diff --git a/nodes.py b/nodes.py index 08f6398..cee8b2f 100644 --- a/nodes.py +++ b/nodes.py @@ -1271,7 +1271,8 @@ class Trellis2MeshTexturing: "texture_alpha_mode": (["OPAQUE","MASK","BLEND"],{"default":"OPAQUE"}), "double_side_material": ("BOOLEAN",{"default":True}), "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}), + "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}), }, } @@ -1281,22 +1282,25 @@ class Trellis2MeshTexturing: 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): - #image = tensor2pil_v2(image) - image = tensor2pil(image) + 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): + 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, + 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 + double_side_material = double_side_material, + max_views = max_views ) diff --git a/trellis2/models/sc_vaes/fdg_vae.py b/trellis2/models/sc_vaes/fdg_vae.py index b1f0bd4..008d9c6 100644 --- a/trellis2/models/sc_vaes/fdg_vae.py +++ b/trellis2/models/sc_vaes/fdg_vae.py @@ -118,7 +118,7 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder): if useTiled: mesh = [Mesh(*tiled_flexible_dual_grid_to_mesh( - coords=h.coords[:, 1:], + coords=v.coords[:, 1:], dual_vertices=v.feats, intersected_flag=i.feats, split_weight=q.feats, @@ -129,7 +129,7 @@ class FlexiDualGridVaeDecoder(SparseUnetVaeDecoder): )) for v, i, q in zip(vertices, intersected, quad_lerp)] else: mesh = [Mesh(*flexible_dual_grid_to_mesh( - coords=h.coords[:, 1:], + coords=v.coords[:, 1:], dual_vertices=v.feats, intersected_flag=i.feats, split_weight=q.feats, diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 5414241..24f2e79 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -1307,13 +1307,22 @@ class Trellis2ImageTo3DPipeline(Pipeline): resolution: int = 1024, texture_size: int = 2048, texture_alpha_mode = 'OPAQUE', - double_side_material = True + double_side_material = True, + max_views = 4 ): mesh = self.preprocess_mesh(mesh) torch.manual_seed(seed) + # Accept either a single PIL image or a list of PIL images (multi-view) + if isinstance(image, (list, tuple)): + images = list(image) + else: + images = [image] + + torch.manual_seed(seed) + self.load_image_cond_model() - cond = self.get_cond(image, resolution) + cond = self.get_cond(images, resolution, max_views = max_views) if not self.keep_models_loaded: self.unload_image_cond_model() diff --git a/trellis2/renderers/pbr_mesh_renderer.py b/trellis2/renderers/pbr_mesh_renderer.py index 876378f..b3b4655 100644 --- a/trellis2/renderers/pbr_mesh_renderer.py +++ b/trellis2/renderers/pbr_mesh_renderer.py @@ -250,9 +250,19 @@ class PbrMeshRenderer: ssaa = self.rendering_options["ssaa"] if mesh.vertices.shape[0] == 0 or mesh.faces.shape[0] == 0: - return edict( - shaded=torch.full((4, resolution, resolution), 0.5, dtype=torch.float32, device=self.device), + out_dict = edict( + normal=torch.zeros((3, resolution, resolution), dtype=torch.float32, device=self.device), + mask=torch.zeros((resolution, resolution), dtype=torch.float32, device=self.device), + base_color=torch.zeros((3, resolution, resolution), dtype=torch.float32, device=self.device), + metallic=torch.zeros((resolution, resolution), dtype=torch.float32, device=self.device), + roughness=torch.zeros((resolution, resolution), dtype=torch.float32, device=self.device), + alpha=torch.zeros((resolution, resolution), dtype=torch.float32, device=self.device), + clay=torch.zeros((resolution, resolution), dtype=torch.float32, device=self.device), ) + for i, k in enumerate(envmap.keys()): + shaded_key = f"shaded_{k}" if k != '' else "shaded" + out_dict[shaded_key] = torch.zeros((3, resolution, resolution), dtype=torch.float32, device=self.device) + return out_dict rays_o, rays_d = utils3d.torch.get_image_rays( extrinsics, intrinsics, resolution * ssaa, resolution * ssaa