Support for Pixal3D MultiView
This commit is contained in:
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user