From c1ed55def8dcd8c07529d230b4af73b882d14d99 Mon Sep 17 00:00:00 2001 From: Bruno Fargnoli Date: Sat, 7 Mar 2026 17:59:59 +0100 Subject: [PATCH] Added the node "Mesh with Voxel Cascade Generator" --- README.md | 1 + nodes.py | 203 +++++++++++- pyproject.toml | 2 +- trellis2/modules/image_feature_extractor.py | 12 +- trellis2/pipelines/trellis2_image_to_3d.py | 349 +++++++++++++++++++- 5 files changed, 558 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index 2858946..4996062 100644 --- a/README.md +++ b/README.md @@ -14,6 +14,7 @@ | Date | Description | | --- | --- | +| **2026-03-07** | Added "Heun" sampler
Added the node "Mesh with Voxel Cascade Generator" | | **2026-03-05** | Added "RK4" and "RK5" samplers
Processing is much slower, so reduce the number of steps | | **2026-03-04** | Sparse Structure Resolution supported up to 128
Experimental for "cascade" pipelines only
Can increase the details | | **2026-02-27** | Added the Wheels for Windows Python 3.13, Torch 2.10.0, CUDA 13.1 | diff --git a/nodes.py b/nodes.py index 7204ffb..81e0c65 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } diff --git a/pyproject.toml b/pyproject.toml index b3dabbe..28527f5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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) diff --git a/trellis2/modules/image_feature_extractor.py b/trellis2/modules/image_feature_extractor.py index bf152c8..6f6a79c 100644 --- a/trellis2/modules/image_feature_extractor.py +++ b/trellis2/modules/image_feature_extractor.py @@ -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() diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 559b198..3f69059 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -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(