Fixed Pixal3D

This commit is contained in:
Bruno Fargnoli
2026-05-21 00:13:29 +02:00
parent e3d592ab35
commit 992782c1fc
7 changed files with 996 additions and 943 deletions
+1
View File
@@ -14,6 +14,7 @@
| Date | Description |
| --- | --- |
| **2026-05-20** | Fixed Pixal3D<br>The node "Mesh with Voxel Advanced Generator" is compatible with Pixal3D |
| **2026-05-13** | Added support for Pixal3D-T model<br>It's not compatible with all nodes<br>Check in the folder example_workflows |
| **2026-04-20** | Recreated all Workflows |
| **2026-04-20** | Added node "Sparse MultiView Generator"<br>Added node "ImageCond MultiView Generator"<br>Added node "Shape MultiView Generator"<br>Added node "Shape Cascade MultiView Generator"<br>Added node "Tex Slat MultiView Generator" |
File diff suppressed because it is too large Load Diff
+80 -42
View File
@@ -61,6 +61,12 @@ class AnyType(str):
any = AnyType("*")
def inpaint_channel(channel, mask_inv, radius=1, mode = cv2.INPAINT_TELEA):
out = cv2.inpaint(channel, mask_inv, radius, mode)
if out.ndim == 2:
out = out[..., None]
return out
def rotate_triton_cache():
"""
Creates a new cache directory and attempts to clean up old ones.
@@ -1278,9 +1284,9 @@ class Trellis2UnWrapAndRasterizer:
# Inpainting: fill gaps (dilation) to prevent black seams at UV boundaries
mask_inv = (~mask).astype(np.uint8)
base_color = cv2.inpaint(base_color, mask_inv, 3, inpainting)
metallic = cv2.inpaint(metallic, mask_inv, 1, inpainting)[..., None]
roughness = cv2.inpaint(roughness, mask_inv, 1, inpainting)[..., None]
alpha = cv2.inpaint(alpha, mask_inv, 1, inpainting)[..., None]
metallic = inpaint_channel(metallic, mask_inv, 1, inpainting)
roughness = inpaint_channel(roughness, mask_inv, 1, inpainting)
alpha = inpaint_channel(alpha, mask_inv, 1, inpainting)
# Create PBR material
# Standard PBR packs Metallic and Roughness into Blue and Green channels
@@ -1435,31 +1441,53 @@ class Trellis2MeshWithVoxelAdvancedGenerator:
if generate_texture_slat:
num_steps = 5
else:
num_steps = 4
pbar = ProgressBar(num_steps)
num_steps = 4
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,
use_tiled=use_tiled_decoder,
pbar=pbar,
sampler=sampler,
fill_holes = fill_holes,
hole_iterations = hole_iterations,
verbose = verbose,
dino_lock = dino_lock,
dino_substeps = dino_substeps,
hole_fill_algorithm=hole_fill_algorithm,
dino_foundation_cap=dino_foundation_cap,
keep_only_shell=keep_only_shell)[0]
if pipeline.isPixal3D:
pipeline.load_moge_model()
if isinstance(images, (list, tuple)):
image = images[0]
else:
image = images
camera_params = pipeline.get_moge_camera_config(image)
if not pipeline.keep_models_loaded:
pipeline.unload_moge_model()
mesh = pipeline.run_pixal3d(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,
generate_texture_slat=generate_texture_slat,
camera_params=camera_params)[0]
else:
pbar = ProgressBar(num_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,
generate_texture_slat=generate_texture_slat,
use_tiled=use_tiled_decoder,
pbar=pbar,
sampler=sampler,
fill_holes = fill_holes,
hole_iterations = hole_iterations,
verbose = verbose,
dino_lock = dino_lock,
dino_substeps = dino_substeps,
hole_fill_algorithm=hole_fill_algorithm,
dino_foundation_cap=dino_foundation_cap,
keep_only_shell=keep_only_shell)[0]
vertices = mesh.vertices.cuda()
faces = mesh.faces.cuda()
@@ -1663,6 +1691,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
"bvh": ("BVH",),
"remove_inner_faces": ("BOOLEAN",{"default":True}),
"inpainting": (["telea","ns"],{"default":"telea"}),
"reorient_vertices":(["None","90 degrees","-90 degrees"],{"default":"90 degrees"}),
}
}
@@ -1672,7 +1701,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, texture_alpha_mode, dual_contouring_resolution, double_side_material, remove_floaters, bake_on_vertices,use_custom_normals,bvh,remove_inner_faces,inpainting):
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, texture_alpha_mode, dual_contouring_resolution, double_side_material, remove_floaters, bake_on_vertices,use_custom_normals,bvh,remove_inner_faces,inpainting,reorient_vertices):
pbar = ProgressBar(5 if not bake_on_vertices else 4)
mesh_copy = copy.deepcopy(mesh)
@@ -1707,10 +1736,7 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
grid_size = np.array(grid_size)
if isinstance(grid_size, np.ndarray):
grid_size = torch.tensor(grid_size, dtype=torch.int32, device=coords.device)
voxel_size = (aabb[1] - aabb[0]) / grid_size
if remove_floaters:
mesh_copy = remove_floater(mesh_copy)
voxel_size = (aabb[1] - aabb[0]) / grid_size
vertices = mesh_copy.vertices
faces = mesh_copy.faces
@@ -1769,18 +1795,21 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
else:
resolution = int(dual_contouring_resolution)
scale = (resolution + 3 * remesh_band) / resolution * scale
print(f"Calculated scale: {scale}")
print('Performing Dual Contouring ...')
# Perform Dual Contouring remeshing (rebuilds topology)
cumesh.init(*CuMesh.remeshing.remesh_narrow_band_dc_quad(
vertices, faces,
center = center,
scale = scale * 1.1, # old calculation : (resolution + 3 * remesh_band) / resolution * scale,
scale = scale,
resolution = resolution,
band = remesh_band,
project_back = remesh_project, # Snaps vertices back to original surface
verbose = True,
remove_inner_faces = remove_inner_faces,
#bvh = bvh,
bvh = bvh,
))
new_vertices, new_faces = cumesh.read()
@@ -1910,8 +1939,12 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
normals_np = out_normals.cpu().numpy()
# Swap Y and Z axes, invert Y (common conversion for GLB compatibility)
vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy()
if reorient_vertices == '90 degrees':
vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy()
elif reorient_vertices == '-90 degrees':
vertices_np[:, 1], vertices_np[:, 2] = -vertices_np[:, 2].copy(), vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = -normals_np[:, 2].copy(), normals_np[:, 1].copy()
# Create mesh with vertex colors using ColorVisuals
if use_custom_normals:
@@ -2021,9 +2054,9 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
mask_inv = (~mask).astype(np.uint8)
base_color = cv2.inpaint(base_color, mask_inv, 3, inpainting)
metallic = cv2.inpaint(metallic, mask_inv, 1, inpainting)[..., None]
roughness = cv2.inpaint(roughness, mask_inv, 1, inpainting)[..., None]
alpha = cv2.inpaint(alpha, mask_inv, 1, inpainting)[..., None]
metallic = inpaint_channel(metallic, mask_inv, 1, inpainting)
roughness = inpaint_channel(roughness, mask_inv, 1, inpainting)
alpha = inpaint_channel(alpha, mask_inv, 1, inpainting)
# Create PBR material
# Standard PBR packs Metallic and Roughness into Blue and Green channels
@@ -2045,8 +2078,13 @@ class Trellis2PostProcessAndUnWrapAndRasterizer:
normals_np = out_normals.cpu().numpy()
# Swap Y and Z axes, invert Y (common conversion for GLB compatibility)
vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy()
if reorient_vertices == '90 degrees':
vertices_np[:, 1], vertices_np[:, 2] = vertices_np[:, 2].copy(), -vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = normals_np[:, 2].copy(), -normals_np[:, 1].copy()
elif reorient_vertices == '-90 degrees':
vertices_np[:, 1], vertices_np[:, 2] = -vertices_np[:, 2].copy(), vertices_np[:, 1].copy()
normals_np[:, 1], normals_np[:, 2] = -normals_np[:, 2].copy(), normals_np[:, 1].copy()
uvs_np[:, 1] = 1 - uvs_np[:, 1] # Flip UV V-coordinate
if use_custom_normals:
@@ -5156,7 +5194,7 @@ class Trellis2SaveImage:
"required": {
"images": ("IMAGE", {"tooltip": "The images to save."}),
"filename_prefix": ("STRING", {"default": "ComfyUI", "tooltip": "The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."}),
"compress_level": ("INT",{"default":4,"min":1,"max":9,"step":1}),
"compress_level": ("INT",{"default":1,"min":1,"max":9,"step":1}),
},
"hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"
+1 -1
View File
@@ -1,3 +1,3 @@
SPCONV_ALGO = 'auto' # 'auto', 'implicit_gemm', 'native'
FLEX_GEMM_ALGO = 'implicit_gemm_splitk' # 'explicit_gemm', 'implicit_gemm', 'implicit_gemm_splitk', 'masked_implicit_gemm', 'masked_implicit_gemm_splitk'
FLEX_GEMM_ALGO = 'masked_implicit_gemm_splitk' # 'explicit_gemm', 'implicit_gemm', 'implicit_gemm_splitk', 'masked_implicit_gemm', 'masked_implicit_gemm_splitk'
FLEX_GEMM_HASHMAP_RATIO = 2.0 # Ratio of hashmap size to input size
+92 -22
View File
@@ -239,14 +239,25 @@ class Trellis2ImageTo3DPipeline(Pipeline):
facebook_model_path = os.path.join(folder_paths.models_dir,"facebook","dinov3-vitl16-pretrain-lvd1689m")
pipeline._pretrained_args['image_cond_model']['args']['model_name'] = facebook_model_path
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["ss"]["model_name"] = facebook_model_path
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["shape_512"]["model_name"] = facebook_model_path
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["shape_1024"]["model_name"] = facebook_model_path
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["tex_1024"]["model_name"] = facebook_model_path
try:
from mmgp import safetensors2 as _mmgp_st2
if 'C64' not in _mmgp_st2._map_to_dtype:
_mmgp_st2._map_to_dtype['C64'] = torch.complex64
_mmgp_st2._map_to_dtype['C128'] = torch.complex128
print("[Pixal3D] Patched mmgp.safetensors2._map_to_dtype with C64/C128")
except (ImportError, AttributeError):
print('mmgp not installed')
# mmgp not installed (comfy_env worker / non-Desktop ComfyUI) — no patch needed.
pass
return pipeline
def load_moge_model(self):
if hasattr(self,'moge_model') and self.moge_model is not None:
return self.moge_model
@@ -526,7 +537,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
# =========================================================================
# Proj mode condition building
# =========================================================================
@torch.no_grad()
def get_proj_cond_ss(
self,
@@ -635,6 +646,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
'neg_cond': {'global': torch.zeros_like(z_global), 'proj': SparseTensor(feats=torch.zeros_like(z_proj_sparse), coords=coords)},
}
@torch.no_grad()
def get_moge_camera_config(self, image):
from ..utils.camera import get_camera_params_wild_moge
@@ -679,7 +691,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
output = output[:, :, :3] * output[:, :, 3:4]
output = Image.fromarray((output * 255).astype(np.uint8))
return output
@torch.no_grad()
def get_cond(
self,
image: Union[torch.Tensor, Image.Image, List[Image.Image]],
@@ -766,20 +779,21 @@ class Trellis2ImageTo3DPipeline(Pipeline):
neg_cond = torch.zeros_like(cond)
return {"cond": cond, "neg_cond": neg_cond}
@torch.no_grad()
def sample_sparse_structure(
self,
cond: dict,
resolution: int,
num_samples: int = 1,
sampler_params: dict = {},
fill_holes: bool = True,
fill_holes: bool = False,
hole_structure: int = 1,
hole_iterations: int = 1,
hole_fill_algorithm: str = "remove_small_holes",
keep_only_shell: bool = True,
verbose: bool = True,
keep_only_shell: bool = False,
verbose: bool = False,
dino_lock: float = 0.0,
dino_substeps: int = 4,
dino_substeps: int = 2,
dino_foundation_cap: float = 0.92
) -> torch.Tensor:
"""
@@ -951,6 +965,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self._cleanup_cuda()
return coords
@torch.no_grad()
def sample_shape_slat(
self,
cond: dict,
@@ -1008,6 +1023,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat
@torch.no_grad()
def sample_shape_slat_cascade(
self,
lr_cond: dict,
@@ -1144,6 +1160,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat, hr_resolution
@torch.no_grad()
def decode_shape_slat(
self,
slat: SparseTensor,
@@ -1178,6 +1195,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return ret
@torch.no_grad()
def sample_tex_slat(
self,
cond: dict,
@@ -1296,6 +1314,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self._cleanup_cuda()
out_mesh = []
for m in meshes:
m.fill_holes()
out_mesh.append(
MeshWithVoxel(
m.vertices, m.faces,
@@ -1315,6 +1334,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self._cleanup_cuda()
out_mesh = []
for m, v in zip(meshes, tex_voxels):
m.fill_holes()
out_mesh.append(
MeshWithVoxel(
m.vertices, m.faces,
@@ -1328,7 +1348,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
)
return out_mesh
@torch.no_grad()
@torch.inference_mode()
def run(
self,
image: Image.Image,
@@ -1701,6 +1721,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
else:
return 'Euler'
@torch.no_grad()
def sample_shape_slat_cascade_advanced(
self,
lr_cond: dict,
@@ -1850,6 +1871,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat, hr_resolution
@torch.no_grad()
def sample_tex_slat_advanced(
self,
cond: dict,
@@ -2085,8 +2107,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return out_mesh
@torch.no_grad()
@torch.inference_mode()
def run_multiview(
self,
front: Image.Image,
@@ -2336,6 +2357,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
else:
return out_mesh
@torch.no_grad()
def sample_sparse_structure_multiview(
self,
conds: dict,
@@ -2543,6 +2565,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return coords
@torch.no_grad()
def sample_shape_slat_multiview(
self,
conds: dict,
@@ -2613,6 +2636,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat
@torch.no_grad()
def sample_shape_slat_cascade_multiview(
self,
lr_conds: dict,
@@ -2767,7 +2791,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat
@torch.no_grad()
def sample_tex_slat_multiview(
self,
conds: dict,
@@ -2854,7 +2878,6 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat
def preprocess_mesh(self, mesh: trimesh.Trimesh) -> trimesh.Trimesh:
"""
Preprocess the input mesh.
@@ -2875,6 +2898,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
mesh.vertices = vertices
return mesh
@torch.no_grad()
def encode_shape_slat(
self,
mesh: trimesh.Trimesh,
@@ -3154,7 +3178,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return textured_mesh, baseColorTexture, metallicRoughnessTexture
@torch.no_grad()
@torch.inference_mode()
def texture_mesh(
self,
mesh: trimesh.Trimesh,
@@ -3239,7 +3263,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
out_mesh, baseColorTexture, metallicRoughnessTexture = self.postprocess_mesh(mesh, pbr_voxel, resolution, texture_size, texture_alpha_mode, double_side_material, bake_on_vertices, use_custom_normals, mesh_cluster_threshold_cone_half_angle_rad, inpainting)
return out_mesh, baseColorTexture, metallicRoughnessTexture
@torch.no_grad()
@torch.inference_mode()
def texture_mesh_multiview(
self,
mesh: trimesh.Trimesh,
@@ -3375,6 +3399,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return coords;
@torch.no_grad()
def sample_mesh_slat(
self,
mesh_slat,
@@ -3470,7 +3495,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
return slat, hr_resolution
@torch.no_grad()
@torch.inference_mode()
def refine_mesh(
self,
mesh: trimesh.Trimesh,
@@ -3631,7 +3656,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
else:
return out_mesh
@torch.no_grad()
@torch.inference_mode()
def run_pixal3d(
self,
image: Image.Image,
@@ -3641,7 +3666,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
sparse_structure_sampler_params: dict = {},
shape_slat_sampler_params: dict = {},
tex_slat_sampler_params: dict = {},
preprocess_image: bool = True,
preprocess_image: bool = False,
return_latent: bool = False,
pipeline_type: Optional[str] = None,
max_num_tokens: int = 49152,
@@ -3667,7 +3692,10 @@ class Trellis2ImageTo3DPipeline(Pipeline):
max_num_tokens (int): The maximum number of tokens to use.
"""
# Check pipeline type
pipeline_type = pipeline_type or self.default_pipeline_type
if pipeline_type == '1536_cascade':
hr_resolution = 1536
else:
hr_resolution = 1024
# Extract camera params
camera_angle_x = camera_params['camera_angle_x']
@@ -3679,35 +3707,51 @@ class Trellis2ImageTo3DPipeline(Pipeline):
torch.manual_seed(seed)
# ---- Stage 1: Sparse Structure (proj) ----
image_cond_model = self.load_pixal3d_image_cond_ss()
cond_ss = self.get_proj_cond_ss(
[image],
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
image_cond_model=image_cond_model
)
del image_cond_model
if not self.keep_models_loaded:
self.unload_pixal3d_image_cond_ss()
ss_res = 32
self.load_sparse_structure_model()
coords = self.sample_sparse_structure(
cond_ss, ss_res,
num_samples, sparse_structure_sampler_params
)
if not self.keep_models_loaded:
self.unload_sparse_structure_model()
del cond_ss
torch.cuda.empty_cache()
# ---- Stage 2: Shape LR 512 (proj) ----
image_cond_model = self.load_pixal3d_image_cond_shape_512()
cond_shape_lr = self.get_proj_cond_shape(
self.image_cond_model_shape_512, [image], coords,
image_cond_model, [image], coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
)
del image_cond_model
if not self.keep_models_loaded:
self.unload_pixal3d_image_cond_shape_512()
self.load_shape_slat_flow_model_512()
lr_slat = self.sample_shape_slat(
cond_shape_lr, self.models['shape_slat_flow_model_512'],
coords, shape_slat_sampler_params
)
if not self.keep_models_loaded:
self.unload_shape_slat_flow_model_512()
del cond_shape_lr
torch.cuda.empty_cache()
# ---- Stage 3a: Upsample LR → HR ----
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
@@ -3716,6 +3760,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
self.models['shape_slat_decoder'].cpu()
self.models['shape_slat_decoder'].low_vram = False
if not self.keep_models_loaded:
self.unload_shape_slat_decoder()
lr_resolution = 512
actual_hr_resolution = hr_resolution
while True:
@@ -3735,18 +3782,26 @@ class Trellis2ImageTo3DPipeline(Pipeline):
torch.cuda.empty_cache()
# ---- Stage 3b: Shape HR (proj) ----
image_cond_model = self.load_pixal3d_image_cond_shape_1024()
cond_shape_hr = self.get_proj_cond_shape(
self.image_cond_model_shape_1024, [image], hr_coords_unique,
image_cond_model, [image], hr_coords_unique,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=actual_grid_res,
)
del image_cond_model
if not self.keep_models_loaded:
self.unload_pixal3d_image_cond_shape_1024()
self.load_shape_slat_flow_model_1024()
noise_hr = SparseTensor(
feats=torch.randn(hr_coords_unique.shape[0], self.models['shape_slat_flow_model_1024'].in_channels).to(self.device),
coords=hr_coords_unique,
)
sampler_params_hr = {**self.shape_slat_sampler_params, **shape_slat_sampler_params}
flow_model_hr = self.models['shape_slat_flow_model_1024']
if self.low_vram:
flow_model_hr.to(self.device)
@@ -3763,23 +3818,38 @@ class Trellis2ImageTo3DPipeline(Pipeline):
std = torch.tensor(self.shape_slat_normalization['std'])[None].to(hr_slat.device)
mean = torch.tensor(self.shape_slat_normalization['mean'])[None].to(hr_slat.device)
shape_slat = hr_slat * std + mean
del flow_model_hr
if not self.keep_models_loaded:
self.unload_shape_slat_flow_model_1024()
del cond_shape_hr, noise_hr, hr_slat, hr_coords_unique
torch.cuda.empty_cache()
if generate_texture_slat:
# ---- Stage 4: Texture (proj) ----
image_cond_model = self.load_pixal3d_image_cond_tex_1024()
tex_grid_res = actual_hr_resolution // 16
cond_tex = self.get_proj_cond_shape(
self.image_cond_model_tex_1024, [image], shape_slat.coords,
image_cond_model, [image], shape_slat.coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=tex_grid_res,
)
if not self.keep_models_loaded:
self.unload_pixal3d_image_cond_tex_1024()
self.load_tex_slat_flow_model_1024()
tex_slat = self.sample_tex_slat(
cond_tex, 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()
del cond_tex
torch.cuda.empty_cache()
@@ -67,7 +67,7 @@ def project_points_to_image_batch(
points_homogeneous = torch.cat([points_3d_batch, ones], dim=-1) # [B, N, 4]
# Compute world to camera transformation matrix
world_to_camera = torch.linalg.inv(transform_matrix) # [B, 4, 4]
world_to_camera = torch.linalg.inv(transform_matrix.float()).to(transform_matrix.dtype) # linalg.inv requires fp32+
# Batch transform to camera coordinate system: [B, N, 4] @ [B, 4, 4]^T -> [B, N, 3]
points_camera = torch.bmm(points_homogeneous, world_to_camera.transpose(-2, -1))[..., :3] # [B, N, 3]
@@ -409,7 +409,7 @@ class DinoV3ProjFeatureExtractor(nn.Module):
# (current behavior, default). >1 enables a streaming wrapper that tiles NAF and
# the projection together, capping the transient HR-feature-map peak. Set transiently
# by the pipeline (see Pixal3DImageTo3DPipeline.get_proj_cond_shape).
self.naf_tile_factor: int = 4
self.naf_tile_factor: int = 1
# proj_channels: the output dimension of proj features
# Without NAF: embed_dim (e.g. 1024)
+1 -1
View File
@@ -19,7 +19,7 @@ def distance_from_fov(camera_angle_x, grid_point, target_point, mesh_scale, imag
distance_x = f_pixels * xw / x_ndc - yw
return {"distance_from_x": float(distance_x), "f_pixels": float(f_pixels)}
def get_camera_params_wild_moge(pil_image, moge_model, device="cuda", mesh_scale=1.0, extend_pixel=0, image_resolution=1024):
def get_camera_params_wild_moge(pil_image, moge_model, device="cuda", mesh_scale=1.0, extend_pixel=0, image_resolution=512):
#pil_image = Image.open(image_path).convert("RGB")
width, height = pil_image.size
image_np = np.array(pil_image).astype(np.float32) / 255.0