"""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()