| |
| """Method 1 — differentiable SEMANTIC DISAMBIGUATION head over frozen FCx top-K modes. |
| |
| Frozen dim-64 cross-venue forecasting encoder -> per-venue bank -> top-N kNN -> K spatial modes |
| (greedy NMS). A SHARED tiny head: (a) predicts the query's room/activity from z_q, (b) scores each |
| mode by agreement between that prediction and the mode's neighbour-semantics (+ mass/spread/sim/rank), |
| softmaxes -> alpha, reads out position = sum_k alpha_k * center_k. Trained on ALL venues' train |
| windows (scene-leave-out) with L_pos + L_sel(closest-mode CE) + L_room + L_act. Realizes the top-3 |
| oracle while keeping the robust frozen latent untouched. Eval vs FCx (cur 3.31, oracle min@3 1.60). |
| """ |
| import json |
| from pathlib import Path |
| import numpy as np, torch, torch.nn as nn, torch.nn.functional as F |
| 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')) |
|
|
| net = Forecaster(cin=12, P=1).to(DEV) |
| net.enc.load_state_dict(torch.load(AN / 'p3_fcx/shared_encoder.pt', map_location=DEV)); net.eval() |
|
|
|
|
| @torch.no_grad() |
| def precompute(V, idx, leave_out): |
| Zb = embed_all(net, V, V.tr); xyz_b = V.xyz[V.tr]; loc_b = V.loc_oh[V.tr]; act_b = V.act[V.tr]; sc_b = V.scene[V.tr] |
| out = {k: [] for k in ['Z', 'C', 'RM', 'AC', 'MA', 'SP', 'SI', 'VA', 'GT', 'LOC', 'ACT', 'LAB']} |
| for i in range(0, len(idx), 256): |
| b = idx[i:i+256]; zq = embed_all(net, V, b); sim = zq @ Zb.t() |
| if leave_out: sim = sim.masked_fill(sc_b[None] == V.scene[b][:, None], -1e9) |
| tv, ti = sim.topk(min(N_TOP, Zb.shape[0]), -1); tw = tv / tv.sum(-1, keepdim=True) |
| txyz, tloc, tact = xyz_b[ti], loc_b[ti], act_b[ti]; B = len(b) |
| remain = torch.ones_like(tw, dtype=torch.bool) |
| cc = torch.zeros(B, K_MODES, 3, device=DEV); rr = torch.zeros(B, K_MODES, RL, device=DEV); aa = torch.zeros(B, K_MODES, A, device=DEV) |
| ma = torch.zeros(B, K_MODES, device=DEV); sp = torch.zeros(B, K_MODES, device=DEV); si = torch.zeros(B, K_MODES, device=DEV) |
| va = torch.zeros(B, K_MODES, dtype=torch.bool, device=DEV) |
| ar = torch.arange(B, device=DEV) |
| for k in range(K_MODES): |
| seed = tw.masked_fill(~remain, -2).argmax(-1); sx = txyz[ar, seed] |
| d = (txyz - sx[:, None]).norm(dim=-1); near = (d < RADIUS) & remain |
| wn = tw * near.float(); ms = wn.sum(-1) |
| cen = (wn[..., None] * txyz).sum(1) / ms.clamp(min=1e-9)[:, None] |
| cc[:, k] = cen; rr[:, k] = (wn[..., None] * tloc).sum(1) / ms.clamp(min=1e-9)[:, None] |
| aa[:, k] = (wn[..., None] * tact).sum(1) / ms.clamp(min=1e-9)[:, None] |
| ma[:, k] = ms; sp[:, k] = ((wn * ((txyz - cen[:, None]).norm(dim=-1) ** 2)).sum(-1) / ms.clamp(min=1e-9)).sqrt() |
| si[:, k] = (tv * near.float()).max(-1).values; va[:, k] = ms > 0; remain = remain & ~near |
| for key, val in [('Z', zq), ('C', cc), ('RM', rr), ('AC', aa), ('MA', ma), ('SP', sp), ('SI', si), |
| ('VA', va), ('GT', V.xyz[b]), ('LOC', V.loc[b]), ('ACT', V.act[b]), ('LAB', V.lab[b])]: |
| out[key].append(val) |
| return {k: torch.cat(v) for k, v in out.items()} |
|
|
|
|
| class Disamb(nn.Module): |
| def __init__(s, zdim=64): |
| super().__init__() |
| s.sem = nn.Sequential(nn.Linear(zdim, 96), nn.GELU(), nn.Linear(96, RL + A)) |
| din = RL + A + RL + A + 4 |
| s.score = nn.Sequential(nn.Linear(din, 64), nn.GELU(), nn.Linear(64, 1)) |
| def forward(s, M): |
| sem = s.sem(M['Z']); qr = sem[:, :RL].softmax(-1); qa = sem[:, RL:].sigmoid() |
| B, K = M['MA'].shape |
| rank = (torch.arange(K, device=M['Z'].device).float() / K)[None, :, None].expand(B, K, 1) |
| feat = torch.cat([qr[:, None].expand(-1, K, -1), qa[:, None].expand(-1, K, -1), M['RM'], M['AC'], |
| M['MA'][..., None], M['SP'][..., None], M['SI'][..., None], rank], -1) |
| sk = s.score(feat).squeeze(-1).masked_fill(~M['VA'], -1e9) |
| alpha = sk.softmax(-1); pos = (alpha[..., None] * M['C']).sum(1) |
| return pos, alpha, sk, sem[:, :RL], sem[:, RL:] |
|
|
|
|
| def sub(M, b): return {k: v[b] for k, v in M.items()} |
|
|
|
|
| def train_head(TR, epochs=120, bs=1024, lr=1e-3, lam_sel=1.0): |
| torch.manual_seed(0) |
| head = Disamb().to(DEV); opt = torch.optim.AdamW(head.parameters(), lr, weight_decay=1e-4) |
| N = len(TR['Z']) |
| dmode = (TR['C'] - TR['GT'][:, None]).norm(dim=-1).masked_fill(~TR['VA'], 1e9) |
| TR['TGT'] = dmode.argmin(-1) |
| for ep in range(epochs): |
| head.train(); perm = torch.randperm(N, device=DEV); tot = [0, 0, 0, 0]; nb = 0 |
| for i in range(0, N - bs + 1, bs): |
| b = perm[i:i+bs]; Mb = sub(TR, b) |
| pos, alpha, sk, rl, al = head(Mb) |
| lpos = F.smooth_l1_loss(pos, Mb['GT']) |
| lsel = F.cross_entropy(sk, Mb['TGT']) |
| lab = Mb['LAB']; hl = lab & (Mb['LOC'] >= 0) |
| lroom = F.cross_entropy(rl[hl], Mb['LOC'][hl]) if hl.any() else rl.new_zeros(()) |
| lact = F.binary_cross_entropy_with_logits(al[lab], Mb['ACT'][lab]) if lab.any() else al.new_zeros(()) |
| loss = lpos + lam_sel * lsel + lroom + lact |
| opt.zero_grad(); loss.backward(); opt.step() |
| for j, v in enumerate([lpos, lsel, lroom, lact]): tot[j] += float(v) |
| nb += 1 |
| if ep % 30 == 0 or ep == epochs - 1: |
| print(f" ep{ep:3d} Lpos={tot[0]/nb:.3f} Lsel={tot[1]/nb:.3f} Lroom={tot[2]/nb:.3f} Lact={tot[3]/nb:.3f}", flush=True) |
| return head |
|
|
|
|
| @torch.no_grad() |
| def evaluate(head, VAL): |
| head.eval() |
| soft, hard, oracle = [], [], [] |
| rt = rp_sem = rp_mode = tp_s = fp_s = fn_s = tp_m = fp_m = fn_m = 0 |
| for v, M in VAL.items(): |
| pos, alpha, sk, rl, al = head(M) |
| hardc = M['C'][torch.arange(len(M['C']), device=DEV), alpha.argmax(-1)] |
| dmode = (M['C'] - M['GT'][:, None]).norm(dim=-1).masked_fill(~M['VA'], 1e9) |
| soft.append((pos - M['GT']).norm(dim=-1)); hard.append((hardc - M['GT']).norm(dim=-1)) |
| oracle.append(dmode.min(-1).values) |
| lab = M['LAB'] & (M['LOC'] >= 0) |
| rt += int(lab.sum()) |
| rp_sem += int((rl.argmax(-1)[lab] == M['LOC'][lab]).sum()) |
| mode_room = (alpha[..., None] * M['RM']).sum(1) |
| rp_mode += int((mode_room.argmax(-1)[lab] == M['LOC'][lab]).sum()) |
| gt = M['ACT'] > 0; labm = M['LAB'] |
| for pr, tp_, fp_, fn_ in [(al.sigmoid() > 0.5, 's', 's', 's')]: |
| pass |
| ps = (al.sigmoid() > 0.5) & labm[:, None]; gtl = gt & labm[:, None] |
| tp_s += int((ps & gtl).sum()); fp_s += int((ps & ~gtl).sum()); fn_s += int((~ps & gtl).sum()) |
| mode_act = (alpha[..., None] * M['AC']).sum(1) |
| pm = (mode_act > 0.5) & labm[:, None] |
| tp_m += int((pm & gtl).sum()); fp_m += int((pm & ~gtl).sum()); fn_m += int((~pm & gtl).sum()) |
| soft = torch.cat(soft); hard = torch.cat(hard); oracle = torch.cat(oracle) |
| f1 = lambda tp, fp, fn: 2*tp/(2*tp+fp+fn) if (2*tp+fp+fn) else 0.0 |
| return dict(soft=float(soft.median()), hard=float(hard.median()), oracle=float(oracle.median()), |
| room_sem=rp_sem/max(rt, 1), room_mode=rp_mode/max(rt, 1), |
| actF1_sem=f1(tp_s, fp_s, fn_s), actF1_mode=f1(tp_m, fp_m, fn_m)) |
|
|
|
|
| def main(): |
| print("precomputing modes (train, scene-leave-out) ...", flush=True) |
| Vs = {v: Venue(v) for v in VENUES} |
| TRs = [precompute(Vs[v], Vs[v].tr, True) for v in VENUES] |
| TR = {k: torch.cat([t[k] for t in TRs]) for k in TRs[0]} |
| print(f" pooled {len(TR['Z']):,} train queries", flush=True) |
| print("precomputing modes (val) ...", flush=True) |
| VAL = {v: precompute(Vs[v], Vs[v].va, False) for v in VENUES} |
| print("training shared disambiguation head ...", flush=True) |
| head = train_head(TR) |
| torch.save(head.state_dict(), AN / 'method1_head.pt') |
| R = evaluate(head, VAL) |
| print("\n=== METHOD 1 (val, macro over windows) ===") |
| print(f" localization: soft={R['soft']:.2f}m hard(argmax mode)={R['hard']:.2f}m oracle(min@5)={R['oracle']:.2f}m") |
| print(f" [FCx baseline: top-1 3.31m, oracle min@3 1.60 / min@5 1.01]") |
| print(f" semantics: room sem-head={R['room_sem']:.3f} room mode-read={R['room_mode']:.3f} " |
| f"actF1 sem-head={R['actF1_sem']:.3f} actF1 mode-read={R['actF1_mode']:.3f}") |
| print(f" [FCx probe: room 0.462 actF1 0.282 | CamFormer: room 0.453 actF1 0.353]") |
| json.dump(R, open(AN / 'method1_metrics.json', 'w'), indent=1) |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|