#!/usr/bin/env python3 """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) # no-head mass-weighted mean e_top0 = (C[:, 0] - gt).norm(dim=-1) # highest-mass single mode 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) # queries with >1 mode 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")