| |
| """Diagnostic eval of the trained Method-1 head: per-venue macro + baselines + selection accuracy, |
| to see whether the head actually disambiguates modes (vs mass-weighting / vs the oracle).""" |
| import numpy as np, torch |
| exec(open('/workspace/ego6d_rag/method1.py').read().split('\ndef main():')[0]) |
|
|
| head = Disamb().to(DEV); head.load_state_dict(torch.load('/workspace/ego6d_rag/analysis/method1_head.pt')); head.eval() |
| Vs = {v: Venue(v) for v in VENUES} |
|
|
| rows = {} |
| with torch.no_grad(): |
| for v in VENUES: |
| M = precompute(Vs[v], Vs[v].va, False) |
| pos, alpha, sk, 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) |
| closest = dmode.argmin(-1); oracle = dmode.min(-1).values; picked = alpha.argmax(-1) |
| 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_top0 = (C[:, 0] - gt).norm(dim=-1) |
| e_soft = (pos - gt).norm(dim=-1); e_hard = (C[ar, picked] - gt).norm(dim=-1) |
| nval = max(int((M['VA'].sum(-1) > 1).sum()), 1) |
| multi = M['VA'].sum(-1) > 1 |
| rows[v] = dict(mass=float(e_mass.median()), top0=float(e_top0.median()), |
| soft=float(e_soft.median()), hard=float(e_hard.median()), oracle=float(oracle.median()), |
| sel_acc=float((picked == closest).float().mean()), |
| sel_acc_multi=float((picked[multi] == closest[multi]).float().mean()), |
| good_mode=float((oracle < 1.0).float().mean()), |
| alpha_ent=float((-(alpha.clamp(min=1e-9).log() * alpha).sum(-1)).mean()), |
| n_modes=float(M['VA'].sum(-1).float().mean())) |
|
|
| mac = lambda k: float(np.mean([rows[v][k] for v in VENUES])) |
| print(f"{'venue':6} {'mass':>5} {'top0':>5} {'soft':>5} {'hard':>5} {'oracle':>6} {'selA':>5} {'selA_m':>6} {'good<1m':>7} {'aEnt':>5} {'nMod':>4}") |
| for v in VENUES: |
| r = rows[v] |
| print(f"{v:6} {r['mass']:5.2f} {r['top0']:5.2f} {r['soft']:5.2f} {r['hard']:5.2f} {r['oracle']:6.2f} " |
| f"{r['sel_acc']:5.2f} {r['sel_acc_multi']:6.2f} {r['good_mode']:7.2f} {r['alpha_ent']:5.2f} {r['n_modes']:4.1f}") |
| print(f"{'MACRO':6} {mac('mass'):5.2f} {mac('top0'):5.2f} {mac('soft'):5.2f} {mac('hard'):5.2f} {mac('oracle'):6.2f} " |
| f"{mac('sel_acc'):5.2f} {mac('sel_acc_multi'):6.2f} {mac('good_mode'):7.2f} {mac('alpha_ent'):5.2f} {mac('n_modes'):4.1f}") |
| print(f"\nmax selection softmax entropy (uniform over 5) = {np.log(5):.2f}") |
| print("sel_acc = P(argmax-alpha mode == closest-to-GT mode); sel_acc_multi = same on queries with >1 mode") |
|
|