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