ego6d_rag / method_final.py
Peanuttoad's picture
Add files using upload-large-folder tool
b8c7534 verified
Raw
History Blame Contribute Delete
7.14 kB
#!/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)