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