import os import torch import torchvision.transforms as transforms from PIL import Image, ImageSequence, ImageOps from pathlib import Path import numpy as np import json import trimesh as Trimesh from tqdm import tqdm import folder_paths import node_helpers import hashlib import cv2 import gc import copy import pymeshlab import cumesh as CuMesh import o_voxel import meshlib.mrmeshnumpy as mrmeshnumpy import meshlib.mrmeshpy as mrmeshpy import nvdiffrast.torch as dr from flex_gemm.ops.grid_sample import grid_sample_3d import comfy.model_management as mm from comfy.utils import load_torch_file, ProgressBar, common_upscale import comfy.utils from .trellis2.pipelines import Trellis2ImageTo3DPipeline from .trellis2.representations import Mesh, MeshWithVoxel script_directory = os.path.dirname(os.path.abspath(__file__)) comfy_path = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) to_pil = transforms.ToPILImage() def pil2tensor(image): return torch.from_numpy(np.array(image).astype(np.float32) / 255.0)[None,] def tensor2pil(image: torch.Tensor) -> Image.Image: """ Accepts either: - (H,W,C) - (1,H,W,C) Returns a PIL RGB/RGBA image depending on channels. """ if isinstance(image, torch.Tensor): t = image.detach().cpu() if t.ndim == 4: # Expect (B,H,W,C); allow only B==1 here if t.shape[0] != 1: raise ValueError(f"tensor2pil expects batch of 1, got batch={t.shape[0]}") t = t[0] elif t.ndim != 3: raise ValueError(f"tensor2pil expects (H,W,C) or (1,H,W,C), got shape={tuple(t.shape)}") arr = (t.numpy() * 255.0).clip(0, 255).astype(np.uint8) return Image.fromarray(arr) raise TypeError(f"tensor2pil expected torch.Tensor, got {type(image)}") def tensor_batch_to_pil_list(images: torch.Tensor, max_views: int = 4) -> list[Image.Image]: """ Converts a ComfyUI IMAGE tensor (B,H,W,C) into a list of PIL images. Caps to max_views for safety. """ if not isinstance(images, torch.Tensor): raise TypeError(f"Expected torch.Tensor for IMAGE, got {type(images)}") if images.ndim == 4: b = int(images.shape[0]) n = min(b, int(max_views)) return [tensor2pil(images[i:i+1]) for i in range(n)] if images.ndim == 3: return [tensor2pil(images)] raise ValueError(f"Unsupported IMAGE tensor shape: {tuple(images.shape)}") def convert_tensor_images_to_pil(images): pil_array = [] for image in images: pil_array.append(tensor2pil(image)) return pil_array def simplify_with_meshlib(vertices, faces, target=1000000): current_faces_num = len(faces) print(f'Current Faces Number: {current_faces_num}') if current_faces_num 1 and img.format not in excluded_formats: output_image = torch.cat(output_images, dim=0) output_mask = torch.cat(output_masks, dim=0) output_image_ori = torch.cat(output_images_ori, dim=0) else: output_image = output_images[0] output_mask = output_masks[0] output_image_ori = output_images_ori[0] return (output_image, output_mask, output_image_ori) @classmethod def IS_CHANGED(s, image): image_path = folder_paths.get_annotated_filepath(image) m = hashlib.sha256() with open(image_path, 'rb') as f: m.update(f.read()) return m.digest().hex() @classmethod def VALIDATE_INPUTS(s, image): if not folder_paths.exists_annotated_filepath(image): return "Invalid image file: {}".format(image) return True class Trellis2SimplifyMesh: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "target_face_num": ("INT",{"default":1000000,"min":1,"max":30000000}), "method": (["Cumesh","Meshlib"],{"default":"Cumesh"}), }, } RETURN_TYPES = ("MESHWITHVOXEL", ) RETURN_NAMES = ("mesh", ) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, target_face_num, method): mesh_copy = copy.deepcopy(mesh) if method=="Cumesh": mesh_copy.simplify_with_cumesh(target = target_face_num) elif method=="Meshlib": mesh_copy.simplify_with_meshlib(target = target_face_num) else: raise Exception("Unknown simplification method") return (mesh_copy,) class Trellis2MeshWithVoxelToTrimesh: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "reorient_vertices":(["None","90 degrees","-90 degrees"],{"default":"90 degrees"}), }, } RETURN_TYPES = ("TRIMESH", ) RETURN_NAMES = ("trimesh", ) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, reorient_vertices): vertices_np = mesh.vertices.cpu().numpy() if reorient_vertices == '90 degrees': vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2], -vertices_np[:, 1] elif reorient_vertices == '-90 degrees': vertices_np[:, 1], vertices_np[:, 2] = -vertices_np[:, 2], vertices_np[:, 1] trimesh = Trimesh.Trimesh( vertices=vertices_np, faces=mesh.faces.cpu().numpy(), process=False ) return (trimesh,) class Trellis2ExportMesh: @classmethod def INPUT_TYPES(s): return { "required": { "trimesh": ("TRIMESH",), "filename_prefix": ("STRING", {"default": "3D/Trellis2"}), "file_format": (["glb", "obj", "ply", "stl", "3mf", "dae"],), }, "optional": { "save_file": ("BOOLEAN", {"default": True}), }, } RETURN_TYPES = ("STRING",) RETURN_NAMES = ("glb_path",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, trimesh, filename_prefix, file_format, save_file=True): full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, folder_paths.get_output_directory()) output_glb_path = Path(full_output_folder, f'{filename}_{counter:05}_.{file_format}') output_glb_path.parent.mkdir(exist_ok=True) if save_file: trimesh.export(output_glb_path, file_type=file_format) relative_path = Path(subfolder) / f'{filename}_{counter:05}_.{file_format}' else: temp_file = Path(full_output_folder, f'hy3dtemp_.{file_format}') trimesh.export(temp_file, file_type=file_format) relative_path = Path(subfolder) / f'hy3dtemp_.{file_format}' return (str(relative_path), ) class Trellis2PostProcessMesh: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "fill_holes": ("BOOLEAN", {"default":True}), "fill_holes_max_perimeter": ("FLOAT",{"default":0.03,"min":0.001,"max":99.999,"step":0.001}), "remove_duplicate_faces": ("BOOLEAN",{"default":True}), "repair_non_manifold_edges": ("BOOLEAN", {"default":True}), "remove_non_manifold_faces": ("BOOLEAN", {"default":True}), "remove_small_connected_components": ("BOOLEAN", {"default":True}), "remove_small_connected_components_size": ("FLOAT", {"default":0.00001,"min":0.00001,"max":9.99999,"step":0.00001}), "unify_faces_orientation": ("BOOLEAN", {"default":True}), "remove_floaters": ("BOOLEAN",{"default":True}), "remove_infinite_vertices": ("BOOLEAN",{"default":True}), }, } RETURN_TYPES = ("MESHWITHVOXEL",) RETURN_NAMES = ("mesh",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, fill_holes, fill_holes_max_perimeter, remove_duplicate_faces, repair_non_manifold_edges, remove_non_manifold_faces, remove_small_connected_components, remove_small_connected_components_size,unify_faces_orientation,remove_floaters,remove_infinite_vertices): mesh_copy = copy.deepcopy(mesh) if remove_floaters: mesh_copy = remove_floater(mesh_copy) if remove_infinite_vertices: mesh_copy = remove_mesh_infinite_vertices(mesh_copy) vertices = mesh_copy.vertices faces = mesh_copy.faces # Move data to GPU vertices = vertices.cuda() faces = faces.cuda() # Initialize CUDA mesh handler cumesh = CuMesh.CuMesh() cumesh.init(vertices, faces) print(f"Current vertices: {cumesh.num_vertices}, faces: {cumesh.num_faces}") if fill_holes: cumesh.fill_holes(max_hole_perimeter=fill_holes_max_perimeter) print(f"After filling holes: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") if remove_duplicate_faces: print('Removing duplicate faces ...') cumesh.remove_duplicate_faces() if repair_non_manifold_edges: print('Repairing non manifold edges ...') cumesh.repair_non_manifold_edges() if remove_non_manifold_faces: print('Removing non manifold faces ...') cumesh.remove_non_manifold_faces() if remove_small_connected_components: print('Removing small connected components ...') cumesh.remove_small_connected_components(remove_small_connected_components_size) if unify_faces_orientation: print('Unifying faces orientation ...') cumesh.unify_face_orientations() print(f"After initial cleanup: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") new_vertices, new_faces = cumesh.read() mesh_copy.vertices = new_vertices.to(mesh_copy.device) mesh_copy.faces = new_faces.to(mesh_copy.device) del cumesh gc.collect() return (mesh_copy,) class Trellis2UnWrapAndRasterizer: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":90.0,"min":0.0,"max":359.9}), "mesh_cluster_refine_iterations": ("INT",{"default":0}), "mesh_cluster_global_iterations": ("INT",{"default":1}), "mesh_cluster_smooth_strength": ("INT",{"default":1}), "texture_size": ("INT",{"default":1024, "min":512, "max":16384}), "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), "double_side_material": ("BOOLEAN",{"default":True}), "bake_on_vertices": ("BOOLEAN",{"default":False}), "use_custom_normals": ("BOOLEAN",{"default":False}), }, } RETURN_TYPES = ("TRIMESH","IMAGE","IMAGE",) RETURN_NAMES = ("trimesh","base_color_texture", "metallic_roughness_texture",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength, texture_size, texture_alpha_mode, double_side_material, bake_on_vertices = False,use_custom_normals=False): mesh_copy = copy.deepcopy(mesh) aabb = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]] vertices = mesh_copy.vertices faces = mesh_copy.faces attr_volume = mesh_copy.attrs coords = mesh_copy.coords attr_layout = mesh_copy.layout voxel_size = mesh_copy.voxel_size mesh_cluster_threshold_cone_half_angle_rad = np.radians(mesh_cluster_threshold_cone_half_angle_rad) # --- Input Normalization (AABB, Voxel Size, Grid Size) --- if isinstance(aabb, (list, tuple)): aabb = np.array(aabb) if isinstance(aabb, np.ndarray): aabb = torch.tensor(aabb, dtype=torch.float32, device=coords.device) # Calculate grid dimensions based on AABB and voxel size if voxel_size is not None: if isinstance(voxel_size, float): voxel_size = [voxel_size, voxel_size, voxel_size] 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) grid_size = ((aabb[1] - aabb[0]) / voxel_size).round().int() else: if isinstance(grid_size, int): grid_size = [grid_size, grid_size, grid_size] 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) voxel_size = (aabb[1] - aabb[0]) / grid_size print(f"Original mesh: {vertices.shape[0]} vertices, {faces.shape[0]} faces") vertices = vertices.cuda() faces = faces.cuda() cumesh = CuMesh.CuMesh() cumesh.init(vertices, faces) # Build BVH for the current mesh to guide remeshing print(f"Building BVH for current mesh...") bvh = CuMesh.cuBVH(vertices, faces) # --- Branch: Bake On Vertices (skip UV unwrapping and texture creation) --- if bake_on_vertices: print('Baking colors on vertices...') out_vertices, out_faces = cumesh.read() out_vertices = out_vertices.cuda() out_faces = out_faces.cuda() cumesh.compute_vertex_normals() out_normals = cumesh.read_vertex_normals() # Sample attributes directly at vertex positions from the voxel grid # No BVH mapping needed - the voxel grid contains all the color information vertex_attrs = grid_sample_3d( attr_volume, torch.cat([torch.zeros_like(coords[:, :1]), coords], dim=-1), shape=torch.Size([1, attr_volume.shape[1], *grid_size.tolist()]), grid=((out_vertices - aabb[0]) / voxel_size).reshape(1, -1, 3), mode='trilinear', ) # Extract base color and alpha per vertex (vertex_attrs shape: N_vertices x C) base_color_idx = attr_layout['base_color'] alpha_idx = attr_layout['alpha'] # Get RGB values and squeeze any extra dimensions to get (N, 3) vertex_colors_rgb = vertex_attrs[..., base_color_idx].cpu().numpy() vertex_colors_rgb = np.squeeze(vertex_colors_rgb) # Remove batch dims if any if vertex_colors_rgb.ndim == 1: vertex_colors_rgb = vertex_colors_rgb[None, :] # Ensure at least 2D vertex_colors_rgb = np.clip(vertex_colors_rgb * 255, 0, 255).astype(np.uint8) # Handle alpha based on texture_alpha_mode if texture_alpha_mode == "OPAQUE": # For OPAQUE mode, use full alpha (255) vertex_alpha = np.full((vertex_colors_rgb.shape[0], 1), 255, dtype=np.uint8) else: vertex_alpha = vertex_attrs[..., alpha_idx].cpu().numpy() vertex_alpha = np.squeeze(vertex_alpha) # Remove batch dims if any vertex_alpha = np.clip(vertex_alpha * 255, 0, 255).astype(np.uint8) # Ensure alpha is 2D with shape (N, 1) if vertex_alpha.ndim == 1: vertex_alpha = vertex_alpha[:, None] # Combine into RGBA vertex_colors_rgba = np.concatenate([vertex_colors_rgb, vertex_alpha], axis=-1) print("Finalizing mesh with vertex colors...") vertices_np = out_vertices.cpu().numpy() faces_np = out_faces.cpu().numpy() normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy() normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy() # Create mesh with vertex colors using ColorVisuals if use_custom_normals: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_normals=normals_np, vertex_colors=vertex_colors_rgba, process=False, ) else: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_colors=vertex_colors_rgba, process=False, ) del cumesh gc.collect() # Return empty placeholder textures for vertex color mode placeholder_texture = pil2tensor(Image.new('RGBA', (1, 1), (0, 0, 0, 0))) return (textured_mesh, placeholder_texture, placeholder_texture,) print('Unwrapping ...') out_vertices, out_faces, out_uvs, out_vmaps = cumesh.uv_unwrap( compute_charts_kwargs={ "threshold_cone_half_angle_rad": mesh_cluster_threshold_cone_half_angle_rad, "refine_iterations": mesh_cluster_refine_iterations, "global_iterations": mesh_cluster_global_iterations, "smooth_strength": mesh_cluster_smooth_strength, }, return_vmaps=True, verbose=True, ) out_vertices = out_vertices.cuda() out_faces = out_faces.cuda() out_uvs = out_uvs.cuda() out_vmaps = out_vmaps.cuda() cumesh.compute_vertex_normals() out_normals = cumesh.read_vertex_normals()[out_vmaps] print("Sampling attributes...") # Setup differentiable rasterizer context ctx = dr.RasterizeCudaContext() # Prepare UV coordinates for rasterization (rendering in UV space) uvs_rast = torch.cat([out_uvs * 2 - 1, torch.zeros_like(out_uvs[:, :1]), torch.ones_like(out_uvs[:, :1])], dim=-1).unsqueeze(0) rast = torch.zeros((1, texture_size, texture_size, 4), device='cuda', dtype=torch.float32) # Rasterize in chunks to save memory for i in range(0, out_faces.shape[0], 100000): rast_chunk, _ = dr.rasterize( ctx, uvs_rast, out_faces[i:i+100000], resolution=[texture_size, texture_size], ) mask_chunk = rast_chunk[..., 3:4] > 0 rast_chunk[..., 3:4] += i # Store face ID in alpha channel rast = torch.where(mask_chunk, rast_chunk, rast) # Mask of valid pixels in texture mask = rast[0, ..., 3] > 0 # Interpolate 3D positions in UV space (finding 3D coord for every texel) pos = dr.interpolate(out_vertices.unsqueeze(0), rast, out_faces)[0][0] valid_pos = pos[mask] # Map these positions back to the *original* high-res mesh to get accurate attributes # This corrects geometric errors introduced by simplification/remeshing _, face_id, uvw = bvh.unsigned_distance(valid_pos, return_uvw=True) orig_tri_verts = vertices[faces[face_id.long()]] # (N_new, 3, 3) valid_pos = (orig_tri_verts * uvw.unsqueeze(-1)).sum(dim=1) # Trilinear sampling from the attribute volume (Color, Material props) attrs = torch.zeros(texture_size, texture_size, attr_volume.shape[1], device='cuda') attrs[mask] = grid_sample_3d( attr_volume, torch.cat([torch.zeros_like(coords[:, :1]), coords], dim=-1), shape=torch.Size([1, attr_volume.shape[1], *grid_size.tolist()]), grid=((valid_pos - aabb[0]) / voxel_size).reshape(1, -1, 3), mode='trilinear', ) # --- Texture Post-Processing & Material Construction --- print("Finalizing mesh...") mask = mask.cpu().numpy() # Extract channels based on layout (BaseColor, Metallic, Roughness, Alpha) base_color = np.clip(attrs[..., attr_layout['base_color']].cpu().numpy() * 255, 0, 255).astype(np.uint8) metallic = np.clip(attrs[..., attr_layout['metallic']].cpu().numpy() * 255, 0, 255).astype(np.uint8) roughness = np.clip(attrs[..., attr_layout['roughness']].cpu().numpy() * 255, 0, 255).astype(np.uint8) alpha = np.clip(attrs[..., attr_layout['alpha']].cpu().numpy() * 255, 0, 255).astype(np.uint8) alpha_mode = texture_alpha_mode # Inpainting: fill gaps (dilation) to prevent black seams at UV boundaries mask_inv = (~mask).astype(np.uint8) base_color = cv2.inpaint(base_color, mask_inv, 3, cv2.INPAINT_TELEA) metallic = cv2.inpaint(metallic, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] roughness = cv2.inpaint(roughness, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] alpha = cv2.inpaint(alpha, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] # Create PBR material # Standard PBR packs Metallic and Roughness into Blue and Green channels baseColorTexture_np = Image.fromarray(np.concatenate([base_color, alpha], axis=-1)) metallicRoughnessTexture_np = Image.fromarray(np.concatenate([np.zeros_like(metallic), roughness, metallic], axis=-1)) material = Trimesh.visual.material.PBRMaterial( baseColorTexture=baseColorTexture_np, baseColorFactor=np.array([255, 255, 255, 255], dtype=np.uint8), metallicRoughnessTexture=metallicRoughnessTexture_np, metallicFactor=1.0, roughnessFactor=1.0, alphaMode=alpha_mode, doubleSided=double_side_material, ) vertices_np = out_vertices.cpu().numpy() faces_np = out_faces.cpu().numpy() uvs_np = out_uvs.cpu().numpy() normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2], -vertices_np[:, 1] normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2], -normals_np[:, 1] uvs_np[:, 1] = 1 - uvs_np[:, 1] # Flip UV V-coordinate if use_custom_normals: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_normals=normals_np, process=False, visual=Trimesh.visual.TextureVisuals(uv=uvs_np,material=material) ) else: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, process=False, visual=Trimesh.visual.TextureVisuals(uv=uvs_np,material=material) ) del cumesh gc.collect() baseColorTexture = pil2tensor(baseColorTexture_np) metallicRoughnessTexture = pil2tensor(metallicRoughnessTexture_np) return (textured_mesh, baseColorTexture, metallicRoughnessTexture, ) class Trellis2MeshWithVoxelAdvancedGenerator: @classmethod def INPUT_TYPES(s): return { "required": { "pipeline": ("TRELLIS2PIPELINE",), "image": ("IMAGE",), "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), "pipeline_type": (["512","1024","1024_cascade","1536_cascade"],{"default":"1024_cascade"}), "sparse_structure_steps": ("INT",{"default":12, "min":1, "max":100},), "sparse_structure_guidance_strength": ("FLOAT",{"default":7.50}), "sparse_structure_guidance_rescale": ("FLOAT",{"default":0.70}), "sparse_structure_rescale_t": ("FLOAT",{"default":5.00}), "shape_steps": ("INT",{"default":12, "min":1, "max":100},), "shape_guidance_strength": ("FLOAT",{"default":7.50}), "shape_guidance_rescale": ("FLOAT",{"default":0.50}), "shape_rescale_t": ("FLOAT",{"default":3.00}), "texture_steps": ("INT",{"default":12, "min":1, "max":100},), "texture_guidance_strength": ("FLOAT",{"default":1.00}), "texture_guidance_rescale": ("FLOAT",{"default":0.00}), "texture_rescale_t": ("FLOAT",{"default":3.00}), "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":128,"step":8}), "generate_texture_slat": ("BOOLEAN", {"default":True}), "sparse_structure_guidance_interval_start": ("FLOAT",{"default":0.30,"min":0.00,"max":1.00,"step":0.01}), "sparse_structure_guidance_interval_end": ("FLOAT",{"default":1.00,"min":0.00,"max":1.00,"step":0.01}), "shape_guidance_interval_start": ("FLOAT",{"default":0.30,"min":0.00,"max":1.00,"step":0.01}), "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}), }, } RETURN_TYPES = ("MESHWITHVOXEL", ) RETURN_NAMES = ("mesh", ) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, pipeline, image, seed, pipeline_type, sparse_structure_steps, sparse_structure_guidance_strength, sparse_structure_guidance_rescale, sparse_structure_rescale_t, 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, max_views, sparse_structure_resolution, generate_texture_slat, sparse_structure_guidance_interval_start, sparse_structure_guidance_interval_end, shape_guidance_interval_start, shape_guidance_interval_end, texture_guidance_interval_start, 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 sparse_structure_guidance_interval = [sparse_structure_guidance_interval_start,sparse_structure_guidance_interval_end] shape_guidance_interval = [shape_guidance_interval_start,shape_guidance_interval_end] texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] sparse_structure_sampler_params = {"steps":sparse_structure_steps,"guidance_strength":sparse_structure_guidance_strength,"guidance_rescale":sparse_structure_guidance_rescale,"guidance_interval":sparse_structure_guidance_interval,"rescale_t":sparse_structure_rescale_t} 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} if generate_texture_slat: num_steps = 5 else: num_steps = 4 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, use_tiled=use_tiled_decoder, pbar=pbar)[0] return (mesh,) class Trellis2PostProcessAndUnWrapAndRasterizer: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":90.0,"min":0.0,"max":359.9}), "mesh_cluster_refine_iterations": ("INT",{"default":0}), "mesh_cluster_global_iterations": ("INT",{"default":1}), "mesh_cluster_smooth_strength": ("INT",{"default":1}), "texture_size": ("INT",{"default":2048, "min":512, "max":16384}), "remesh": ("BOOLEAN",{"default":True}), "remesh_band": ("FLOAT",{"default":1.0}), "remesh_project": ("FLOAT",{"default":0.0}), "target_face_num": ("INT",{"default":2000000,"min":1,"max":16000000}), "simplify_method": (["Cumesh","Meshlib"],{"default":"Cumesh"}), "fill_holes": ("BOOLEAN", {"default":True}), "fill_holes_max_perimeter": ("FLOAT",{"default":0.03,"min":0.001,"max":99.999,"step":0.001}), "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), "dual_contouring_resolution": (["Auto","128","256","512","1024","2048"],{"default":"512"}), "double_side_material": ("BOOLEAN",{"default":True}), "remove_floaters": ("BOOLEAN",{"default":True}), "bake_on_vertices": ("BOOLEAN",{"default":False}), "use_custom_normals":("BOOLEAN",{"default":False}), }, } RETURN_TYPES = ("TRIMESH","IMAGE","IMAGE",) RETURN_NAMES = ("trimesh","base_color_texture","metallic_roughness_texture",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, mesh_cluster_threshold_cone_half_angle_rad, mesh_cluster_refine_iterations, mesh_cluster_global_iterations, mesh_cluster_smooth_strength, texture_size, remesh, remesh_band, remesh_project, target_face_num, simplify_method, fill_holes, fill_holes_max_perimeter, texture_alpha_mode, dual_contouring_resolution, double_side_material, remove_floaters, bake_on_vertices=False,use_custom_normals=False): pbar = ProgressBar(5 if not bake_on_vertices else 4) mesh_copy = copy.deepcopy(mesh) aabb = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]] attr_volume = mesh_copy.attrs coords = mesh_copy.coords attr_layout = mesh_copy.layout voxel_size = mesh_copy.voxel_size mesh_cluster_threshold_cone_half_angle_rad = np.radians(mesh_cluster_threshold_cone_half_angle_rad) # --- Input Normalization (AABB, Voxel Size, Grid Size) --- if isinstance(aabb, (list, tuple)): aabb = np.array(aabb) if isinstance(aabb, np.ndarray): aabb = torch.tensor(aabb, dtype=torch.float32, device=coords.device) # Calculate grid dimensions based on AABB and voxel size if voxel_size is not None: if isinstance(voxel_size, float): voxel_size = [voxel_size, voxel_size, voxel_size] 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) grid_size = ((aabb[1] - aabb[0]) / voxel_size).round().int() else: if isinstance(grid_size, int): grid_size = [grid_size, grid_size, grid_size] 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) voxel_size = (aabb[1] - aabb[0]) / grid_size if remove_floaters: mesh_copy = remove_floater(mesh_copy) vertices = mesh_copy.vertices faces = mesh_copy.faces vertices = vertices.cuda() faces = faces.cuda() # Initialize CUDA mesh handler cumesh = CuMesh.CuMesh() cumesh.init(vertices, faces) print(f"Current vertices: {cumesh.num_vertices}, faces: {cumesh.num_faces}") # --- Initial Mesh Cleaning --- # Fills holes as much as we can before processing if fill_holes: cumesh.fill_holes(max_hole_perimeter=fill_holes_max_perimeter) print(f"After filling holes: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") # Build BVH for the current mesh to guide remeshing print(f"Building BVH for current mesh...") bvh = CuMesh.cuBVH(vertices, faces) pbar.update(1) print("Cleaning mesh...") # --- Branch 1: Standard Pipeline (Simplification & Cleaning) --- if not remesh: if simplify_method == 'Cumesh': cumesh.simplify(target_face_num * 3, verbose=True) elif simplify_method == 'Meshlib': # GPU -> CPU -> Meshlib -> CPU -> GPU v, f = cumesh.read() new_vertices, new_faces = simplify_with_meshlib(v.cpu().numpy(), f.cpu().numpy(), target_face_num) cumesh.init(torch.from_numpy(new_vertices).float().cuda(), torch.from_numpy(new_faces).int().cuda()) cumesh.remove_duplicate_faces() cumesh.repair_non_manifold_edges() cumesh.remove_small_connected_components(1e-5) if fill_holes: cumesh.fill_holes(max_hole_perimeter=fill_holes_max_perimeter) if simplify_method == 'Cumesh': cumesh.simplify(target_face_num, verbose=True) elif simplify_method == 'Meshlib': # GPU -> CPU -> Meshlib -> CPU -> GPU v, f = cumesh.read() new_vertices, new_faces = simplify_with_meshlib(v.cpu().numpy(), f.cpu().numpy(), target_face_num) cumesh.init(torch.from_numpy(new_vertices).float().cuda(), torch.from_numpy(new_faces).int().cuda()) cumesh.remove_duplicate_faces() cumesh.repair_non_manifold_edges() cumesh.remove_small_connected_components(1e-5) if fill_holes: cumesh.fill_holes(max_hole_perimeter=fill_holes_max_perimeter) print(f"After initial cleanup: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") # Step 2: Unify face orientations cumesh.unify_face_orientations() # --- Branch 2: Remeshing Pipeline --- else: center = aabb.mean(dim=0) scale = (aabb[1] - aabb[0]).max().item() if dual_contouring_resolution == "Auto": resolution = grid_size.max().item() print(f"Dual Contouring resolution: {resolution}") else: resolution = int(dual_contouring_resolution) print('Performing Dual Contouring ...') # Perform Dual Contouring remeshing (rebuilds topology) cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc( vertices, faces, center = center, scale = scale, # old calculation : (resolution + 3 * remesh_band) / resolution * scale, resolution = resolution, band = remesh_band, project_back = remesh_project, # Snaps vertices back to original surface verbose = True, bvh = bvh, )) print(f"After remeshing: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") # Step 2: Unify face orientations #cumesh.unify_face_orientations() if simplify_method == 'Cumesh': cumesh.simplify(target_face_num, verbose=True) elif simplify_method == 'Meshlib': # GPU -> CPU -> Meshlib -> CPU -> GPU v, f = cumesh.read() new_vertices, new_faces = simplify_with_meshlib(v.cpu().numpy(), f.cpu().numpy(), target_face_num) cumesh.init(torch.from_numpy(new_vertices).float().cuda(), torch.from_numpy(new_faces).int().cuda()) print(f"After simplifying: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") pbar.update(1) # --- Branch: Bake On Vertices (skip UV unwrapping and texture creation) --- if bake_on_vertices: print('Baking colors on vertices...') out_vertices, out_faces = cumesh.read() out_vertices = out_vertices.cuda() out_faces = out_faces.cuda() cumesh.compute_vertex_normals() out_normals = cumesh.read_vertex_normals() # Map vertex positions back to original mesh for accurate attribute sampling # Use BVH to find the closest point on original mesh surface for more accurate colors _, face_id, uvw = bvh.unsigned_distance(out_vertices, return_uvw=True) orig_tri_verts = vertices[faces[face_id.long()]] # (N_verts, 3, 3) mapped_pos = (orig_tri_verts * uvw.unsqueeze(-1)).sum(dim=1) # Sample attributes at mapped positions from the voxel grid vertex_attrs = grid_sample_3d( attr_volume, torch.cat([torch.zeros_like(coords[:, :1]), coords], dim=-1), shape=torch.Size([1, attr_volume.shape[1], *grid_size.tolist()]), grid=((mapped_pos - aabb[0]) / voxel_size).reshape(1, -1, 3), mode='trilinear', ) # Extract base color and alpha per vertex (vertex_attrs shape: N_vertices x C) base_color_idx = attr_layout['base_color'] alpha_idx = attr_layout['alpha'] # Get RGB values and squeeze any extra dimensions to get (N, 3) vertex_colors_rgb = vertex_attrs[..., base_color_idx].cpu().numpy() vertex_colors_rgb = np.squeeze(vertex_colors_rgb) # Remove batch dims if any if vertex_colors_rgb.ndim == 1: vertex_colors_rgb = vertex_colors_rgb[None, :] # Ensure at least 2D vertex_colors_rgb = np.clip(vertex_colors_rgb * 255, 0, 255).astype(np.uint8) # Handle alpha based on texture_alpha_mode if texture_alpha_mode == "OPAQUE": # For OPAQUE mode, use full alpha (255) vertex_alpha = np.full((vertex_colors_rgb.shape[0], 1), 255, dtype=np.uint8) else: vertex_alpha = vertex_attrs[..., alpha_idx].cpu().numpy() vertex_alpha = np.squeeze(vertex_alpha) # Remove batch dims if any vertex_alpha = np.clip(vertex_alpha * 255, 0, 255).astype(np.uint8) # Ensure alpha is 2D with shape (N, 1) if vertex_alpha.ndim == 1: vertex_alpha = vertex_alpha[:, None] # Combine into RGBA vertex_colors_rgba = np.concatenate([vertex_colors_rgb, vertex_alpha], axis=-1) print("Finalizing mesh with vertex colors...") pbar.update(1) vertices_np = out_vertices.cpu().numpy() faces_np = out_faces.cpu().numpy() normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy() normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy() # Create mesh with vertex colors using ColorVisuals if use_custom_normals: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_normals=normals_np, vertex_colors=vertex_colors_rgba, process=False, ) else: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_colors=vertex_colors_rgba, process=False, ) del cumesh gc.collect() # Return empty placeholder textures for vertex color mode placeholder_texture = pil2tensor(Image.new('RGBA', (1, 1), (0, 0, 0, 0))) return (textured_mesh, placeholder_texture, placeholder_texture,) # --- Standard texture baking path --- print('Unwrapping ...') out_vertices, out_faces, out_uvs, out_vmaps = cumesh.uv_unwrap( compute_charts_kwargs={ "threshold_cone_half_angle_rad": mesh_cluster_threshold_cone_half_angle_rad, "refine_iterations": mesh_cluster_refine_iterations, "global_iterations": mesh_cluster_global_iterations, "smooth_strength": mesh_cluster_smooth_strength, }, return_vmaps=True, verbose=True, ) pbar.update(1) out_vertices = out_vertices.cuda() out_faces = out_faces.cuda() out_uvs = out_uvs.cuda() out_vmaps = out_vmaps.cuda() cumesh.compute_vertex_normals() out_normals = cumesh.read_vertex_normals()[out_vmaps] print("Sampling attributes...") # Setup differentiable rasterizer context ctx = dr.RasterizeCudaContext() # Prepare UV coordinates for rasterization (rendering in UV space) uvs_rast = torch.cat([out_uvs * 2 - 1, torch.zeros_like(out_uvs[:, :1]), torch.ones_like(out_uvs[:, :1])], dim=-1).unsqueeze(0) rast = torch.zeros((1, texture_size, texture_size, 4), device='cuda', dtype=torch.float32) # Rasterize in chunks to save memory for i in range(0, out_faces.shape[0], 100000): rast_chunk, _ = dr.rasterize( ctx, uvs_rast, out_faces[i:i+100000], resolution=[texture_size, texture_size], ) mask_chunk = rast_chunk[..., 3:4] > 0 rast_chunk[..., 3:4] += i # Store face ID in alpha channel rast = torch.where(mask_chunk, rast_chunk, rast) # Mask of valid pixels in texture mask = rast[0, ..., 3] > 0 # Interpolate 3D positions in UV space (finding 3D coord for every texel) pos = dr.interpolate(out_vertices.unsqueeze(0), rast, out_faces)[0][0] valid_pos = pos[mask] # Map these positions back to the *original* high-res mesh to get accurate attributes # This corrects geometric errors introduced by simplification/remeshing _, face_id, uvw = bvh.unsigned_distance(valid_pos, return_uvw=True) orig_tri_verts = vertices[faces[face_id.long()]] # (N_new, 3, 3) valid_pos = (orig_tri_verts * uvw.unsqueeze(-1)).sum(dim=1) # Trilinear sampling from the attribute volume (Color, Material props) attrs = torch.zeros(texture_size, texture_size, attr_volume.shape[1], device='cuda') attrs[mask] = grid_sample_3d( attr_volume, torch.cat([torch.zeros_like(coords[:, :1]), coords], dim=-1), shape=torch.Size([1, attr_volume.shape[1], *grid_size.tolist()]), grid=((valid_pos - aabb[0]) / voxel_size).reshape(1, -1, 3), mode='trilinear', ) # --- Texture Post-Processing & Material Construction --- print("Finalizing mesh...") pbar.update(1) mask = mask.cpu().numpy() # Extract channels based on layout (BaseColor, Metallic, Roughness, Alpha) base_color = np.clip(attrs[..., attr_layout['base_color']].cpu().numpy() * 255, 0, 255).astype(np.uint8) metallic = np.clip(attrs[..., attr_layout['metallic']].cpu().numpy() * 255, 0, 255).astype(np.uint8) roughness = np.clip(attrs[..., attr_layout['roughness']].cpu().numpy() * 255, 0, 255).astype(np.uint8) alpha = np.clip(attrs[..., attr_layout['alpha']].cpu().numpy() * 255, 0, 255).astype(np.uint8) alpha_mode = texture_alpha_mode # Inpainting: fill gaps (dilation) to prevent black seams at UV boundaries mask_inv = (~mask).astype(np.uint8) base_color = cv2.inpaint(base_color, mask_inv, 3, cv2.INPAINT_TELEA) metallic = cv2.inpaint(metallic, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] roughness = cv2.inpaint(roughness, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] alpha = cv2.inpaint(alpha, mask_inv, 1, cv2.INPAINT_TELEA)[..., None] # Create PBR material # Standard PBR packs Metallic and Roughness into Blue and Green channels baseColorTexture_np = Image.fromarray(np.concatenate([base_color, alpha], axis=-1)) metallicRoughnessTexture_np = Image.fromarray(np.concatenate([np.zeros_like(metallic), roughness, metallic], axis=-1)) material = Trimesh.visual.material.PBRMaterial( baseColorTexture=baseColorTexture_np, baseColorFactor=np.array([255, 255, 255, 255], dtype=np.uint8), metallicRoughnessTexture=metallicRoughnessTexture_np, metallicFactor=1.0, roughnessFactor=1.0, alphaMode=alpha_mode, doubleSided=double_side_material, ) vertices_np = out_vertices.cpu().numpy() faces_np = out_faces.cpu().numpy() uvs_np = out_uvs.cpu().numpy() normals_np = out_normals.cpu().numpy() # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy() normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy() uvs_np[:, 1] = 1 - uvs_np[:, 1] # Flip UV V-coordinate if use_custom_normals: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, vertex_normals=normals_np, process=False, visual=Trimesh.visual.TextureVisuals(uv=uvs_np,material=material) ) else: textured_mesh = Trimesh.Trimesh( vertices=vertices_np, faces=faces_np, process=False, visual=Trimesh.visual.TextureVisuals(uv=uvs_np,material=material) ) pbar.update(1) del cumesh gc.collect() baseColorTexture = pil2tensor(baseColorTexture_np) metallicRoughnessTexture = pil2tensor(metallicRoughnessTexture_np) return (textured_mesh, baseColorTexture, metallicRoughnessTexture,) class Trellis2Remesh: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "remesh_band": ("FLOAT",{"default":1.0}), "remesh_project": ("FLOAT",{"default":0.0}), "fill_holes": ("BOOLEAN", {"default":True}), "fill_holes_max_perimeter": ("FLOAT",{"default":0.03,"min":0.001,"max":99.999,"step":0.001}), "dual_contouring_resolution": (["Auto","128","256","512","1024","2048"],{"default":"Auto"}), "remove_floaters": ("BOOLEAN",{"default":True}), }, } RETURN_TYPES = ("MESHWITHVOXEL",) RETURN_NAMES = ("mesh",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, remesh_band, remesh_project, fill_holes, fill_holes_max_perimeter, dual_contouring_resolution, remove_floaters): mesh_copy = copy.deepcopy(mesh) if remove_floaters: mesh_copy = remove_floater(mesh_copy) aabb = [[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]] vertices = mesh_copy.vertices faces = mesh_copy.faces attr_volume = mesh_copy.attrs coords = mesh_copy.coords attr_layout = mesh_copy.layout voxel_size = mesh_copy.voxel_size # --- Input Normalization (AABB, Voxel Size, Grid Size) --- if isinstance(aabb, (list, tuple)): aabb = np.array(aabb) if isinstance(aabb, np.ndarray): aabb = torch.tensor(aabb, dtype=torch.float32, device='cuda') # Calculate grid dimensions based on AABB and voxel size if voxel_size is not None: if isinstance(voxel_size, float): voxel_size = [voxel_size, voxel_size, voxel_size] 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='cuda') grid_size = ((aabb[1] - aabb[0]) / voxel_size).round().int() else: if isinstance(grid_size, int): grid_size = [grid_size, grid_size, grid_size] 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='cuda') voxel_size = (aabb[1] - aabb[0]) / grid_size # Move data to GPU vertices = vertices.cuda() faces = faces.cuda() # Initialize CUDA mesh handler cumesh = CuMesh.CuMesh() cumesh.init(vertices, faces) print(f"Current vertices: {cumesh.num_vertices}, faces: {cumesh.num_faces}") # --- Initial Mesh Cleaning --- # Fills holes as much as we can before processing if fill_holes: cumesh.fill_holes(max_hole_perimeter=fill_holes_max_perimeter) print(f"After filling holes: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") vertices, faces = cumesh.read() # Build BVH for the current mesh to guide remeshing print(f"Building BVH for current mesh...") bvh = CuMesh.cuBVH(vertices, faces) print("Cleaning mesh...") center = aabb.mean(dim=0) scale = (aabb[1] - aabb[0]).max().item() if dual_contouring_resolution == "Auto": resolution = grid_size.max().item() print(f"Dual Contouring resolution: {resolution}") else: resolution = int(dual_contouring_resolution) print('Performing Dual Contouring ...') # Perform Dual Contouring remeshing (rebuilds topology) cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc( vertices, faces, center = center, scale = scale, # old calculation (resolution + 3 * remesh_band) / resolution * scale, resolution = resolution, band = remesh_band, project_back = remesh_project, # Snaps vertices back to original surface verbose = True, bvh = bvh, )) print(f"After remeshing: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces") # Step 2: Unify face orientations #cumesh.unify_face_orientations() new_vertices, new_faces = cumesh.read() mesh_copy.vertices = new_vertices.to(mesh_copy.device) mesh_copy.faces = new_faces.to(mesh_copy.device) del cumesh gc.collect() return (mesh_copy,) class Trellis2MeshTexturing: @classmethod def INPUT_TYPES(s): return { "required": { "pipeline": ("TRELLIS2PIPELINE",), "image": ("IMAGE",), "trimesh": ("TRIMESH",), "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), "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}), "resolution": ([512,1024],{"default":1024}), "texture_size": ("INT",{"default":2048,"min":512,"max":16384}), "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), "double_side_material": ("BOOLEAN",{"default":True}), "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}), "max_views": ("INT", {"default": 4, "min": 1, "max": 16}), }, } RETURN_TYPES = ("TRIMESH","IMAGE","IMAGE",) RETURN_NAMES = ("trimesh","base_color_texture","metallic_roughness_texture",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, pipeline, image, trimesh, seed, texture_steps, texture_guidance_strength, texture_guidance_rescale, texture_rescale_t, resolution, texture_size, texture_alpha_mode, double_side_material, texture_guidance_interval_start, texture_guidance_interval_end, max_views,): images = tensor_batch_to_pil_list(image, max_views=max_views) image_in = images[0] if len(images) == 1 else images #image = tensor2pil(image) texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] 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} textured_mesh, baseColorTexture_np, metallicRoughnessTexture_np = pipeline.texture_mesh(mesh=trimesh, image=image_in, seed=seed, tex_slat_sampler_params = tex_slat_sampler_params, resolution = resolution, texture_size = texture_size, texture_alpha_mode = texture_alpha_mode, double_side_material = double_side_material, max_views = max_views, ) baseColorTexture = pil2tensor(baseColorTexture_np) metallicRoughnessTexture = pil2tensor(metallicRoughnessTexture_np) return (textured_mesh, baseColorTexture, metallicRoughnessTexture, ) class Trellis2LoadMesh: @classmethod def INPUT_TYPES(s): return { "required": { "glb_path": ("STRING", {"default": "", "tooltip": "The glb path with mesh to load."}), } } RETURN_TYPES = ("TRIMESH",) RETURN_NAMES = ("trimesh",) OUTPUT_TOOLTIPS = ("The glb model with mesh to texturize.",) FUNCTION = "load" CATEGORY = "Trellis2Wrapper" DESCRIPTION = "Loads a glb model from the given path." def load(self, glb_path): if not os.path.exists(glb_path): glb_path = os.path.join(folder_paths.get_input_directory(), glb_path) trimesh = Trimesh.load(glb_path, force="mesh") return (trimesh,) class Trellis2PreProcessImage: @classmethod def INPUT_TYPES(s): return { "required": { "image": ("IMAGE",), } } RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("image",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" def process(self, image): image = tensor2pil(image) image = self.preprocess_image(image) image = pil2tensor(image) return (image,) def preprocess_image(self, input: Image.Image) -> Image.Image: """ Preprocess the input image. """ # if has alpha channel, use it directly; otherwise, remove background has_alpha = False if input.mode == 'RGBA': alpha = np.array(input)[:, :, 3] if not np.all(alpha == 255): has_alpha = True max_size = max(input.size) scale = min(1, 2048 / max_size) if scale < 1: input = input.resize((int(input.width * scale), int(input.height * scale)), Image.Resampling.LANCZOS) # if has_alpha: # output = input # else: # input = input.convert('RGB') # if self.low_vram: # self.rembg_model.to(self.device) # output = self.rembg_model(input) # if self.low_vram: # self.rembg_model.cpu() output = input output_np = np.array(output) alpha = output_np[:, :, 3] bbox = np.argwhere(alpha > 0.8 * 255) bbox = np.min(bbox[:, 1]), np.min(bbox[:, 0]), np.max(bbox[:, 1]), np.max(bbox[:, 0]) center = (bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2 size = max(bbox[2] - bbox[0], bbox[3] - bbox[1]) size = int(size * 1) bbox = center[0] - size // 2, center[1] - size // 2, center[0] + size // 2, center[1] + size // 2 output = output.crop(bbox) # type: ignore output = np.array(output).astype(np.float32) / 255 output = output[:, :, :3] * output[:, :, 3:4] output = Image.fromarray((output * 255).astype(np.uint8)) 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.50}), "shape_guidance_rescale": ("FLOAT",{"default":0.50}), "shape_rescale_t": ("FLOAT",{"default":3.00}), "texture_steps": ("INT",{"default":12, "min":1, "max":100},), "texture_guidance_strength": ("FLOAT",{"default":1.00}), "texture_guidance_rescale": ("FLOAT",{"default":0.00}), "texture_rescale_t": ("FLOAT",{"default":3.00}), "max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}), "generate_texture_slat": ("BOOLEAN", {"default":True}), "downsampling":([16,32,64],{"default":16}), "shape_guidance_interval_start": ("FLOAT",{"default":0.30,"min":0.00,"max":1.00,"step":0.01}), "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}), }, } 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, downsampling, shape_guidance_interval_start, shape_guidance_interval_end, texture_guidance_interval_start, texture_guidance_interval_end, use_tiled_decoder): image = tensor2pil(image) shape_guidance_interval = [shape_guidance_interval_start,shape_guidance_interval_end] texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end] 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, use_tiled=use_tiled_decoder)[0] return (mesh,) class Trellis2PostProcess2: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "fill_holes": ("BOOLEAN", {"default":True}), "fix_normals": ("BOOLEAN", {"default":False}), "fix_face_orientation": ("BOOLEAN", {"default":True}), "remove_duplicate_faces": ("BOOLEAN",{"default":True}), }, } RETURN_TYPES = ("MESHWITHVOXEL",) RETURN_NAMES = ("mesh",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, fill_holes, fix_normals, fix_face_orientation, remove_duplicate_faces,): mesh_copy = copy.deepcopy(mesh) vertices_np = mesh_copy.vertices.cpu().numpy() faces_np = mesh_copy.faces.cpu().numpy() trimesh = Trimesh.Trimesh(vertices=vertices_np,faces=faces_np) print(f"Initial mesh: {len(trimesh.faces)} faces") print(f"Is winding consistent? {trimesh.is_winding_consistent}") if fix_normals: print('Fixing normals ...') trimesh.fix_normals() if fix_face_orientation: if trimesh.is_watertight: print('Mesh is watertight, fixing inversion ...') Trimesh.repair.fix_inversion(trimesh) else: print('Mesh is not watertight, cannot fix inversion') if remove_duplicate_faces: print('Removing duplicate faces ...') trimesh.remove_duplicate_faces() if fill_holes: print('Filling holes ...') trimesh.fill_holes() new_vertices = torch.from_numpy(trimesh.vertices).float() new_faces = torch.from_numpy(trimesh.faces).int() mesh_copy.vertices = new_vertices.to(mesh_copy.device) mesh_copy.faces = new_faces.to(mesh_copy.device) del trimesh gc.collect() return (mesh_copy,) class Trellis2OvoxelExportToGLB: @classmethod def INPUT_TYPES(s): return { "required": { "mesh": ("MESHWITHVOXEL",), "resolution": ([512,1024],{"default":1024}), "texture_size": ([512,1024,2048,4096],{"default":2048}), "target_face_num": ("INT",{"default":2000000,"min":500,"max":16000000}), }, } RETURN_TYPES = ("TRIMESH",) RETURN_NAMES = ("trimesh",) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, mesh, resolution, texture_size, target_face_num): mesh_copy = copy.deepcopy(mesh) glb = o_voxel.postprocess.to_glb( vertices=mesh_copy.vertices, faces=mesh_copy.faces, attr_volume=mesh_copy.attrs, coords=mesh_copy.coords, attr_layout=mesh_copy.layout, grid_size=resolution, aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], decimation_target=target_face_num, texture_size=texture_size, remesh=True, remesh_band=1, remesh_project=0, use_tqdm=True, ) return (glb,) class Trellis2TrimeshToMeshWithVoxel: @classmethod def INPUT_TYPES(s): return { "required": { "trimesh": ("TRIMESH",), "resolution": ([512,1024],{"default":1024}), }, } RETURN_TYPES = ("MESHWITHVOXEL", ) RETURN_NAMES = ("mesh", ) FUNCTION = "process" CATEGORY = "Trellis2Wrapper" OUTPUT_NODE = True def process(self, trimesh, resolution): mesh_copy = trimesh.copy() mvoxel = self.get_voxelmesh_from_trimesh(mesh_copy, resolution) return (mvoxel,) def get_voxelmesh_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() del voxel_indices del dual_vertices del intersected gc.collect() pbr_attr_layout = { 'base_color': slice(0, 3), 'metallic': slice(3, 4), 'roughness': slice(4, 5), 'alpha': slice(5, 6), } mvoxel = MeshWithVoxel( vertices, faces, origin = [-0.5, -0.5, -0.5], voxel_size = 1 / resolution, coords = coords, attrs = None, voxel_shape = None, layout=pbr_attr_layout ) return mvoxel NODE_CLASS_MAPPINGS = { "Trellis2LoadModel": Trellis2LoadModel, "Trellis2MeshWithVoxelGenerator": Trellis2MeshWithVoxelGenerator, "Trellis2LoadImageWithTransparency": Trellis2LoadImageWithTransparency, "Trellis2SimplifyMesh": Trellis2SimplifyMesh, "Trellis2MeshWithVoxelToTrimesh": Trellis2MeshWithVoxelToTrimesh, "Trellis2ExportMesh": Trellis2ExportMesh, "Trellis2PostProcessMesh": Trellis2PostProcessMesh, "Trellis2UnWrapAndRasterizer": Trellis2UnWrapAndRasterizer, "Trellis2MeshWithVoxelAdvancedGenerator": Trellis2MeshWithVoxelAdvancedGenerator, "Trellis2PostProcessAndUnWrapAndRasterizer": Trellis2PostProcessAndUnWrapAndRasterizer, "Trellis2Remesh": Trellis2Remesh, "Trellis2MeshTexturing": Trellis2MeshTexturing, "Trellis2LoadMesh": Trellis2LoadMesh, "Trellis2PreProcessImage": Trellis2PreProcessImage, "Trellis2MeshRefiner": Trellis2MeshRefiner, "Trellis2PostProcess2": Trellis2PostProcess2, "Trellis2OvoxelExportToGLB": Trellis2OvoxelExportToGLB, "Trellis2TrimeshToMeshWithVoxel": Trellis2TrimeshToMeshWithVoxel, } NODE_DISPLAY_NAME_MAPPINGS = { "Trellis2LoadModel": "Trellis2 - LoadModel", "Trellis2MeshWithVoxelGenerator": "Trellis2 - Mesh With Voxel Generator", "Trellis2LoadImageWithTransparency": "Trellis2 - Load Image with Transparency", "Trellis2SimplifyMesh": "Trellis2 - Simplify Mesh", "Trellis2MeshWithVoxelToTrimesh": "Trellis2 - Mesh With Voxel To Trimesh", "Trellis2ExportMesh": "Trellis2 - Export Mesh", "Trellis2PostProcessMesh": "Trellis2 - PostProcess Mesh", "Trellis2UnWrapAndRasterizer": "Trellis2 - UV Unwrap and Rasterize", "Trellis2MeshWithVoxelAdvancedGenerator": "Trellis2 - Mesh With Voxel Advanced Generator", "Trellis2PostProcessAndUnWrapAndRasterizer": "Trellis2 - Post Process/UnWrap and Rasterize", "Trellis2Remesh": "Trellis2 - Remesh", "Trellis2MeshTexturing": "Trellis2 - Mesh Texturing", "Trellis2LoadMesh": "Trellis2 - Load Mesh", "Trellis2PreProcessImage": "Trellis2 - PreProcess Image", "Trellis2MeshRefiner": "Trellis2 - Mesh Refiner", "Trellis2PostProcess2": "Trellis2 - PostProcess Mesh 2", "Trellis2OvoxelExportToGLB": "Trellis2 - Ovoxel Export to GLB", "Trellis2TrimeshToMeshWithVoxel": "Trellis2 - Trimesh to Mesh with Voxel", }