File size: 3,970 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
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
"""Shape metrics against ShapeNetCore.v2.PC15k GT, following MinD-3D's tools/get_cd.py protocol
(2048 points, pc_norm), plus a shared orientation search so methods with different canonical frames
are compared fairly."""
import os
import glob
import json
import argparse
import collections

import numpy as np
import trimesh
from scipy.spatial import cKDTree
from scipy.optimize import linear_sum_assignment

N_POINTS = 2048
FSCORE_TAU = 0.02


def pc_norm(pc):
    pc = pc - pc.mean(0)
    return pc / (2 * np.max(np.linalg.norm(pc, axis=1)))


def rot_y(deg):
    t = np.deg2rad(deg)
    c, s = np.cos(t), np.sin(t)
    return np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])


UP_MAPS = {
    "y_up": np.eye(3),
    "z_up": np.array([[1, 0, 0], [0, 0, 1], [0, -1, 0]], dtype=np.float64),
}
CANDIDATES = [(up, az, rot_y(az) @ m) for up, m in UP_MAPS.items() for az in range(0, 360, 45)]


def chamfer(a, b):
    d_ab = cKDTree(b).query(a)[0]
    d_ba = cKDTree(a).query(b)[0]
    return np.mean(d_ab ** 2) + np.mean(d_ba ** 2), d_ab, d_ba


def fscore(d_ab, d_ba, tau):
    p = np.mean(d_ab < tau)
    r = np.mean(d_ba < tau)
    return 0.0 if p + r == 0 else 2 * p * r / (p + r)


def emd(a, b):
    cost = np.linalg.norm(a[:, None] - b[None], axis=-1)
    r, c = linear_sum_assignment(cost)
    return cost[r, c].mean()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--gen_dir", required=True)
    parser.add_argument("--gt_root", default="/home/hubin/data/ShapeNetCore.v2.PC15k/ShapeNetCore.v2.PC15k")
    parser.add_argument("--test_list", default="/home/hubin/data/fMRI-Shape/annotations/core_test_list.txt")
    parser.add_argument("--out", required=True)
    parser.add_argument("--seed", type=int, default=0)
    args = parser.parse_args()

    gt_paths = {os.path.basename(p)[:-4]: p for p in glob.glob(f"{args.gt_root}/*/_/*/*.npy")}
    ids = [l.strip() for l in open(args.test_list) if l.strip()]
    rows = []
    for obj in ids:
        cat, uid = obj.split("/")
        gen = [p for p in glob.glob(os.path.join(args.gen_dir, f"{cat}_{uid}.*")) if p.endswith((".ply", ".obj"))]
        if uid not in gt_paths or not gen:
            continue
        rng = np.random.default_rng(args.seed)
        gt = np.load(gt_paths[uid])
        gt = pc_norm(gt[rng.choice(len(gt), N_POINTS, replace=False)].astype(np.float64))
        mesh = trimesh.load(gen[0], force="mesh")
        pts = trimesh.sample.sample_surface(mesh, N_POINTS, seed=args.seed)[0].astype(np.float64)
        pts = pc_norm(pts)

        best = None
        for up, az, R in CANDIDATES:
            cd, d_ab, d_ba = chamfer(pts @ R.T, gt)
            if best is None or cd < best[0]:
                best = (cd, up, az, R, d_ab, d_ba)
        cd, up, az, R, d_ab, d_ba = best
        rows.append({
            "id": obj, "cat": cat, "cd": cd, "emd": emd(pts @ R.T, gt),
            "fscore": fscore(d_ab, d_ba, FSCORE_TAU), "up": up, "azimuth": az,
        })
        print(f"{obj} CD={cd * 1e3:.3f}e-3 EMD={rows[-1]['emd']:.4f} F={rows[-1]['fscore']:.3f} ({up},{az})", flush=True)

    by_cat = collections.defaultdict(list)
    for r in rows:
        by_cat[r["cat"]].append(r)
    summary = {
        "n": len(rows),
        "cd_x1e3": float(np.mean([r["cd"] for r in rows]) * 1e3),
        "emd": float(np.mean([r["emd"] for r in rows])),
        "fscore@0.02": float(np.mean([r["fscore"] for r in rows])),
        "per_category": {c: {"n": len(v), "cd_x1e3": float(np.mean([r["cd"] for r in v]) * 1e3),
                             "emd": float(np.mean([r["emd"] for r in v]))} for c, v in sorted(by_cat.items())},
        "chosen_orientation": collections.Counter(f"{r['up']}/{r['azimuth']}" for r in rows).most_common(),
    }
    with open(args.out, "w") as f:
        json.dump({"summary": summary, "rows": rows}, f, indent=1)
    print(json.dumps({k: v for k, v in summary.items() if k != "per_category"}, indent=1))


if __name__ == "__main__":
    main()