mind3d-trellis2 / code /upper_bound_trellis2.py
jamie33's picture
MinD-3D + TRELLIS.2 sub-01 experiments: code, metrics, logs, report
87e9895 verified
Raw History Blame Contribute Delete
3.67 kB
"""Oracle upper bound: feed GT stimulus frames to TRELLIS.2 and export shape meshes."""
import os
os.environ.setdefault("OPENCV_IO_ENABLE_OPENEXR", "1")
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import sys
import json
import time
import argparse
import numpy as np
import torch
import trimesh
import imageio.v3 as iio
from PIL import Image
sys.path.insert(0, "/home/hubin/trellis_work/TRELLIS.2")
from trellis2.pipelines import Trellis2ImageTo3DPipeline
def load_frame(video_path, frame_idx):
frames = iio.imread(video_path, index=frame_idx)
return Image.fromarray(frames)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--weights", default="/home/hubin/trellis_work/weights/TRELLIS.2-4B-merged")
parser.add_argument("--config_file", default="pipeline_local.json")
parser.add_argument("--test_list", default="/home/hubin/data/fMRI-Shape/annotations/core_test_list.txt")
parser.add_argument("--stimuli_dir", default="/home/hubin/data/fMRI-Shape/stimuli_test/stimuli")
parser.add_argument("--frame_idx", type=int, default=24)
parser.add_argument("--pipeline_type", default="512", choices=["512", "1024"])
parser.add_argument("--out_dir", default="/home/hubin/trellis_work/outputs/upper_bound_512_f24")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--limit", type=int, default=0)
parser.add_argument("--shard", type=int, default=0)
parser.add_argument("--num_shards", type=int, default=1)
args = parser.parse_args()
ids = [l.strip() for l in open(args.test_list) if l.strip()]
ids = ids[args.shard::args.num_shards]
if args.limit:
ids = ids[:args.limit]
os.makedirs(os.path.join(args.out_dir, "meshes"), exist_ok=True)
os.makedirs(os.path.join(args.out_dir, "inputs"), exist_ok=True)
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(args.weights, config_file=args.config_file)
pipeline.low_vram = False
pipeline.cuda()
res = int(args.pipeline_type)
ss_res = {512: 32, 1024: 64}[res]
flow_key = f"shape_slat_flow_model_{res}"
stats = []
for i, obj in enumerate(ids):
name = obj.replace("/", "_")
mesh_path = os.path.join(args.out_dir, "meshes", f"{name}.ply")
if os.path.exists(mesh_path):
continue
t0 = time.time()
image = load_frame(os.path.join(args.stimuli_dir, f"{obj}.mp4"), args.frame_idx)
image = pipeline.preprocess_image(image)
image.save(os.path.join(args.out_dir, "inputs", f"{name}.png"))
torch.manual_seed(args.seed)
with torch.no_grad():
cond = pipeline.get_cond([image], res)
coords = pipeline.sample_sparse_structure(cond, ss_res, 1)
shape_slat = pipeline.sample_shape_slat(cond, pipeline.models[flow_key], coords)
meshes, _ = pipeline.decode_shape_slat(shape_slat, res)
mesh = meshes[0]
mesh.fill_holes()
trimesh.Trimesh(
vertices=mesh.vertices.detach().cpu().numpy(),
faces=mesh.faces.detach().cpu().numpy(),
process=False,
).export(mesh_path)
dt = time.time() - t0
stats.append({"id": obj, "time": dt, "n_voxels": int(coords.shape[0]),
"n_verts": int(mesh.vertices.shape[0])})
print(f"[{i + 1}/{len(ids)}] {obj} {dt:.1f}s voxels={coords.shape[0]} verts={mesh.vertices.shape[0]}", flush=True)
torch.cuda.empty_cache()
with open(os.path.join(args.out_dir, f"stats_shard{args.shard}.json"), "w") as f:
json.dump(stats, f, indent=1)
if __name__ == "__main__":
main()