Implemented "Tiled" Decoder + Updated Cumesh and O_voxel
This commit is contained in:
+4
-1
@@ -1 +1,4 @@
|
||||
__pycache__
|
||||
__pycache__
|
||||
Cumesh/*
|
||||
o_voxel/*
|
||||
FlexGEMM/*
|
||||
@@ -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,)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Reference in New Issue
Block a user