diff --git a/nodes.py b/nodes.py index 13cc28e..836881f 100644 --- a/nodes.py +++ b/nodes.py @@ -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 diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 3b79d57..4d889d3 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -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 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 4caa732..4101353 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/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 6ac4b00..3d17c2b 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/flex_gemm-0.0.1-cp312-cp312-linux_x86_64.whl b/wheels/Linux/Torch291/flex_gemm-0.0.1-cp312-cp312-linux_x86_64.whl index 1b45b8a..ee6dc8c 100644 Binary files a/wheels/Linux/Torch291/flex_gemm-0.0.1-cp312-cp312-linux_x86_64.whl and b/wheels/Linux/Torch291/flex_gemm-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 b395a0f..51cc7ba 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 2b989f2..fe23cae 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 f9c6ee3..348d799 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/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl b/wheels/Windows/Torch280/cumesh-0.0.1-cp311-cp311-win_amd64.whl index 4f253fe..8814679 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 e4044b2..1d0e5f6 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