Added 2048_cascade + increased precision in some nodes + added missing guidance_interval in some nodes

Preparing the code to process by chunk
This commit is contained in:
Bruno Fargnoli
2026-01-02 14:44:39 +01:00
parent 868bde0889
commit 0f5c7c305a
2 changed files with 169 additions and 38 deletions
+86 -32
View File
@@ -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]
+83 -6
View File
@@ -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