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