Files
visualbruno-ComfyUI-Trellis2/nodes.py
T
Bruno Fargnoli ecf938f220 Fixed progress_bar in Fill_Holes + udpated cumesh + added new nodes
Added "Remesh with Quad" node
Added "Batch Simplify Mesh and Export" node
2026-02-08 20:14:48 +01:00

2655 lines
113 KiB
Python

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()
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
any = AnyType("*")
def parse_string_to_int_list(number_string):
"""
Parses a string containing comma-separated numbers into a list of integers.
Args:
number_string: A string containing comma-separated numbers (e.g., "20000,10000,5000").
Returns:
A list of integers parsed from the input string.
Returns an empty list if the input string is empty or None.
"""
if not number_string:
return []
try:
# Split the string by comma and convert each part to an integer
int_list = [int(num.strip()) for num in number_string.split(',')]
return int_list
except ValueError as e:
print(f"Error converting string to integer: {e}. Please ensure all values are valid numbers.")
return []
def reset_cuda():
# Force garbage collection of Python objects
gc.collect()
# Clear PyTorch CUDA cache
torch.cuda.empty_cache()
# Synchronize to ensure all GPU operations complete
torch.cuda.synchronize()
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<target:
return
settings = mrmeshpy.DecimateSettings()
faces_to_delete = current_faces_num - target
settings.maxDeletedFaces = faces_to_delete
settings.packMesh = True
print('Generating Meshlib Mesh ...')
mesh = mrmeshnumpy.meshFromFacesVerts(faces, vertices)
print('Packing Optimally ...')
mesh.packOptimally()
print('Decimating ...')
mrmeshpy.decimateMesh(mesh, settings)
new_vertices = mrmeshnumpy.getNumpyVerts(mesh)
new_faces = mrmeshnumpy.getNumpyFaces(mesh.topology)
print(f"Reduced faces, resulting in {len(new_vertices)} vertices and {len(new_faces)} faces")
return new_vertices, new_faces
def remove_floater(mesh):
print('Removing floater ...')
faces = mesh.faces.cpu().numpy()
print(f"Current faces: {len(faces)}")
mesh_set = pymeshlab.MeshSet()
mesh_pymeshlab = pymeshlab.Mesh(vertex_matrix=mesh.vertices.cpu().numpy(), face_matrix=faces)
mesh_set.add_mesh(mesh_pymeshlab, "converted_mesh")
mesh_set = pymeshlab_remove_floater(mesh_set)
mesh_pymeshlab = mesh_set.current_mesh()
new_faces = mesh_pymeshlab.face_matrix()
print(f"After removing floater: {len(new_faces)}")
new_vertices = torch.from_numpy(mesh_pymeshlab.vertex_matrix()).contiguous().float()
new_faces = torch.from_numpy(new_faces).contiguous().int()
mesh.vertices = new_vertices
mesh.faces = new_faces
return mesh
def remove_floater2(vertices, faces):
print('Removing floater ...')
#faces = faces.cpu().numpy()
print(f"Current faces: {len(faces)}")
mesh_set = pymeshlab.MeshSet()
mesh_pymeshlab = pymeshlab.Mesh(vertex_matrix=vertices, face_matrix=faces)
mesh_set.add_mesh(mesh_pymeshlab, "converted_mesh")
mesh_set = pymeshlab_remove_floater(mesh_set)
mesh_pymeshlab = mesh_set.current_mesh()
new_faces = mesh_pymeshlab.face_matrix()
print(f"After removing floater: {len(new_faces)}")
new_vertices = mesh_pymeshlab.vertex_matrix()
return new_vertices, new_faces
def remove_mesh_infinite_vertices(mesh):
print('Removing infinite vertices ...')
vertices = mesh.vertices.cpu().numpy()
faces = mesh.faces.cpu().numpy()
trimesh = Trimesh.Trimesh(vertices=vertices,faces=faces)
print(f"Original vertex count: {len(trimesh.vertices)}")
# Remove anything outside a reasonable bounding box
limit = 1e10
valid_mask = (np.abs(trimesh.vertices) < limit).all(axis=1)
trimesh.update_vertices(valid_mask)
# Removing vertices can leave "degenerate" faces or orphan nodes
trimesh.update_faces(trimesh.nondegenerate_faces())
trimesh.remove_unreferenced_vertices()
print(f"Cleaned vertex count: {len(trimesh.vertices)}")
new_vertices = torch.from_numpy(trimesh.vertices).float()
new_faces = torch.from_numpy(trimesh.faces).int()
mesh.vertices = new_vertices
mesh.faces = new_faces
return mesh
def pymeshlab_remove_floater(mesh: pymeshlab.MeshSet):
mesh.apply_filter("compute_selection_by_small_disconnected_components_per_face",
nbfaceratio=0.005)
mesh.apply_filter("compute_selection_transfer_face_to_vertex", inclusive=False)
mesh.apply_filter("meshing_remove_selected_vertices_and_faces")
return mesh
def _batched_unsigned_distance(bvh, positions, batch_size=100000, return_uvw=False):
"""
Batch unsigned_distance queries to avoid GPU kernel timeout on large meshes.
When processing high-resolution textures (e.g., 2048x2048 = ~4M pixels) on complex
meshes, a single BVH query can cause GPU watchdog timeout. This function splits
the query into smaller batches.
Args:
bvh: The BVH structure from cumesh
positions: (N, 3) tensor of query positions
batch_size: Maximum number of queries per batch (default 100K, matching
the rasterization chunk size used elsewhere in this file)
return_uvw: Whether to return barycentric coordinates
Returns:
Same as bvh.unsigned_distance()
"""
import torch
N = positions.shape[0]
if N <= batch_size:
return bvh.unsigned_distance(positions, return_uvw=return_uvw)
distances_list = []
face_id_list = []
uvw_list = [] if return_uvw else None
for i in range(0, N, batch_size):
end = min(i + batch_size, N)
d, f, u = bvh.unsigned_distance(positions[i:end], return_uvw=return_uvw)
distances_list.append(d)
face_id_list.append(f)
if return_uvw:
uvw_list.append(u)
return (
torch.cat(distances_list),
torch.cat(face_id_list),
torch.cat(uvw_list) if return_uvw else None
)
class Trellis2LoadModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"modelname": (["TRELLIS.2-4B"],),
"backend": (["flash_attn","xformers"],{"default":"xformers"}),
"device": (["cpu","cuda"],{"default":"cuda"}),
"low_vram": ("BOOLEAN",{"default":True}),
"keep_models_loaded": ("BOOLEAN", {"default":True}),
},
}
RETURN_TYPES = ("TRELLIS2PIPELINE", )
RETURN_NAMES = ("pipeline", )
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, modelname, backend, device, low_vram, keep_models_loaded):
os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1'
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" # Can save GPU memory
#os.environ["FLEX_GEMM_AUTOTUNE_CACHE_PATH"] = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'autotune_cache.json')
#os.environ["FLEX_GEMM_AUTOTUNER_VERBOSE"] = '1'
os.environ['ATTN_BACKEND'] = backend
reset_cuda()
torch.backends.cudnn.benchmark = False
download_path = os.path.join(folder_paths.models_dir,"microsoft")
model_path = os.path.join(download_path, modelname)
hf_model_name = f"microsoft/{modelname}"
if not os.path.exists(model_path):
print(f"Downloading model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=hf_model_name,
local_dir=model_path,
local_dir_use_symlinks=False,
)
dinov3_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m","model.safetensors")
if not os.path.exists(dinov3_model_path):
raise Exception("Facebook Dinov3 model not found in models/facebook/dinov3-vitl16-pretrain-lvd1689m folder")
trellis_image_large_path = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16.safetensors")
if not os.path.exists(trellis_image_large_path):
print('Trellis-Image-Large ss_dec_conv3d_16l8_fp16 files not found. Trying to download the files from huggingface ...')
import requests
url = "https://huggingface.co/microsoft/TRELLIS-image-large/resolve/main/ckpts/ss_dec_conv3d_16l8_fp16.json?download=true"
filename = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16.json")
path = Path(filename)
path.parent.mkdir(parents=True, exist_ok=True)
response = requests.get(url)
if response.status_code == 200:
with open(filename, "wb") as f:
f.write(response.content)
print("Download ss_dec_conv3d_16l8_fp16.json complete!")
else:
raise Exception("Cannot download Trellis-Image-Large file ss_dec_conv3d_16l8_fp16.json")
url = "https://huggingface.co/microsoft/TRELLIS-image-large/resolve/main/ckpts/ss_dec_conv3d_16l8_fp16.safetensors?download=true"
filename = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16.safetensors")
response = requests.get(url)
if response.status_code == 200:
with open(filename, "wb") as f:
f.write(response.content)
print("Download ss_dec_conv3d_16l8_fp16.safetensors complete!")
else:
raise Exception("Cannot download Trellis-Image-Large file ss_dec_conv3d_16l8_fp16.safetensors")
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path, keep_models_loaded = keep_models_loaded)
pipeline.low_vram = low_vram
if device=="cuda":
if low_vram:
pipeline.cuda()
else:
pipeline.to(device)
else:
pipeline.to(device)
return (pipeline,)
class Trellis2MeshWithVoxelGenerator:
@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},),
"shape_steps": ("INT",{"default":12, "min":1, "max":100},),
"texture_steps": ("INT",{"default":12, "min":1, "max":100},),
"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}),
"use_tiled_decoder": ("BOOLEAN", {"default":True}),
},
}
RETURN_TYPES = ("MESHWITHVOXEL", "BVH", )
RETURN_NAMES = ("mesh", "bvh", )
FUNCTION = "process"
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, use_tiled_decoder):
reset_cuda()
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
sparse_structure_sampler_params = {"steps":sparse_structure_steps}
shape_slat_sampler_params = {"steps":shape_steps}
tex_slat_sampler_params = {"steps":texture_steps}
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]
vertices = mesh.vertices.cuda()
faces = mesh.faces.cuda()
# Build BVH for the current mesh to guide remeshing
if generate_texture_slat:
print("Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
bvh.vertices = vertices.detach().clone()
bvh.faces = faces.detach().clone()
else:
print("Not building BVH : only used for texturing")
bvh = None
return (mesh, bvh,)
class Trellis2LoadImageWithTransparency:
@classmethod
def INPUT_TYPES(s):
input_dir = folder_paths.get_input_directory()
files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
files = folder_paths.filter_files_content_types(files, ["image"])
return {"required":
{"image": (sorted(files), {"image_upload": True})},
}
CATEGORY = "Trellis2Wrapper"
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", )
RETURN_NAMES = ("image", "mask", "image_with_alpha")
FUNCTION = "load_image"
def load_image(self, image):
image_path = folder_paths.get_annotated_filepath(image)
img = node_helpers.pillow(Image.open, image_path)
output_images = []
output_masks = []
output_images_ori = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
i = node_helpers.pillow(ImageOps.exif_transpose, i)
output_images_ori.append(pil2tensor(i))
if i.mode == 'I':
i = i.point(lambda i: i * (1 / 255))
image = i.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
if 'A' in i.getbands():
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
elif i.mode == 'P' and 'transparency' in i.info:
mask = np.array(i.convert('RGBA').getchannel('A')).astype(np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 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 Trellis2SimplifyTrimesh:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"trimesh": ("TRIMESH",),
"target_face_num": ("INT",{"default":1000000,"min":1,"max":30000000}),
"method": (["Cumesh","Meshlib"],{"default":"Cumesh"}),
},
}
RETURN_TYPES = ("TRIMESH", )
RETURN_NAMES = ("trimesh", )
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, trimesh, target_face_num, method):
mesh_copy = copy.deepcopy(trimesh)
if method=="Cumesh":
cumesh = CuMesh.CuMesh()
cumesh.init(torch.from_numpy(mesh_copy.vertices).float().cuda(), torch.from_numpy(mesh_copy.faces).int().cuda())
cumesh.simplify(target_face_num, verbose=True)
new_vertices, new_faces = cumesh.read()
mesh_copy.vertices = new_vertices.cpu().numpy()
mesh_copy.faces = new_faces.cpu().numpy()
del cumesh
elif method=="Meshlib":
new_vertices, new_faces = simplify_with_meshlib(mesh_copy.vertices, mesh_copy.faces, target = target_face_num)
mesh_copy.vertices = new_vertices
mesh_copy.faces = new_faces
else:
raise Exception("Unknown simplification method")
return (mesh_copy,)
class Trellis2ProgressiveSimplify:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"max_edge_length": ("FLOAT",{"default":0.00,"min":0.00,"max":99999.99,"step":0.01}),
"max_triangle_aspect_ratio": ("FLOAT",{"default":20.00,"min":0.01,"max":99999.99,"step":0.01}),
"strategy": (["Minimal Error First","Shortest Edge First"],{"default":"Minimal Error First"}),
"stabilizer": ("FLOAT",{"default":0.000001,"min":0.0,"max":0.999999,"step":0.000001}),
"touch_near_boundary_edges": ("BOOLEAN",{"default":True}),
"optimize_vertex_positions": ("BOOLEAN",{"default":True}),
"angle_based_weights": ("BOOLEAN",{"default":False}),
},
"optional": {
"trimesh": ("TRIMESH",),
"mesh": ("MESHWITHVOXEL",),
}
}
RETURN_TYPES = ("TRIMESH", "MESHWITHVOXEL",)
RETURN_NAMES = ("trimesh", "mesh", )
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, max_edge_length, max_triangle_aspect_ratio, strategy, stabilizer, touch_near_boundary_edges, optimize_vertex_positions, angle_based_weights, trimesh = None, mesh = None):
if trimesh is not None:
trimesh = copy.deepcopy(trimesh)
vertices = trimesh.vertices
faces = trimesh.faces
vertices, faces = self.simplify(vertices, faces, max_edge_length, max_triangle_aspect_ratio, strategy, stabilizer, touch_near_boundary_edges, optimize_vertex_positions, angle_based_weights)
trimesh.vertices = vertices
trimesh.faces = faces
if mesh is not None:
mesh = copy.deepcopy(mesh)
vertices = mesh.vertices.cpu().numpy()
faces = mesh.faces.cpu().numpy()
vertices, faces = self.simplify(vertices, faces, max_edge_length, max_triangle_aspect_ratio, strategy, stabilizer, touch_near_boundary_edges, optimize_vertex_positions, angle_based_weights)
mesh.vertices = torch.from_numpy(vertices).float()
mesh.faces = torch.from_numpy(faces).int()
return (trimesh, mesh)
def simplify(self, vertices, faces, max_edge_length, max_triangle_aspect_ratio, strategy, stabilizer, touch_near_boundary_edges, optimize_vertex_positions, angle_based_weights):
current_faces_num = len(faces)
print(f'Current Faces Number: {current_faces_num}')
settings = mrmeshpy.DecimateSettings()
if strategy == "Minimal Error First":
settings.strategy = mrmeshpy.DecimateStrategy.MinimizeError
else:
settings.strategy = mrmeshpy.DecimateStrategy.ShortestEdgeFirst
settings.maxTriangleAspectRatio = max_triangle_aspect_ratio
settings.stabilizer = stabilizer
settings.touchNearBdEdges = touch_near_boundary_edges
settings.optimizeVertexPos = optimize_vertex_positions
settings.angleWeightedDistToPlane = angle_based_weights
settings.packMesh = True
print('Generating Meshlib Mesh ...')
mesh = mrmeshnumpy.meshFromFacesVerts(faces, vertices)
if max_edge_length == 0.0:
max_edge_length = 2.0
# for edge_id in mesh.topology.allValidEdges():
# edge_len = mesh.computeEdgeLen(edge_id)
# if edge_len > max_edge_length:
# max_edge_length = edge_len
# print(f"Calculated Max Edge Length: {max_edge_length}")
settings.maxEdgeLen = max_edge_length
settings.maxError = max_edge_length / 1000
print('Packing Optimally ...')
mesh.packOptimally()
print('Decimating ...')
mrmeshpy.decimateMesh(mesh, settings)
new_vertices = mrmeshnumpy.getNumpyVerts(mesh)
new_faces = mrmeshnumpy.getNumpyFaces(mesh.topology)
print(f"Reduced faces, resulting in {len(new_vertices)} vertices and {len(new_faces)} faces")
del mesh
gc.collect()
return new_vertices, new_faces
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":False}),
"fill_holes_max_perimeter": ("FLOAT",{"default":0.03,"min":0.001,"max":99.999,"step":0.001}),
"remove_duplicate_faces": ("BOOLEAN",{"default":False}),
"repair_non_manifold_edges": ("BOOLEAN", {"default":False}),
"remove_non_manifold_faces": ("BOOLEAN", {"default":False}),
"remove_small_connected_components": ("BOOLEAN", {"default":False}),
"remove_small_connected_components_size": ("FLOAT", {"default":0.00001,"min":0.00001,"max":9.99999,"step":0.00001}),
"unify_faces_orientation": ("BOOLEAN", {"default":False}),
"remove_floaters": ("BOOLEAN",{"default":False}),
"remove_infinite_vertices": ("BOOLEAN",{"default":False}),
"merge_vertices": ("BOOLEAN",{"default":False}),
"merge_distance": ("FLOAT",{"default":0.0010,"min":0.0001,"max":999.9999,"step":0.0001}),
"remove_nan_vertices": ("BOOLEAN",{"default":False}),
},
}
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,
merge_vertices,
merge_distance,
remove_nan_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()
if merge_vertices or remove_nan_vertices:
import open3d
open3d_mesh = open3d.geometry.TriangleMesh()
open3d_mesh.vertices = open3d.utility.Vector3dVector(vertices.cpu().numpy())
open3d_mesh.triangles = open3d.utility.Vector3iVector(faces.cpu().numpy().astype(np.int32))
# NaN check
print('Removing NaN vertices ...')
verts = np.asarray(open3d_mesh.vertices)
if np.any(np.isnan(verts)) or np.any(np.isinf(verts)):
print('NaN found. Cleaning them ...')
verts = np.nan_to_num(verts, nan=0.0, posinf=0.0, neginf=0.0)
open3d_mesh.vertices = open3d.utility.Vector3dVector(verts)
open3d_mesh = open3d_mesh.remove_duplicated_vertices()
open3d_mesh = open3d_mesh.remove_duplicated_triangles()
open3d_mesh = open3d_mesh.remove_degenerate_triangles()
open3d_mesh = open3d_mesh.remove_unreferenced_vertices()
#bbox = open3d_mesh.get_axis_aligned_bounding_box()
#max_extent = np.max(bbox.get_extent())
#safe_merge_distance = max_extent * 0.0005 # More conservative
#print(f"Auto-calculated merge distance: {safe_merge_distance:.6f}")
if merge_vertices:
# Merge and cleanup
open3d_mesh = open3d_mesh.merge_close_vertices(merge_distance)
open3d_mesh = open3d_mesh.remove_duplicated_vertices()
open3d_mesh = open3d_mesh.remove_duplicated_triangles()
open3d_mesh = open3d_mesh.remove_degenerate_triangles()
open3d_mesh = open3d_mesh.remove_unreferenced_vertices()
# Proper normal computation sequence
open3d_mesh.compute_triangle_normals()
open3d_mesh.compute_vertex_normals()
open3d_mesh.normalize_normals()
open3d_mesh.orient_triangles() # Orient based on computed normals
open3d_mesh.compute_vertex_normals() # Recompute after orientation
# Gentler smoothing
open3d_mesh = open3d_mesh.filter_smooth_taubin(number_of_iterations=3)
open3d_mesh.compute_vertex_normals() # Final recompute
cumesh.init(torch.from_numpy(np.asarray(open3d_mesh.vertices)).cuda().float(), torch.from_numpy(np.asarray(open3d_mesh.triangles)).cuda().int())
del open3d_mesh
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":60.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":4096, "min":512, "max":16384}),
"texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}),
"double_side_material": ("BOOLEAN",{"default":False}),
"bake_on_vertices": ("BOOLEAN",{"default":False}),
"use_custom_normals": ("BOOLEAN",{"default":False}),
"bvh": ("BVH",),
}
}
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,use_custom_normals,bvh):
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
# if bvh == None:
# print(f"Building BVH for current mesh...")
# bvh = CuMesh.cuBVH(vertices, faces)
# bvh.vertices = vertices
# bvh.faces = 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 = bvh.vertices[bvh.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, 1, 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": 12345, "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":6.50}),
"sparse_structure_guidance_rescale": ("FLOAT",{"default":0.20}),
"sparse_structure_rescale_t": ("FLOAT",{"default":4.00}),
"shape_steps": ("INT",{"default":12, "min":1, "max":100},),
"shape_guidance_strength": ("FLOAT",{"default":6.50}),
"shape_guidance_rescale": ("FLOAT",{"default":0.20}),
"shape_rescale_t": ("FLOAT",{"default":4.00}),
"texture_steps": ("INT",{"default":12, "min":1, "max":100},),
"texture_guidance_strength": ("FLOAT",{"default":3.00}),
"texture_guidance_rescale": ("FLOAT",{"default":0.20}),
"texture_rescale_t": ("FLOAT",{"default":3.00}),
"max_num_tokens": ("INT",{"default":999999,"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.10,"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.10,"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.00,"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","BVH", )
RETURN_NAMES = ("mesh", "bvh", )
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):
reset_cuda()
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]
vertices = mesh.vertices.cuda()
faces = mesh.faces.cuda()
if generate_texture_slat:
# Build BVH for the current mesh to guide remeshing
print("Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
bvh.vertices = vertices.detach().clone()
bvh.faces = faces.detach().clone()
else:
print("Not building BVH : only used for texturing")
bvh = None
return (mesh,bvh,)
class Trellis2PostProcessAndUnWrapAndRasterizer:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESHWITHVOXEL",),
"mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60.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":4096, "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}),
"texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}),
"dual_contouring_resolution": (["Auto","128","256","512","1024","2048"],{"default":"1024"}),
"double_side_material": ("BOOLEAN",{"default":False}),
"remove_floaters": ("BOOLEAN",{"default":True}),
"bake_on_vertices": ("BOOLEAN",{"default":False}),
"use_custom_normals":("BOOLEAN",{"default":False}),
"bvh": ("BVH",),
"remove_inner_faces": ("BOOLEAN",{"default":True}),
}
}
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, texture_alpha_mode, dual_contouring_resolution, double_side_material, remove_floaters, bake_on_vertices,use_custom_normals,bvh,remove_inner_faces):
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")
# vertices, faces = cumesh.read()
# BVH is coming from MeshWithVoxel Generator node
# print(f"Building BVH for current mesh...")
# bvh = CuMesh.cuBVH(vertices, faces)
# bvh.vertices = vertices
# bvh.faces = 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
print('Unifying faces orientation ...')
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 * 1.1, # 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,
remove_inner_faces = remove_inner_faces,
#bvh = bvh,
))
new_vertices, new_faces = cumesh.read()
if remove_floaters:
new_vertices, new_faces = remove_floater2(new_vertices.cpu().numpy(),new_faces.cpu().numpy())
new_vertices = torch.from_numpy(new_vertices).contiguous().float().cuda()
new_faces = torch.from_numpy(new_faces).contiguous().int().cuda()
cumesh.init(new_vertices, new_faces)
print(f"After remeshing: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces")
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)
if fill_holes:
new_vertices, new_faces = cumesh.read()
meshlib_mesh = mrmeshnumpy.meshFromFacesVerts(new_faces.detach().clone().cpu().numpy(), new_vertices.detach().clone().cpu().numpy())
hole_edges = meshlib_mesh.topology.findHoleRepresentiveEdges()
holes_filled = 0
nb_holes = len(hole_edges)
print(f"{nb_holes} holes found")
if nb_holes>0:
progress_bar_holes = tqdm(total=nb_holes,desc="Filling holes")
for e in hole_edges:
params = mrmeshpy.FillHoleParams()
params.metric = mrmeshpy.getUniversalMetric(meshlib_mesh)
mrmeshpy.fillHole(meshlib_mesh, e, params)
holes_filled += 1
progress_bar_holes.update(1)
progress_bar_holes.close()
new_vertices = mrmeshnumpy.getNumpyVerts(meshlib_mesh)
new_faces = mrmeshnumpy.getNumpyFaces(meshlib_mesh.topology)
del meshlib_mesh
gc.collect()
cumesh.init(torch.from_numpy(new_vertices).float().to(coords.device), torch.from_numpy(new_faces).int().to(coords.device))
# --- 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 = bvh.vertices[bvh.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, 1, 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}),
"remove_inner_faces": ("BOOLEAN",{"default":False}),
}
}
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, remove_inner_faces):
reset_cuda()
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()
del cumesh
gc.collect()
# Build BVH for the current mesh to guide remeshing
#print(f"Building BVH for current mesh...")
#bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
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)
vertices, faces = CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces,
center = center,
scale = scale * 1.1, # 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,
remove_inner_faces = remove_inner_faces,
#bvh = bvh,
)
if remove_floaters:
vertices, faces = remove_floater2(vertices.cpu().numpy(),faces.cpu().numpy())
vertices = torch.from_numpy(vertices).contiguous().float()
faces = torch.from_numpy(faces).contiguous().int()
print(f"After remeshing: {len(vertices)} vertices, {len(faces)} faces")
mesh_copy.vertices = vertices.to(mesh_copy.device)
mesh_copy.faces = faces.to(mesh_copy.device)
return (mesh_copy,)
class Trellis2ReconstructMesh:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESHWITHVOXEL",),
"remesh_band": ("FLOAT",{"default":1.0}),
"resolution": ([128,256,512,1024,2048],{"default":512}),
}
}
RETURN_TYPES = ("MESHWITHVOXEL",)
RETURN_NAMES = ("mesh",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, mesh, remesh_band, resolution):
reset_cuda()
mesh_copy = copy.deepcopy(mesh)
vertices = mesh_copy.vertices.cuda()
faces = mesh_copy.faces.cuda()
# Perform Dual Contouring remeshing (rebuilds topology)
print('Reconstructing mesh ...')
vertices, faces = CuMesh.remeshing.reconstruct_mesh_dc(vertices, faces, resolution, verbose=True)
print(f"After reconstruction: {len(vertices)} vertices, {len(faces)} faces")
mesh_copy.vertices = vertices.to(mesh_copy.device)
mesh_copy.faces = faces.to(mesh_copy.device)
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":3.0}),
"texture_guidance_rescale": ("FLOAT",{"default":0.2}),
"texture_rescale_t": ("FLOAT",{"default":3.0}),
"resolution": ([512,1024],{"default":1024}),
"texture_size": ("INT",{"default":4096,"min":512,"max":16384}),
"texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}),
"double_side_material": ("BOOLEAN",{"default":False}),
"texture_guidance_interval_start": ("FLOAT",{"default":0.00,"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}),
"bake_on_vertices": ("BOOLEAN",{"default":False}),
"use_custom_normals": ("BOOLEAN",{"default":False}),
"mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60.0,"min":0.0,"max":359.9}),
},
}
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,bake_on_vertices,use_custom_normals,mesh_cluster_threshold_cone_half_angle_rad):
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,
bake_on_vertices = bake_on_vertices,
use_custom_normals = use_custom_normals
)
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",),
"padding": ("INT",{"default":0,"min":0,"max":1024}),
"remove_background": ("BOOLEAN",{"default":False}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
def process(self, image, padding, remove_background):
image = tensor2pil(image)
if remove_background:
from rembg import remove
image = remove(image)
image = self.preprocess_image(image)
if padding>0:
border = (int(padding), int(padding), int(padding), int(padding))
fill_color = self.parse_fill_for_image("0,0,0,255", image)
image = ImageOps.expand(image,border=border,fill=fill_color)
image = pil2tensor(image)
return (image,)
def parse_fill_for_image(self, fill: str, img):
values = [int(x.strip()) for x in fill.split(",")]
if img.mode in ("L", "P"):
return values[0]
if img.mode == "RGB":
return tuple(values[:3])
if img.mode == "RGBA":
return tuple(values[:4])
raise ValueError(f"Unsupported image mode: {img.mode}")
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": 12345, "min": 0, "max": 0x7fffffff}),
"resolution": ([512,1024,1536],{"default":1024}),
"shape_steps": ("INT",{"default":12, "min":1, "max":100},),
"shape_guidance_strength": ("FLOAT",{"default":6.50}),
"shape_guidance_rescale": ("FLOAT",{"default":0.20}),
"shape_rescale_t": ("FLOAT",{"default":4.00}),
"texture_steps": ("INT",{"default":12, "min":1, "max":100},),
"texture_guidance_strength": ("FLOAT",{"default":3.00}),
"texture_guidance_rescale": ("FLOAT",{"default":0.20}),
"texture_rescale_t": ("FLOAT",{"default":3.00}),
"max_num_tokens": ("INT",{"default":999999,"min":0,"max":999999}),
"generate_texture_slat": ("BOOLEAN", {"default":True}),
"downsampling":([16,32,64],{"default":16}),
"shape_guidance_interval_start": ("FLOAT",{"default":0.10,"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.00,"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}),
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
},
}
RETURN_TYPES = ("MESHWITHVOXEL", "BVH", )
RETURN_NAMES = ("mesh", "bvh", )
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,
max_views):
images = tensor_batch_to_pil_list(image, max_views=max_views)
image_in = images[0] if len(images) == 1 else images
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_in, 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, max_views = max_views)[0]
vertices = mesh.vertices.cuda()
faces = mesh.faces.cuda()
# Build BVH for the current mesh to guide remeshing
if generate_texture_slat:
print("Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
bvh.vertices = vertices.detach().clone()
bvh.faces = faces.detach().clone()
else:
print('Not building BVH, only used for texturing')
bvh = None
return (mesh, bvh,)
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
class Trellis2Continue:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"input_1": (any,),
"input_2": (any,),
},
}
RETURN_TYPES = (any, any, )
RETURN_NAMES = ("output_1", "output_2", )
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, input_1, input_2):
return (input_1, input_2,)
class Trellis2MeshWithVoxelToMeshlibMesh:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESHWITHVOXEL",),
},
}
RETURN_TYPES = ("MESHLIB_MESH", )
RETURN_NAMES = ("meshlib_mesh",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, mesh):
meshlib_mesh = mrmeshnumpy.meshFromFacesVerts(mesh.faces.cpu().numpy(), mesh.vertices.cpu().numpy())
return (meshlib_mesh,)
class Trellis2FillHolesWithMeshlib:
"""Fill all holes in a mesh"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESHWITHVOXEL",),
},
}
RETURN_TYPES = ("MESHWITHVOXEL", "INT")
RETURN_NAMES = ("mesh", "holes_filled")
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
DESCRIPTION = "Fill all holes in a mesh using optimal triangulation."
def process(self, mesh):
import meshlib.mrmeshpy as mrmeshpy
mesh_copy = copy.deepcopy(mesh)
mesh = mrmeshnumpy.meshFromFacesVerts(mesh_copy.faces.detach().clone().cpu().numpy(), mesh_copy.vertices.detach().clone().cpu().numpy())
hole_edges = mesh.topology.findHoleRepresentiveEdges()
holes_filled = 0
nb_holes = len(hole_edges)
print(f"{nb_holes} holes found")
if nb_holes>0:
progress_bar = tqdm(total=nb_holes,desc="Filling holes")
pbar = ProgressBar(nb_holes)
for e in hole_edges:
params = mrmeshpy.FillHoleParams()
params.metric = mrmeshpy.getUniversalMetric(mesh)
mrmeshpy.fillHole(mesh, e, params)
holes_filled += 1
progress_bar.update(1)
pbar.update(1)
progress_bar.close()
new_vertices = mrmeshnumpy.getNumpyVerts(mesh)
new_faces = mrmeshnumpy.getNumpyFaces(mesh.topology)
del mesh
gc.collect()
mesh_copy.vertices = torch.from_numpy(new_vertices).float().to(mesh_copy.device)
mesh_copy.faces = torch.from_numpy(new_faces).int().to(mesh_copy.device)
return (mesh_copy, holes_filled)
class Trellis2SmoothNormals:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"trimesh": ("TRIMESH",),
},
}
RETURN_TYPES = ("TRIMESH",)
RETURN_NAMES = ("trimesh",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
def process(self, trimesh):
new_mesh = trimesh.copy()
new_mesh.vertex_normals = Trimesh.smoothing.get_vertices_normals(new_mesh)
return (new_mesh,)
class Trellis2RemeshWithQuad:
@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":False}),
"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}),
"remove_inner_faces": ("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, remove_inner_faces):
reset_cuda()
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()
del cumesh
gc.collect()
# Build BVH for the current mesh to guide remeshing
#print(f"Building BVH for current mesh...")
#bvh = CuMesh.cuBVH(vertices.detach().clone(), faces.detach().clone())
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)
vertices, faces = CuMesh.remeshing.remesh_narrow_band_dc_quad(
vertices, faces,
center = center,
scale = scale * 1.1, # 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,
remove_inner_faces = remove_inner_faces,
#bvh = bvh,
)
if remove_floaters:
vertices, faces = remove_floater2(vertices.cpu().numpy(),faces.cpu().numpy())
vertices = torch.from_numpy(vertices).contiguous().float()
faces = torch.from_numpy(faces).contiguous().int()
print(f"After remeshing: {len(vertices)} vertices, {len(faces)} faces")
mesh_copy.vertices = vertices.to(mesh_copy.device)
mesh_copy.faces = faces.to(mesh_copy.device)
return (mesh_copy,)
class Trellis2BatchSimplifyMeshAndExport:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mesh": ("MESHWITHVOXEL",),
"target_face_num": ("STRING",{"default":"2000000,1000000,500000,100000,50000,10000,5000,2500,1000"}),
"method": (["Cumesh","Meshlib"],{"default":"Cumesh"}),
"fill_holes":("BOOLEAN",{"default":True}),
"reorient_vertices":(["None","90 degrees","-90 degrees"],{"default":"90 degrees"}),
"filename_prefix":("STRING",),
"file_format": (["glb", "obj", "ply", "stl", "3mf", "dae"],),
},
}
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("lst_glb_path", )
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, mesh, target_face_num, method, fill_holes, reorient_vertices, filename_prefix, file_format):
lst_output_mesh = []
list_of_faces = parse_string_to_int_list(target_face_num)
if len(list_of_faces)>0:
cumesh = CuMesh.CuMesh()
mesh_copy = copy.deepcopy(mesh)
for target_nbfaces in list_of_faces:
print(f"Processing at {target_nbfaces} ...")
vertices = mesh_copy.vertices.detach().clone().cpu().numpy()
faces = mesh_copy.faces.detach().clone().cpu().numpy()
if method=="Cumesh":
cumesh.init(torch.from_numpy(vertices).float().cuda(), torch.from_numpy(faces).int().cuda())
cumesh.simplify(target_nbfaces, verbose=True)
vertices, faces = cumesh.read()
vertices = vertices.cpu().numpy()
faces = faces.cpu().numpy()
elif method=="Meshlib":
vertices, faces = simplify_with_meshlib(vertices, faces, target_nbfaces)
else:
raise Exception("Unknown simplification method")
if fill_holes:
import meshlib.mrmeshpy as mrmeshpy
mmesh = mrmeshnumpy.meshFromFacesVerts(faces, vertices)
hole_edges = mmesh.topology.findHoleRepresentiveEdges()
nb_holes = len(hole_edges)
print(f"{nb_holes} holes found")
if nb_holes>0:
progress_bar = tqdm(total=nb_holes,desc="Filling holes")
for e in hole_edges:
params = mrmeshpy.FillHoleParams()
params.metric = mrmeshpy.getUniversalMetric(mmesh)
mrmeshpy.fillHole(mmesh, e, params)
progress_bar.update(1)
progress_bar.close()
vertices = mrmeshnumpy.getNumpyVerts(mmesh)
faces = mrmeshnumpy.getNumpyFaces(mmesh.topology)
del mmesh
gc.collect()
if reorient_vertices == '90 degrees':
vertices[:, 1], vertices[:, 2] = vertices[:, 2], -vertices[:, 1]
elif reorient_vertices == '-90 degrees':
vertices[:, 1], vertices[:, 2] = -vertices[:, 2], vertices[:, 1]
trimesh = Trimesh.Trimesh(
vertices=vertices,
faces=faces,
process=False
)
filename_prefix_with_nbfaces = f"{filename_prefix}_{target_nbfaces}"
full_output_folder, filename, counter, subfolder, filename_prefix_with_nbfaces = folder_paths.get_save_image_path(filename_prefix_with_nbfaces, 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)
trimesh.export(output_glb_path, file_type=file_format)
lst_output_mesh.append(str(output_glb_path))
del trimesh
del cumesh
del mesh_copy
return (lst_output_mesh,)
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,
"Trellis2SimplifyTrimesh": Trellis2SimplifyTrimesh,
"Trellis2Continue": Trellis2Continue,
"Trellis2ProgressiveSimplify": Trellis2ProgressiveSimplify,
"Trellis2ReconstructMesh": Trellis2ReconstructMesh,
"Trellis2MeshWithVoxelToMeshlibMesh": Trellis2MeshWithVoxelToMeshlibMesh,
"Trellis2FillHolesWithMeshlib": Trellis2FillHolesWithMeshlib,
"Trellis2SmoothNormals": Trellis2SmoothNormals,
"Trellis2RemeshWithQuad": Trellis2RemeshWithQuad,
"Trellis2BatchSimplifyMeshAndExport": Trellis2BatchSimplifyMeshAndExport,
}
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",
"Trellis2SimplifyTrimesh": "Trellis2 - Simplify Trimesh",
"Trellis2Continue": "Trellis2 - Continue",
"Trellis2ProgressiveSimplify": "Trellis2 - Progressive Simplify",
"Trellis2ReconstructMesh": "Trellis2 - Reconstruct Mesh",
"Trellis2MeshWithVoxelToMeshlibMesh": "Trellis2 - Mesh with Voxel to Meshlib Mesh",
"Trellis2FillHolesWithMeshlib": "Trellis2 - Fill Holes with Meshlib",
"Trellis2SmoothNormals": "Trellis2 - Smooth Normals",
"Trellis2RemeshWithQuad": "Trellis2 - Remesh With Quad",
"Trellis2BatchSimplifyMeshAndExport": "Trellis2 - Batch Simplify Mesh And Export",
}