Added the node "Trellis2ReconstructMesh" + "Multiple Images" support for "Mesh Refiner" node

This commit is contained in:
Bruno Fargnoli
2026-01-27 17:28:00 +01:00
parent 1577b5c748
commit fa93c4c067
11 changed files with 265 additions and 71 deletions
+249 -64
View File
@@ -352,10 +352,14 @@ class Trellis2MeshWithVoxelGenerator:
faces = mesh.faces.cuda()
# Build BVH for the current mesh to guide remeshing
print(f"Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
bvh.vertices = vertices
bvh.faces = faces
if generate_texture_slat:
print("Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
bvh.vertices = vertices
bvh.faces = faces
else:
print("Not building BVH : only used for texturing")
bvh = None
return (mesh, bvh,)
@@ -504,7 +508,92 @@ class Trellis2SimplifyTrimesh:
else:
raise Exception("Unknown simplification method")
return (mesh_copy,)
return (mesh_copy,)
class Trellis2ProgressiveSimplify:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"max_edge_length": ("FLOAT",{"default":2.64,"min":0.01,"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.maxEdgeLen = max_edge_length
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.maxError = max_edge_length / 1000
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")
del mesh
gc.collect()
return new_vertices, new_faces
class Trellis2MeshWithVoxelToTrimesh:
@classmethod
@@ -588,6 +677,9 @@ class Trellis2PostProcessMesh:
"unify_faces_orientation": ("BOOLEAN", {"default":True}),
"remove_floaters": ("BOOLEAN",{"default":True}),
"remove_infinite_vertices": ("BOOLEAN",{"default":True}),
"merge_vertices": ("BOOLEAN",{"default":True}),
"merge_distance": ("FLOAT",{"default":0.0010,"min":0.0001,"max":999.9999,"step":0.0001}),
"remove_nan_vertices": ("BOOLEAN",{"default":True}),
},
}
@@ -597,13 +689,28 @@ class Trellis2PostProcessMesh:
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):
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)
mesh_copy = remove_mesh_infinite_vertices(mesh_copy)
vertices = mesh_copy.vertices
faces = mesh_copy.faces
@@ -639,10 +746,54 @@ class Trellis2PostProcessMesh:
if unify_faces_orientation:
print('Unifying faces orientation ...')
cumesh.unify_face_orientations()
cumesh.unify_face_orientations()
print(f"After initial cleanup: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces")
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()
@@ -668,10 +819,8 @@ class Trellis2UnWrapAndRasterizer:
"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}),
},
"optional":{
"bvh": ("BVH",),
"use_custom_normals": ("BOOLEAN",{"default":False}),
"bvh": ("BVH",),
}
}
@@ -681,7 +830,7 @@ class Trellis2UnWrapAndRasterizer:
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,bvh=None):
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]]
@@ -728,11 +877,11 @@ class Trellis2UnWrapAndRasterizer:
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
# 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:
@@ -1028,11 +1177,15 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
vertices = mesh.vertices.cuda()
faces = mesh.faces.cuda()
# Build BVH for the current mesh to guide remeshing
print(f"Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
bvh.vertices = vertices
bvh.faces = faces
if generate_texture_slat:
# Build BVH for the current mesh to guide remeshing
print("Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
bvh.vertices = vertices
bvh.faces = faces
else:
print("Not building BVH : only used for texturing")
bvh = None
return (mesh,bvh,)
@@ -1060,10 +1213,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
"remove_floaters": ("BOOLEAN",{"default":True}),
"bake_on_vertices": ("BOOLEAN",{"default":False}),
"use_custom_normals":("BOOLEAN",{"default":False}),
},
"optional":{
"bvh": ("BVH",),
"bvh": ("BVH",),
}
}
@@ -1073,7 +1223,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
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,bvh=None):
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,use_custom_normals,bvh):
pbar = ProgressBar(5 if not bake_on_vertices else 4)
mesh_copy = copy.deepcopy(mesh)
@@ -1128,14 +1278,14 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
# 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")
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
if bvh == None:
print(f"Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
bvh.vertices = vertices
bvh.faces = faces
# 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)
@@ -1175,6 +1325,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
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 ---
@@ -1193,7 +1344,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces,
center = center,
scale = scale, # old calculation : (resolution + 3 * remesh_band) / resolution * scale,
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
@@ -1204,7 +1355,8 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
print(f"After remeshing: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces")
# Step 2: Unify face orientations
#cumesh.unify_face_orientations()
print('Unifying faces orientation ...')
cumesh.unify_face_orientations()
if simplify_method == 'Cumesh':
cumesh.simplify(target_face_num, verbose=True)
@@ -1445,9 +1597,6 @@ class Trellis2Remesh:
"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}),
},
"optional":{
"bvh": ("BVH",),
}
}
@@ -1457,7 +1606,7 @@ class Trellis2Remesh:
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, mesh, remesh_band, remesh_project, fill_holes, fill_holes_max_perimeter, dual_contouring_resolution, remove_floaters, bvh=None):
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:
@@ -1512,11 +1661,13 @@ class Trellis2Remesh:
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
if bvh == None:
print(f"Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
print(f"Building BVH for current mesh...")
bvh = CuMesh.cuBVH(vertices, faces)
print("Cleaning mesh...")
center = aabb.mean(dim=0)
@@ -1530,32 +1681,58 @@ class Trellis2Remesh:
print('Performing Dual Contouring ...')
# Perform Dual Contouring remeshing (rebuilds topology)
cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces = CuMesh.remeshing.remesh_narrow_band_dc(
vertices, faces,
center = center,
scale = scale, # old calculation (resolution + 3 * remesh_band) / resolution * scale,
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,
bvh = bvh,
))
)
print(f"After remeshing: {cumesh.num_vertices} vertices, {cumesh.num_faces} faces")
print(f"After remeshing: {len(vertices)} vertices, {len(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()
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):
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):
@@ -1746,6 +1923,7 @@ class Trellis2MeshRefiner:
"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}),
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
},
}
@@ -1771,9 +1949,11 @@ class Trellis2MeshRefiner:
shape_guidance_interval_end,
texture_guidance_interval_start,
texture_guidance_interval_end,
use_tiled_decoder):
use_tiled_decoder,
max_views):
image = tensor2pil(image)
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]
@@ -1781,7 +1961,7 @@ class Trellis2MeshRefiner:
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"guidance_interval":texture_guidance_interval,"rescale_t":texture_rescale_t}
mesh = pipeline.refine_mesh(mesh = trimesh, image=image, seed=seed, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, resolution = resolution, max_num_tokens = max_num_tokens, generate_texture_slat=generate_texture_slat, downsampling=downsampling, use_tiled=use_tiled_decoder)[0]
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()
@@ -1974,7 +2154,7 @@ class Trellis2Continue:
OUTPUT_NODE = True
def process(self, input_1, input_2):
return (input_1, input_2,)
return (input_1, input_2,)
NODE_CLASS_MAPPINGS = {
"Trellis2LoadModel": Trellis2LoadModel,
@@ -1997,7 +2177,10 @@ NODE_CLASS_MAPPINGS = {
"Trellis2TrimeshToMeshWithVoxel": Trellis2TrimeshToMeshWithVoxel,
"Trellis2SimplifyTrimesh": Trellis2SimplifyTrimesh,
"Trellis2Continue": Trellis2Continue,
"Trellis2ProgressiveSimplify": Trellis2ProgressiveSimplify,
"Trellis2ReconstructMesh": Trellis2ReconstructMesh,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"Trellis2LoadModel": "Trellis2 - LoadModel",
@@ -2020,4 +2203,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"Trellis2TrimeshToMeshWithVoxel": "Trellis2 - Trimesh to Mesh with Voxel",
"Trellis2SimplifyTrimesh": "Trellis2 - Simplify Trimesh",
"Trellis2Continue": "Trellis2 - Continue",
"Trellis2ProgressiveSimplify": "Trellis2 - Progressive Simplify",
"Trellis2ReconstructMesh": "Trellis2 - Reconstruct Mesh",
}
+2 -1
View File
@@ -2,4 +2,5 @@ meshlib
requests
pymeshlab
opencv-python
scipy
scipy
open3d
+14 -6
View File
@@ -1165,6 +1165,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
Returns:
SparseTensor: The encoded structured latent.
"""
print('Converting mesh to flexible dual grid ...')
vertices = torch.from_numpy(mesh.vertices).float()
faces = torch.from_numpy(mesh.faces).long()
@@ -1522,6 +1523,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
) -> SparseTensor:
# Upsample
self.load_shape_slat_decoder()
print('Decoding mesh slat ...')
if self.low_vram:
self.models['shape_slat_decoder'].to(self.device)
self.models['shape_slat_decoder'].low_vram = True
@@ -1535,7 +1537,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self.unload_shape_slat_decoder()
#downsampling = 16
lr_resolution = hr_resolution
lr_resolution = 512
# if hr_resolution == 512:
# downsampling = 16
# elif hr_resolution == 1024:
@@ -1546,7 +1548,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
while True:
quant_coords = torch.cat([
hr_coords[:, :1],
((hr_coords[:, 1:] + 0.5) / 512 * (lr_resolution // downsampling)).int(),
((hr_coords[:, 1:] + 0.5) / lr_resolution * (hr_resolution // downsampling)).int(),
], dim=1)
coords = quant_coords.unique(dim=0)
num_tokens = coords.shape[0]
@@ -1598,7 +1600,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
def refine_mesh(
self,
mesh: trimesh.Trimesh,
image: Image.Image,
image,
seed: int = 42,
shape_slat_sampler_params: dict = {},
tex_slat_sampler_params: dict = {},
@@ -1608,16 +1610,22 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return_latent = False,
downsampling = 16,
use_tiled: bool = True,
max_views: int = 4,
):
mesh = self.preprocess_mesh(mesh)
torch.manual_seed(seed)
self.load_image_cond_model()
if resolution == 512:
cond = self.get_cond(image, 512)
if isinstance(image, (list, tuple)):
images = list(image)
else:
cond = self.get_cond(image, 1024)
images = [image]
if resolution == 512:
cond = self.get_cond(images, 512, max_views = max_views)
else:
cond = self.get_cond(images, 1024, max_views = max_views)
if not self.keep_models_loaded:
self.unload_image_cond_model()