Can bypass "texture slat generation" for "Mesh Only" + Updated CuMesh package

This commit is contained in:
Bruno Fargnoli
2025-12-26 17:03:29 +01:00
parent 5db44f5a55
commit f9534d8b21
10 changed files with 87 additions and 52 deletions
+10 -7
View File
@@ -208,6 +208,7 @@ class Trellis2MeshWithVoxelGenerator:
"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}),
"generate_texture_slat": ("BOOLEAN", {"default":True}),
},
}
@@ -217,7 +218,7 @@ class Trellis2MeshWithVoxelGenerator:
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, pipeline, image, seed, pipeline_type, sparse_structure_steps, shape_steps, texture_steps, max_num_tokens, max_views, sparse_structure_resolution):
def process(self, pipeline, image, seed, pipeline_type, sparse_structure_steps, shape_steps, texture_steps, max_num_tokens, max_views, sparse_structure_resolution, generate_texture_slat):
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
@@ -225,7 +226,7 @@ class Trellis2MeshWithVoxelGenerator:
shape_slat_sampler_params = {"steps":shape_steps}
tex_slat_sampler_params = {"steps":texture_steps}
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = 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)[0]
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = 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)[0]
return (mesh,)
@@ -681,6 +682,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
"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}),
"generate_texture_slat": ("BOOLEAN", {"default":True}),
},
}
@@ -704,7 +706,8 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
texture_rescale_t,
max_num_tokens,
max_views,
sparse_structure_resolution):
sparse_structure_resolution,
generate_texture_slat):
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
@@ -713,7 +716,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
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}
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = 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)[0]
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = 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)[0]
return (mesh,)
@@ -1009,7 +1012,7 @@ class Trellis2Remesh:
if isinstance(aabb, (list, tuple)):
aabb = np.array(aabb)
if isinstance(aabb, np.ndarray):
aabb = torch.tensor(aabb, dtype=torch.float32, device=coords.device)
aabb = torch.tensor(aabb, dtype=torch.float32, device='cuda')
# Calculate grid dimensions based on AABB and voxel size
if voxel_size is not None:
@@ -1018,7 +1021,7 @@ class Trellis2Remesh:
if isinstance(voxel_size, (list, tuple)):
voxel_size = np.array(voxel_size)
if isinstance(voxel_size, np.ndarray):
voxel_size = torch.tensor(voxel_size, dtype=torch.float32, device=coords.device)
voxel_size = torch.tensor(voxel_size, dtype=torch.float32, device='cuda')
grid_size = ((aabb[1] - aabb[0]) / voxel_size).round().int()
else:
if isinstance(grid_size, int):
@@ -1026,7 +1029,7 @@ class Trellis2Remesh:
if isinstance(grid_size, (list, tuple)):
grid_size = np.array(grid_size)
if isinstance(grid_size, np.ndarray):
grid_size = torch.tensor(grid_size, dtype=torch.int32, device=coords.device)
grid_size = torch.tensor(grid_size, dtype=torch.int32, device='cuda')
voxel_size = (aabb[1] - aabb[0]) / grid_size
# Move data to GPU
+77 -45
View File
@@ -762,24 +762,45 @@ class Trellis2ImageTo3DPipeline(Pipeline):
meshes, subs = self.decode_shape_slat(shape_slat, resolution)
if self.low_vram:
self._cleanup_cuda()
tex_voxels = self.decode_tex_slat(tex_slat, subs)
if self.low_vram:
self._cleanup_cuda()
out_mesh = []
for m, v in zip(meshes, tex_voxels):
m.fill_holes()
out_mesh.append(
MeshWithVoxel(
m.vertices, m.faces,
origin = [-0.5, -0.5, -0.5],
voxel_size = 1 / resolution,
coords = v.coords[:, 1:],
attrs = v.feats,
voxel_shape = torch.Size([*v.shape, *v.spatial_shape]),
layout=self.pbr_attr_layout
if tex_slat is None:
if self.low_vram:
self._cleanup_cuda()
out_mesh = []
for m in meshes:
m.fill_holes()
out_mesh.append(
MeshWithVoxel(
m.vertices, m.faces,
origin = [-0.5, -0.5, -0.5],
voxel_size = 1 / resolution,
coords = None,
attrs = None,
voxel_shape = None,
layout=self.pbr_attr_layout
)
)
)
return out_mesh
return out_mesh
else:
tex_voxels = self.decode_tex_slat(tex_slat, subs)
if self.low_vram:
self._cleanup_cuda()
out_mesh = []
for m, v in zip(meshes, tex_voxels):
m.fill_holes()
out_mesh.append(
MeshWithVoxel(
m.vertices, m.faces,
origin = [-0.5, -0.5, -0.5],
voxel_size = 1 / resolution,
coords = v.coords[:, 1:],
attrs = v.feats,
voxel_shape = torch.Size([*v.shape, *v.spatial_shape]),
layout=self.pbr_attr_layout
)
)
return out_mesh
@torch.no_grad()
def run(
@@ -795,7 +816,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
pipeline_type: Optional[str] = None,
max_num_tokens: int = 49152,
sparse_structure_resolution: int = 32,
max_views: int = 4
max_views: int = 4,
generate_texture_slat = True
) -> List[MeshWithVoxel]:
"""
Run the pipeline.
@@ -875,12 +897,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
if not self.keep_models_loaded:
self.unload_shape_slat_flow_model_512()
self.unload_tex_slat_flow_model_1024()
self.load_tex_slat_flow_model_512()
tex_slat = self.sample_tex_slat(
cond_512, self.models['tex_slat_flow_model_512'],
shape_slat, tex_slat_sampler_params
)
if generate_texture_slat:
self.unload_tex_slat_flow_model_1024()
self.load_tex_slat_flow_model_512()
tex_slat = self.sample_tex_slat(
cond_512, self.models['tex_slat_flow_model_512'],
shape_slat, tex_slat_sampler_params
)
if not self.keep_models_loaded:
self.unload_tex_slat_flow_model_512()
@@ -897,12 +920,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
if not self.keep_models_loaded:
self.unload_shape_slat_flow_model_1024()
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 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()
@@ -923,12 +947,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self.unload_shape_slat_flow_model_512()
self.unload_shape_slat_flow_model_1024()
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 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()
@@ -947,21 +972,28 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self.unload_shape_slat_flow_model_512()
self.unload_shape_slat_flow_model_1024()
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 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()
self.unload_tex_slat_flow_model_1024()
torch.cuda.empty_cache()
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
if generate_texture_slat:
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
else:
out_mesh = self.decode_latent(shape_slat, None, res)
torch.cuda.empty_cache()
if return_latent:
return out_mesh, (shape_slat, tex_slat, res)
if generate_texture_slat:
return out_mesh, (shape_slat, tex_slat, res)
else:
return out_mesh, (shape_slat, None, res)
else:
return out_mesh