| |
| """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]) |
| 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) |
| 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)) |
| 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) |
| feat = torch.stack([fr, fa, M['CFA'], M['SP'], M['SI'], rank], -1) |
| 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() |
|
|