Support for Pixal3D MultiView

This commit is contained in:
Bruno Fargnoli
2026-09-07 14:05:53 +02:00
parent fd69b62175
commit 41170c6a23
11 changed files with 5010 additions and 169 deletions
+1
View File
@@ -14,6 +14,7 @@
| Date | Description |
| --- | --- |
| **2026-09-07** | Add support for Pixal3D MultiView |
| **2026-07-31** | Added new nodes "Smooth Mesh with PyMeshlab" and "Smooth Trimesh with PyMeshlab" |
| **2026-06-02** | Added new node "Render MultiView (Nvdiffrast)"<br>Thanks GiusTex |
| **2026-05-26** | Added new nodes used for the projection<br>Check the example Projection_Blender_Qwen_XViews |
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -976,7 +976,7 @@
"widget_ue_connectable": {}
},
"widgets_values": [
"TencentARC/Pixal3D-T",
"TencentARC/Pixal3D",
"flash_attn",
"cuda",
true,
File diff suppressed because it is too large Load Diff
@@ -212,7 +212,7 @@
"widget_ue_connectable": {}
},
"widgets_values": [
"TencentARC/Pixal3D-T",
"TencentARC/Pixal3D",
"flash_attn",
"cuda",
true,
+518 -146
View File
@@ -334,12 +334,26 @@ def _batched_unsigned_distance(bvh, positions, batch_size=100000, return_uvw=Fal
torch.cat(uvw_list) if return_uvw else None
)
def check_pixal3d_mv_pipeline(pipeline):
"""
The multi-view conditioning only lines up with the *_mv denoisers.
Feeding it to the single-view checkpoints silently produces garbage rather
than an error, so refuse it here: the pipeline has to have been loaded with
pixal3d_multiview on (pipeline_mv.json).
"""
if not getattr(pipeline, 'isPixal3DMV', False):
raise Exception(
'pixal3d_mv_views needs the Pixal3D multi-view weights. Turn on '
'"pixal3d_multiview" in Trellis2 - LoadModel (loads pipeline_mv.json / ckpts/*_mv).')
class Trellis2LoadModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"modelname": (["microsoft/TRELLIS.2-4B","visualbruno/TRELLIS.2-4B-FP8","TencentARC/Pixal3D-T"],{"default":"microsoft/TRELLIS.2-4B"}),
"modelname": (["microsoft/TRELLIS.2-4B","visualbruno/TRELLIS.2-4B-FP8","TencentARC/Pixal3D"],{"default":"microsoft/TRELLIS.2-4B"}),
"backend": (["flash_attn","xformers","sdpa","flash_attn_3"],{"default":"flash_attn"}),
"device": (["cpu","cuda"],{"default":"cuda"}),
"low_vram": ("BOOLEAN",{"default":True}),
@@ -347,6 +361,7 @@ class Trellis2LoadModel:
"conv_backend": (["spconv","torchsparse","flex_gemm"],{"default":"flex_gemm"}),
"sparse_backend": (["xformers","flash_attn"],{"default":"flash_attn"}),
"use_reconviagen": ("BOOLEAN",{"default":False}),
"pixal3d_multiview": ("BOOLEAN",{"default":False,"tooltip":"Pixal3D only: load the multi-view denoisers (pipeline_mv.json / ckpts/*_mv) instead of the single-view ones"}),
#"naf_chunk_size":(["None","144","208","272","336","400","464","528","592","656","720","784","848","912","976","1024"],{"default":"None"}),
}
}
@@ -357,7 +372,7 @@ class Trellis2LoadModel:
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, modelname, backend, device, low_vram, keep_models_loaded, conv_backend, sparse_backend, use_reconviagen):
def process(self, modelname, backend, device, low_vram, keep_models_loaded, conv_backend, sparse_backend, use_reconviagen, pixal3d_multiview = False):
import requests
os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1'
@@ -384,7 +399,18 @@ class Trellis2LoadModel:
local_dir=model_path,
local_dir_use_symlinks=False,
)
elif pixal3d_multiview and not os.path.exists(os.path.join(model_path,'pipeline_mv.json')):
# A Pixal3D folder downloaded before the multi-view release has no
# pipeline_mv.json / ckpts/*_mv. snapshot_download skips what is already
# there, so this only pulls the missing multi-view files.
print(f"Multi-view weights missing in {model_path}, downloading them ...")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=modelname,
local_dir=model_path,
local_dir_use_symlinks=False,
)
reconviagen_pipeline_file = os.path.join(folder_paths.models_dir,'microsoft','TRELLIS.2-4B','reconviagen_pipeline.json')
if not os.path.exists(reconviagen_pipeline_file):
source_reconviagen_pipeline_file = os.path.join(script_directory,'reconviagen_pipeline.json')
@@ -440,7 +466,7 @@ class Trellis2LoadModel:
else:
raise Exception("Cannot download Trellis-Image-Large file ss_dec_conv3d_16l8_fp16.safetensors")
if use_reconviagen and modelname == 'TencentARC/Pixal3D-T':
if use_reconviagen and modelname == 'TencentARC/Pixal3D':
raise Exception('Model TencentARC/Pixal3D-T is not compatible with ReconViaGen')
if use_reconviagen:
@@ -516,10 +542,13 @@ class Trellis2LoadModel:
use_fp8 = False
isPixal3D = False
if modelname == "TencentARC/Pixal3D-T":
if modelname == "TencentARC/Pixal3D":
isPixal3D = True
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path, keep_models_loaded = keep_models_loaded, use_fp8=use_fp8, use_reconviagen=use_reconviagen, isPixal3D = isPixal3D)
if pixal3d_multiview and not isPixal3D:
raise Exception('pixal3d_multiview only applies to TencentARC/Pixal3D')
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(model_path, keep_models_loaded = keep_models_loaded, use_fp8=use_fp8, use_reconviagen=use_reconviagen, isPixal3D = isPixal3D, isPixal3DMV = pixal3d_multiview)
pipeline.low_vram = low_vram
# if naf_chunk_size == "None":
@@ -4473,7 +4502,8 @@ class Trellis2SparseGenerator:
},
"optional":{
"image":("IMAGE",),
"moge_camera_config":("MOGE_CAM_CONFIG",)
"moge_camera_config":("MOGE_CAM_CONFIG",),
"pixal3d_mv_views":("PIXAL3D_MV_VIEWS",)
}
}
@@ -4502,8 +4532,9 @@ class Trellis2SparseGenerator:
dino_foundation_cap,
keep_only_shell,
image = None,
moge_camera_config = None
):
moge_camera_config = None,
pixal3d_mv_views = None
):
self.seed_all(seed)
sparse_structure_guidance_interval = [sparse_structure_guidance_interval_start,sparse_structure_guidance_interval_end]
@@ -4514,31 +4545,44 @@ class Trellis2SparseGenerator:
pipeline.sparse_structure_sampler = getattr(samplers, f"Flow{sparse_sampler_prefix}GuidanceIntervalSampler")(**args['sparse_structure_sampler']['args'])
if pipeline.isPixal3D:
if image is not None:
images = tensor_batch_to_pil_list(image, max_views=16)
images = list(images)
else:
raise Exception('Image is required for Pixal3D')
if moge_camera_config is not None:
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
else:
raise Exception('MoGe Camera Config is required for Pixal3D')
if pixal3d_mv_views is not None:
check_pixal3d_mv_pipeline(pipeline)
image_cond_model = pipeline.load_pixal3d_image_cond_ss()
image_cond = pipeline.get_proj_cond_ss(
image=images,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
image_cond_model=image_cond_model
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_ss()
image_cond_model = pipeline.load_pixal3d_mv_image_cond_ss()
image_cond = pipeline.get_proj_cond_ss_mv(
pixal3d_mv_views,
image_cond_model=image_cond_model
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_mv_image_cond_ss()
else:
if image is not None:
images = tensor_batch_to_pil_list(image, max_views=16)
images = list(images)
else:
raise Exception('Image is required for Pixal3D')
if moge_camera_config is not None:
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
else:
raise Exception('MoGe Camera Config is required for Pixal3D')
image_cond_model = pipeline.load_pixal3d_image_cond_ss()
image_cond = pipeline.get_proj_cond_ss(
image=images,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
image_cond_model=image_cond_model
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_ss()
pipeline.load_sparse_structure_model()
@@ -4595,7 +4639,8 @@ class Trellis2ShapeGenerator:
{
"image": ("IMAGE",),
"moge_camera_config": ("MOGE_CAM_CONFIG",),
}
"pixal3d_mv_views": ("PIXAL3D_MV_VIEWS",),
}
}
RETURN_TYPES = ("SHAPE_SLAT", "INT", "TRELLIS2PIPELINE",)
@@ -4618,43 +4663,56 @@ class Trellis2ShapeGenerator:
dino_substeps,
dino_foundation_cap,
image = None,
moge_camera_config = None
moge_camera_config = None,
pixal3d_mv_views = None
):
shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end]
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end]
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
args = pipeline._pretrained_args
shape_sampler_prefix = pipeline.GetSamplerName(shape_sampler)
pipeline.shape_slat_sampler = getattr(samplers, f"Flow{shape_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
pipeline.shape_slat_sampler = getattr(samplers, f"Flow{shape_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
if resolution == 512:
pipeline.unload_shape_slat_flow_model_1024()
if pipeline.isPixal3D:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
if pixal3d_mv_views is not None:
check_pixal3d_mv_pipeline(pipeline)
image_cond_model = pipeline.load_pixal3d_mv_image_cond_shape_512()
image_cond = pipeline.get_proj_cond_shape_mv(
image_cond_model, pixal3d_mv_views, coords,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_mv_image_cond_shape_512()
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_512()
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_512()
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_512()
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_512()
pipeline.load_shape_slat_flow_model_512()
@@ -4674,29 +4732,41 @@ class Trellis2ShapeGenerator:
pipeline.unload_shape_slat_flow_model_512()
if pipeline.isPixal3D:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
if pixal3d_mv_views is not None:
check_pixal3d_mv_pipeline(pipeline)
image_cond_model = pipeline.load_pixal3d_mv_image_cond_shape_1024()
image_cond = pipeline.get_proj_cond_shape_mv(
image_cond_model, pixal3d_mv_views, coords,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_mv_image_cond_shape_1024()
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_1024()
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_1024()
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_1024()
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_1024()
pipeline.load_shape_slat_flow_model_1024()
@@ -4742,7 +4812,8 @@ class Trellis2ShapeCascadeGenerator:
{
"image": ("IMAGE",),
"moge_camera_config": ("MOGE_CAM_CONFIG",),
}
"pixal3d_mv_views": ("PIXAL3D_MV_VIEWS",),
}
}
RETURN_TYPES = ("SHAPE_SLAT","INT","TRELLIS2PIPELINE","INT",)
@@ -4765,23 +4836,24 @@ class Trellis2ShapeCascadeGenerator:
dino_substeps,
dino_foundation_cap,
image = None,
moge_camera_config = None
moge_camera_config = None,
pixal3d_mv_views = None
):
shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end]
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
shape_guidance_interval = [shape_guidance_interval_start, shape_guidance_interval_end]
shape_slat_sampler_params = {"steps":shape_steps,"guidance_strength":shape_guidance_strength,"guidance_rescale":shape_guidance_rescale,"guidance_interval":shape_guidance_interval,"rescale_t":shape_rescale_t}
args = pipeline._pretrained_args
shape_sampler_prefix = pipeline.GetSamplerName(shape_sampler)
pipeline.shape_slat_sampler = getattr(samplers, f"Flow{shape_sampler_prefix}GuidanceIntervalSampler")(**args['shape_slat_sampler']['args'])
slat, hr_resolution, num_tokens = self.sample(pipeline, shape_slat, from_resolution, to_resolution, sparse_structure_resolution, max_num_tokens, image_cond, shape_slat_sampler_params, verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config)
slat, hr_resolution, num_tokens = self.sample(pipeline, shape_slat, from_resolution, to_resolution, sparse_structure_resolution, max_num_tokens, image_cond, shape_slat_sampler_params, verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config, pixal3d_mv_views)
if not pipeline.keep_models_loaded:
pipeline.unload_shape_slat_flow_model_1024()
return (slat, hr_resolution, pipeline, num_tokens,)
def sample(self, pipeline, slat, lr_resolution, resolution, sparse_structure_resolution, max_num_tokens, cond, sampler_params, verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config):
def sample(self, pipeline, slat, lr_resolution, resolution, sparse_structure_resolution, max_num_tokens, cond, sampler_params, verbose, dino_lock, dino_substeps, dino_foundation_cap, image, moge_camera_config, pixal3d_mv_views = None):
# Upsample
pipeline.load_shape_slat_decoder()
if pipeline.low_vram:
@@ -4821,32 +4893,45 @@ class Trellis2ShapeCascadeGenerator:
break
if pipeline.isPixal3D:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_1024()
actual_grid_res = hr_resolution // 16
cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=actual_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_1024()
if pixal3d_mv_views is not None:
check_pixal3d_mv_pipeline(pipeline)
image_cond_model = pipeline.load_pixal3d_mv_image_cond_shape_1024()
cond = pipeline.get_proj_cond_shape_mv(
image_cond_model, pixal3d_mv_views, coords,
grid_resolution_override=actual_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_mv_image_cond_shape_1024()
else:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_shape_1024()
cond = pipeline.get_proj_cond_shape(
image_cond_model, images, coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=actual_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_shape_1024()
pipeline.load_shape_slat_flow_model_1024()
flow_model = pipeline.models['shape_slat_flow_model_1024']
@@ -4915,7 +5000,8 @@ class Trellis2TexSlatGenerator:
"image": ("IMAGE",),
"moge_camera_config": ("MOGE_CAM_CONFIG",),
"from_resolution": ("INT",),
}
"pixal3d_mv_views": ("PIXAL3D_MV_VIEWS",),
}
}
RETURN_TYPES = ("TEXTURE_SLAT", "TRELLIS2PIPELINE",)
@@ -4939,7 +5025,8 @@ class Trellis2TexSlatGenerator:
dino_foundation_cap,
image = None,
moge_camera_config = None,
from_resolution = None
from_resolution = None,
pixal3d_mv_views = None
):
texture_guidance_interval = [texture_guidance_interval_start,texture_guidance_interval_end]
@@ -4968,32 +5055,45 @@ class Trellis2TexSlatGenerator:
pipeline.unload_tex_slat_flow_model_512()
if pipeline.isPixal3D:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_tex_1024()
tex_grid_res = from_resolution // 16
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, shape_slat.coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=tex_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_tex_1024()
if pixal3d_mv_views is not None:
check_pixal3d_mv_pipeline(pipeline)
image_cond_model = pipeline.load_pixal3d_mv_image_cond_tex_1024()
image_cond = pipeline.get_proj_cond_shape_mv(
image_cond_model, pixal3d_mv_views, shape_slat.coords,
grid_resolution_override=tex_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_mv_image_cond_tex_1024()
else:
images = tensor_batch_to_pil_list(image, max_views=16)
image_in = images[0] if len(images) == 1 else images
if isinstance(image_in, (list, tuple)):
images = list(image_in)
else:
images = [image_in]
camera_angle_x = moge_camera_config['camera_angle_x']
distance = moge_camera_config['distance']
mesh_scale = moge_camera_config['mesh_scale']
image_cond_model = pipeline.load_pixal3d_image_cond_tex_1024()
image_cond = pipeline.get_proj_cond_shape(
image_cond_model, images, shape_slat.coords,
camera_angle_x=camera_angle_x,
distance=distance,
mesh_scale=mesh_scale,
grid_resolution_override=tex_grid_res,
)
if not pipeline.keep_models_loaded:
pipeline.unload_pixal3d_image_cond_tex_1024()
pipeline.load_tex_slat_flow_model_1024()
@@ -7880,6 +7980,274 @@ class Trellis2SelectImagesForMultiView:
return (front_image, back_image, left_image, right_image, )
def comfy_images_to_rgba_pils(images, masks=None, invert_mask=False, remove_background=False,
max_views=16):
"""
Turn a ComfyUI IMAGE batch into the RGBA PIL views the Pixal3D MV path expects.
Alpha is taken, in order of preference, from an explicit MASK input, from the
image's own alpha channel, or from rembg. The views are never cropped or
rescaled here: transforms.json describes the framing as given, so changing it
would break the correspondence between the pixels and the cameras.
"""
if not isinstance(images, torch.Tensor):
raise TypeError(f"Expected torch.Tensor for IMAGE, got {type(images)}")
if images.ndim == 3:
images = images.unsqueeze(0)
if images.ndim != 4:
raise ValueError(f"Unsupported IMAGE tensor shape: {tuple(images.shape)}")
V = min(int(images.shape[0]), int(max_views))
if V < int(images.shape[0]):
print(f"[Pixal3D MV] Warning: {images.shape[0]} views given, using the first {V}")
if masks is not None:
if masks.ndim == 2:
masks = masks.unsqueeze(0)
if masks.shape[0] != images.shape[0]:
raise ValueError(
f"got {images.shape[0]} images but {masks.shape[0]} masks")
out = []
for i in range(V):
img = images[i]
alpha = None
if masks is not None:
alpha = masks[i].detach().cpu().float().clamp(0, 1)
if invert_mask:
alpha = 1.0 - alpha
if alpha.shape != img.shape[:2]:
alpha = F.interpolate(
alpha[None, None], size=tuple(img.shape[:2]),
mode='bilinear', align_corners=False)[0, 0]
elif img.shape[-1] == 4:
a = img[..., 3].detach().cpu().float().clamp(0, 1)
# A fully opaque alpha channel counts as no mask, as in preprocess_image.
if not torch.all(a >= 254.0 / 255.0):
alpha = a
rgb = img[..., :3].detach().cpu().float().clamp(0, 1)
pil = Image.fromarray((rgb.numpy() * 255.0).astype(np.uint8), mode='RGB')
if alpha is None:
if not remove_background:
raise ValueError(
f"view {i} has no alpha channel and no mask. Give it a MASK, an "
f"RGBA image, or turn remove_background on.")
from rembg import remove
pil = remove(pil).convert('RGBA')
else:
pil = pil.convert('RGBA')
pil.putalpha(Image.fromarray((alpha.numpy() * 255.0).astype(np.uint8), mode='L'))
out.append(pil)
return out
def pixal3d_views_to_preview(views):
"""The alpha-premultiplied 512px views, as a ComfyUI IMAGE batch, for previewing."""
size = min(views['images'].keys())
return views['images'][size][0].permute(0, 2, 3, 1).contiguous().cpu()
class Trellis2Pixal3DMultiViewConfig:
"""
Build the Pixal3D multi-view conditioning bundle from a batch of posed views.
The first image is the MAIN view and its pose must be the canonical front view
(azimuth 0, elevation 0), because every other view is placed relative to it.
Azimuth / elevation follow the same convention as
Trellis2RenderMultiViewNvdiffrast: 0/90/180/270 -> front/left/back/right, and
positive elevation is above the object.
fov / distance / mesh_scale describe the camera the views were SHOT with, and
the views are never cropped or rescaled here, so they have to match how the
images are actually framed. Getting the distance wrong scales the whole
projection: at 10% too near, the surface of the model samples the background
instead of itself, which shows up as washed-out texture long before the mesh
suffers. `framing` picks how the distance is derived:
auto measure the object in the alpha channel and fit the distance to
it. Works whatever the margin is; the default.
pixal3d_rig the rig the multi-view weights were trained on, 10% margin
(the shipped example: fov 20 deg, distance 3.1192).
fill_frame object touches the frame edges. This is what
Trellis2PreProcessImage produces and what
Trellis2FovMoGeCameraConfig assumes, so it is right for views
you cropped yourself and wrong for raw renders.
camera_config take the distance from the wired moge_camera_config.
A distance widget above 0 always wins. A wired moge_camera_config supplies fov
and mesh_scale; its distance is only used by framing = camera_config.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"azimuths": ("STRING", {"default": "0,90,180,270"}),
"elevations": ("STRING", {"default": "0,0,0,0"}),
"fov": ("FLOAT", {"default": 20.0, "min": 0.001, "max": 179.999, "step": 0.001}),
"fov_unit": (["deg", "rad"], {"default": "deg"}),
"mesh_scale": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 9.9, "step": 0.1}),
"distance": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 99.99, "step": 0.0001,
"tooltip": "0 = derive it from `framing`. Above 0 always wins."}),
"remove_background": ("BOOLEAN", {"default": False}),
"invert_mask": ("BOOLEAN", {"default": False}),
"framing": (["auto", "pixal3d_rig", "fill_frame", "camera_config"],
{"default": "auto",
"tooltip": "How to derive the camera distance. auto = fit it to the "
"object in the alpha channel; pixal3d_rig = the 10% margin "
"the MV weights were trained on; fill_frame = object touches "
"the frame edges (what Trellis2PreProcessImage produces)."}),
},
"optional": {
"masks": ("MASK",),
"moge_camera_config": ("MOGE_CAM_CONFIG",),
}
}
RETURN_TYPES = ("PIXAL3D_MV_VIEWS", "MOGE_CAM_CONFIG", "IMAGE",)
RETURN_NAMES = ("pixal3d_mv_views", "moge_camera_config", "preview",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, images, azimuths, elevations, fov, fov_unit, mesh_scale, distance,
remove_background, invert_mask, framing="auto", masks=None,
moge_camera_config=None):
from .trellis2.utils import mv_camera
az_list = Trellis2ImagesToViewConfigs()._parse_angles(azimuths)
el_list = Trellis2ImagesToViewConfigs()._parse_angles(elevations)
if not az_list or not el_list:
raise Exception("azimuths and elevations are required")
if len(az_list) != len(el_list):
raise Exception("azimuths and elevations must have the same amount of values")
if moge_camera_config is not None:
# A wired camera supplies the lens; the distance is settled below, since
# its own distance assumes the single-view path's cropped framing.
camera_angle_x = float(moge_camera_config['camera_angle_x'])
mesh_scale = float(moge_camera_config.get('mesh_scale', 1.0))
else:
camera_angle_x = float(fov) if fov_unit == "rad" else math.radians(float(fov))
pils = comfy_images_to_rgba_pils(
images, masks=masks, invert_mask=invert_mask,
remove_background=remove_background, max_views=len(az_list),
)
if len(pils) != len(az_list):
raise Exception(
f"got {len(pils)} images but {len(az_list)} azimuth/elevation pairs")
fill = mv_camera.measure_object_fill(pils)
fill_frame = mv_camera.camera_distance_for_fov(camera_angle_x, mesh_scale)
print(f"[Pixal3D MV] views fill {100 * fill:.1f}% of the frame "
f"(fill_frame distance would be {fill_frame:.4f})")
if distance > 0:
print(f"[Pixal3D MV] framing: distance {distance:.4f} set explicitly")
elif framing == "auto":
distance = mv_camera.camera_distance_for_extent(
camera_angle_x, fill * 512 / 2, mesh_scale)
print(f"[Pixal3D MV] framing=auto: distance {distance:.4f} fitted to the views")
elif framing == "pixal3d_rig":
distance = fill_frame * mv_camera.PIXAL3D_RIG_MARGIN
print(f"[Pixal3D MV] framing=pixal3d_rig: distance {distance:.4f} "
f"({mv_camera.PIXAL3D_RIG_MARGIN}x fill_frame, the trained 10% margin)")
elif framing == "fill_frame":
distance = fill_frame
print(f"[Pixal3D MV] framing=fill_frame: distance {distance:.4f}")
elif framing == "camera_config":
if moge_camera_config is None:
raise Exception("framing=camera_config needs a moge_camera_config input")
distance = float(moge_camera_config['distance'])
print(f"[Pixal3D MV] framing=camera_config: distance {distance:.4f}")
else:
raise Exception(f"unknown framing {framing!r}")
# The projection scales with distance, so a mismatch here quietly makes every
# grid point sample the wrong pixel -- loudest in the texture stage.
fitted = mv_camera.camera_distance_for_extent(camera_angle_x, fill * 512 / 2, mesh_scale)
if abs(distance - fitted) / fitted > 0.03:
print(f"[Pixal3D MV] Warning: distance {distance:.4f} does not match how the views "
f"are framed ({fitted:.4f} would). The projection is off by "
f"{100 * abs(distance - fitted) / fitted:.1f}%, which washes out the texture. "
f"Try framing=auto.")
views = mv_camera.build_views_from_angles(
pils, az_list, el_list, camera_angle_x,
mesh_scale=mesh_scale, distance=distance,
)
mv_camera.check_main_view(views)
cam_config = {'camera_angle_x': camera_angle_x,
'distance': float(distance),
'mesh_scale': float(mesh_scale)}
return (views, cam_config, pixal3d_views_to_preview(views),)
class Trellis2Pixal3DLoadMultiViewFolder:
"""
Load a Pixal3D multi-view input folder (the inference_mv.py format).
<folder_path>/
transforms.json mesh_scale + per-frame file_path / transform_matrix
view00_azim000.png RGBA alpha is used as the mask when present
...
transform_matrix is a 4x4 camera-to-world in the Blender/NeRF convention (Z-up
world, each camera looking along its own -Z with its own +Y up) and
camera_angle_x is the horizontal fov in radians, given per frame or once at the
top level. Frame 0 is the main view.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"folder_path": ("STRING", {"default": ""}),
"num_views": ("INT", {"default": 0, "min": 0, "max": 64,
"tooltip": "0 = every frame in transforms.json"}),
"remove_background": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("PIXAL3D_MV_VIEWS", "MOGE_CAM_CONFIG", "IMAGE",)
RETURN_NAMES = ("pixal3d_mv_views", "moge_camera_config", "preview",)
FUNCTION = "process"
CATEGORY = "Trellis2Wrapper"
OUTPUT_NODE = True
def process(self, folder_path, num_views, remove_background):
from .trellis2.utils import mv_camera
if not os.path.isdir(folder_path):
raise Exception(f"Folder not found: {folder_path}")
if not os.path.exists(os.path.join(folder_path, 'transforms.json')):
raise Exception(f"No transforms.json in {folder_path}")
rembg = None
if remove_background:
from rembg import remove
rembg = lambda im: remove(im.convert('RGB'))
views = mv_camera.load_views_from_dir(
folder_path,
num_views=None if num_views <= 0 else int(num_views),
rembg=rembg,
)
mv_camera.check_main_view(views)
cam_config = {'camera_angle_x': float(views['camera_angle_x'][0, 0]),
'distance': float(views['camera_distance'][0, 0]),
'mesh_scale': float(views['mesh_scale'])}
return (views, cam_config, pixal3d_views_to_preview(views),)
NODE_CLASS_MAPPINGS = {
"Trellis2LoadModel": Trellis2LoadModel,
"Trellis2MeshWithVoxelGenerator": Trellis2MeshWithVoxelGenerator,
@@ -7956,6 +8324,8 @@ NODE_CLASS_MAPPINGS = {
"Trellis2SelectImagesForMultiView": Trellis2SelectImagesForMultiView,
"Trellis2SmoothMeshWithPyMeshlab": Trellis2SmoothMeshWithPyMeshlab,
"Trellis2SmoothTrimeshWithPyMeshlab": Trellis2SmoothTrimeshWithPyMeshlab,
"Trellis2Pixal3DMultiViewConfig": Trellis2Pixal3DMultiViewConfig,
"Trellis2Pixal3DLoadMultiViewFolder": Trellis2Pixal3DLoadMultiViewFolder,
}
@@ -8035,4 +8405,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"Trellis2SelectImagesForMultiView": "Trellis2 - Select Images For MultiView",
"Trellis2SmoothMeshWithPyMeshlab": "Trellis2 - Smooth Mesh With PyMeshlab",
"Trellis2SmoothTrimeshWithPyMeshlab": "Trellis2 - Smooth Trimesh With PyMeshlab",
"Trellis2Pixal3DMultiViewConfig": "Trellis2 - Pixal3D MultiView Config",
"Trellis2Pixal3DLoadMultiViewFolder": "Trellis2 - Pixal3D Load MultiView Folder",
}
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "trellis2"
description = "ComfyUI Wrapper for Microsoft Trellis.2 - Native and Compact Structured Latents for 3D Generation"
version = "1.0.26"
version = "1.0.27"
license = {file = "LICENSE"}
# classifiers = [
# # For OS-independent nodes (works on all operating systems)
+58 -4
View File
@@ -278,10 +278,64 @@ class VarLenTensor:
raise ValueError(f"Unsupported reduce operation: {op}")
if dim is None or 0 in dim:
return red
red = torch.segment_reduce(red, reduce=op, lengths=self.seqlen)
return red
return red
# Defensive cache validation: in the cascade Shape-SLat sampler's
# CFG-rescale path (classifier_free_guidance_mixin.std_pos = ...),
# x_0_pos is derived from x_t via an elementwise chain that
# propagates _spatial_cache by reference (see replace() /
# __merge_sparse_cache below). If the inherited cached `seqlen`
# was computed at a different scale than x_0_pos.feats, the
# invariant sum(seqlen) == feats.shape[0] no longer holds.
# CUDA's segment_reduce kernel silently consumes the mismatch;
# the CPU path (MPS fallback) raises RuntimeError. Recompute
# lengths from coords (authoritative) and evict the stale
# cache entries so subsequent reads stay correct.
lengths = self.seqlen
n_data = red.shape[0]
if int(lengths.sum().item()) != n_data:
fresh = None
coords = getattr(self, 'coords', None)
if coords is not None and coords.shape[0] == n_data:
batch_size = int(coords[:, 0].max().item()) + 1 if coords.shape[0] > 0 else 1
fresh = torch.bincount(coords[:, 0].long(), minlength=batch_size).to(
dtype=torch.long, device=red.device
)
if fresh is None or int(fresh.sum().item()) != n_data:
fresh = torch.tensor(
[l.stop - l.start for l in self.layout],
dtype=torch.long, device=red.device,
)
if hasattr(self, '_spatial_cache') and hasattr(self, '_scale'):
try:
scale_key = str(self._scale)
slot = self._spatial_cache.get(scale_key, {})
for k in ('seqlen', 'cum_seqlen', 'batch_boardcast_map', 'layout'):
slot.pop(k, None)
except Exception:
pass
elif hasattr(self, '_cache'):
try:
for k in ('seqlen', 'cum_seqlen', 'batch_boardcast_map'):
self._cache.pop(k, None)
except Exception:
pass
lengths = fresh
if int(lengths.sum().item()) != n_data:
raise RuntimeError(
f"VarLenTensor.reduce: cannot reconcile seqlen "
f"sum({int(lengths.sum().item())}) with data.size(0)={n_data}. "
f"layout has {len(self.layout)} segments."
)
# torch.segment_reduce has no Metal kernel; PyTorch's auto fallback
# is racy on cascade-sized workloads. Run explicitly on CPU and
# copy back. No-op on CUDA / CPU.
if red.device.type == 'mps':
reduced = torch.segment_reduce(red.cpu(), reduce=op, lengths=lengths.cpu())
return reduced.to(red.device)
return torch.segment_reduce(red, reduce=op, lengths=lengths)
def mean(self, dim: Optional[Union[int, Tuple[int,...]]] = None, keepdim: bool = False) -> torch.Tensor:
return self.reduce(op='mean', dim=dim, keepdim=keepdim)
+199 -6
View File
@@ -32,7 +32,13 @@ def build_pixal3d_image_cond_model(config: dict):
from ..trainers.flow_matching.mixins.image_conditioned_proj import DinoV3ProjFeatureExtractor
model = DinoV3ProjFeatureExtractor(**config)
model.eval()
return model
return model
def build_pixal3d_mv_image_cond_model(config: dict):
from ..trainers.flow_matching.mixins.image_conditioned_proj import DinoV3ProjMultiViewFeatureExtractor
model = DinoV3ProjMultiViewFeatureExtractor(**config)
model.eval()
return model
def pil2tensor(image):
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0)[None,]
@@ -142,8 +148,16 @@ class Trellis2ImageTo3DPipeline(Pipeline):
"use_naf_upsample": True,
"naf_target_size": 1024,
},
}
}
# Multi-view variants of the four stages. Same DINOv3 backbone and grid
# resolutions as above -- only the extractor class and the fusion differ.
# "average" keeps the fused feature shape identical to single-view, which is
# why the *_mv denoisers need no architectural change.
self.PIXAL3D_MV_IMAGE_COND_CONFIGS = {
key: {**cfg, "multiview_fusion": "average"}
for key, cfg in self.PIXAL3D_IMAGE_COND_CONFIGS.items()
}
def switch_samplers(self, sampler_type: str = "euler"):
"""Dynamically switches the sampler instances based on user selection."""
self._sampler_prefix = "Euler"
@@ -186,7 +200,7 @@ class Trellis2ImageTo3DPipeline(Pipeline):
torch.cuda.empty_cache()
@classmethod
def from_pretrained(cls, path: str, config_file: str = "pipeline.json", keep_models_loaded = True, use_fp8 = False, use_reconviagen = False, isPixal3D = False) -> "Trellis2ImageTo3DPipeline":
def from_pretrained(cls, path: str, config_file: str = "pipeline.json", keep_models_loaded = True, use_fp8 = False, use_reconviagen = False, isPixal3D = False, isPixal3DMV = False) -> "Trellis2ImageTo3DPipeline":
"""
Load a pretrained model.
@@ -197,6 +211,10 @@ class Trellis2ImageTo3DPipeline(Pipeline):
config_file = "reconviagen_pipeline.json"
elif use_fp8:
config_file = "pipeline_fp8.json"
elif isPixal3DMV:
# The multi-view denoisers live next to the single-view ones under
# ckpts/*_mv in the same repo; this config file is what points at them.
config_file = "pipeline_mv.json"
pipeline = super().from_pretrained(path, config_file)
args = pipeline._pretrained_args
@@ -233,7 +251,8 @@ class Trellis2ImageTo3DPipeline(Pipeline):
pipeline.last_processing = ''
pipeline.use_fp8 = use_fp8
pipeline.isPixal3D = isPixal3D
pipeline.isPixal3DMV = isPixal3DMV
if not isPixal3D:
pipeline._pretrained_args['models']['sparse_structure_decoder'] = os.path.join(folder_paths.models_dir,"microsoft","TRELLIS-image-large","ckpts","ss_dec_conv3d_16l8_fp16")
@@ -245,6 +264,9 @@ class Trellis2ImageTo3DPipeline(Pipeline):
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["shape_1024"]["model_name"] = facebook_model_path
pipeline.PIXAL3D_IMAGE_COND_CONFIGS["tex_1024"]["model_name"] = facebook_model_path
for _cfg in pipeline.PIXAL3D_MV_IMAGE_COND_CONFIGS.values():
_cfg["model_name"] = facebook_model_path
try:
from mmgp import safetensors2 as _mmgp_st2
if 'C64' not in _mmgp_st2._map_to_dtype:
@@ -345,7 +367,50 @@ class Trellis2ImageTo3DPipeline(Pipeline):
if hasattr(self,'pixal3d_image_cond_tex_1024') and self.pixal3d_image_cond_tex_1024 is not None:
del self.pixal3d_image_cond_tex_1024
self.pixal3d_image_cond_tex_1024 = None
self._cleanup_cuda()
self._cleanup_cuda()
# ---- Pixal3D multi-view image cond models (DinoV3ProjMultiViewFeatureExtractor) ----
def _load_pixal3d_mv_image_cond(self, stage: str):
attr = f'pixal3d_mv_image_cond_{stage}'
model = getattr(self, attr, None)
if model is not None:
return model
print(f'Loading Pixal3D MultiView Image Cond {stage} Model ...')
model = build_pixal3d_mv_image_cond_model(self.PIXAL3D_MV_IMAGE_COND_CONFIGS[stage])
setattr(self, attr, model)
return model
def _unload_pixal3d_mv_image_cond(self, stage: str):
attr = f'pixal3d_mv_image_cond_{stage}'
if getattr(self, attr, None) is not None:
setattr(self, attr, None)
self._cleanup_cuda()
def load_pixal3d_mv_image_cond_ss(self):
return self._load_pixal3d_mv_image_cond('ss')
def unload_pixal3d_mv_image_cond_ss(self):
self._unload_pixal3d_mv_image_cond('ss')
def load_pixal3d_mv_image_cond_shape_512(self):
return self._load_pixal3d_mv_image_cond('shape_512')
def unload_pixal3d_mv_image_cond_shape_512(self):
self._unload_pixal3d_mv_image_cond('shape_512')
def load_pixal3d_mv_image_cond_shape_1024(self):
return self._load_pixal3d_mv_image_cond('shape_1024')
def unload_pixal3d_mv_image_cond_shape_1024(self):
self._unload_pixal3d_mv_image_cond('shape_1024')
def load_pixal3d_mv_image_cond_tex_1024(self):
return self._load_pixal3d_mv_image_cond('tex_1024')
def unload_pixal3d_mv_image_cond_tex_1024(self):
self._unload_pixal3d_mv_image_cond('tex_1024')
def load_sparse_structure_model(self):
if self.models['sparse_structure_flow_model'] is None:
@@ -646,6 +711,134 @@ class Trellis2ImageTo3DPipeline(Pipeline):
'neg_cond': {'global': torch.zeros_like(z_global), 'proj': SparseTensor(feats=torch.zeros_like(z_proj_sparse), coords=coords)},
}
# ------------------------------------------------------------------
# Pixal3D multi-view proj conditioning
#
# Same cascade as the single-view Pixal3D path (SS -> Shape 512 -> Shape 1024
# -> Tex 1024); the only difference is the image condition. Instead of one
# image at the canonical front view it takes V views with explicit c2w
# matrices and lets DinoV3ProjMultiViewFeatureExtractor project all of them
# into the shared 3D grid, averaging the per-view features.
#
# Because the fusion is "average", the fused z_proj has exactly the
# single-view shape [B, R^3, C], so every per-block proj_linear / proj
# cross-attn weight of the *_mv denoisers stays compatible. View 0 is the
# main view and its calc_mat_0 == F by construction, so V=1 degenerates to
# the single-view path.
#
# `views` is the bundle built by trellis2.utils.mv_camera.build_views().
# ------------------------------------------------------------------
@staticmethod
def _mv_image_for(views: dict, image_cond_model: nn.Module) -> torch.Tensor:
"""
Pick the pre-decoded [1, V, 3, H, W] batch matching this stage input size.
Unlike the single-view extractor (which takes PIL and resizes internally), the
multi-view one takes tensors, so the caller has to supply the right resolution:
the SS / shape-512 stages run at 512 and the 1024 stages at 1024.
"""
size = int(image_cond_model.image_size)
if size not in views['images']:
raise KeyError(
f"No view images at resolution {size}; got {sorted(views['images'])}. "
f"Build them with mv_camera.build_views(image_sizes=...).")
return views['images'][size]
def _run_mv_extractor(self, image_cond_model: nn.Module, views: dict):
device = self.device
image = self._mv_image_for(views, image_cond_model).to(device)
cam_angle = views['camera_angle_x'].to(device=device, dtype=torch.float32)
dist = views['camera_distance'].to(device=device, dtype=torch.float32)
transform_matrix = views['transform_matrix'].to(device=device, dtype=torch.float32)
# mesh_scale is per-object, i.e. [B]; the extractor expands it over views.
scale = torch.as_tensor(views['mesh_scale'], dtype=torch.float32, device=device).reshape(-1)
return image_cond_model(
image,
camera_angle_x=cam_angle,
distance=dist,
mesh_scale=scale,
transform_matrix=transform_matrix,
)
@torch.no_grad()
def get_proj_cond_ss_mv(
self,
views: dict,
image_cond_model: nn.Module = None,
) -> dict:
"""Multi-view proj conditioning for the sparse structure stage (dense grid)."""
print('Getting MultiView Proj Image Cond ...')
if image_cond_model is None:
image_cond_model = self.load_pixal3d_mv_image_cond_ss()
# The MV cond models are built lazily, after pipeline.to(), so they always
# start on CPU -- move them in regardless of mode and only offload again
# when low_vram is on.
image_cond_model.to(self.device)
z_global, z_proj = self._run_mv_extractor(image_cond_model, views)
if self.low_vram:
image_cond_model.cpu()
return {
'cond': {'global': z_global, 'proj': z_proj},
'neg_cond': {'global': torch.zeros_like(z_global), 'proj': torch.zeros_like(z_proj)},
}
@torch.no_grad()
def get_proj_cond_shape_mv(
self,
image_cond_model: nn.Module,
views: dict,
coords: torch.Tensor,
grid_resolution_override: int = None,
) -> dict:
"""Multi-view proj conditioning for the shape / texture stages (sparse tokens)."""
print('Getting MultiView Projected Image Cond ...')
device = self.device
# See get_proj_cond_ss_mv: these models are built after pipeline.to().
image_cond_model.to(device)
# The HR cascade grid resolution floats with the token budget (1536/16=96 and
# down), so it can differ from what this stage was trained at; swap the grid.
orig_grid_res = image_cond_model.grid_resolution
override = (grid_resolution_override is not None
and grid_resolution_override != orig_grid_res)
if override:
image_cond_model.grid_resolution = grid_resolution_override
image_cond_model.proj_grid = image_cond_model.proj_grid.__class__(
grid_resolution=grid_resolution_override,
image_resolution=image_cond_model.proj_grid.image_resolution,
).to(device)
z_global, z_proj = self._run_mv_extractor(image_cond_model, views)
B = z_global.shape[0]
grid_res = image_cond_model.grid_resolution
b_idx = coords[:, 0].long()
x_idx = coords[:, 1].long()
y_idx = coords[:, 2].long()
z_idx = coords[:, 3].long()
z_proj_grid = z_proj.reshape(B, grid_res, grid_res, grid_res, -1)
z_proj_sparse = z_proj_grid[b_idx, x_idx, y_idx, z_idx]
z_proj_st = SparseTensor(feats=z_proj_sparse, coords=coords)
if override:
image_cond_model.grid_resolution = orig_grid_res
image_cond_model.proj_grid = image_cond_model.proj_grid.__class__(
grid_resolution=orig_grid_res,
image_resolution=image_cond_model.proj_grid.image_resolution,
).to(device)
if self.low_vram:
image_cond_model.cpu()
return {
'cond': {'global': z_global, 'proj': z_proj_st},
'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
@@ -337,6 +337,102 @@ class ProjGrid(nn.Module):
return vis_images
class ProjGridMV(ProjGrid):
"""
Multi-view variant of ProjGrid.
Identical to ProjGrid except it ALLOWS an explicit per-view transform_matrix
(calc_mat) to be passed in for projection (the base ProjGrid asserts it is
None and always uses the fixed front-view matrix). This is used by the
multi-view feature extractor to project the SAME 3D grid through each view's
relative pose (calc_mat_i = F @ inv(C_0) @ C_i).
When transform_matrix is None it falls back to the front-view + distance
behavior, so it is a strict superset of ProjGrid.
"""
def forward(
self,
features_map: torch.Tensor,
camera_angle_x: torch.Tensor,
distance: torch.Tensor,
mesh_scale: torch.Tensor,
transform_matrix: Optional[torch.Tensor] = None,
BHWC: bool = True,
) -> torch.Tensor:
if BHWC:
B, H, W, C = features_map.shape
else:
B, C, H, W = features_map.shape
grid_points = self.grid_points
grid_points = grid_points.expand(B, -1, -1)
grid_points = grid_points / mesh_scale.unsqueeze(-1).unsqueeze(-1) / 2 # Scale alignment
if transform_matrix is None:
transform_matrix = self.front_view_transform_matrix
transform_matrix = transform_matrix.expand(B, -1, -1).clone()
transform_matrix[:, 1, 3] = -distance # Set camera distance
# else: use the provided per-view calc_mat directly
image_points, depth, valid_mask = project_points_to_image_batch(
grid_points, transform_matrix, camera_angle_x, self.image_resolution
)
image_points_norm = (image_points + 0.5) / self.image_resolution * 2 - 1
if BHWC:
features_map = features_map.permute(0, 3, 1, 2) # [B, C, H, W]
x = sample_features(features_map, image_points_norm) # [B, C, K]
x = x.permute(0, 2, 1) # [B, K, C]
return x
def compute_relative_calc_mat(
transform_matrix: torch.Tensor,
distance: torch.Tensor,
front_view_transform_matrix: torch.Tensor,
) -> torch.Tensor:
"""
Compute the per-view projection matrix (calc_mat) that maps each view into
the coordinate frame where the MAIN view (index 0) is snapped to the fixed
front-view pose F (with F's camera distance set to the main view's distance).
calc_mat_i = F @ inv(C_0) @ C_i
where C_i = transform_matrix[:, i] (c2w). For i == 0, calc_mat_0 == F exactly,
which reproduces the single-view behavior.
Args:
transform_matrix: [B, V, 4, 4] c2w matrices (index 0 = main view).
distance: [B, V] camera distances (only index 0 used for F).
front_view_transform_matrix: [4, 4] the canonical front-view c2w.
Returns:
calc_mat: [B, V, 4, 4]
"""
B, V = transform_matrix.shape[:2]
device = transform_matrix.device
# Fixed front matrix with main-view distance in the translation slot.
F_mat = front_view_transform_matrix.to(device).unsqueeze(0).expand(B, -1, -1).clone() # [B,4,4]
F_mat[:, 1, 3] = -distance[:, 0]
F_mat = F_mat.unsqueeze(1) # [B,1,4,4]
C0 = transform_matrix[:, 0:1] # [B,1,4,4]
# Do the matrix math in fp32 for numerical stability (inv is sensitive).
with torch.amp.autocast('cuda', enabled=False):
C0f = C0.float().expand(B, V, 4, 4).reshape(B * V, 4, 4)
Cif = transform_matrix.float().reshape(B * V, 4, 4)
Ff = F_mat.float().expand(B, V, 4, 4).reshape(B * V, 4, 4)
rel = torch.bmm(torch.linalg.inv(C0f), Cif) # inv(C_0) @ C_i
calc = torch.bmm(Ff, rel) # F @ rel
calc_mat = calc.reshape(B, V, 4, 4)
return calc_mat
# =============================================================================
# DINOv3 Feature Extractor with Projection
# =============================================================================
@@ -479,11 +575,21 @@ class DinoV3ProjFeatureExtractor(nn.Module):
hidden_states = self.model.embeddings(image, bool_masked_pos=None)
position_embeddings = self.model.rope_embeddings(image)
for layer_module in self.model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
# transformers < 5
if hasattr(self.model,'layer'):
for layer_module in self.model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
elif hasattr(self.model,'model') and hasattr(self.model.model,'layer'): # transformers >= 5
for layer_module in self.model.model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
else:
raise Exception("Cannot extract features")
return F.layer_norm(hidden_states, hidden_states.shape[-1:])
@@ -781,6 +887,165 @@ class DinoV3ProjFeatureExtractor(nn.Module):
)
# =============================================================================
# Multi-View DINOv3 Feature Extractor with Projection
# =============================================================================
class DinoV3ProjMultiViewFeatureExtractor(DinoV3ProjFeatureExtractor):
"""
Multi-view DINOv3 feature extractor with view-aligned projection.
Reuses DinoV3ProjFeatureExtractor's DINOv3 backbone + NAF, but:
- accepts multi-view inputs: image [B, V, 3, H, W], camera_angle_x [B, V],
distance [B, V], mesh_scale [B], transform_matrix [B, V, 4, 4] (c2w, index0=main).
- computes per-view calc_mat = F @ inv(C_0) @ C_i and projects the SAME 3D grid
through each view's relative pose (uses ProjGridMV).
- fuses the V views according to `multiview_fusion`:
* "average": z_proj -> [B, R^3, C] (mean over V, matches single-view);
z_global -> [B, 1+num_reg, C] (mean over V).
* "attention": z_proj -> [B, R^3, V, C] (keep V dim for per-voxel attn);
z_global -> [B, V*(1+num_reg), C] (concat over V as longer KV).
proj_channels is inherited (embed_dim, or embed_dim*2 with NAF), so downstream
per-block proj_linear / proj cross-attn weights are compatible with single-view.
"""
def __init__(self, *args, multiview_fusion: str = "average", **kwargs):
super().__init__(*args, **kwargs)
assert multiview_fusion in ("average", "attention"), \
f"multiview_fusion must be average or attention, got {multiview_fusion}"
self.multiview_fusion = multiview_fusion
# Replace proj_grid with the multi-view variant that accepts calc_mat.
self.proj_grid = ProjGridMV(
grid_resolution=self.grid_resolution,
image_resolution=self.image_size,
)
def _project_single_view(self, z_patchtokens_spatial, image_for_naf,
camera_angle_x_v, distance_v, mesh_scale_v, calc_mat_v):
"""Project one view of DINOv3 features to the 3D grid. Returns [b, R^3, C]."""
z_proj_lr = self.proj_grid(
z_patchtokens_spatial, camera_angle_x_v, distance_v, mesh_scale_v, calc_mat_v
)
if not self.use_naf_upsample:
return z_proj_lr # [b, R^3, D]
self._load_naf()
lr_features_bchw = z_patchtokens_spatial.permute(0, 3, 1, 2)
# Honour the same NAF tiling knob the single-view path uses (see
# DinoV3ProjFeatureExtractor.naf_tile_factor). 1 = un-tiled.
K = getattr(self, 'naf_tile_factor', 1) or 1
if K <= 1:
hr_features = self.naf_model(image_for_naf, lr_features_bchw, self.naf_target_size)
z_proj_hr = self.proj_grid(
hr_features, camera_angle_x_v, distance_v, mesh_scale_v, calc_mat_v, BHWC=False
)
del hr_features
else:
z_proj_hr = self._proj_naf_tiled(
image_for_naf, lr_features_bchw,
camera_angle_x_v, distance_v, mesh_scale_v, calc_mat_v,
tile_factor=int(K),
)
return torch.cat([z_proj_lr, z_proj_hr], dim=-1) # [b, R^3, 2D]
def forward(
self,
image: torch.Tensor,
camera_angle_x: Optional[torch.Tensor] = None,
distance: Optional[torch.Tensor] = None,
mesh_scale: Optional[torch.Tensor] = None,
transform_matrix: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Args:
image: [B, V, 3, H, W]
camera_angle_x: [B, V]
distance: [B, V]
mesh_scale: [B]
transform_matrix: [B, V, 4, 4] (c2w, index 0 = main view)
Returns:
(z_global, z_proj) - shapes depend on multiview_fusion (see class doc).
"""
assert image.ndim == 5, f"Expected image [B,V,3,H,W], got {tuple(image.shape)}"
B, V = image.shape[:2]
if camera_angle_x is None or distance is None or mesh_scale is None or transform_matrix is None:
raise ValueError("camera_angle_x, distance, mesh_scale, transform_matrix must be provided")
# Compute per-view calc_mat = F @ inv(C_0) @ C_i -> [B, V, 4, 4]
calc_mat = compute_relative_calc_mat(
transform_matrix, distance, self.proj_grid.front_view_transform_matrix
)
# Flatten B*V and run DINOv3 once.
img_flat = image.reshape(B * V, *image.shape[2:]) # [B*V, 3, H, W]
if self.use_naf_upsample:
image_for_naf = img_flat.clone()
img_norm = self.transform(img_flat)
with torch.no_grad():
z = self.extract_features(img_norm)
z_clstoken = z[:, 0:1]
num_reg = getattr(self.model.config, 'num_register_tokens', 4)
z_regtokens = z[:, 1:1 + num_reg]
z_patchtokens = z[:, 1 + num_reg:]
z_patchtokens_spatial = z_patchtokens.reshape(
B * V, self.patch_number, self.patch_number, -1
)
camera_angle_x_flat = camera_angle_x.reshape(B * V)
distance_flat = distance.reshape(B * V)
# mesh_scale is shared per object -> expand to per-view then flatten.
mesh_scale_flat = mesh_scale[:, None].expand(B, V).reshape(B * V)
calc_mat_flat = calc_mat.reshape(B * V, 4, 4)
# Project each view (loop over V to bound peak memory).
# A per-view projection is [B, R^3, C], which at R=64 / C=2048 is 2 GiB
# in fp32. Averaging accumulates in place so peak stays at one view
# instead of growing with V (a list + stack would cost 2*V*2 GiB, i.e.
# ~40 GiB at V=9, which both starves the denoiser and makes the peak of
# a step depend on N -- something the elastic memory controller models
# purely from token count and therefore mispredicts).
fuse_by_average = self.multiview_fusion == "average"
z_proj_acc = None # average: running sum
per_view_proj = [] if not fuse_by_average else None # attention: keep V
for v in range(V):
idx = torch.arange(v, B * V, V, device=z.device)
img_naf_v = image_for_naf[idx] if self.use_naf_upsample else None
z_view = self._project_single_view(
z_patchtokens_spatial[idx],
img_naf_v,
camera_angle_x_flat[idx],
distance_flat[idx],
mesh_scale_flat[idx],
calc_mat_flat[idx],
)
if fuse_by_average:
if z_proj_acc is None:
# Without NAF the projection is a permuted view, so make the
# accumulator contiguous before adding into it in place.
z_proj_acc = z_view.contiguous()
else:
z_proj_acc.add_(z_view)
del z_view
else:
per_view_proj.append(z_view)
# z_global per view: [B, V, 1+num_reg, D]
z_global_all = torch.cat([z_clstoken, z_regtokens], dim=1) # [B*V, 1+num_reg, D]
z_global_all = z_global_all.reshape(B, V, z_global_all.shape[-2], z_global_all.shape[-1])
if fuse_by_average:
z_proj = z_proj_acc.div_(V) # [B, R^3, C]
z_global = z_global_all.mean(dim=1) # [B, 1+num_reg, D]
else: # attention
z_proj = torch.stack(per_view_proj, dim=2) # [B, R^3, V, C]
z_global = z_global_all.reshape(B, V * z_global_all.shape[2], z_global_all.shape[3])
return z_global, z_proj
# =============================================================================
# DINOv3 + VAE Gated Feature Extractor with Projection
# =============================================================================
@@ -890,11 +1155,23 @@ class DinoV3VaeProjFeatureExtractor(nn.Module):
image = image.to(self.dino_model.embeddings.patch_embeddings.weight.dtype)
hidden_states = self.dino_model.embeddings(image, bool_masked_pos=None)
position_embeddings = self.dino_model.rope_embeddings(image)
for layer_module in self.dino_model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
# transformers < 5
if hasattr(self.dino_model,'layer'):
for layer_module in self.dino_model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
elif hasattr(self.dino_model,'model') and hasattr(self.dino_model.model,'layer'): # transformers >= 5
for layer_module in self.dino_model.model.layer:
hidden_states = layer_module(
hidden_states,
position_embeddings=position_embeddings,
)
else:
raise Exception("Cannot extract features")
return F.layer_norm(hidden_states, hidden_states.shape[-1:])
@torch.no_grad()
+373
View File
@@ -0,0 +1,373 @@
"""
Camera / view-bundle helpers for the Pixal3D multi-view path.
Pixal3D's multi-view conditioning wants the same "dataset-style" bundle that
inference_mv.py builds from a `transforms.json` directory:
{
'images': {image_size: [1, V, 3, S, S]}, alpha-premultiplied, in [0,1]
'camera_angle_x': [1, V] horizontal fov, radians
'camera_distance': [1, V] camera distance (norm of the c2w translation)
'transform_matrix': [1, V, 4, 4] camera-to-world, index 0 = MAIN view
'mesh_scale': float
'view_names': [str] * V
}
`transform_matrix` follows the Blender/NeRF convention used by the training
renders: a world that is Z-up, each camera looking along its own -Z with its own
+Y as up. Frame 0 is the main view and should be the canonical front view -- a
camera at (0, -d, 0) -- because every other view is placed relative to it
(calc_mat_i = F @ inv(C_0) @ C_i, see compute_relative_calc_mat).
Azimuth / elevation here match the convention used by
Trellis2RenderMultiViewNvdiffrast, which works in the model's own Y-up mesh
frame:
eye_mesh = d * [cos(el)sin(az), sin(el), cos(el)cos(az)]
so az=0/90/180/270 -> front/left/back/right and el=+/-90 -> top/bottom. ProjGrid
rotates the mesh frame into the Blender frame with (x, y, z) -> (x, -z, y), so
the same camera in Blender coordinates sits at
eye = d * [cos(el)sin(az), -cos(el)cos(az), sin(el)]
which is exactly (0, -d, 0) at az=el=0 -- the canonical front view.
"""
import os
import json
import math
from typing import *
import numpy as np
import torch
from PIL import Image
from .camera import compute_f_pixels, distance_from_fov
# Blender-frame world up (Z-up).
_WORLD_UP = (0.0, 0.0, 1.0)
# The Pixal3D training rig leaves a 10% margin around the object: the shipped
# assets/mv_images/example/transforms.json has distance 3.1192050 at a 20 deg fov,
# and the fill-the-frame distance for that fov is 2.8356409 -- exactly 1.1x smaller.
# Multi-view renders (and the outputs of most multi-view diffusion models) come
# framed this way, unlike the single-view path, where preprocess_image crops the
# object until it fills the frame.
PIXAL3D_RIG_MARGIN = 1.1
def camera_distance_for_extent(camera_angle_x: float, half_extent_px: float,
mesh_scale: float = 1.0, image_resolution: int = 512) -> float:
"""
Distance at which the unit box half-extent projects to `half_extent_px` pixels.
This is the one relation that has to hold between the views and the cameras: the
projection maps the [-0.5, 0.5]^3 grid (scaled by mesh_scale) into each image, so
if the object is drawn smaller than the camera says, every grid point samples too
far out -- the surface ends up reading the background.
"""
if half_extent_px <= 0:
raise ValueError("half_extent_px must be positive")
f_pixels = compute_f_pixels(camera_angle_x, image_resolution)
return float(f_pixels * 0.5 / (mesh_scale * half_extent_px))
def camera_distance_for_fov(camera_angle_x: float, mesh_scale: float = 1.0,
image_resolution: int = 512, extend_pixel: int = 0) -> float:
"""
Distance at which a unit object of `mesh_scale` exactly fills the frame.
Same formula the single-view Pixal3D path uses (Trellis2FovMoGeCameraConfig /
get_camera_params_wild_moge). Note that the single-view path only gets away with
it because preprocess_image crops the object to fill the frame first; un-cropped
multi-view renders usually want PIXAL3D_RIG_MARGIN times this.
"""
grid_point = torch.tensor([-1.0, 0.0, 0.0])
target_point = torch.tensor([0 - extend_pixel, image_resolution - 1 + extend_pixel])
return float(distance_from_fov(
camera_angle_x, grid_point, target_point, mesh_scale, image_resolution
)["distance_from_x"])
def measure_object_fill(images: List[Image.Image], alpha_threshold: float = 0.8) -> float:
"""
Fraction of the frame the object occupies, as the largest silhouette extent
across the views.
Mirrors the framing convention preprocess_image establishes for a single view:
a square crop of side max(bbox_w, bbox_h), so the object's *longest* axis is what
maps to the frame. Taking the max over views is the tightest rig that still
contains every view.
"""
best = 0.0
for im in images:
alpha = np.array(im.convert('RGBA').getchannel(3))
ys, xs = np.nonzero(alpha > alpha_threshold * 255)
if len(xs) == 0:
continue
w = int(xs.max() - xs.min() + 1)
h = int(ys.max() - ys.min() + 1)
best = max(best, max(w, h) / max(im.size))
if best <= 0:
raise ValueError("every view is empty; cannot measure the framing")
return best
def blender_c2w_from_azimuth_elevation(azimuth_deg: float, elevation_deg: float,
distance: float) -> np.ndarray:
"""
Build a 4x4 camera-to-world matrix in the Blender/NeRF convention Pixal3D uses.
Args:
azimuth_deg: 0 = front, 90 = left, 180 = back, 270 = right.
elevation_deg: positive = above the object.
distance: camera distance from the origin.
Returns:
[4, 4] float64 c2w matrix. Columns are (right, up, back, eye); the camera
looks along -back.
"""
az = math.radians(float(azimuth_deg))
el = math.radians(float(elevation_deg))
eye = np.array([
math.cos(el) * math.sin(az),
-math.cos(el) * math.cos(az),
math.sin(el),
], dtype=np.float64) * float(distance)
# +Z of the camera points away from the object (the camera looks along -Z).
back = eye / (np.linalg.norm(eye) + 1e-12)
world_up = np.array(_WORLD_UP, dtype=np.float64)
if abs(float(np.dot(back, world_up))) > 0.999:
# Looking straight down / up: the world up is degenerate, so pick the
# in-plane reference that the surrounding elevations converge to. At the
# top pole (back_z > 0) that limit is +Y, at the bottom pole it is -Y;
# using one fixed vector for both rolls one of the two views by 180 deg.
world_up = np.array([0.0, 1.0 if back[2] > 0 else -1.0, 0.0], dtype=np.float64)
right = np.cross(world_up, back)
right = right / (np.linalg.norm(right) + 1e-12)
up = np.cross(back, right)
up = up / (np.linalg.norm(up) + 1e-12)
c2w = np.eye(4, dtype=np.float64)
c2w[:3, 0] = right
c2w[:3, 1] = up
c2w[:3, 2] = back
c2w[:3, 3] = eye
return c2w
def to_cond_tensor(image: Image.Image, image_size: int) -> torch.Tensor:
"""
Turn an RGBA view into a conditioning tensor the way training read its views:
LANCZOS resize, then premultiply by alpha so the background is black.
"""
image = image.convert('RGBA').resize((image_size, image_size), Image.Resampling.LANCZOS)
alpha = torch.tensor(np.array(image.getchannel(3))).float() / 255.0
rgb = torch.tensor(np.array(image.convert('RGB'))).permute(2, 0, 1).float() / 255.0
return rgb * alpha.unsqueeze(0)
def build_views(
images: List[Image.Image],
transform_matrix: Union[np.ndarray, torch.Tensor, List],
camera_angle_x: Union[float, Sequence[float]],
mesh_scale: float = 1.0,
image_sizes: Sequence[int] = (512, 1024),
view_names: Optional[List[str]] = None,
) -> dict:
"""
Assemble the views bundle the MV conditioning path consumes.
Args:
images: V RGBA PIL images. Alpha is used as the object mask, so views that
carry no real mask must be matted BEFORE they get here (the views are
never cropped or rescaled -- the framing has to be the framing the
cameras describe).
transform_matrix: [V, 4, 4] c2w matrices, index 0 = main view.
camera_angle_x: horizontal fov in radians, one value or one per view.
mesh_scale: per-object mesh scale.
image_sizes: the stage resolutions to pre-decode (512 for SS / shape-512,
1024 for the two 1024 stages).
view_names: optional labels, used only for logging.
"""
V = len(images)
if V == 0:
raise ValueError("build_views needs at least one view")
tm = torch.as_tensor(np.asarray(transform_matrix, dtype=np.float32), dtype=torch.float32)
if tm.ndim != 3 or tm.shape[-2:] != (4, 4):
raise ValueError(f"transform_matrix must be [V, 4, 4], got {tuple(tm.shape)}")
if tm.shape[0] != V:
raise ValueError(f"got {V} images but {tm.shape[0]} transform matrices")
tm = tm[None] # [1, V, 4, 4]
if isinstance(camera_angle_x, (int, float)):
cax = torch.full((1, V), float(camera_angle_x), dtype=torch.float32)
else:
cax = torch.tensor([float(a) for a in camera_angle_x], dtype=torch.float32)[None]
if cax.shape[1] != V:
raise ValueError(f"got {V} images but {cax.shape[1]} camera_angle_x values")
# Derive the distance from the pose rather than trusting a separate field, so it
# can never disagree with transform_matrix.
camera_distance = torch.norm(tm[:, :, :3, 3], dim=-1) # [1, V]
bundle_images = {
int(size): torch.stack([to_cond_tensor(im, int(size)) for im in images], dim=0)[None]
for size in image_sizes # [1, V, 3, S, S]
}
if view_names is None:
view_names = [f"view{i:02d}" for i in range(V)]
print(f"[Pixal3D MV] V={V} ({', '.join(view_names)})")
print(f"[Pixal3D MV] fov={math.degrees(float(cax[0, 0])):.2f}deg, "
f"distance={float(camera_distance[0, 0]):.4f}, mesh_scale={float(mesh_scale):.4f}")
return {
'images': bundle_images,
'camera_angle_x': cax,
'camera_distance': camera_distance,
'transform_matrix': tm,
'mesh_scale': float(mesh_scale),
'view_names': list(view_names),
}
def build_views_from_angles(
images: List[Image.Image],
azimuths: Sequence[float],
elevations: Sequence[float],
camera_angle_x: float,
mesh_scale: float = 1.0,
distance: Optional[float] = None,
image_sizes: Sequence[int] = (512, 1024),
) -> dict:
"""
Build the views bundle for an orbit described by azimuth / elevation angles.
`distance` defaults to the framing distance implied by the fov and mesh scale,
the same one the single-view path derives from MoGe.
"""
if len(images) != len(azimuths) or len(images) != len(elevations):
raise ValueError(
f"images ({len(images)}), azimuths ({len(azimuths)}) and elevations "
f"({len(elevations)}) must have the same length")
if distance is None:
distance = camera_distance_for_fov(camera_angle_x, mesh_scale)
transform_matrix = np.stack([
blender_c2w_from_azimuth_elevation(a, e, distance)
for a, e in zip(azimuths, elevations)
], axis=0)
view_names = [f"azim{int(round(a)) % 360:03d}_elev{int(round(e)):+03d}"
for a, e in zip(azimuths, elevations)]
return build_views(
images, transform_matrix, camera_angle_x,
mesh_scale=mesh_scale, image_sizes=image_sizes, view_names=view_names,
)
def load_rgba(path: str, rembg: Optional[Callable] = None) -> Tuple[Image.Image, bool]:
"""
Read one view as RGBA, matting it first if it does not already carry a mask.
Returns (image, was_matted). The alpha test matches preprocess_image: a fully
opaque alpha channel counts as no mask.
"""
image = Image.open(path)
alpha = np.array(image.getchannel(3)) if image.mode == 'RGBA' else None
if alpha is not None and not np.all(alpha == 255):
return image.convert('RGBA'), False
if rembg is None:
raise ValueError(f"{path} has no alpha channel and no matting model was given")
return rembg(image).convert('RGBA'), True
def load_views_from_dir(
views_dir: str,
num_views: Optional[int] = None,
rembg: Optional[Callable] = None,
image_sizes: Sequence[int] = (512, 1024),
) -> dict:
"""
Load a `transforms.json` view directory, the input format of inference_mv.py.
<views_dir>/
transforms.json mesh_scale + per-frame file_path / transform_matrix
view00_azim000.png RGBA alpha is used as the mask if present
...
`camera_angle_x` may be given per frame or once at the top level.
"""
with open(os.path.join(views_dir, 'transforms.json')) as f:
meta = json.load(f)
frames = meta['frames']
if num_views is not None:
if num_views > len(frames):
raise ValueError(f"num_views {num_views} > {len(frames)} views in {views_dir}")
frames = frames[:num_views]
def camera_angle_x_of(frame):
for src in (frame, meta):
if 'camera_angle_x' in src:
return float(src['camera_angle_x'])
raise KeyError(f"camera_angle_x missing for {frame.get('file_path')}")
paths = [os.path.join(views_dir, fr['file_path']) for fr in frames]
loaded = [load_rgba(p, rembg) for p in paths]
rgba = [im for im, _ in loaded]
matted = sum(was_matted for _, was_matted in loaded)
if matted:
print(f"[Pixal3D MV] matted {matted}/{len(paths)} view(s) that had no alpha channel")
views = build_views(
rgba,
np.array([fr['transform_matrix'] for fr in frames], dtype=np.float32),
[camera_angle_x_of(fr) for fr in frames],
mesh_scale=float(meta.get('mesh_scale', 1.0)),
image_sizes=image_sizes,
view_names=[fr.get('name', os.path.splitext(fr['file_path'])[0]) for fr in frames],
)
print(f"[Pixal3D MV] loaded from {views_dir}")
return views
def check_main_view(views: dict, atol: float = 1e-4) -> float:
"""
Warn if frame 0 is not the canonical front view, and return the deviation.
The extractor maps every view through calc_mat_i = F @ inv(C_0) @ C_i, so the
main view is always snapped onto F -- but if C_0 is not itself a front view,
the whole rig gets rotated relative to the object and the generated mesh comes
out in a different frame than the models were trained for.
"""
from ..trainers.flow_matching.mixins.image_conditioned_proj import ProjGridMV
F_mat = ProjGridMV(grid_resolution=2, image_resolution=64).front_view_transform_matrix.clone()
F_mat[1, 3] = -views['camera_distance'][0, 0]
err = float((views['transform_matrix'][0, 0] - F_mat).abs().max())
if err > atol:
print(f"[Pixal3D MV] Warning: main view (frame 0) is not the canonical front view "
f"(max deviation {err:.3e}). The result will be posed in that view's frame.")
else:
print(f"[Pixal3D MV] main view == canonical front view (max err {err:.1e})")
return err
def views_to(views: dict, device) -> dict:
"""Move every tensor in a views bundle to `device` (images included)."""
out = dict(views)
out['images'] = {k: v.to(device) for k, v in views['images'].items()}
for k in ('camera_angle_x', 'camera_distance', 'transform_matrix'):
out[k] = views[k].to(device)
return out