#!/usr/bin/env python3 """Method 1b (refinement) — mass-prior gating + CamFormer dual disambiguation signal. Forecasting encoder retrieves modes (robust geometry). For each mode, compute the mean CamFormer (semantic) embedding of its neighbours; score modes by agreement with the query's CamFormer embedding (an ORTHOGONAL encoder, unlike the circular forecasting-semantics of Method 1). Readout is mass-gated: alpha = softmax(beta*log(mass) + score), so the head can only IMPROVE on the mass-weighted mean (3.27). Semantic head sits on the CamFormer embedding (stronger semantics). Eval per-venue macro vs baselines. """ 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_cf.py').read().split('\ndef main():')[0]) # Venue, embed_all, cf_embed_all, Forecaster, CamFormer, IN_LEN, RL, A, DEV 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')) 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 precompute(V, idx, leave_out): Zb = embed_all(fnet, V, V.tr); CFb = cf_embed_all(cnet, V, V.tr) # forecasting + CamFormer banks xyz_b, loc_b, act_b, sc_b = V.xyz[V.tr], V.loc_oh[V.tr], V.act[V.tr], V.scene[V.tr] o = {k: [] for k in ['CFQ', 'C', 'RM', 'AC', 'MA', 'SP', 'SI', 'CFA', 'VA', 'GT', 'LOC', 'ACT', 'LAB']} for i in range(0, len(idx), 256): b = idx[i:i+256]; zq = embed_all(fnet, V, b); cfq = cf_embed_all(cnet, 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, tcf = xyz_b[ti], loc_b[ti], act_b[ti], CFb[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); 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) cfa = 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] 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 cfsig = F.normalize((wn[..., None] * tcf).sum(1) / ms.clamp(min=1e-9)[:, None], dim=-1) cfa[:, k] = (cfq * cfsig).sum(-1); va[:, k] = ms > 0; remain = remain & ~near for key, val in [('CFQ', cfq), ('C', cc), ('RM', rr), ('AC', aa), ('MA', ma), ('SP', sp), ('SI', si), ('CFA', cfa), ('VA', va), ('GT', V.xyz[b]), ('LOC', V.loc[b]), ('ACT', V.act[b]), ('LAB', V.lab[b])]: o[key].append(val) return {k: torch.cat(v) for k, v in o.items()} class DisambB(nn.Module): def __init__(s): super().__init__() s.sem = nn.Sequential(nn.Linear(384, 128), nn.GELU(), nn.Linear(128, RL + A)) # semantics on CamFormer emb s.score = nn.Sequential(nn.Linear(6, 32), nn.GELU(), nn.Linear(32, 1)) s.beta = nn.Parameter(torch.tensor(1.0)) def forward(s, M): sem = s.sem(M['CFQ']); qr = sem[:, :RL].softmax(-1); qa = sem[:, RL:].sigmoid() B, K = M['MA'].shape; rank = (torch.arange(K, device=M['CFQ'].device).float() / K)[None, :].expand(B, K) fr = (qr[:, None] * M['RM']).sum(-1); fa = (qa[:, None] * M['AC']).sum(-1) # forecasting-neighbour agreement feat = torch.stack([fr, fa, M['CFA'], M['SP'], M['SI'], rank], -1) # [B,K,6] score = s.score(feat).squeeze(-1) logits = (s.beta * torch.log(M['MA'] + 1e-6) + score).masked_fill(~M['VA'], -1e9) alpha = logits.softmax(-1); pos = (alpha[..., None] * M['C']).sum(1) return pos, alpha, logits, 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): torch.manual_seed(0); head = DisambB().to(DEV); opt = torch.optim.AdamW(head.parameters(), lr, weight_decay=1e-4) N = len(TR['CFQ']); TR['TGT'] = (TR['C'] - TR['GT'][:, None]).norm(dim=-1).masked_fill(~TR['VA'], 1e9).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): Mb = sub(TR, perm[i:i+bs]); pos, alpha, logits, rl, al = head(Mb) lpos = F.smooth_l1_loss(pos, Mb['GT']); lsel = F.cross_entropy(logits, 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 + lsel + lroom + lact; opt.zero_grad(); loss.backward(); opt.step() for j, vv in enumerate([lpos, lsel, lroom, lact]): tot[j] += float(vv) 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} beta={float(head.beta):.2f}", flush=True) return head @torch.no_grad() def evaluate(head, Vs): head.eval(); rows = {} for v in VENUES: M = precompute(Vs[v], Vs[v].va, False); pos, alpha, logits, rl, al = head(M) C, gt = M['C'], M['GT']; ar = torch.arange(len(C), device=DEV) dmode = (C - gt[:, None]).norm(dim=-1).masked_fill(~M['VA'], 1e9) massw = M['MA'] / M['MA'].sum(-1, keepdim=True).clamp(min=1e-9) e_mass = ((massw[..., None] * C).sum(1) - gt).norm(dim=-1) e_soft = (pos - gt).norm(dim=-1); e_hard = (C[ar, alpha.argmax(-1)] - gt).norm(dim=-1) lab = M['LAB'] & (M['LOC'] >= 0); rt = max(int(lab.sum()), 1) room = int((rl.argmax(-1)[lab] == M['LOC'][lab]).sum()) / rt gt_a = M['ACT'] > 0; lm = M['LAB']; pr = (al.sigmoid() > 0.5) & lm[:, None]; gl = gt_a & lm[:, None] tp, fp, fn = int((pr & gl).sum()), int((pr & ~gl).sum()), int((~pr & gl).sum()) rows[v] = dict(mass=float(e_mass.median()), soft=float(e_soft.median()), hard=float(e_hard.median()), oracle=float(dmode.min(-1).values.median()), sel_acc=float((alpha.argmax(-1) == dmode.argmin(-1)).float().mean()), room=room, actF1=2*tp/(2*tp+fp+fn) if (2*tp+fp+fn) else 0.0) return rows def main(): Vs = {v: Venue(v) for v in VENUES} print("precompute train (CamFormer+forecasting modes, scene-leave-out) ...", flush=True) 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['CFQ']):,}; training mass-gated dual head ...", flush=True) head = train_head(TR); torch.save(head.state_dict(), AN / 'method1b_head.pt') rows = evaluate(head, Vs); mac = lambda k: float(np.mean([rows[v][k] for v in VENUES])) print(f"\n{'venue':6} {'mass':>5} {'soft':>5} {'hard':>5} {'oracle':>6} {'selA':>5} {'room':>5} {'actF1':>6}") for v in VENUES: r = rows[v]; print(f"{v:6} {r['mass']:5.2f} {r['soft']:5.2f} {r['hard']:5.2f} {r['oracle']:6.2f} {r['sel_acc']:5.2f} {r['room']:5.3f} {r['actF1']:6.3f}") print(f"{'MACRO':6} {mac('mass'):5.2f} {mac('soft'):5.2f} {mac('hard'):5.2f} {mac('oracle'):6.2f} {mac('sel_acc'):5.2f} {mac('room'):5.3f} {mac('actF1'):6.3f}") print(f"\n[refs] FCx top-1 3.31 | Method1 soft 3.43/hard 3.77 selA 0.37 | mass-baseline 3.27 | oracle ~1.0") print(f"[sem refs] FCx probe room .462/actF1 .282 | CamFormer room .453/actF1 .353") json.dump({'macro': {k: mac(k) for k in rows[VENUES[0]]}, 'per_venue': rows}, open(AN / 'method1b_metrics.json', 'w'), indent=1) if __name__ == '__main__': main()