diff --git a/example_workflows/TextureMesh.json b/example_workflows/TextureMesh.json new file mode 100644 index 0000000..5eb9d10 --- /dev/null +++ b/example_workflows/TextureMesh.json @@ -0,0 +1,454 @@ +{ + "id": "9b1321c3-d206-474a-bf17-498b285fb45b", + "revision": 0, + "last_node_id": 12, + "last_link_id": 17, + "nodes": [ + { + "id": 1, + "type": "Trellis2LoadImageWithTransparency", + "pos": [ + 805.1940845502396, + 1197.36441919367 + ], + "size": [ + 392.34356316210614, + 466.5324210306162 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [] + }, + { + "name": "mask", + "type": "MASK", + "links": [] + }, + { + "name": "image_with_alpha", + "type": "IMAGE", + "links": [ + 15 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "e7f9b30df7a09bcedf1c955e176754c73f983254", + "Node name for S&R": "Trellis2LoadImageWithTransparency", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "Image_04298_2048.png", + "image" + ] + }, + { + "id": 3, + "type": "Trellis2MeshTexturing", + "pos": [ + 1308.169514665661, + 1217.733755068405 + ], + "size": [ + 314.03515625, + 314 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "pipeline", + "type": "TRELLIS2PIPELINE", + "link": 1 + }, + { + "name": "image", + "type": "IMAGE", + "link": 15 + }, + { + "name": "trimesh", + "type": "TRIMESH", + "link": 8 + } + ], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 11 + ] + }, + { + "name": "base_color_texture", + "type": "IMAGE", + "links": [ + 9 + ] + }, + { + "name": "metallic_roughness_texture", + "type": "IMAGE", + "links": [ + 10 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "20651320d20ed56a52168c1bd3f29d6886c4be06", + "Node name for S&R": "Trellis2MeshTexturing", + "widget_ue_connectable": {} + }, + "widgets_values": [ + 810478905, + "randomize", + 12, + 1, + 0, + 3, + 1024, + 2048, + "OPAQUE", + true + ] + }, + { + "id": 9, + "type": "PreviewImage", + "pos": [ + 1317.9756230852583, + 1591.0694221632657 + ], + "size": [ + 309.20638371079303, + 279.8094875474301 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 9 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.5.1", + "Node name for S&R": "PreviewImage", + "widget_ue_connectable": {} + }, + "widgets_values": [] + }, + { + "id": 12, + "type": "Preview3D", + "pos": [ + 2027.3925330734914, + 840.4085320617074 + ], + "size": [ + 960.9769683217987, + 1025.7966066516265 + ], + "flags": {}, + "order": 7, + "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": 12 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.5.1", + "Node name for S&R": "Preview3D", + "widget_ue_connectable": {}, + "Last Time Model File": "TexturedMesh_00001_.glb", + "Scene Config": { + "showGrid": true, + "backgroundColor": "#282828", + "backgroundImage": "", + "backgroundRenderMode": "tiled" + }, + "Camera Config": { + "cameraType": "perspective", + "fov": 35, + "state": { + "position": { + "x": 0.12885373643844386, + "y": 4.965867390532585, + "z": 8.22624562737759 + }, + "target": { + "x": 0, + "y": 2.5, + "z": 0 + }, + "zoom": 1, + "cameraType": "perspective" + } + }, + "Light Config": { + "intensity": 3 + } + }, + "widgets_values": [ + "TexturedMesh_00001_.glb", + "" + ] + }, + { + "id": 2, + "type": "Trellis2LoadModel", + "pos": [ + 881.0563207451194, + 957.3003110075967 + ], + "size": [ + 301.7261859434295, + 154 + ], + "flags": {}, + "order": 1, + "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": 8, + "type": "Trellis2LoadMesh", + "pos": [ + 794.1090714719911, + 1780.9456209238792 + ], + "size": [ + 428.6309297171438, + 63.28766538369291 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "links": [ + 8 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "20651320d20ed56a52168c1bd3f29d6886c4be06", + "Node name for S&R": "Trellis2LoadMesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "C:\\Git\\ComfyUI\\output\\DwarfSorcerer_Hy20_Hy3D_00001_.obj" + ] + }, + { + "id": 10, + "type": "PreviewImage", + "pos": [ + 1656.7729796760464, + 1590.5886284895685 + ], + "size": [ + 316.12846064390465, + 278.5596651672149 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 10 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.5.1", + "Node name for S&R": "PreviewImage", + "widget_ue_connectable": {} + }, + "widgets_values": [] + }, + { + "id": 11, + "type": "Trellis2ExportMesh", + "pos": [ + 1686.1918418492774, + 1215.162062552115 + ], + "size": [ + 270, + 106 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "trimesh", + "type": "TRIMESH", + "link": 11 + } + ], + "outputs": [ + { + "name": "glb_path", + "type": "STRING", + "links": [ + 12 + ] + } + ], + "properties": { + "aux_id": "visualbruno/ComfyUI-Trellis2", + "ver": "20651320d20ed56a52168c1bd3f29d6886c4be06", + "Node name for S&R": "Trellis2ExportMesh", + "widget_ue_connectable": {} + }, + "widgets_values": [ + "TexturedMesh", + "glb", + true + ] + } + ], + "links": [ + [ + 1, + 2, + 0, + 3, + 0, + "TRELLIS2PIPELINE" + ], + [ + 8, + 8, + 0, + 3, + 2, + "TRIMESH" + ], + [ + 9, + 3, + 1, + 9, + 0, + "IMAGE" + ], + [ + 10, + 3, + 2, + 10, + 0, + "IMAGE" + ], + [ + 11, + 3, + 0, + 11, + 0, + "TRIMESH" + ], + [ + 12, + 11, + 0, + 12, + 2, + "STRING" + ], + [ + 15, + 1, + 2, + 3, + 1, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "workflowRendererVersion": "LG", + "ue_links": [], + "ds": { + "scale": 0.6303940863128702, + "offset": [ + -558.8013640114224, + -570.1544608868662 + ] + }, + "links_added_by_ue": [], + "frontendVersion": "1.34.9", + "VHS_latentpreview": false, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index e456d80..e372879 100644 --- a/nodes.py +++ b/nodes.py @@ -31,6 +31,8 @@ from .trellis2.pipelines import Trellis2ImageTo3DPipeline script_directory = os.path.dirname(os.path.abspath(__file__)) comfy_path = os.path.dirname(os.path.dirname(os.path.dirname(__file__))) +to_pil = transforms.ToPILImage() + def pil2tensor(image): return torch.from_numpy(np.array(image).astype(np.float32) / 255.0)[None,] @@ -54,7 +56,7 @@ def tensor2pil(image: torch.Tensor) -> Image.Image: 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)}") + 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]: """ @@ -1082,7 +1084,78 @@ class Trellis2Remesh: del cumesh gc.collect() - return (mesh,) + return (mesh,) + +class Trellis2MeshTexturing: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "pipeline": ("TRELLIS2PIPELINE",), + "image": ("IMAGE",), + "trimesh": ("TRIMESH",), + "seed": ("INT", {"default": 0, "min": 0, "max": 0x7fffffff}), + "texture_steps": ("INT",{"default":12, "min":1, "max":100},), + "texture_guidance_strength": ("FLOAT",{"default":1.0}), + "texture_guidance_rescale": ("FLOAT",{"default":0.0}), + "texture_rescale_t": ("FLOAT",{"default":3.0}), + "resolution": ([512,1024],{"default":1024}), + "texture_size": ("INT",{"default":2048,"min":512,"max":16384}), + "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), + "double_side_material": ("BOOLEAN",{"default":True}), + }, + } + + 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): + #image = tensor2pil_v2(image) + image = tensor2pil(image) + tex_slat_sampler_params = {"steps":texture_steps,"guidance_strength":texture_guidance_strength,"guidance_rescale":texture_guidance_rescale,"rescale_t":texture_rescale_t} + + textured_mesh, baseColorTexture_np, metallicRoughnessTexture_np = pipeline.texture_mesh(mesh=trimesh, + image=image, + 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 + ) + + + 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,) NODE_CLASS_MAPPINGS = { "Trellis2LoadModel": Trellis2LoadModel, @@ -1096,6 +1169,8 @@ NODE_CLASS_MAPPINGS = { "Trellis2MeshWithVoxelAdvancedGenerator": Trellis2MeshWithVoxelAdvancedGenerator, "Trellis2PostProcessAndUnWrapAndRasterizer": Trellis2PostProcessAndUnWrapAndRasterizer, "Trellis2Remesh": Trellis2Remesh, + "Trellis2MeshTexturing": Trellis2MeshTexturing, + "Trellis2LoadMesh": Trellis2LoadMesh, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -1110,4 +1185,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "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", } \ No newline at end of file diff --git a/trellis2/pipelines/base.py b/trellis2/pipelines/base.py index 59ddffa..bc9bd2f 100644 --- a/trellis2/pipelines/base.py +++ b/trellis2/pipelines/base.py @@ -46,6 +46,8 @@ class Pipeline: # except Exception as e: # _models[k] = models.from_pretrained(v) + _models['shape_slat_encoder'] = None + new_pipeline = cls(_models) new_pipeline._pretrained_args = args return new_pipeline diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index dc73155..3254535 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -14,6 +14,12 @@ from .. import models import gc import os import folder_paths +import trimesh +import o_voxel +import cumesh +import nvdiffrast.torch as dr +import cv2 +import flex_gemm class Trellis2ImageTo3DPipeline(Pipeline): @@ -243,7 +249,20 @@ class Trellis2ImageTo3DPipeline(Pipeline): if self.models['tex_slat_flow_model_1024'] is not None: del self.models['tex_slat_flow_model_1024'] self.models['tex_slat_flow_model_1024'] = None - gc.collect() + gc.collect() + + def load_shape_slat_encoder(self): + if self.models['shape_slat_encoder'] is None: + print('Loading Shape Slat Encoder model ...') + self.models['shape_slat_encoder'] = models.from_pretrained(f"{self.path}/ckpts/shape_enc_next_dc_f16c32_fp16") + self.models['shape_slat_encoder'].eval() + self.models['shape_slat_encoder'].to(self._device) + + def unload_shape_slat_encoder(self): + if self.models['shape_slat_encoder'] is not None: + del self.models['shape_slat_encoder'] + self.models['shape_slat_encoder'] = None + gc.collect() def to(self, device: torch.device) -> None: self._device = device @@ -638,7 +657,7 @@ class Trellis2ImageTo3DPipeline(Pipeline): def decode_tex_slat( self, slat: SparseTensor, - subs: List[SparseTensor], + subs: List[SparseTensor] = None, ) -> SparseTensor: """ Decode the structured latent. @@ -654,7 +673,12 @@ class Trellis2ImageTo3DPipeline(Pipeline): if self.low_vram: self.models['tex_slat_decoder'].to(self.device) - ret = self.models['tex_slat_decoder'](slat, guide_subs=subs) * 0.5 + 0.5 + + if subs is None: + ret = self.models['tex_slat_decoder'](slat) * 0.5 + 0.5 + else: + ret = self.models['tex_slat_decoder'](slat, guide_subs=subs) * 0.5 + 0.5 + if self.low_vram: self.models['tex_slat_decoder'].cpu() @@ -879,3 +903,214 @@ class Trellis2ImageTo3DPipeline(Pipeline): return out_mesh, (shape_slat, tex_slat, res) else: return out_mesh + + def preprocess_mesh(self, mesh: trimesh.Trimesh) -> trimesh.Trimesh: + """ + Preprocess the input mesh. + """ + vertices = mesh.vertices + vertices_min = vertices.min(axis=0) + vertices_max = vertices.max(axis=0) + center = (vertices_min + vertices_max) / 2 + scale = 0.99999 / (vertices_max - vertices_min).max() + vertices = (vertices - center) * scale + tmp = vertices[:, 1].copy() + vertices[:, 1] = -vertices[:, 2] + vertices[:, 2] = tmp + assert np.all(vertices >= -0.5) and np.all(vertices <= 0.5), 'vertices out of range' + return trimesh.Trimesh(vertices=vertices, faces=mesh.faces, process=False) + + def encode_shape_slat( + self, + mesh: trimesh.Trimesh, + resolution: int = 1024, + ) -> SparseTensor: + """ + Encode the meshes to structured latent. + + Args: + mesh (trimesh.Trimesh): The mesh to encode. + resolution (int): The resolution of mesh + + Returns: + SparseTensor: The encoded structured latent. + """ + 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, + ) + + vertices = SparseTensor( + feats=dual_vertices * resolution - voxel_indices, + coords=torch.cat([torch.zeros_like(voxel_indices[:, 0:1]), voxel_indices], dim=-1) + ).to(self.device) + intersected = vertices.replace(intersected).to(self.device) + + self.load_shape_slat_encoder() + + if self.low_vram: + self.models['shape_slat_encoder'].to(self.device) + shape_slat = self.models['shape_slat_encoder'](vertices, intersected) + if self.low_vram: + self.models['shape_slat_encoder'].cpu() + + if not self.keep_models_loaded: + self.unload_shape_slat_encoder() + + return shape_slat + + def postprocess_mesh( + self, + mesh: trimesh.Trimesh, + pbr_voxel: SparseTensor, + resolution: int = 1024, + texture_size: int = 1024, + texture_alpha_mode = 'OPAQUE', + double_side_material = True + ): + vertices = mesh.vertices + faces = mesh.faces + normals = mesh.vertex_normals + vertices_torch = torch.from_numpy(vertices).float().cuda() + faces_torch = torch.from_numpy(faces).int().cuda() + if hasattr(mesh, 'visual') and hasattr(mesh.visual, 'uv') and mesh.visual.uv is not None: + uvs = mesh.visual.uv.copy() + uvs[:, 1] = 1 - uvs[:, 1] + uvs_torch = torch.from_numpy(uvs).float().cuda() + else: + _cumesh = cumesh.CuMesh() + _cumesh.init(vertices_torch, faces_torch) + print('Unwrapping mesh ...') + vertices_torch, faces_torch, uvs_torch, vmap = _cumesh.uv_unwrap(return_vmaps=True) + vertices_torch = vertices_torch.cuda() + faces_torch = faces_torch.cuda() + uvs_torch = uvs_torch.cuda() + vertices = vertices_torch.cpu().numpy() + faces = faces_torch.cpu().numpy() + uvs = uvs_torch.cpu().numpy() + normals = normals[vmap.cpu().numpy()] + + # rasterize + print('Finalizing mesh ...') + ctx = dr.RasterizeCudaContext() + uvs_torch = torch.cat([uvs_torch * 2 - 1, torch.zeros_like(uvs_torch[:, :1]), torch.ones_like(uvs_torch[:, :1])], dim=-1).unsqueeze(0) + rast, _ = dr.rasterize( + ctx, uvs_torch, faces_torch, + resolution=[texture_size, texture_size], + ) + mask = rast[0, ..., 3] > 0 + pos = dr.interpolate(vertices_torch.unsqueeze(0), rast, faces_torch)[0][0] + + attrs = torch.zeros(texture_size, texture_size, pbr_voxel.shape[1], device=self.device) + attrs[mask] = flex_gemm.ops.grid_sample.grid_sample_3d( + pbr_voxel.feats, + pbr_voxel.coords, + shape=torch.Size([*pbr_voxel.shape, *pbr_voxel.spatial_shape]), + grid=((pos[mask] + 0.5) * resolution).reshape(1, -1, 3), + mode='trilinear', + ) + + # construct mesh + mask = mask.cpu().numpy() + base_color = np.clip(attrs[..., self.pbr_attr_layout['base_color']].cpu().numpy() * 255, 0, 255).astype(np.uint8) + metallic = np.clip(attrs[..., self.pbr_attr_layout['metallic']].cpu().numpy() * 255, 0, 255).astype(np.uint8) + roughness = np.clip(attrs[..., self.pbr_attr_layout['roughness']].cpu().numpy() * 255, 0, 255).astype(np.uint8) + alpha = np.clip(attrs[..., self.pbr_attr_layout['alpha']].cpu().numpy() * 255, 0, 255).astype(np.uint8) + + # extend + mask = (~mask).astype(np.uint8) + base_color = cv2.inpaint(base_color, mask, 3, cv2.INPAINT_TELEA) + metallic = cv2.inpaint(metallic, mask, 1, cv2.INPAINT_TELEA)[..., None] + roughness = cv2.inpaint(roughness, mask, 1, cv2.INPAINT_TELEA)[..., None] + alpha = cv2.inpaint(alpha, mask, 1, cv2.INPAINT_TELEA)[..., None] + + baseColorTexture = Image.fromarray(np.concatenate([base_color, alpha], axis=-1)) + metallicRoughnessTexture = Image.fromarray(np.concatenate([np.zeros_like(metallic), roughness, metallic], axis=-1)) + + material = trimesh.visual.material.PBRMaterial( + baseColorTexture=baseColorTexture, + baseColorFactor=np.array([255, 255, 255, 255], dtype=np.uint8), + metallicRoughnessTexture=metallicRoughnessTexture, + metallicFactor=1.0, + roughnessFactor=1.0, + alphaMode=texture_alpha_mode, + doubleSided=True, + ) + + # Swap Y and Z axes, invert Y (common conversion for GLB compatibility) + vertices[:, 1], vertices[:, 2] = vertices[:, 2], -vertices[:, 1] + normals[:, 1], normals[:, 2] = normals[:, 2], -normals[:, 1] + uvs[:, 1] = 1 - uvs[:, 1] # Flip UV V-coordinate + + textured_mesh = trimesh.Trimesh( + vertices=vertices, + faces=faces, + vertex_normals=normals, + process=False, + visual=trimesh.visual.TextureVisuals(uv=uvs, material=material) + ) + + return textured_mesh, baseColorTexture, metallicRoughnessTexture + + @torch.no_grad() + def texture_mesh( + self, + mesh: trimesh.Trimesh, + image: Image.Image, + seed: int = 42, + tex_slat_sampler_params: dict = {}, + resolution: int = 1024, + texture_size: int = 2048, + texture_alpha_mode = 'OPAQUE', + double_side_material = True + ): + mesh = self.preprocess_mesh(mesh) + torch.manual_seed(seed) + + self.load_image_cond_model() + cond = self.get_cond(image, resolution) + + if not self.keep_models_loaded: + self.unload_image_cond_model() + + shape_slat = self.encode_shape_slat(mesh, resolution) + + if resolution==512: + self.unload_tex_slat_flow_model_1024() + self.load_tex_slat_flow_model_512() + tex_model = self.models['tex_slat_flow_model_512'] + + tex_slat = self.sample_tex_slat( + cond, tex_model, + shape_slat, tex_slat_sampler_params + ) + + if not self.keep_models_loaded: + self.unload_tex_slat_flow_model_512() + else: + self.unload_tex_slat_flow_model_512() + self.load_tex_slat_flow_model_1024() + tex_model = self.models['tex_slat_flow_model_1024'] + + tex_slat = self.sample_tex_slat( + cond, tex_model, + shape_slat, tex_slat_sampler_params + ) + + if not self.keep_models_loaded: + self.unload_shape_slat_flow_model_1024() + + torch.cuda.empty_cache() + pbr_voxel = self.decode_tex_slat(tex_slat) + 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 \ No newline at end of file