#!/usr/bin/env python3 """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) # semantics from selected mode 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()