File size: 3,668 Bytes
87e9895
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
"""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()