Added the node "Mesh with Voxel Cascade Generator"
This commit is contained in:
@@ -14,6 +14,7 @@
|
||||
|
||||
| Date | Description |
|
||||
| --- | --- |
|
||||
| **2026-03-07** | Added "Heun" sampler<br>Added the node "Mesh with Voxel Cascade Generator" |
|
||||
| **2026-03-05** | Added "RK4" and "RK5" samplers<br>Processing is much slower, so reduce the number of steps |
|
||||
| **2026-03-04** | Sparse Structure Resolution supported up to 128<br>Experimental for "cascade" pipelines only<br>Can increase the details |
|
||||
| **2026-02-27** | Added the Wheels for Windows Python 3.13, Torch 2.10.0, CUDA 13.1 |
|
||||
|
||||
@@ -3123,7 +3123,204 @@ class Trellis2LaplacianSmoothingWithOpen3d:
|
||||
mesh_copy.vertices = torch.from_numpy(new_vertices).float().to(mesh_copy.device)
|
||||
mesh_copy.faces = torch.from_numpy(new_faces).int().to(mesh_copy.device)
|
||||
|
||||
return (mesh_copy,)
|
||||
return (mesh_copy,)
|
||||
|
||||
class Trellis2UnWrapTrimesh:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"trimesh": ("TRIMESH",),
|
||||
"mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60.0,"min":0.0,"max":359.9}),
|
||||
"mesh_cluster_refine_iterations": ("INT",{"default":0}),
|
||||
"mesh_cluster_global_iterations": ("INT",{"default":1}),
|
||||
"mesh_cluster_smooth_strength": ("INT",{"default":1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRIMESH", )
|
||||
RETURN_NAMES = ("trimesh", )
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, trimesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength):
|
||||
mesh_cluster_threshold_cone_half_angle_rad = np.radians(mesh_cluster_threshold_cone_half_angle_rad)
|
||||
|
||||
mesh_copy = trimesh.copy()
|
||||
|
||||
vertices = torch.from_numpy(mesh_copy.vertices).float().cuda()
|
||||
faces = torch.from_numpy(mesh_copy.faces).int().cuda()
|
||||
|
||||
cumesh = CuMesh.CuMesh()
|
||||
cumesh.init(vertices, faces)
|
||||
|
||||
out_vertices, out_faces, out_uvs = cumesh.uv_unwrap(
|
||||
compute_charts_kwargs={
|
||||
"threshold_cone_half_angle_rad": mesh_cluster_threshold_cone_half_angle_rad,
|
||||
"refine_iterations": mesh_cluster_refine_iterations,
|
||||
"global_iterations": mesh_cluster_global_iterations,
|
||||
"smooth_strength": mesh_cluster_smooth_strength,
|
||||
},
|
||||
return_vmaps=False,
|
||||
verbose=True,
|
||||
)
|
||||
|
||||
del cumesh
|
||||
|
||||
mesh_copy.vertices = out_vertices.cpu().numpy()
|
||||
mesh_copy.faces = out_faces.cpu().numpy()
|
||||
mesh_copy.visual.uv = out_uvs.cpu().numpy()
|
||||
|
||||
return (mesh_copy,)
|
||||
|
||||
class Trellis2MeshWithVoxelCascadeGenerator:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"pipeline": ("TRELLIS2PIPELINE",),
|
||||
"image": ("IMAGE",),
|
||||
"seed": ("INT", {"default": 12345, "min": 0, "max": 0x7fffffff}),
|
||||
"pipeline_type": (["1024_cascade","1536_cascade"],{"default":"1024_cascade"}),
|
||||
"sparse_structure_steps": ("INT",{"default":12, "min":1, "max":100},),
|
||||
"sparse_structure_guidance_strength": ("FLOAT",{"default":6.50,"min":0.00,"max":99.99,"step":0.01}),
|
||||
"sparse_structure_guidance_rescale": ("FLOAT",{"default":0.05,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"sparse_structure_rescale_t": ("FLOAT",{"default":4.00,"min":0.00,"max":9.99,"step":0.01}),
|
||||
"sparse_structure_sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}),
|
||||
"sparse_structure_resolution": ("INT", {"default":32,"min":32,"max":128,"step":4}),
|
||||
"sparse_structure_guidance_interval_start": ("FLOAT",{"default":0.10,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"sparse_structure_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"low_res_shape_steps": ("INT",{"default":12, "min":1, "max":100},),
|
||||
"low_res_shape_guidance_strength": ("FLOAT",{"default":6.50,"min":0.00,"max":99.99,"step":0.01}),
|
||||
"low_res_shape_guidance_rescale": ("FLOAT",{"default":0.05,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"low_res_shape_rescale_t": ("FLOAT",{"default":4.00,"min":0.00,"max":9.99,"step":0.01}),
|
||||
"low_res_shape_sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}),
|
||||
"low_res_shape_guidance_interval_start": ("FLOAT",{"default":0.10,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"low_res_shape_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"high_res_shape_steps": ("INT",{"default":12, "min":1, "max":100},),
|
||||
"high_res_shape_guidance_strength": ("FLOAT",{"default":6.50,"min":0.00,"max":99.99,"step":0.01}),
|
||||
"high_res_shape_guidance_rescale": ("FLOAT",{"default":0.05,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"high_res_shape_rescale_t": ("FLOAT",{"default":4.00,"min":0.00,"max":9.99,"step":0.01}),
|
||||
"high_res_shape_sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}),
|
||||
"high_res_shape_guidance_interval_start": ("FLOAT",{"default":0.10,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"high_res_shape_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"generate_texture_slat": ("BOOLEAN", {"default":True}),
|
||||
"texture_steps": ("INT",{"default":12, "min":1, "max":100},),
|
||||
"texture_guidance_strength": ("FLOAT",{"default":6.50,"min":0.00,"max":99.99,"step":0.01}),
|
||||
"texture_guidance_rescale": ("FLOAT",{"default":0.05,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"texture_rescale_t": ("FLOAT",{"default":4.00,"min":0.00,"max":9.99,"step":0.01}),
|
||||
"texture_sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}),
|
||||
"texture_guidance_interval_start": ("FLOAT",{"default":0.00,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"texture_guidance_interval_end": ("FLOAT",{"default":0.90,"min":0.00,"max":1.00,"step":0.01}),
|
||||
"max_num_tokens": ("INT",{"default":999999,"min":0,"max":999999}),
|
||||
"use_tiled_decoder": ("BOOLEAN", {"default":True}),
|
||||
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MESHWITHVOXEL","BVH", )
|
||||
RETURN_NAMES = ("mesh", "bvh", )
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, pipeline, image, seed, pipeline_type,
|
||||
# sparse
|
||||
sparse_structure_steps,
|
||||
sparse_structure_guidance_strength,
|
||||
sparse_structure_guidance_rescale,
|
||||
sparse_structure_rescale_t,
|
||||
sparse_structure_sampler,
|
||||
sparse_structure_resolution,
|
||||
sparse_structure_guidance_interval_start,
|
||||
sparse_structure_guidance_interval_end,
|
||||
# low res shape
|
||||
low_res_shape_steps,
|
||||
low_res_shape_guidance_strength,
|
||||
low_res_shape_guidance_rescale,
|
||||
low_res_shape_rescale_t,
|
||||
low_res_shape_sampler,
|
||||
low_res_shape_guidance_interval_start,
|
||||
low_res_shape_guidance_interval_end,
|
||||
# high res shape
|
||||
high_res_shape_steps,
|
||||
high_res_shape_guidance_strength,
|
||||
high_res_shape_guidance_rescale,
|
||||
high_res_shape_rescale_t,
|
||||
high_res_shape_sampler,
|
||||
high_res_shape_guidance_interval_start,
|
||||
high_res_shape_guidance_interval_end,
|
||||
# texture,
|
||||
generate_texture_slat,
|
||||
texture_steps,
|
||||
texture_guidance_strength,
|
||||
texture_guidance_rescale,
|
||||
texture_rescale_t,
|
||||
texture_sampler,
|
||||
texture_guidance_interval_start,
|
||||
texture_guidance_interval_end,
|
||||
# others
|
||||
max_num_tokens,
|
||||
use_tiled_decoder,
|
||||
max_views
|
||||
):
|
||||
|
||||
reset_cuda()
|
||||
|
||||
images = tensor_batch_to_pil_list(image, max_views=max_views)
|
||||
image_in = images[0] if len(images) == 1 else images
|
||||
|
||||
sparse_structure_guidance_interval = [sparse_structure_guidance_interval_start,sparse_structure_guidance_interval_end]
|
||||
low_res_shape_guidance_interval = [low_res_shape_guidance_interval_start, low_res_shape_guidance_interval_end]
|
||||
high_res_shape_guidance_interval = [high_res_shape_guidance_interval_start, high_res_shape_guidance_interval_end]
|
||||
texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end]
|
||||
|
||||
sparse_structure_sampler_params = {"steps":sparse_structure_steps,"guidance_strength":sparse_structure_guidance_strength,"guidance_rescale":sparse_structure_guidance_rescale,"guidance_interval":sparse_structure_guidance_interval,"rescale_t":sparse_structure_rescale_t}
|
||||
low_res_shape_slat_sampler_params = {"steps":low_res_shape_steps,"guidance_strength":low_res_shape_guidance_strength,"guidance_rescale":low_res_shape_guidance_rescale,"guidance_interval":low_res_shape_guidance_interval,"rescale_t":low_res_shape_rescale_t}
|
||||
high_res_shape_slat_sampler_params = {"steps":high_res_shape_steps,"guidance_strength":high_res_shape_guidance_strength,"guidance_rescale":high_res_shape_guidance_rescale,"guidance_interval":high_res_shape_guidance_interval,"rescale_t":high_res_shape_rescale_t}
|
||||
tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"guidance_interval":texture_guidance_interval,"rescale_t":texture_rescale_t}
|
||||
|
||||
if generate_texture_slat:
|
||||
num_steps = 5
|
||||
else:
|
||||
num_steps = 4
|
||||
|
||||
pbar = ProgressBar(num_steps)
|
||||
|
||||
mesh = pipeline.run_cascade(image=image_in,
|
||||
seed=seed,
|
||||
pipeline_type=pipeline_type,
|
||||
sparse_structure_sampler_params = sparse_structure_sampler_params,
|
||||
low_res_shape_slat_sampler_params = low_res_shape_slat_sampler_params,
|
||||
high_res_shape_slat_sampler_params = high_res_shape_slat_sampler_params,
|
||||
tex_slat_sampler_params = tex_slat_sampler_params,
|
||||
max_num_tokens = max_num_tokens,
|
||||
sparse_structure_resolution = sparse_structure_resolution,
|
||||
max_views = max_views,
|
||||
generate_texture_slat=generate_texture_slat,
|
||||
use_tiled=use_tiled_decoder,
|
||||
pbar=pbar,
|
||||
sparse_structure_sampler = sparse_structure_sampler,
|
||||
low_res_shape_sampler = low_res_shape_sampler,
|
||||
high_res_shape_sampler = high_res_shape_sampler,
|
||||
tex_sampler = texture_sampler
|
||||
)[0]
|
||||
|
||||
vertices = mesh.vertices.cuda()
|
||||
faces = mesh.faces.cuda()
|
||||
|
||||
if generate_texture_slat:
|
||||
# Build BVH for the current mesh to guide remeshing
|
||||
print("Building BVH for current mesh...")
|
||||
bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
|
||||
bvh.vertices = vertices.detach().clone()
|
||||
bvh.faces = faces.detach().clone()
|
||||
else:
|
||||
print("Not building BVH : only used for texturing")
|
||||
bvh = None
|
||||
|
||||
return (mesh,bvh,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -3161,6 +3358,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"Trellis2StringSelector": Trellis2StringSelector,
|
||||
"Trellis2FillHolesWithCuMesh": Trellis2FillHolesWithCuMesh,
|
||||
"Trellis2LaplacianSmoothingWithOpen3d": Trellis2LaplacianSmoothingWithOpen3d,
|
||||
"Trellis2UnWrapTrimesh": Trellis2UnWrapTrimesh,
|
||||
"Trellis2MeshWithVoxelCascadeGenerator": Trellis2MeshWithVoxelCascadeGenerator,
|
||||
}
|
||||
|
||||
|
||||
@@ -3199,4 +3398,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Trellis2StringSelector": "Trellis2 - String Selector",
|
||||
"Trellis2FillHolesWithCuMesh": "Trellis2 - Fill Holes with CuMesh",
|
||||
"Trellis2LaplacianSmoothingWithOpen3d": "Trellis2 - Laplacian Smoothing (using open3d)",
|
||||
"Trellis2UnWrapTrimesh": "Trellis2 - UnWrap Trimesh",
|
||||
"Trellis2MeshWithVoxelCascadeGenerator": "Trellis2 - Mesh With Voxel Cascade Generator"
|
||||
}
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "trellis2"
|
||||
description = "ComfyUI Wrapper for Microsoft Trellis.2 - Native and Compact Structured Latents for 3D Generation"
|
||||
version = "1.0.14"
|
||||
version = "1.0.15"
|
||||
license = {file = "LICENSE"}
|
||||
# classifiers = [
|
||||
# # For OS-independent nodes (works on all operating systems)
|
||||
|
||||
@@ -107,12 +107,12 @@ class DinoV3FeatureExtractor:
|
||||
elif isinstance(image, list):
|
||||
assert all(isinstance(i, Image.Image) for i in image), "Image list should be list of PIL images"
|
||||
# We resize the images only if they are bigger than self.image_size
|
||||
image = [
|
||||
i.resize((self.image_size, self.image_size), Image.LANCZOS)
|
||||
if max(i.size) > self.image_size else i
|
||||
for i in image
|
||||
]
|
||||
#image = [i.resize((self.image_size, self.image_size), Image.LANCZOS) for i in image]
|
||||
# image = [
|
||||
# i.resize((self.image_size, self.image_size), Image.LANCZOS)
|
||||
# if max(i.size) > self.image_size else i
|
||||
# for i in image
|
||||
# ]
|
||||
image = [i.resize((self.image_size, self.image_size), Image.LANCZOS) for i in image]
|
||||
image = [np.array(i.convert('RGB')).astype(np.float32) / 255 for i in image]
|
||||
image = [torch.from_numpy(i).permute(2, 0, 1).float() for i in image]
|
||||
image = torch.stack(image).cuda()
|
||||
|
||||
@@ -118,7 +118,6 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
|
||||
args = self._pretrained_args
|
||||
self.sparse_structure_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['sparse_structure_sampler']['args'])
|
||||
# Re-instantiate the samplers using the new prefix but keeping original args
|
||||
self.shape_slat_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
|
||||
self.tex_slat_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['tex_slat_sampler']['args'])
|
||||
|
||||
@@ -1181,6 +1180,354 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
return out_mesh, (shape_slat, None, res)
|
||||
else:
|
||||
return out_mesh
|
||||
|
||||
def GetSamplerName(self, sampler):
|
||||
if sampler == 'euler':
|
||||
return 'Euler'
|
||||
elif sampler == 'rk4':
|
||||
return 'RK4'
|
||||
elif sampler == 'rk5':
|
||||
return 'RK5'
|
||||
elif sampler == 'heun':
|
||||
return 'Heun'
|
||||
else:
|
||||
return 'Euler'
|
||||
|
||||
def sample_shape_slat_cascade_advanced(
|
||||
self,
|
||||
lr_cond: dict,
|
||||
cond: dict,
|
||||
flow_model_lr,
|
||||
flow_model,
|
||||
lr_resolution: int,
|
||||
resolution: int,
|
||||
coords: torch.Tensor,
|
||||
low_res_sampler_params: dict = {},
|
||||
high_res_sampler_params: dict = {},
|
||||
max_num_tokens: int = 999999,
|
||||
sparse_structure_resolution: int = 32,
|
||||
low_res_sampler_name: str = 'euler',
|
||||
high_res_sampler_name: str = 'euler',
|
||||
) -> SparseTensor:
|
||||
"""
|
||||
Sample structured latent with the given conditioning.
|
||||
|
||||
Args:
|
||||
cond (dict): The conditioning information.
|
||||
coords (torch.Tensor): The coordinates of the sparse structure.
|
||||
sampler_params (dict): Additional parameters for the sampler.
|
||||
"""
|
||||
# LR
|
||||
|
||||
if self.low_vram:
|
||||
lr_cond = self._cond_to(lr_cond, self.device)
|
||||
cond = self._cond_to(cond, self.device)
|
||||
|
||||
coords_dev = coords.to(self.device)
|
||||
# Sample structured latent
|
||||
noise = SparseTensor(
|
||||
feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device),
|
||||
coords=coords_dev,
|
||||
)
|
||||
sampler_params = {**self.shape_slat_sampler_params, **low_res_sampler_params}
|
||||
if self.low_vram:
|
||||
flow_model_lr.to(self.device)
|
||||
|
||||
args = self._pretrained_args
|
||||
sparse_sampler_prefix = self.GetSamplerName(low_res_sampler_name)
|
||||
self.shape_slat_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
|
||||
|
||||
slat = self.shape_slat_sampler.sample(
|
||||
flow_model_lr,
|
||||
noise,
|
||||
**lr_cond,
|
||||
**sampler_params,
|
||||
verbose=True,
|
||||
tqdm_desc="Sampling shape SLat (LR)",
|
||||
).samples
|
||||
if self.low_vram:
|
||||
flow_model_lr.cpu()
|
||||
self._cleanup_cuda()
|
||||
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device)
|
||||
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device)
|
||||
slat = slat * std + mean
|
||||
|
||||
del coords_dev
|
||||
if self.low_vram:
|
||||
lr_cond = self._cond_cpu(lr_cond)
|
||||
self._cleanup_cuda()
|
||||
|
||||
# Upsample
|
||||
self.load_shape_slat_decoder()
|
||||
if self.low_vram:
|
||||
self.models['shape_slat_decoder'].to(self.device)
|
||||
self.models['shape_slat_decoder'].low_vram = True
|
||||
hr_coords = self.models['shape_slat_decoder'].upsample(slat, upsample_times=4)
|
||||
if self.low_vram:
|
||||
self.models['shape_slat_decoder'].cpu()
|
||||
self.models['shape_slat_decoder'].low_vram = False
|
||||
hr_resolution = resolution
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_decoder()
|
||||
|
||||
ratio = (sparse_structure_resolution / 32)
|
||||
|
||||
while True:
|
||||
quant_coords = torch.cat([
|
||||
hr_coords[:, :1],
|
||||
((hr_coords[:, 1:] + 0.5) / (lr_resolution * ratio) * (hr_resolution // 16)).int(),
|
||||
], dim=1)
|
||||
coords = quant_coords.unique(dim=0)
|
||||
num_tokens = coords.shape[0]
|
||||
if num_tokens < max_num_tokens:
|
||||
if hr_resolution != resolution:
|
||||
print(f"Due to the limited number of tokens, the resolution is reduced to {hr_resolution}.")
|
||||
print(f"Num Tokens: {num_tokens}")
|
||||
break
|
||||
hr_resolution -= 128
|
||||
if hr_resolution < 1024 and resolution >= 1024:
|
||||
print(f"Num Tokens: {num_tokens}")
|
||||
hr_resolution = 1024
|
||||
break
|
||||
if hr_resolution < 512:
|
||||
print(f"Num Tokens: {num_tokens}")
|
||||
hr_resolution = 512
|
||||
break
|
||||
|
||||
coords_dev = coords.to(self.device)
|
||||
# Sample structured latent
|
||||
noise = SparseTensor(
|
||||
feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device),
|
||||
coords=coords_dev,
|
||||
)
|
||||
|
||||
sampler_params = {**self.shape_slat_sampler_params, **high_res_sampler_params}
|
||||
|
||||
sparse_sampler_prefix = self.GetSamplerName(high_res_sampler_name)
|
||||
self.shape_slat_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
|
||||
|
||||
if self.low_vram:
|
||||
flow_model.to(self.device)
|
||||
slat = self.shape_slat_sampler.sample(
|
||||
flow_model,
|
||||
noise,
|
||||
**cond,
|
||||
**sampler_params,
|
||||
verbose=True,
|
||||
tqdm_desc="Sampling shape SLat (HR)",
|
||||
).samples
|
||||
if self.low_vram:
|
||||
flow_model.cpu()
|
||||
self._cleanup_cuda()
|
||||
|
||||
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device)
|
||||
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device)
|
||||
slat = slat * std + mean
|
||||
|
||||
del coords_dev
|
||||
if self.low_vram:
|
||||
cond = self._cond_cpu(cond)
|
||||
self._cleanup_cuda()
|
||||
|
||||
return slat, hr_resolution
|
||||
|
||||
def sample_tex_slat_advanced(
|
||||
self,
|
||||
cond: dict,
|
||||
flow_model,
|
||||
shape_slat: SparseTensor,
|
||||
sampler_params: dict = {},
|
||||
sampler_name: str = 'euler',
|
||||
) -> SparseTensor:
|
||||
"""
|
||||
Sample structured latent with the given conditioning.
|
||||
|
||||
Args:
|
||||
cond (dict): The conditioning information.
|
||||
shape_slat (SparseTensor): The structured latent for shape
|
||||
sampler_params (dict): Additional parameters for the sampler.
|
||||
"""
|
||||
if self.low_vram:
|
||||
cond = self._cond_to(cond, self.device)
|
||||
# Sample structured latent
|
||||
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(shape_slat.device)
|
||||
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(shape_slat.device)
|
||||
shape_slat = (shape_slat - mean) / std
|
||||
|
||||
in_channels = flow_model.in_channels if isinstance(flow_model, nn.Module) else flow_model[0].in_channels
|
||||
noise = shape_slat.replace(feats=torch.randn(shape_slat.coords.shape[0], in_channels - shape_slat.feats.shape[1]).to(self.device))
|
||||
sampler_params = {**self.tex_slat_sampler_params, **sampler_params}
|
||||
|
||||
args = self._pretrained_args
|
||||
sparse_sampler_prefix = self.GetSamplerName(sampler_name)
|
||||
self.tex_slat_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['tex_slat_sampler']['args'])
|
||||
|
||||
if self.low_vram:
|
||||
flow_model.to(self.device)
|
||||
slat = self.tex_slat_sampler.sample(
|
||||
flow_model,
|
||||
noise,
|
||||
concat_cond=shape_slat,
|
||||
**cond,
|
||||
**sampler_params,
|
||||
verbose=True,
|
||||
tqdm_desc="Sampling texture SLat",
|
||||
).samples
|
||||
if self.low_vram:
|
||||
flow_model.cpu()
|
||||
self._cleanup_cuda()
|
||||
|
||||
std = torch.tensor(self.tex_slat_normalization['std'])[None].to(slat.device)
|
||||
mean = torch.tensor(self.tex_slat_normalization['mean'])[None].to(slat.device)
|
||||
slat = slat * std + mean
|
||||
|
||||
if self.low_vram:
|
||||
cond = self._cond_cpu(cond)
|
||||
self._cleanup_cuda()
|
||||
return slat
|
||||
|
||||
@torch.no_grad()
|
||||
def run_cascade(
|
||||
self,
|
||||
image: Image.Image,
|
||||
num_samples: int = 1,
|
||||
seed: int = 42,
|
||||
sparse_structure_sampler_params: dict = {},
|
||||
low_res_shape_slat_sampler_params: dict = {},
|
||||
high_res_shape_slat_sampler_params: dict = {},
|
||||
tex_slat_sampler_params: dict = {},
|
||||
pipeline_type: str = '1024_cascade',
|
||||
max_num_tokens: int = 999999,
|
||||
sparse_structure_resolution: int = 32,
|
||||
generate_texture_slat = True,
|
||||
use_tiled: bool = True,
|
||||
pbar = None,
|
||||
sparse_structure_sampler = 'euler',
|
||||
low_res_shape_sampler = 'euler',
|
||||
high_res_shape_sampler = 'euler',
|
||||
tex_sampler = 'euler',
|
||||
max_views: int = 4
|
||||
) -> List[MeshWithVoxel]:
|
||||
|
||||
if isinstance(image, (list, tuple)):
|
||||
images = list(image)
|
||||
else:
|
||||
images = [image]
|
||||
|
||||
seed_all(seed)
|
||||
|
||||
# Get Image Cond
|
||||
self.load_image_cond_model()
|
||||
# Multi-view conditioning happens inside get_cond()
|
||||
cond_512 = self.get_cond(images, 512, max_views = max_views)
|
||||
cond_1024 = self.get_cond(images, 1024, max_views = max_views) if pipeline_type != '512' else None
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_image_cond_model()
|
||||
|
||||
args = self._pretrained_args
|
||||
|
||||
#
|
||||
#self.shape_slat_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
|
||||
#self.tex_slat_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['tex_slat_sampler']['args'])
|
||||
|
||||
# Sampling Sparse Structure
|
||||
sparse_sampler_prefix = self.GetSamplerName(sparse_structure_sampler)
|
||||
self.sparse_structure_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['sparse_structure_sampler']['args'])
|
||||
self.load_sparse_structure_model()
|
||||
coords = self.sample_sparse_structure(
|
||||
cond_512, sparse_structure_resolution,
|
||||
num_samples, sparse_structure_sampler_params
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_sparse_structure_model()
|
||||
|
||||
# Sampling Shape
|
||||
if pipeline_type == '1024_cascade':
|
||||
self.load_shape_slat_flow_model_512()
|
||||
self.load_shape_slat_flow_model_1024()
|
||||
shape_slat, res = self.sample_shape_slat_cascade_advanced(
|
||||
cond_512, cond_1024,
|
||||
self.models['shape_slat_flow_model_512'], self.models['shape_slat_flow_model_1024'],
|
||||
512, 1024,
|
||||
coords, low_res_shape_slat_sampler_params, high_res_shape_slat_sampler_params,
|
||||
max_num_tokens,
|
||||
sparse_structure_resolution,
|
||||
low_res_shape_sampler, high_res_shape_sampler
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat_advanced(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params, tex_sampler
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
elif pipeline_type == '1536_cascade':
|
||||
self.load_shape_slat_flow_model_512()
|
||||
self.load_shape_slat_flow_model_1024()
|
||||
shape_slat, res = self.sample_shape_slat_cascade_advanced(
|
||||
cond_512, cond_1024,
|
||||
self.models['shape_slat_flow_model_512'], self.models['shape_slat_flow_model_1024'],
|
||||
512, 1536,
|
||||
coords, low_res_shape_slat_sampler_params, high_res_shape_slat_sampler_params,
|
||||
max_num_tokens,
|
||||
sparse_structure_resolution,
|
||||
low_res_shape_sampler, high_res_shape_sampler
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat_advanced(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params, tex_sampler
|
||||
)
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
if generate_texture_slat:
|
||||
out_mesh = self.decode_latent(shape_slat, tex_slat, res, use_tiled=use_tiled)
|
||||
else:
|
||||
out_mesh = self.decode_latent(shape_slat, None, res, use_tiled=use_tiled)
|
||||
torch.cuda.empty_cache()
|
||||
pbar.update(1)
|
||||
|
||||
return out_mesh
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def run_multiview(
|
||||
|
||||
Reference in New Issue
Block a user