ego6d_rag / method1_eval.py
Peanuttoad's picture
Add files using upload-large-folder tool
b8c7534 verified
Raw
History Blame Contribute Delete
2.86 kB
#!/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")