Can bypass "texture slat generation" for "Mesh Only" + Updated CuMesh package
This commit is contained in:
@@ -208,6 +208,7 @@ class Trellis2MeshWithVoxelGenerator:
|
||||
"max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}),
|
||||
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
|
||||
"sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":64,"step":8}),
|
||||
"generate_texture_slat": ("BOOLEAN", {"default":True}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -217,7 +218,7 @@ class Trellis2MeshWithVoxelGenerator:
|
||||
CATEGORY = "Trellis2Wrapper"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def process(self, pipeline, image, seed, pipeline_type, sparse_structure_steps, shape_steps, texture_steps, max_num_tokens, max_views, sparse_structure_resolution):
|
||||
def process(self, pipeline, image, seed, pipeline_type, sparse_structure_steps, shape_steps, texture_steps, max_num_tokens, max_views, sparse_structure_resolution, generate_texture_slat):
|
||||
images = tensor_batch_to_pil_list(image, max_views=max_views)
|
||||
image_in = images[0] if len(images) == 1 else images
|
||||
|
||||
@@ -225,7 +226,7 @@ class Trellis2MeshWithVoxelGenerator:
|
||||
shape_slat_sampler_params = {"steps":shape_steps}
|
||||
tex_slat_sampler_params = {"steps":texture_steps}
|
||||
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views)[0]
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views, generate_texture_slat = generate_texture_slat)[0]
|
||||
|
||||
return (mesh,)
|
||||
|
||||
@@ -681,6 +682,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
|
||||
"max_num_tokens": ("INT",{"default":49152,"min":0,"max":999999}),
|
||||
"max_views": ("INT", {"default": 4, "min": 1, "max": 16}),
|
||||
"sparse_structure_resolution": ("INT", {"default":32,"min":8,"max":64,"step":8}),
|
||||
"generate_texture_slat": ("BOOLEAN", {"default":True}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -704,7 +706,8 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
|
||||
texture_rescale_t,
|
||||
max_num_tokens,
|
||||
max_views,
|
||||
sparse_structure_resolution):
|
||||
sparse_structure_resolution,
|
||||
generate_texture_slat):
|
||||
|
||||
images = tensor_batch_to_pil_list(image, max_views=max_views)
|
||||
image_in = images[0] if len(images) == 1 else images
|
||||
@@ -713,7 +716,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
|
||||
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"rescale_t":shape_rescale_t}
|
||||
tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"rescale_t":texture_rescale_t}
|
||||
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views)[0]
|
||||
mesh = pipeline.run(image=image_in, seed=seed, pipeline_type=pipeline_type, sparse_structure_sampler_params = sparse_structure_sampler_params, shape_slat_sampler_params = shape_slat_sampler_params, tex_slat_sampler_params = tex_slat_sampler_params, max_num_tokens = max_num_tokens, sparse_structure_resolution = sparse_structure_resolution, max_views = max_views, generate_texture_slat=generate_texture_slat)[0]
|
||||
|
||||
return (mesh,)
|
||||
|
||||
@@ -1009,7 +1012,7 @@ class Trellis2Remesh:
|
||||
if isinstance(aabb, (list, tuple)):
|
||||
aabb = np.array(aabb)
|
||||
if isinstance(aabb, np.ndarray):
|
||||
aabb = torch.tensor(aabb, dtype=torch.float32, device=coords.device)
|
||||
aabb = torch.tensor(aabb, dtype=torch.float32, device='cuda')
|
||||
|
||||
# Calculate grid dimensions based on AABB and voxel size
|
||||
if voxel_size is not None:
|
||||
@@ -1018,7 +1021,7 @@ class Trellis2Remesh:
|
||||
if isinstance(voxel_size, (list, tuple)):
|
||||
voxel_size = np.array(voxel_size)
|
||||
if isinstance(voxel_size, np.ndarray):
|
||||
voxel_size = torch.tensor(voxel_size, dtype=torch.float32, device=coords.device)
|
||||
voxel_size = torch.tensor(voxel_size, dtype=torch.float32, device='cuda')
|
||||
grid_size = ((aabb[1] - aabb[0]) / voxel_size).round().int()
|
||||
else:
|
||||
if isinstance(grid_size, int):
|
||||
@@ -1026,7 +1029,7 @@ class Trellis2Remesh:
|
||||
if isinstance(grid_size, (list, tuple)):
|
||||
grid_size = np.array(grid_size)
|
||||
if isinstance(grid_size, np.ndarray):
|
||||
grid_size = torch.tensor(grid_size, dtype=torch.int32, device=coords.device)
|
||||
grid_size = torch.tensor(grid_size, dtype=torch.int32, device='cuda')
|
||||
voxel_size = (aabb[1] - aabb[0]) / grid_size
|
||||
|
||||
# Move data to GPU
|
||||
|
||||
@@ -762,24 +762,45 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
meshes, subs = self.decode_shape_slat(shape_slat, resolution)
|
||||
if self.low_vram:
|
||||
self._cleanup_cuda()
|
||||
tex_voxels = self.decode_tex_slat(tex_slat, subs)
|
||||
if self.low_vram:
|
||||
self._cleanup_cuda()
|
||||
out_mesh = []
|
||||
for m, v in zip(meshes, tex_voxels):
|
||||
m.fill_holes()
|
||||
out_mesh.append(
|
||||
MeshWithVoxel(
|
||||
m.vertices, m.faces,
|
||||
origin = [-0.5, -0.5, -0.5],
|
||||
voxel_size = 1 / resolution,
|
||||
coords = v.coords[:, 1:],
|
||||
attrs = v.feats,
|
||||
voxel_shape = torch.Size([*v.shape, *v.spatial_shape]),
|
||||
layout=self.pbr_attr_layout
|
||||
|
||||
if tex_slat is None:
|
||||
if self.low_vram:
|
||||
self._cleanup_cuda()
|
||||
out_mesh = []
|
||||
for m in meshes:
|
||||
m.fill_holes()
|
||||
out_mesh.append(
|
||||
MeshWithVoxel(
|
||||
m.vertices, m.faces,
|
||||
origin = [-0.5, -0.5, -0.5],
|
||||
voxel_size = 1 / resolution,
|
||||
coords = None,
|
||||
attrs = None,
|
||||
voxel_shape = None,
|
||||
layout=self.pbr_attr_layout
|
||||
)
|
||||
)
|
||||
)
|
||||
return out_mesh
|
||||
return out_mesh
|
||||
|
||||
else:
|
||||
tex_voxels = self.decode_tex_slat(tex_slat, subs)
|
||||
if self.low_vram:
|
||||
self._cleanup_cuda()
|
||||
out_mesh = []
|
||||
for m, v in zip(meshes, tex_voxels):
|
||||
m.fill_holes()
|
||||
out_mesh.append(
|
||||
MeshWithVoxel(
|
||||
m.vertices, m.faces,
|
||||
origin = [-0.5, -0.5, -0.5],
|
||||
voxel_size = 1 / resolution,
|
||||
coords = v.coords[:, 1:],
|
||||
attrs = v.feats,
|
||||
voxel_shape = torch.Size([*v.shape, *v.spatial_shape]),
|
||||
layout=self.pbr_attr_layout
|
||||
)
|
||||
)
|
||||
return out_mesh
|
||||
|
||||
@torch.no_grad()
|
||||
def run(
|
||||
@@ -795,7 +816,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
pipeline_type: Optional[str] = None,
|
||||
max_num_tokens: int = 49152,
|
||||
sparse_structure_resolution: int = 32,
|
||||
max_views: int = 4
|
||||
max_views: int = 4,
|
||||
generate_texture_slat = True
|
||||
) -> List[MeshWithVoxel]:
|
||||
"""
|
||||
Run the pipeline.
|
||||
@@ -875,12 +897,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
self.load_tex_slat_flow_model_512()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_512, self.models['tex_slat_flow_model_512'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
self.load_tex_slat_flow_model_512()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_512, self.models['tex_slat_flow_model_512'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
@@ -897,12 +920,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
@@ -923,12 +947,13 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
@@ -947,21 +972,28 @@ class Trellis2ImageTo3DPipeline(Pipeline):
|
||||
self.unload_shape_slat_flow_model_512()
|
||||
self.unload_shape_slat_flow_model_1024()
|
||||
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
if generate_texture_slat:
|
||||
self.unload_tex_slat_flow_model_512()
|
||||
self.load_tex_slat_flow_model_1024()
|
||||
tex_slat = self.sample_tex_slat(
|
||||
cond_1024, self.models['tex_slat_flow_model_1024'],
|
||||
shape_slat, tex_slat_sampler_params
|
||||
)
|
||||
|
||||
if not self.keep_models_loaded:
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
self.unload_tex_slat_flow_model_1024()
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
|
||||
if generate_texture_slat:
|
||||
out_mesh = self.decode_latent(shape_slat, tex_slat, res)
|
||||
else:
|
||||
out_mesh = self.decode_latent(shape_slat, None, res)
|
||||
torch.cuda.empty_cache()
|
||||
if return_latent:
|
||||
return out_mesh, (shape_slat, tex_slat, res)
|
||||
if generate_texture_slat:
|
||||
return out_mesh, (shape_slat, tex_slat, res)
|
||||
else:
|
||||
return out_mesh, (shape_slat, None, res)
|
||||
else:
|
||||
return out_mesh
|
||||
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Reference in New Issue
Block a user