Implemented "Tiled" Decoder + Updated Cumesh and O_voxel

This commit is contained in:
Bruno Fargnoli
2026-01-05 19:03:24 +01:00
parent ee7edd2f53
commit 2fd9df85fa
17 changed files with 63 additions and 24 deletions
+4 -1
View File
@@ -1 +1,4 @@
__pycache__
__pycache__
Cumesh/*
o_voxel/*
FlexGEMM/*
+11 -6
View File
@@ -285,6 +285,7 @@ class Trellis2MeshWithVoxelGenerator:
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
"sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":128,"step":8}),
"generate_texture_slat": ("BOOLEAN", {"default":True}),
"use_tiled_decoder": ("BOOLEAN", {"default":True}),
},
}
@@ -294,7 +295,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, generate_texture_slat):
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, use_tiled_decoder):
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
@@ -309,7 +310,7 @@ class Trellis2MeshWithVoxelGenerator:
pbar = ProgressBar(num_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, generate_texture_slat = generate_texture_slat, pbar=pbar)[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, use_tiled=use_tiled_decoder, pbar=pbar)[0]
return (mesh,)
@@ -785,6 +786,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
"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}),
"use_tiled_decoder": ("BOOLEAN", {"default":True}),
},
}
@@ -815,7 +817,8 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
shape_guidance_interval_start,
shape_guidance_interval_end,
texture_guidance_interval_start,
texture_guidance_interval_end):
texture_guidance_interval_end,
use_tiled_decoder):
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
@@ -835,7 +838,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
pbar = ProgressBar(num_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, generate_texture_slat=generate_texture_slat, pbar=pbar)[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, use_tiled=use_tiled_decoder, pbar=pbar)[0]
return (mesh,)
@@ -1411,6 +1414,7 @@ class Trellis2MeshRefiner:
"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}),
"use_tiled_decoder": ("BOOLEAN", {"default":True}),
},
}
@@ -1435,7 +1439,8 @@ class Trellis2MeshRefiner:
shape_guidance_interval_start,
shape_guidance_interval_end,
texture_guidance_interval_start,
texture_guidance_interval_end):
texture_guidance_interval_end,
use_tiled_decoder):
image = tensor2pil(image)
@@ -1445,7 +1450,7 @@ class Trellis2MeshRefiner:
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]
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, use_tiled=use_tiled_decoder)[0]
return (mesh,)
+25 -8
View File
@@ -18,7 +18,7 @@ from .sparse_unet_vae import (
chunked_apply,
)
from ...representations import Mesh
from o_voxel.convert import flexible_dual_grid_to_mesh
from o_voxel.convert import flexible_dual_grid_to_mesh, tiled_flexible_dual_grid_to_mesh
class FlexiDualGridVaeEncoder(SparseUnetVaeEncoder):
@@ -81,7 +81,7 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder):
def set_resolution(self, resolution: int) -> None:
self.resolution = resolution
def forward(self, x: sp.SparseTensor, gt_intersected: sp.SparseTensor = None, **kwargs):
def forward(self, x: sp.SparseTensor, gt_intersected: sp.SparseTensor = None, useTiled: bool = True, **kwargs):
decoded = super().forward(x, **kwargs)
if self.training:
h, subs_gt, subs = decoded
@@ -115,11 +115,28 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder):
vertices = h.replace((1 + 2 * self.voxel_margin) * F.sigmoid(h.feats[..., 0:3]) - self.voxel_margin)
intersected = h.replace(h.feats[..., 3:6] > 0)
quad_lerp = h.replace(F.softplus(h.feats[..., 6:7]))
mesh = [Mesh(*flexible_dual_grid_to_mesh(
h.coords[:, 1:], v.feats, i.feats, q.feats,
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],
grid_size=self.resolution,
train=False
)) for v, i, q in zip(vertices, intersected, quad_lerp)]
if useTiled:
mesh = [Mesh(*tiled_flexible_dual_grid_to_mesh(
coords=h.coords[:, 1:],
dual_vertices=v.feats,
intersected_flag=i.feats,
split_weight=q.feats,
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],
grid_size=self.resolution,
tile_size=128, # Adjust this based on your GPU VRAM
train=False
)) for v, i, q in zip(vertices, intersected, quad_lerp)]
else:
mesh = [Mesh(*flexible_dual_grid_to_mesh(
coords=h.coords[:, 1:],
dual_vertices=v.feats,
intersected_flag=i.feats,
split_weight=q.feats,
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],
grid_size=self.resolution,
train=False
)) for v, i, q in zip(vertices, intersected, quad_lerp)]
out_list[0] = mesh
return out_list[0] if len(out_list) == 1 else tuple(out_list)
@@ -608,6 +608,7 @@ class SparseUnetVaeDecoder(nn.Module):
h = block(h, subdiv=guide_subs[i] if guide_subs is not None else None)
else:
h = block(h)
torch.cuda.empty_cache()
if self.low_vram:
def fused_finalize(t):
t = t.type(x.dtype)
+22 -9
View File
@@ -579,7 +579,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
**lr_cond,
**sampler_params,
verbose=True,
tqdm_desc="Sampling shape SLat",
tqdm_desc="Sampling shape SLat (LR)",
).samples
if self.low_vram:
flow_model_lr.cpu()
@@ -641,7 +641,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
**cond,
**sampler_params,
verbose=True,
tqdm_desc="Sampling shape SLat",
tqdm_desc="Sampling shape SLat (HR)",
).samples
if self.low_vram:
flow_model.cpu()
@@ -662,6 +662,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self,
slat: SparseTensor,
resolution: int,
use_tiled: bool = True,
) -> Tuple[List[Mesh], List[SparseTensor]]:
"""
Decode the structured latent.
@@ -680,7 +681,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
if self.low_vram:
self.models['shape_slat_decoder'].to(self.device)
self.models['shape_slat_decoder'].low_vram = True
ret = self.models['shape_slat_decoder'](slat, return_subs=True)
ret = self.models['shape_slat_decoder'](slat, return_subs=True, useTiled=use_tiled)
if self.low_vram:
self.models['shape_slat_decoder'].cpu()
self.models['shape_slat_decoder'].low_vram = False
@@ -782,6 +783,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
shape_slat: SparseTensor,
tex_slat: SparseTensor,
resolution: int,
use_tiled: bool = True,
) -> List[MeshWithVoxel]:
"""
Decode the latent codes.
@@ -791,7 +793,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
tex_slat (SparseTensor): The structured latent for texture.
resolution (int): The resolution of the output.
"""
meshes, subs = self.decode_shape_slat(shape_slat, resolution)
meshes, subs = self.decode_shape_slat(shape_slat, resolution, use_tiled=use_tiled)
if self.low_vram:
self._cleanup_cuda()
@@ -850,6 +852,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
sparse_structure_resolution: int = 32,
max_views: int = 4,
generate_texture_slat = True,
use_tiled: bool = True,
pbar = None
) -> List[MeshWithVoxel]:
"""
@@ -1110,9 +1113,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
torch.cuda.empty_cache()
if generate_texture_slat:
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
out_mesh = self.decode_latent(shape_slat, tex_slat, res, use_tiled=use_tiled)
else:
out_mesh = self.decode_latent(shape_slat, None, res)
out_mesh = self.decode_latent(shape_slat, None, res, use_tiled=use_tiled)
torch.cuda.empty_cache()
pbar.update(1)
if return_latent:
@@ -1214,7 +1217,16 @@ class Trellis2ImageTo3DPipeline(Pipeline):
_cumesh = cumesh.CuMesh()
_cumesh.init(vertices_torch, faces_torch)
print('Unwrapping mesh ...')
vertices_torch, faces_torch, uvs_torch, vmap = _cumesh.uv_unwrap(return_vmaps=True)
vertices_torch, faces_torch, uvs_torch, vmap = _cumesh.uv_unwrap(
compute_charts_kwargs={
"threshold_cone_half_angle_rad": np.radians(90.0),
"refine_iterations": 0,
"global_iterations": 1,
"smooth_strength": 1,
},
return_vmaps=True,
verbose=True,
)
vertices_torch = vertices_torch.cuda()
faces_torch = faces_torch.cuda()
uvs_torch = uvs_torch.cuda()
@@ -1465,6 +1477,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
generate_texture_slat = True,
return_latent = False,
downsampling = 16,
use_tiled: bool = True,
):
mesh = self.preprocess_mesh(mesh)
torch.manual_seed(seed)
@@ -1562,9 +1575,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
torch.cuda.empty_cache()
if generate_texture_slat:
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
out_mesh = self.decode_latent(shape_slat, tex_slat, res, use_tiled=use_tiled)
else:
out_mesh = self.decode_latent(shape_slat, None, res)
out_mesh = self.decode_latent(shape_slat, None, res, use_tiled=use_tiled)
torch.cuda.empty_cache()
if return_latent: