diff --git a/.gitignore b/.gitignore index ed8ebf5..b50249e 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,4 @@ -__pycache__ \ No newline at end of file +__pycache__ +Cumesh/* +o_voxel/* +FlexGEMM/* \ No newline at end of file diff --git a/nodes.py b/nodes.py index f47854b..08f6398 100644 --- a/nodes.py +++ b/nodes.py @@ -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,) diff --git a/trellis2/models/sc_vaes/fdg_vae.py b/trellis2/models/sc_vaes/fdg_vae.py index 1bb1f55..b1f0bd4 100644 --- a/trellis2/models/sc_vaes/fdg_vae.py +++ b/trellis2/models/sc_vaes/fdg_vae.py @@ -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) diff --git a/trellis2/models/sc_vaes/sparse_unet_vae.py b/trellis2/models/sc_vaes/sparse_unet_vae.py index ad38c50..d6d6dcd 100644 --- a/trellis2/models/sc_vaes/sparse_unet_vae.py +++ b/trellis2/models/sc_vaes/sparse_unet_vae.py @@ -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) diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 9d8e091..5414241 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -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: diff --git a/wheels/Linux/Torch270/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl b/wheels/Linux/Torch270/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl index 2aafb24..b559846 100644 Binary files a/wheels/Linux/Torch270/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl and b/wheels/Linux/Torch270/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl differ diff --git a/wheels/Linux/Torch270/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl b/wheels/Linux/Torch270/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl index 4d4e40d..a2d531c 100644 Binary files a/wheels/Linux/Torch270/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl and b/wheels/Linux/Torch270/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl differ diff --git a/wheels/Linux/Torch291/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl b/wheels/Linux/Torch291/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl index 3d17c2b..91d8d09 100644 Binary files a/wheels/Linux/Torch291/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl and b/wheels/Linux/Torch291/cumesh-0.0.1-cp312-cp312-linux_x86_64.whl differ diff --git a/wheels/Linux/Torch291/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl b/wheels/Linux/Torch291/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl index 51cc7ba..d94a133 100644 Binary files a/wheels/Linux/Torch291/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl and b/wheels/Linux/Torch291/o_voxel-0.0.1-cp312-cp312-linux_x86_64.whl differ diff --git a/wheels/Windows/Torch270/cumesh-0.0.1-cp311-cp311-win_amd64.whl b/wheels/Windows/Torch270/cumesh-0.0.1-cp311-cp311-win_amd64.whl index 52870a8..8f57690 100644 Binary files a/wheels/Windows/Torch270/cumesh-0.0.1-cp311-cp311-win_amd64.whl and b/wheels/Windows/Torch270/cumesh-0.0.1-cp311-cp311-win_amd64.whl differ diff --git a/wheels/Windows/Torch270/cumesh-0.0.1-cp312-cp312-win_amd64.whl b/wheels/Windows/Torch270/cumesh-0.0.1-cp312-cp312-win_amd64.whl index dea85bb..1f2a0f3 100644 Binary files a/wheels/Windows/Torch270/cumesh-0.0.1-cp312-cp312-win_amd64.whl and b/wheels/Windows/Torch270/cumesh-0.0.1-cp312-cp312-win_amd64.whl differ diff --git a/wheels/Windows/Torch270/o_voxel-0.0.1-cp311-cp311-win_amd64.whl b/wheels/Windows/Torch270/o_voxel-0.0.1-cp311-cp311-win_amd64.whl index fa8f6cf..b81052c 100644 Binary files a/wheels/Windows/Torch270/o_voxel-0.0.1-cp311-cp311-win_amd64.whl and b/wheels/Windows/Torch270/o_voxel-0.0.1-cp311-cp311-win_amd64.whl differ diff --git a/wheels/Windows/Torch270/o_voxel-0.0.1-cp312-cp312-win_amd64.whl b/wheels/Windows/Torch270/o_voxel-0.0.1-cp312-cp312-win_amd64.whl index 6ff3238..09c5c88 100644 Binary files a/wheels/Windows/Torch270/o_voxel-0.0.1-cp312-cp312-win_amd64.whl and b/wheels/Windows/Torch270/o_voxel-0.0.1-cp312-cp312-win_amd64.whl differ diff --git a/wheels/Windows/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl b/wheels/Windows/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl index c3fb21d..e7eeaca 100644 Binary files a/wheels/Windows/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl and b/wheels/Windows/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl differ diff --git a/wheels/Windows/Torch280/cumesh-0.0.1-cp312-cp312-win_amd64.whl b/wheels/Windows/Torch280/cumesh-0.0.1-cp312-cp312-win_amd64.whl index a32f5d0..a5eddd0 100644 Binary files a/wheels/Windows/Torch280/cumesh-0.0.1-cp312-cp312-win_amd64.whl and b/wheels/Windows/Torch280/cumesh-0.0.1-cp312-cp312-win_amd64.whl differ diff --git a/wheels/Windows/Torch280/o_voxel-0.0.1-cp311-cp311-win_amd64.whl b/wheels/Windows/Torch280/o_voxel-0.0.1-cp311-cp311-win_amd64.whl index 942f173..24906eb 100644 Binary files a/wheels/Windows/Torch280/o_voxel-0.0.1-cp311-cp311-win_amd64.whl and b/wheels/Windows/Torch280/o_voxel-0.0.1-cp311-cp311-win_amd64.whl differ diff --git a/wheels/Windows/Torch280/o_voxel-0.0.1-cp312-cp312-win_amd64.whl b/wheels/Windows/Torch280/o_voxel-0.0.1-cp312-cp312-win_amd64.whl index 40ca21c..2a535ee 100644 Binary files a/wheels/Windows/Torch280/o_voxel-0.0.1-cp312-cp312-win_amd64.whl and b/wheels/Windows/Torch280/o_voxel-0.0.1-cp312-cp312-win_amd64.whl differ