#!/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")