| import argparse |
| import os |
| import torch |
|
|
| os.environ.setdefault("SPCONV_ALGO", "native") |
| os.environ.setdefault("ATTN_BACKEND", "flash_attn") |
|
|
| from PIL import Image |
|
|
| from trellis.pipelines import ( |
| DVDImageToVoxelPipeline, |
| TrellisImageTo3DPipeline, |
| as_voxel_output, |
| export_cubified_voxels, |
| run_image_stage2_from_dvd_voxels, |
| ) |
| from trellis.utils import postprocessing_utils |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description="DVD voxel editing followed by TRELLIS stage 2.") |
| parser.add_argument("--target-image", required=True, help="Target image condition for voxel editing.") |
| parser.add_argument( |
| "--voxel-coords", |
| required=True, |
| help="Existing voxel coords in DVD convention. Supports .npy, .pt, and .pth.", |
| ) |
| parser.add_argument("--dvd-config", default="ckpts/dvd_img_BSP_ft.json", help="BSP fine-tuned DVD model config JSON.") |
| parser.add_argument( |
| "--dvd-checkpoint", |
| default="ckpts/dvd_img_BSP_ft.safetensors", |
| help="BSP fine-tuned DVD safetensors checkpoint.", |
| ) |
| parser.add_argument("--output-dir", default="example_results", help="Directory for generated assets.") |
| parser.add_argument("--resolution", type=int, default=64, help="Voxel grid resolution for loaded coords.") |
| parser.add_argument("--seed", type=int, default=0) |
| parser.add_argument("--device", default="cuda") |
| return parser.parse_args() |
|
|
|
|
| def load_voxel_coords(path): |
| import numpy as np |
| import torch |
|
|
| ext = os.path.splitext(path)[1].lower() |
| if ext == ".npy": |
| data = np.load(path) |
| elif ext in {".pt", ".pth"}: |
| data = torch.load(path, map_location="cpu") |
| if isinstance(data, dict): |
| for key in ("coords", "voxels", "samples"): |
| if key in data: |
| data = data[key] |
| break |
| else: |
| raise ValueError(f"Unsupported voxel coord file extension: {ext}") |
| return torch.as_tensor(data) |
|
|
|
|
| def main(): |
| args = parse_args() |
| os.makedirs(args.output_dir, exist_ok=True) |
|
|
| target_image = Image.open(args.target_image) |
| target_name = os.path.splitext(os.path.basename(args.target_image))[0] |
| voxel_name = os.path.splitext(os.path.basename(args.voxel_coords))[0] |
| name = f"{target_name}_{voxel_name}" |
|
|
| dvd = DVDImageToVoxelPipeline.from_files( |
| args.dvd_config, |
| args.dvd_checkpoint, |
| resolution=args.resolution, |
| device=args.device, |
| ) |
| voxels = as_voxel_output(load_voxel_coords(args.voxel_coords), resolution=args.resolution) |
|
|
| |
| |
|
|
| keep_mask = torch.ones_like(voxels.samples) |
| keep_mask[...,32:,:] = 0 |
| keep_mask = keep_mask.bool() |
|
|
| edited_voxels = dvd.edit_voxels(target_image, voxels, keep_mask=keep_mask, seed=args.seed) |
| export_cubified_voxels(edited_voxels, os.path.join(args.output_dir, f"{name}_edited_voxels.glb")) |
|
|
| trellis = TrellisImageTo3DPipeline.from_pretrained("microsoft/TRELLIS-image-large") |
| trellis.to(args.device) |
| outputs = run_image_stage2_from_dvd_voxels(trellis, target_image, edited_voxels, seed=args.seed) |
|
|
| glb = postprocessing_utils.to_glb(outputs["gaussian"][0], outputs["mesh"][0]) |
| glb.export(os.path.join(args.output_dir, f"{name}_edited.glb")) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|