#!/usr/bin/env python3 """Final refinement: (A) RELIABILITY-GATED temporal smoothing for localization — apply soft temporal HMM only where retrieval is trustworthy (deployable proxy = mean top-mode retrieval mass), else fall back to the mass-weighted mean; (B) DUAL semantic probe on [FCx (geometry) + CamFormer (semantic)] concatenated embeddings for room+activity. Reports vs all baselines.""" import json from pathlib import Path import numpy as np, pandas as pd, torch, torch.nn as nn, torch.nn.functional as F exec(open('/workspace/ego6d_rag/train_cf.py').read().split('\ndef main():')[0]) # Venue, embed_all, cf_embed_all, Forecaster, CamFormer, RL, A, DEV AN = Path('/workspace/ego6d_rag/analysis') N_TOP, K_MODES, RADIUS, SIGMA = 128, 5, 1.5, 0.75 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() cnet = CamFormer().to(DEV); cnet.load_state_dict(torch.load(AN / 'p3_cf/camformer.pt', map_location=DEV)); cnet.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, sid): t0 = np.full(len(gt), np.nan) for s in np.unique(sid): sub = LW[LW.scene == V.names[int(s)]] if not len(sub): continue cen = sub[['cx', 'cy', 'cz']].values; tt = sub.t0.values for j in np.where(sid == s)[0]: t0[j] = tt[int(((cen - gt[j]) ** 2).sum(1).argmin())] return t0 def fwd_bwd(c, ma, va, sigma): T, K, _ = c.shape; e = np.where(va, ma, 0.) + 1e-9; e = e / e.sum(1, keepdims=True) A = []; a = np.zeros((T, K)); a[0] = e[0] for t in range(1, T): d2 = ((c[t-1][:, None] - c[t][None]) ** 2).sum(-1); At = np.exp(-d2 / (2*sigma*sigma)); At[~va[t-1]] = 0 At = At / (At.sum(1, keepdims=True) + 1e-9); A.append(At); a[t] = e[t] * (a[t-1] @ At); a[t] /= (a[t].sum() + 1e-9) b = np.ones((T, K)) for t in range(T-2, -1, -1): b[t] = (A[t] * (e[t+1]*b[t+1])[None]).sum(1); b[t] /= (b[t].sum() + 1e-9) g = a*b; g = g / (g.sum(1, keepdims=True) + 1e-9); return (g[..., None]*c).sum(1) # ---------- PART A: gated temporal localization ---------- loc = {} for v in VENUES: V = Venue(v); C, MA, VA, GT, SID = val_modes(V); t0 = recover_t0(V, GT, SID) em, es, eo = [], [], [] for s in np.unique(SID): idx = np.where(SID == s)[0]; o = idx[np.argsort(t0[idx])]; c, ma, va, gt = C[o], MA[o], VA[o], GT[o] soft = fwd_bwd(c, ma, va, SIGMA); mw = ma / ma.sum(1, keepdims=True).clip(min=1e-9); pm = (mw[..., None]*c).sum(1) em += list(np.linalg.norm(pm-gt, axis=1)); es += list(np.linalg.norm(soft-gt, axis=1)) dm = np.linalg.norm(c-gt[:, None], axis=-1); dm[~va] = 1e9; eo += list(dm.min(1)) loc[v] = dict(mass=float(np.median(em)), soft=float(np.median(es)), oracle=float(np.median(eo)), proxy=float(np.mean(MA[:, 0]))) # deployable reliability proxy = mean top-mode mass # ---------- PART B: dual semantic probe [FCx | CamFormer] ---------- @torch.no_grad() def dual_emb(net_pair, V, idx): return torch.cat([embed_all(net_pair[0], V, idx), cf_embed_all(net_pair[1], V, idx)], -1) # [N, 64+384] def dual_probe(): Vs = {v: Venue(v) for v in VENUES}; rows = {} for v in VENUES: V = Vs[v]; Ztr = dual_emb((fnet, cnet), V, V.tr).detach(); lp = torch.where(V.lab[V.tr])[0] Zl, lo, ac = Ztr[lp], V.loc[V.tr][lp], V.act[V.tr][lp]; hl = lo >= 0 p = nn.Sequential(nn.Linear(Ztr.shape[1], 128), nn.GELU(), nn.Linear(128, RL+A)).to(DEV) opt = torch.optim.AdamW(p.parameters(), 1e-3, weight_decay=1e-4) with torch.enable_grad(): for _ in range(400): o = p(Zl); ll = F.cross_entropy(o[hl, :RL], lo[hl]) if hl.any() else o.new_zeros(()) la = F.binary_cross_entropy_with_logits(o[:, RL:], ac); opt.zero_grad(); (ll+la).backward(); opt.step() Zv = dual_emb((fnet, cnet), V, V.va); o = p(Zv); lab = V.lab[V.va] & (V.loc[V.va] >= 0) rt = max(int(lab.sum()), 1); room = int((o[lab, :RL].argmax(-1) == V.loc[V.va][lab]).sum())/rt gt = V.act[V.va] > 0; lm = V.lab[V.va]; pr = (o[:, RL:].sigmoid() > 0.5) & lm[:, None]; gl = gt & lm[:, None] tp, fp, fn = int((pr & gl).sum()), int((pr & ~gl).sum()), int((~pr & gl).sum()) rows[v] = dict(room=room, actF1=2*tp/(2*tp+fp+fn) if (2*tp+fp+fn) else 0.0) return rows sem = dual_probe() mac = lambda d, k: float(np.mean([d[v][k] for v in VENUES])) # gate on proxy: try thresholds, pick best macro; also report oracle-gate upper bound prox = np.array([loc[v]['proxy'] for v in VENUES]) def gated(thr): return float(np.mean([loc[v]['soft'] if loc[v]['proxy'] >= thr else loc[v]['mass'] for v in VENUES])) best_thr = min(np.unique(prox), key=gated) oracle_gate = float(np.mean([min(loc[v]['mass'], loc[v]['soft']) for v in VENUES])) print(f"{'venue':6}{'mass':>6}{'soft':>6}{'oracle':>7}{'proxy':>7} gate->{'use':>5}") for v in VENUES: use = 'soft' if loc[v]['proxy'] >= best_thr else 'mass' print(f"{v:6}{loc[v]['mass']:6.2f}{loc[v]['soft']:6.2f}{loc[v]['oracle']:7.2f}{loc[v]['proxy']:7.3f} {use:>9}") print(f"\nLOCALIZATION macro: mass-mean {mac(loc,'mass'):.2f} | soft(ungated) {mac(loc,'soft'):.2f} | " f"PROXY-GATED(thr={best_thr:.3f}) {gated(best_thr):.2f} | oracle-gate {oracle_gate:.2f} | per-window oracle {mac(loc,'oracle'):.2f}") print(f" refs: FCx top-1 3.31 | chance 3.95 | Method1/1b 3.4-3.7 (failed)") print(f"\nSEMANTICS macro (dual FCx|CamFormer probe): room {mac(sem,'room'):.3f} actF1 {mac(sem,'actF1'):.3f}") print(f" refs: FCx room .462/actF1 .282 | CamFormer room .453/actF1 .353") json.dump({'loc': loc, 'sem': sem, 'gate_thr': float(best_thr), 'gated_macro': gated(best_thr), 'oracle_gate': oracle_gate}, open(AN / 'method_final.json', 'w'), indent=1)