diff --git a/nodes.py b/nodes.py index 8791c97..1abdde8 100644 --- a/nodes.py +++ b/nodes.py @@ -143,6 +143,44 @@ def pymeshlab_remove_floater(mesh: pymeshlab.MeshSet): mesh.apply_filter("compute_selection_transfer_face_to_vertex", inclusive=False) mesh.apply_filter("meshing_remove_selected_vertices_and_faces") return mesh + +def _batched_unsigned_distance(bvh, positions, batch_size=100000, return_uvw=False): + """ + Batch unsigned_distance queries to avoid GPU kernel timeout on large meshes. + When processing high-resolution textures (e.g., 2048x2048 = ~4M pixels) on complex + meshes, a single BVH query can cause GPU watchdog timeout. This function splits + the query into smaller batches. + Args: + bvh: The BVH structure from cumesh + positions: (N, 3) tensor of query positions + batch_size: Maximum number of queries per batch (default 100K, matching + the rasterization chunk size used elsewhere in this file) + return_uvw: Whether to return barycentric coordinates + Returns: + Same as bvh.unsigned_distance() + """ + import torch + N = positions.shape[0] + if N <= batch_size: + return bvh.unsigned_distance(positions, return_uvw=return_uvw) + + distances_list = [] + face_id_list = [] + uvw_list = [] if return_uvw else None + + for i in range(0, N, batch_size): + end = min(i + batch_size, N) + d, f, u = bvh.unsigned_distance(positions[i:end], return_uvw=return_uvw) + distances_list.append(d) + face_id_list.append(f) + if return_uvw: + uvw_list.append(u) + + return ( + torch.cat(distances_list), + torch.cat(face_id_list), + torch.cat(uvw_list) if return_uvw else None + ) class Trellis2LoadModel: @classmethod @@ -239,13 +277,13 @@ class Trellis2MeshWithVoxelGenerator: "pipeline": ("TRELLIS2PIPELINE",), "image": ("IMAGE",), "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), - "pipeline_type": (["512","1024","1024_cascade","1536_cascade"],{"default":"1024_cascade"}), + "pipeline_type": (["512","1024","1024_cascade","1536_cascade","2048_cascade"],{"default":"1024_cascade"}), "sparse_structure_steps": ("INT",{"default":12, "min":1, "max":100},), "shape_steps": ("INT",{"default":12, "min":1, "max":100},), "texture_steps": ("INT",{"default":12, "min":1, "max":100},), "max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}), "max_views": ("INT", {"default": 4, "min": 1, "max": 16}), - "sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":64,"step":8}), + "sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":128,"step":8}), "generate_texture_slat": ("BOOLEAN", {"default":True}), }, } @@ -717,29 +755,29 @@ class Trellis2MeshWithVoxelAdvancedGenerator: "pipeline": ("TRELLIS2PIPELINE",), "image": ("IMAGE",), "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), - "pipeline_type": (["512","1024","1024_cascade","1536_cascade"],{"default":"1024_cascade"}), + "pipeline_type": (["512","1024","1024_cascade","1536_cascade","2048_cascade"],{"default":"1024_cascade"}), "sparse_structure_steps": ("INT",{"default":12, "min":1, "max":100},), - "sparse_structure_guidance_strength": ("FLOAT",{"default":7.5}), - "sparse_structure_guidance_rescale": ("FLOAT",{"default":0.7}), - "sparse_structure_rescale_t": ("FLOAT",{"default":5.0}), + "sparse_structure_guidance_strength": ("FLOAT",{"default":7.50}), + "sparse_structure_guidance_rescale": ("FLOAT",{"default":0.70}), + "sparse_structure_rescale_t": ("FLOAT",{"default":5.00}), "shape_steps": ("INT",{"default":12, "min":1, "max":100},), - "shape_guidance_strength": ("FLOAT",{"default":7.5}), - "shape_guidance_rescale": ("FLOAT",{"default":0.5}), - "shape_rescale_t": ("FLOAT",{"default":3.0}), + "shape_guidance_strength": ("FLOAT",{"default":7.50}), + "shape_guidance_rescale": ("FLOAT",{"default":0.50}), + "shape_rescale_t": ("FLOAT",{"default":3.00}), "texture_steps": ("INT",{"default":12, "min":1, "max":100},), - "texture_guidance_strength": ("FLOAT",{"default":1.0}), - "texture_guidance_rescale": ("FLOAT",{"default":0.0}), - "texture_rescale_t": ("FLOAT",{"default":3.0}), + "texture_guidance_strength": ("FLOAT",{"default":1.00}), + "texture_guidance_rescale": ("FLOAT",{"default":0.00}), + "texture_rescale_t": ("FLOAT",{"default":3.00}), "max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}), "max_views": ("INT", {"default": 4, "min": 1, "max": 16}), - "sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":64,"step":8}), + "sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":128,"step":8}), "generate_texture_slat": ("BOOLEAN", {"default":True}), - "sparse_structure_guidance_interval_start": ("FLOAT",{"default":0.3,"min":0.0,"max":1.0}), - "sparse_structure_guidance_interval_end": ("FLOAT",{"default":1.0,"min":0.0,"max":1.0}), - "shape_guidance_interval_start": ("FLOAT",{"default":0.3,"min":0.0,"max":1.0}), - "shape_guidance_interval_end": ("FLOAT",{"default":1.0,"min":0.0,"max":1.0}), - "texture_guidance_interval_start": ("FLOAT",{"default":0.6,"min":0.0,"max":1.0}), - "texture_guidance_interval_end": ("FLOAT",{"default":0.9,"min":0.0,"max":1.0}), + "sparse_structure_guidance_interval_start": ("FLOAT",{"default":0.30,"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}), + "shape_guidance_interval_start": ("FLOAT",{"default":0.30,"min":0.00,"max":1.00,"step":0.01}), + "shape_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}), + "texture_guidance_interval_start": ("FLOAT",{"default":0.60,"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}), }, } @@ -1208,7 +1246,9 @@ class Trellis2MeshTexturing: "resolution": ([512,1024],{"default":1024}), "texture_size": ("INT",{"default":2048,"min":512,"max":16384}), "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), - "double_side_material": ("BOOLEAN",{"default":True}), + "double_side_material": ("BOOLEAN",{"default":True}), + "texture_guidance_interval_start": ("FLOAT",{"default":0.60,"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}), }, } @@ -1218,10 +1258,13 @@ class Trellis2MeshTexturing: CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True - def process(self, pipeline, image, trimesh, seed, texture_steps, texture_guidance_strength, texture_guidance_rescale, texture_rescale_t, resolution, texture_size, texture_alpha_mode, double_side_material): + def process(self, pipeline, image, trimesh, seed, texture_steps, texture_guidance_strength, texture_guidance_rescale, texture_rescale_t, resolution, texture_size, texture_alpha_mode, double_side_material, texture_guidance_interval_start, texture_guidance_interval_end): #image = tensor2pil_v2(image) image = tensor2pil(image) - tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"rescale_t":texture_rescale_t} + + texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] + + 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} textured_mesh, baseColorTexture_np, metallicRoughnessTexture_np = pipeline.texture_mesh(mesh=trimesh, image=image, @@ -1296,7 +1339,7 @@ class Trellis2PreProcessImage: if not np.all(alpha == 255): has_alpha = True max_size = max(input.size) - scale = min(1, 1024 / max_size) + scale = min(1, 2048 / max_size) if scale < 1: input = input.resize((int(input.width * scale), int(input.height * scale)), Image.Resampling.LANCZOS) # if has_alpha: @@ -1334,16 +1377,20 @@ class Trellis2MeshRefiner: "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), "resolution": ([512,1024,1536],{"default":1024}), "shape_steps": ("INT",{"default":12, "min":1, "max":100},), - "shape_guidance_strength": ("FLOAT",{"default":7.5}), - "shape_guidance_rescale": ("FLOAT",{"default":0.5}), - "shape_rescale_t": ("FLOAT",{"default":3.0}), + "shape_guidance_strength": ("FLOAT",{"default":7.50}), + "shape_guidance_rescale": ("FLOAT",{"default":0.50}), + "shape_rescale_t": ("FLOAT",{"default":3.00}), "texture_steps": ("INT",{"default":12, "min":1, "max":100},), - "texture_guidance_strength": ("FLOAT",{"default":1.0}), - "texture_guidance_rescale": ("FLOAT",{"default":0.0}), - "texture_rescale_t": ("FLOAT",{"default":3.0}), + "texture_guidance_strength": ("FLOAT",{"default":1.00}), + "texture_guidance_rescale": ("FLOAT",{"default":0.00}), + "texture_rescale_t": ("FLOAT",{"default":3.00}), "max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}), "generate_texture_slat": ("BOOLEAN", {"default":True}), "downsampling":([16,32,64],{"default":16}), + "shape_guidance_interval_start": ("FLOAT",{"default":0.30,"min":0.00,"max":1.00,"step":0.01}), + "shape_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}), + "texture_guidance_interval_start": ("FLOAT",{"default":0.60,"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}), }, } @@ -1364,12 +1411,19 @@ class Trellis2MeshRefiner: texture_rescale_t, max_num_tokens, generate_texture_slat, - downsampling): + downsampling, + shape_guidance_interval_start, + shape_guidance_interval_end, + texture_guidance_interval_start, + texture_guidance_interval_end): image = tensor2pil(image) - shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"rescale_t":shape_rescale_t} - tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"rescale_t":texture_rescale_t} + shape_guidance_interval = [shape_guidance_interval_start,shape_guidance_interval_end] + texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] + + shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":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} mesh = pipeline.refine_mesh(mesh = trimesh, image=image, seed=seed, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, resolution = resolution, max_num_tokens = max_num_tokens, generate_texture_slat=generate_texture_slat, downsampling=downsampling)[0] diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index e88dbb3..3c44acb 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -80,7 +80,7 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.tex_slat_normalization = tex_slat_normalization self.image_cond_model = image_cond_model self.rembg_model = rembg_model - self.low_vram = low_vram + self._low_vram = low_vram self.default_pipeline_type = default_pipeline_type self.pbr_attr_layout = { 'base_color': slice(0, 3), @@ -90,6 +90,19 @@ class Trellis2ImageTo3DPipeline(Pipeline): } self._device = 'cpu' + @property + def low_vram(self) -> bool: + return self._low_vram + + @low_vram.setter + def low_vram(self, value: bool): + self._low_vram = value + for m in self.models.values(): + if hasattr(m, 'low_vram'): + m.low_vram = value + if hasattr(self, 'image_cond_model') and hasattr(self.image_cond_model, 'low_vram'): + self.image_cond_model.low_vram = value + def _cond_to(self, cond: dict, device: torch.device) -> dict: # Move only tensors; keep other items unchanged return {k: (v.to(device) if torch.is_tensor(v) else v) for k, v in cond.items()} @@ -162,6 +175,8 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.models['sparse_structure_decoder'] = models.from_pretrained(self._pretrained_args['models']['sparse_structure_decoder']) self.models['sparse_structure_decoder'].eval() self.models['sparse_structure_decoder'].to(self._device) + if hasattr(self.models['sparse_structure_decoder'], 'low_vram'): + self.models['sparse_structure_decoder'].low_vram = self.low_vram def unload_sparse_structure_model(self): if self.models['sparse_structure_flow_model']: @@ -218,6 +233,8 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.models['tex_slat_decoder'] = models.from_pretrained(f"{self.path}/{self._pretrained_args['models']['tex_slat_decoder']}") self.models['tex_slat_decoder'].eval() self.models['tex_slat_decoder'].to(self._device) + if hasattr(self.models['tex_slat_decoder'], 'low_vram'): + self.models['tex_slat_decoder'].low_vram = self.low_vram def unload_tex_slat_decoder(self): if self.models['tex_slat_decoder'] is not None: @@ -231,6 +248,8 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.models['shape_slat_decoder'] = models.from_pretrained(f"{self.path}/{self._pretrained_args['models']['shape_slat_decoder']}") self.models['shape_slat_decoder'].eval() self.models['shape_slat_decoder'].to(self._device) + if hasattr(self.models['shape_slat_decoder'], 'low_vram'): + self.models['shape_slat_decoder'].low_vram = self.low_vram def unload_shape_slat_decoder(self): if self.models['shape_slat_decoder'] is not None: @@ -270,6 +289,8 @@ class Trellis2ImageTo3DPipeline(Pipeline): self.models['shape_slat_encoder'] = models.from_pretrained(f"{self.path}/ckpts/shape_enc_next_dc_f16c32_fp16") self.models['shape_slat_encoder'].eval() self.models['shape_slat_encoder'].to(self._device) + if hasattr(self.models['shape_slat_encoder'], 'low_vram'): + self.models['shape_slat_encoder'].low_vram = self.low_vram def unload_shape_slat_encoder(self): if self.models['shape_slat_encoder'] is not None: @@ -592,11 +613,17 @@ class Trellis2ImageTo3DPipeline(Pipeline): ], dim=1) coords = quant_coords.unique(dim=0) num_tokens = coords.shape[0] - if num_tokens < max_num_tokens or hr_resolution == 1024: + 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}.") break hr_resolution -= 128 + if hr_resolution < 1024 and resolution >= 1024: + hr_resolution = 1024 + break + if hr_resolution < 512: + hr_resolution = 512 + break coords_dev = coords.to(self.device) # Sample structured latent @@ -866,9 +893,9 @@ class Trellis2ImageTo3DPipeline(Pipeline): # 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 + # 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 not self.keep_models_loaded: self.unload_image_cond_model() @@ -957,6 +984,50 @@ class Trellis2ImageTo3DPipeline(Pipeline): if not self.keep_models_loaded: self.unload_tex_slat_flow_model_1024() + elif pipeline_type == '2048_cascade': + self.load_shape_slat_flow_model_512() + self.load_shape_slat_flow_model_1024() + shape_slat, res = self.sample_shape_slat_cascade( + cond_512, cond_1024, + self.models['shape_slat_flow_model_512'], self.models['shape_slat_flow_model_1024'], + 512, 2048, + coords, shape_slat_sampler_params, + max_num_tokens + ) + 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( + cond_1024, self.models['tex_slat_flow_model_1024'], + shape_slat, tex_slat_sampler_params + ) + if not self.keep_models_loaded: + self.unload_tex_slat_flow_model_1024() + elif pipeline_type == '4096_cascade': + self.load_shape_slat_flow_model_512() + self.load_shape_slat_flow_model_1024() + shape_slat, res = self.sample_shape_slat_cascade( + cond_512, cond_1024, + self.models['shape_slat_flow_model_512'], self.models['shape_slat_flow_model_1024'], + 512, 4096, + coords, shape_slat_sampler_params, + max_num_tokens + ) + 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( + cond_1024, self.models['tex_slat_flow_model_1024'], + shape_slat, tex_slat_sampler_params + ) + 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() @@ -1282,11 +1353,17 @@ class Trellis2ImageTo3DPipeline(Pipeline): ], dim=1) coords = quant_coords.unique(dim=0) num_tokens = coords.shape[0] - if num_tokens < max_num_tokens or hr_resolution == 1024: + 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}.") break hr_resolution -= 128 + if hr_resolution < 1024 and resolution >= 1024: + hr_resolution = 1024 + break + if hr_resolution < 512: + hr_resolution = 512 + break coords_dev = coords.to(self.device) # Sample structured latent