File size: 5,535 Bytes
b8c7534
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Refinement that CAN capture the headroom: temporal Viterbi over the top-K modes.

Val windows are ~5s apart (dense trajectory) -> the correct mode is the one CONSISTENT IN TIME (you
can't teleport between modes). Per val scene, order windows by time and run Viterbi: emission =
log(retrieval mass of mode k), transition j->k = -||center_j - center_k||^2 / (2 sigma^2). This is the
one disambiguation signal ORTHOGONAL to window content (Methods 1/1b failed because their cue was
entangled with the same ambiguity). Compare vs mass-weighted mean / independent mass-argmax / oracle."""
import json
from pathlib import Path
import numpy as np, pandas as pd, torch
exec(open('/workspace/ego6d_rag/train_fc.py').read().split('\ndef main():')[0])
AN = Path('/workspace/ego6d_rag/analysis')
N_TOP, K_MODES, RADIUS = 128, 5, 1.5
VENUES = sorted(p.stem for p in (AN / 'p2').glob('Loc_*.npz'))
LW = pd.read_parquet(AN / 'p1/labels_windows.parquet')

fnet = Forecaster(cin=12, P=1).to(DEV); fnet.enc.load_state_dict(torch.load(AN / 'p3_fcx/shared_encoder.pt', map_location=DEV)); fnet.eval()


@torch.no_grad()
def val_modes(V):
    Zb = embed_all(fnet, V, V.tr); xyz_b = V.xyz[V.tr]
    C = []; MA = []; VA = []; GT = []; SC = []
    for i in range(0, len(V.va), 256):
        b = V.va[i:i+256]; zq = embed_all(fnet, V, b); sim = zq @ Zb.t()
        tv, ti = sim.topk(min(N_TOP, Zb.shape[0]), -1); tw = tv / tv.sum(-1, keepdim=True)
        txyz = xyz_b[ti]; B = len(b); ar = torch.arange(B, device=DEV); remain = torch.ones_like(tw, dtype=torch.bool)
        cc = torch.zeros(B, K_MODES, 3, device=DEV); ma = torch.zeros(B, K_MODES, device=DEV); va = torch.zeros(B, K_MODES, dtype=torch.bool, device=DEV)
        for k in range(K_MODES):
            seed = tw.masked_fill(~remain, -2).argmax(-1); sx = txyz[ar, seed]
            near = ((txyz - sx[:, None]).norm(dim=-1) < RADIUS) & remain; wn = tw * near.float(); ms = wn.sum(-1)
            cc[:, k] = (wn[..., None] * txyz).sum(1) / ms.clamp(min=1e-9)[:, None]; ma[:, k] = ms; va[:, k] = ms > 0; remain = remain & ~near
        C.append(cc); MA.append(ma); VA.append(va); GT.append(V.xyz[b]); SC.append(V.scene[b])
    return (torch.cat(C).cpu().numpy(), torch.cat(MA).cpu().numpy(), torch.cat(VA).cpu().numpy(),
            torch.cat(GT).cpu().numpy(), torch.cat(SC).cpu().numpy())


def recover_t0(V, gt, scene_id):                       # match each val window to labels_windows by scene+nearest centroid
    t0 = np.full(len(gt), np.nan)
    for sid in np.unique(scene_id):
        name = V.names[int(sid)]; sub = LW[(LW.scene == name)]
        if not len(sub): continue
        cen = sub[['cx', 'cy', 'cz']].values; tt = sub.t0.values; idx = np.where(scene_id == sid)[0]
        for j in idx:
            t0[j] = tt[int(((cen - gt[j]) ** 2).sum(1).argmin())]
    return t0


def viterbi(centers, logmass, valid, sigma):
    T, K, _ = centers.shape
    emis = np.where(valid, logmass, -1e9)
    dp = emis[0].copy(); bp = np.zeros((T, K), int)
    for t in range(1, T):
        d2 = ((centers[t-1][:, None, :] - centers[t][None, :, :]) ** 2).sum(-1)   # [Kj,Kk]
        prev = dp[:, None] + (-d2 / (2 * sigma * sigma))
        prev[~valid[t-1]] = -1e18
        bp[t] = prev.argmax(0); dp = emis[t] + prev.max(0)
    path = np.zeros(T, int); path[-1] = int(dp.argmax())
    for t in range(T-1, 0, -1): path[t-1] = bp[t, path[t]]
    return path


def run(sigma):
    rows = {}
    for v in VENUES:
        V = Venue(v); C, MA, VA, GT, SID = val_modes(V); t0 = recover_t0(V, GT, SID)
        logmass = np.log(MA + 1e-6)
        e_mass, e_argmax, e_vit, e_oracle = [], [], [], []
        for sid in np.unique(SID):
            idx = np.where(SID == sid)[0]; order = idx[np.argsort(t0[idx])]
            c, lm, va, gt = C[order], logmass[order], VA[order], GT[order]
            path = viterbi(c, lm, va, sigma)
            massw = MA[order] / MA[order].sum(1, keepdims=True).clip(min=1e-9)
            pm = (massw[..., None] * c).sum(1)
            am = lm.copy(); am[~va] = -1e9; amax = am.argmax(1)
            ar = np.arange(len(order))
            e_mass += list(np.linalg.norm(pm - gt, axis=1))
            e_argmax += list(np.linalg.norm(c[ar, amax] - gt, axis=1))
            e_vit += list(np.linalg.norm(c[ar, path] - gt, axis=1))
            dmode = np.linalg.norm(c - gt[:, None], axis=-1); dmode[~va] = 1e9
            e_oracle += list(dmode.min(1))
        rows[v] = dict(mass=float(np.median(e_mass)), argmax=float(np.median(e_argmax)),
                       viterbi=float(np.median(e_vit)), oracle=float(np.median(e_oracle)),
                       vit_r1=float(np.mean(np.array(e_vit) < 1.0)), mass_r1=float(np.mean(np.array(e_mass) < 1.0)))
    return rows


for sigma in [1.0, 2.0, 4.0]:
    rows = run(sigma); mac = lambda k: float(np.mean([rows[v][k] for v in VENUES]))
    print(f"\n=== sigma={sigma}m ===  {'venue':6}{'mass':>6}{'argmax':>7}{'viterbi':>8}{'oracle':>7}{'vitR@1':>7}")
    for v in VENUES:
        r = rows[v]; print(f"  {v:6} {r['mass']:6.2f}{r['argmax']:7.2f}{r['viterbi']:8.2f}{r['oracle']:7.2f}{r['vit_r1']:7.2f}")
    print(f"  {'MACRO':6} {mac('mass'):6.2f}{mac('argmax'):7.2f}{mac('viterbi'):8.2f}{mac('oracle'):7.2f}{mac('vit_r1'):7.2f}  (mass R@1 {mac('mass_r1'):.2f})")
    json.dump({'sigma': sigma, 'macro': {k: mac(k) for k in rows[VENUES[0]]}, 'per_venue': rows}, open(AN / f'temporal_sig{sigma}.json', 'w'), indent=1)
print("\n[refs] FCx top-1 3.31 | mass-mean 3.27 | Method1/1b ~3.4-3.7 (failed) | oracle ~1.0")