From d7649d7799b6d6a9ac27c90663ccd5ba37a6bc0a Mon Sep 17 00:00:00 2001 From: Bruno Fargnoli Date: Sat, 27 Dec 2025 12:39:44 +0100 Subject: [PATCH] Added node for "Mesh Refiner" with an example --- example_workflows/MeshRefiner.json | 518 +++++++++++++++++++++ nodes.py | 54 ++- trellis2/pipelines/trellis2_image_to_3d.py | 217 ++++++++- 3 files changed, 787 insertions(+), 2 deletions(-) create mode 100644 example_workflows/MeshRefiner.json diff --git a/example_workflows/MeshRefiner.json b/example_workflows/MeshRefiner.json new file mode 100644 index 0000000..b75d57d --- /dev/null +++ b/example_workflows/MeshRefiner.json @@ -0,0 +1,518 @@ +{ + "id": "bdcdd41a-9d6d-48d5-8de6-ab8903d3b227", + "revision": 0, + "last_node_id": 13, + "last_link_id": 19, + "nodes": [ + { + "id": 1, + "type": "Trellis2LoadModel", + "pos": [ + 268.89245296723817, + 497.0522341473507 + ], + "size": [ + 301.7261859434295, + 154 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "links": [ + 1 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "c5d66aa46bf1bbf903fd4fe9ff086ac774bb7e03", + "Node name for S&R": "Trellis2LoadModel", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "TRELLIS.2-4B", + "flash_attn", + "cuda", + true, + true + ] + }, + { + "id": 13, + "type": "Trellis2Remesh", + "pos": [ + 1023.4554969911898, + 623.7648065627299 + ], + "size": [ + 316.630859375, + 154 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 19 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 18 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2Remesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 1, + 0, + true, + 0.03, + "Auto" + ] + }, + { + "id": 10, + "type": "Trellis2SimplifyMesh", + "pos": [ + 1393.5945383749734, + 624.72625748105 + ], + "size": [ + 270, + 82 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 18 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 13 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2SimplifyMesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 2000000, + "Cumesh" + ] + }, + { + "id": 9, + "type": "Trellis2MeshWithVoxelToTrimesh", + "pos": [ + 1710.8561337529898, + 626.6489172660999 + ], + "size": [ + 342.2886366432093, + 26 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "link": 13 + } + ], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 9 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2MeshWithVoxelToTrimesh", + "widget_ue_connectable": {} + }, + "widgets_values": [] + }, + { + "id": 7, + "type": "Preview3D", + "pos": [ + 1375.9044312035774, + 767.9746810560968 + ], + "size": [ + 716.1079334929157, + 767.2763682808798 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "camera_info", + "shape": 7, + "type": "LOAD3D_CAMERA", + "link": null + }, + { + "name": "bg_image", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "model_file", + "type": "STRING", + "widget": { + "name": "model_file" + }, + "link": 6 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.6.0", + "Node name for S&R": "Preview3D", + "widget_ue_connectable": {}, + "Last Time Model File": "DwarfWarrior_Refined_1536_00001_.glb", + "Scene Config": { + "showGrid": true, + "backgroundColor": "#282828", + "backgroundImage": "", + "backgroundRenderMode": "tiled" + }, + "Camera Config": { + "cameraType": "perspective", + "fov": 35, + "state": { + "position": { + "x": -0.12117550423665688, + "y": 2.8575581491112874, + "z": 8.963607314880118 + }, + "target": { + "x": 0, + "y": 2.5, + "z": 0 + }, + "zoom": 1, + "cameraType": "perspective" + } + }, + "Light Config": { + "intensity": 3 + }, + "Model Config": { + "upDirection": "original", + "materialMode": "original" + } + }, + "widgets_values": [ + "DwarfWarrior_Refined_1536_00001_.glb", + "" + ] + }, + { + "id": 3, + "type": "Trellis2LoadMesh", + "pos": [ + 71.861984424133, + 734.9986287230716 + ], + "size": [ + 489.9682548146561, + 63.28771379401087 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 2 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2LoadMesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "C:\\Git\\ComfyUI\\output\\DwarfWarrior_Hy20_Hy3D_00002_.obj" + ] + }, + { + "id": 2, + "type": "Trellis2LoadImageWithTransparency", + "pos": [ + -30.719285579198413, + 866.7073090204125 + ], + "size": [ + 605.6781257083595, + 694.7687444555902 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [] + }, + { + "name": "mask", + "type": "MASK", + "links": [] + }, + { + "name": "image_with_alpha", + "type": "IMAGE", + "links": [ + 3 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "e7f9b30df7a09bcedf1c955e176754c73f983254", + "Node name for S&R": "Trellis2LoadImageWithTransparency", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "DwarfWarrior02.png", + "image" + ] + }, + { + "id": 4, + "type": "Trellis2MeshRefiner", + "pos": [ + 657.7388525725853, + 713.8478219573461 + ], + "size": [ + 313.0738153551299, + 390.7720378587203 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 1 + }, + { + "name": "trimesh", + "type": "TRIMESH", + "link": 2 + }, + { + "name": "image", + "type": "IMAGE", + "link": 3 + } + ], + "outputs": [ + { + "name": "mesh", + "type": "MESHWITHVOXEL", + "links": [ + 19 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2MeshRefiner", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 12345, + "fixed", + 1536, + 12, + 7.5, + 0.5, + 3, + 12, + 1, + 0, + 3, + 50000, + false + ] + }, + { + "id": 6, + "type": "Trellis2ExportMesh", + "pos": [ + 1031.915640139208, + 861.0380028567 + ], + "size": [ + 270, + 106 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "link": 9 + } + ], + "outputs": [ + { + "name": "glb_path", + "type": "STRING", + "links": [ + 6 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "87f48b0726b5d5e58b744bbac632331cd7a34103", + "Node name for S&R": "Trellis2ExportMesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "DwarfWarrior_Refined_1536", + "glb", + true + ] + } + ], + "links": [ + [ + 1, + 1, + 0, + 4, + 0, + "TRELLIS2PIPELINE" + ], + [ + 2, + 3, + 0, + 4, + 1, + "TRIMESH" + ], + [ + 3, + 2, + 2, + 4, + 2, + "IMAGE" + ], + [ + 6, + 6, + 0, + 7, + 2, + "STRING" + ], + [ + 9, + 9, + 0, + 6, + 0, + "TRIMESH" + ], + [ + 13, + 10, + 0, + 9, + 0, + "MESHWITHVOXEL" + ], + [ + 18, + 13, + 0, + 10, + 0, + "MESHWITHVOXEL" + ], + [ + 19, + 4, + 0, + 13, + 0, + "MESHWITHVOXEL" + ] + ], + "groups": [], + "config": {}, + "extra": { + "workflowRendererVersion": "LG", + "ue_links": [], + "ds": { + "scale": 0.6934334949441421, + "offset": [ + 160.98892492810913, + -350.72740088339407 + ] + }, + "links_added_by_ue": [], + "frontendVersion": "1.35.9", + "VHS_latentpreview": false, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index 2aafd74..47b72a8 100644 --- a/nodes.py +++ b/nodes.py @@ -1222,7 +1222,57 @@ class Trellis2PreProcessImage: output = np.array(output).astype(np.float32) / 255 output = output[:, :, :3] * output[:, :, 3:4] output = Image.fromarray((output * 255).astype(np.uint8)) - return output + return output + +class Trellis2MeshRefiner: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "pipeline": ("TRELLIS2PIPELINE",), + "trimesh": ("TRIMESH",), + "image": ("IMAGE",), + "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), + "resolution": ([512,1024,1536],{"default":1024}), + "shape_steps": ("INT",{"default":12, "min":1, "max":100},), + "shape_guidance_strength": ("FLOAT",{"default":7.5}), + "shape_guidance_rescale": ("FLOAT",{"default":0.5}), + "shape_rescale_t": ("FLOAT",{"default":3.0}), + "texture_steps": ("INT",{"default":12, "min":1, "max":100},), + "texture_guidance_strength": ("FLOAT",{"default":1.0}), + "texture_guidance_rescale": ("FLOAT",{"default":0.0}), + "texture_rescale_t": ("FLOAT",{"default":3.0}), + "max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}), + "generate_texture_slat": ("BOOLEAN", {"default":True}), + }, + } + + RETURN_TYPES = ("MESHWITHVOXEL", ) + RETURN_NAMES = ("mesh", ) + FUNCTION = "process" + CATEGORY = "Trellis2Wrapper" + OUTPUT_NODE = True + + def process(self, pipeline, trimesh, image, seed, resolution, + shape_steps, + shape_guidance_strength, + shape_guidance_rescale, + shape_rescale_t, + texture_steps, + texture_guidance_strength, + texture_guidance_rescale, + texture_rescale_t, + max_num_tokens, + generate_texture_slat): + + image = tensor2pil(image) + + 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.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)[0] + + return (mesh,) NODE_CLASS_MAPPINGS = { "Trellis2LoadModel": Trellis2LoadModel, @@ -1239,6 +1289,7 @@ NODE_CLASS_MAPPINGS = { "Trellis2MeshTexturing": Trellis2MeshTexturing, "Trellis2LoadMesh": Trellis2LoadMesh, "Trellis2PreProcessImage": Trellis2PreProcessImage, + "Trellis2MeshRefiner": Trellis2MeshRefiner, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -1256,4 +1307,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "Trellis2MeshTexturing": "Trellis2 - Mesh Texturing", "Trellis2LoadMesh": "Trellis2 - Load Mesh", "Trellis2PreProcessImage": "Trellis2 - PreProcess Image", + "Trellis2MeshRefiner": "Trellis2 - Mesh Refiner", } \ No newline at end of file diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 4d889d3..10f8344 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -1206,4 +1206,219 @@ class Trellis2ImageTo3DPipeline(Pipeline): torch.cuda.empty_cache() out_mesh, baseColorTexture, metallicRoughnessTexture = self.postprocess_mesh(mesh, pbr_voxel, resolution, texture_size, texture_alpha_mode, double_side_material) - return out_mesh, baseColorTexture, metallicRoughnessTexture \ No newline at end of file + return out_mesh, baseColorTexture, metallicRoughnessTexture + + def get_coords_from_trimesh(self, mesh, resolution): + vertices = torch.from_numpy(mesh.vertices).float() + faces = torch.from_numpy(mesh.faces).long() + + voxel_indices, dual_vertices, intersected = o_voxel.convert.mesh_to_flexible_dual_grid( + vertices.cpu(), faces.cpu(), + grid_size=resolution, + aabb=[[-0.5,-0.5,-0.5],[0.5,0.5,0.5]], + face_weight=1.0, + boundary_weight=0.2, + regularization_weight=1e-2, + timing=True, + ) + + coords = torch.cat([torch.zeros_like(voxel_indices[:, 0:1]), voxel_indices], dim=-1) + coords = coords.cpu() + + print(coords) + + del voxel_indices + del dual_vertices + del intersected + + if self.low_vram: + self._cleanup_cuda() + + return coords; + + def sample_mesh_slat( + self, + mesh_slat, + cond: dict, + flow_model, + resolution: int, + sampler_params: dict = {}, + max_num_tokens: int = 49152, + ) -> SparseTensor: + # Upsample + self.load_shape_slat_decoder() + if self.low_vram: + self.models['shape_slat_decoder'].to(self.device) + self.models['shape_slat_decoder'].low_vram = True + hr_coords = self.models['shape_slat_decoder'].upsample(mesh_slat, upsample_times=4) + if self.low_vram: + self.models['shape_slat_decoder'].cpu() + self.models['shape_slat_decoder'].low_vram = False + hr_resolution = resolution + + if not self.keep_models_loaded: + self.unload_shape_slat_decoder() + + while True: + quant_coords = torch.cat([ + hr_coords[:, :1], + ((hr_coords[:, 1:] + 0.5) / hr_resolution * (hr_resolution // 16)).int(), + ], dim=1) + coords = quant_coords.unique(dim=0) + num_tokens = coords.shape[0] + if num_tokens < max_num_tokens or hr_resolution == 1024: + if hr_resolution != resolution: + print(f"Due to the limited number of tokens, the resolution is reduced to {hr_resolution}.") + break + hr_resolution -= 128 + + coords_dev = coords.to(self.device) + # Sample structured latent + noise = SparseTensor( + feats=torch.randn(coords.shape[0], flow_model.in_channels, device=self.device), + coords=coords_dev, + ) + sampler_params = {**self.shape_slat_sampler_params, **sampler_params} + if self.low_vram: + flow_model.to(self.device) + slat = self.shape_slat_sampler.sample( + flow_model, + noise, + **cond, + **sampler_params, + verbose=True, + tqdm_desc="Sampling shape SLat", + ).samples + if self.low_vram: + flow_model.cpu() + self._cleanup_cuda() + + std = torch.tensor(self.shape_slat_normalization['std'])[None].to(slat.device) + mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(slat.device) + slat = slat * std + mean + + del coords_dev + if self.low_vram: + cond = self._cond_cpu(cond) + self._cleanup_cuda() + + return slat, hr_resolution + + @torch.no_grad() + def refine_mesh( + self, + mesh: trimesh.Trimesh, + image: Image.Image, + seed: int = 42, + shape_slat_sampler_params: dict = {}, + tex_slat_sampler_params: dict = {}, + resolution: int = 1024, + max_num_tokens = 50000, + generate_texture_slat = True, + return_latent = False + ): + mesh = self.preprocess_mesh(mesh) + torch.manual_seed(seed) + + self.load_image_cond_model() + + if resolution == 512: + cond = self.get_cond(image, 512) + else: + cond = self.get_cond(image, 1024) + + if not self.keep_models_loaded: + self.unload_image_cond_model() + + mesh_slat = self.encode_shape_slat(mesh, resolution) + + if resolution==512: + self.unload_shape_slat_flow_model_1024() + self.load_shape_slat_flow_model_512() + shape_slat, res = self.sample_mesh_slat( + mesh_slat, + cond, + self.models['shape_slat_flow_model_512'], + 512, + shape_slat_sampler_params, + max_num_tokens + ) + + if not self.keep_models_loaded: + self.unload_shape_slat_flow_model_512() + + 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, 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() + elif resolution == 1024: + self.unload_shape_slat_flow_model_512() + self.load_shape_slat_flow_model_1024() + shape_slat, res = self.sample_mesh_slat( + mesh_slat, + cond, + self.models['shape_slat_flow_model_1024'], + 1024, + shape_slat_sampler_params, + max_num_tokens + ) + + if not self.keep_models_loaded: + self.unload_shape_slat_flow_model_1024() + + 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, 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() + elif resolution == 1536: + self.unload_shape_slat_flow_model_512() + self.load_shape_slat_flow_model_1024() + shape_slat, res = self.sample_mesh_slat( + mesh_slat, + cond, + self.models['shape_slat_flow_model_1024'], + 1536, + shape_slat_sampler_params, + max_num_tokens + ) + + if not self.keep_models_loaded: + self.unload_shape_slat_flow_model_1024() + + 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, 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() + + torch.cuda.empty_cache() + 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: + if generate_texture_slat: + return out_mesh, (shape_slat, tex_slat, res) + else: + return out_mesh, (shape_slat, None, res) + else: + return out_mesh \ No newline at end of file