ego6d_rag / method1.py
Peanuttoad's picture
Add files using upload-large-folder tool
b236bda verified
Raw
History Blame Contribute Delete
8.66 kB
#!/usr/bin/env python3
"""Method 1 — differentiable SEMANTIC DISAMBIGUATION head over frozen FCx top-K modes.
Frozen dim-64 cross-venue forecasting encoder -> per-venue bank -> top-N kNN -> K spatial modes
(greedy NMS). A SHARED tiny head: (a) predicts the query's room/activity from z_q, (b) scores each
mode by agreement between that prediction and the mode's neighbour-semantics (+ mass/spread/sim/rank),
softmaxes -> alpha, reads out position = sum_k alpha_k * center_k. Trained on ALL venues' train
windows (scene-leave-out) with L_pos + L_sel(closest-mode CE) + L_room + L_act. Realizes the top-3
oracle while keeping the robust frozen latent untouched. Eval vs FCx (cur 3.31, oracle min@3 1.60).
"""
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_fc.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'))
net = Forecaster(cin=12, P=1).to(DEV)
net.enc.load_state_dict(torch.load(AN / 'p3_fcx/shared_encoder.pt', map_location=DEV)); net.eval()
@torch.no_grad()
def precompute(V, idx, leave_out):
Zb = embed_all(net, V, V.tr); xyz_b = V.xyz[V.tr]; loc_b = V.loc_oh[V.tr]; act_b = V.act[V.tr]; sc_b = V.scene[V.tr]
out = {k: [] for k in ['Z', 'C', 'RM', 'AC', 'MA', 'SP', 'SI', 'VA', 'GT', 'LOC', 'ACT', 'LAB']}
for i in range(0, len(idx), 256):
b = idx[i:i+256]; zq = embed_all(net, 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 = xyz_b[ti], loc_b[ti], act_b[ti]; B = len(b)
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)
va = torch.zeros(B, K_MODES, dtype=torch.bool, device=DEV)
ar = torch.arange(B, 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; va[:, k] = ms > 0; remain = remain & ~near
for key, val in [('Z', zq), ('C', cc), ('RM', rr), ('AC', aa), ('MA', ma), ('SP', sp), ('SI', si),
('VA', va), ('GT', V.xyz[b]), ('LOC', V.loc[b]), ('ACT', V.act[b]), ('LAB', V.lab[b])]:
out[key].append(val)
return {k: torch.cat(v) for k, v in out.items()}
class Disamb(nn.Module):
def __init__(s, zdim=64):
super().__init__()
s.sem = nn.Sequential(nn.Linear(zdim, 96), nn.GELU(), nn.Linear(96, RL + A))
din = RL + A + RL + A + 4
s.score = nn.Sequential(nn.Linear(din, 64), nn.GELU(), nn.Linear(64, 1))
def forward(s, M):
sem = s.sem(M['Z']); qr = sem[:, :RL].softmax(-1); qa = sem[:, RL:].sigmoid()
B, K = M['MA'].shape
rank = (torch.arange(K, device=M['Z'].device).float() / K)[None, :, None].expand(B, K, 1)
feat = torch.cat([qr[:, None].expand(-1, K, -1), qa[:, None].expand(-1, K, -1), M['RM'], M['AC'],
M['MA'][..., None], M['SP'][..., None], M['SI'][..., None], rank], -1)
sk = s.score(feat).squeeze(-1).masked_fill(~M['VA'], -1e9)
alpha = sk.softmax(-1); pos = (alpha[..., None] * M['C']).sum(1)
return pos, alpha, sk, 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, lam_sel=1.0):
torch.manual_seed(0)
head = Disamb().to(DEV); opt = torch.optim.AdamW(head.parameters(), lr, weight_decay=1e-4)
N = len(TR['Z'])
dmode = (TR['C'] - TR['GT'][:, None]).norm(dim=-1).masked_fill(~TR['VA'], 1e9)
TR['TGT'] = dmode.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):
b = perm[i:i+bs]; Mb = sub(TR, b)
pos, alpha, sk, rl, al = head(Mb)
lpos = F.smooth_l1_loss(pos, Mb['GT'])
lsel = F.cross_entropy(sk, 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 + lam_sel * lsel + lroom + lact
opt.zero_grad(); loss.backward(); opt.step()
for j, v in enumerate([lpos, lsel, lroom, lact]): tot[j] += float(v)
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}", flush=True)
return head
@torch.no_grad()
def evaluate(head, VAL):
head.eval()
soft, hard, oracle = [], [], []
rt = rp_sem = rp_mode = tp_s = fp_s = fn_s = tp_m = fp_m = fn_m = 0
for v, M in VAL.items():
pos, alpha, sk, rl, al = head(M)
hardc = M['C'][torch.arange(len(M['C']), device=DEV), alpha.argmax(-1)]
dmode = (M['C'] - M['GT'][:, None]).norm(dim=-1).masked_fill(~M['VA'], 1e9)
soft.append((pos - M['GT']).norm(dim=-1)); hard.append((hardc - M['GT']).norm(dim=-1))
oracle.append(dmode.min(-1).values)
lab = M['LAB'] & (M['LOC'] >= 0)
rt += int(lab.sum())
rp_sem += int((rl.argmax(-1)[lab] == M['LOC'][lab]).sum())
mode_room = (alpha[..., None] * M['RM']).sum(1) # semantics from selected mode
rp_mode += int((mode_room.argmax(-1)[lab] == M['LOC'][lab]).sum())
gt = M['ACT'] > 0; labm = M['LAB']
for pr, tp_, fp_, fn_ in [(al.sigmoid() > 0.5, 's', 's', 's')]:
pass
ps = (al.sigmoid() > 0.5) & labm[:, None]; gtl = gt & labm[:, None]
tp_s += int((ps & gtl).sum()); fp_s += int((ps & ~gtl).sum()); fn_s += int((~ps & gtl).sum())
mode_act = (alpha[..., None] * M['AC']).sum(1)
pm = (mode_act > 0.5) & labm[:, None]
tp_m += int((pm & gtl).sum()); fp_m += int((pm & ~gtl).sum()); fn_m += int((~pm & gtl).sum())
soft = torch.cat(soft); hard = torch.cat(hard); oracle = torch.cat(oracle)
f1 = lambda tp, fp, fn: 2*tp/(2*tp+fp+fn) if (2*tp+fp+fn) else 0.0
return dict(soft=float(soft.median()), hard=float(hard.median()), oracle=float(oracle.median()),
room_sem=rp_sem/max(rt, 1), room_mode=rp_mode/max(rt, 1),
actF1_sem=f1(tp_s, fp_s, fn_s), actF1_mode=f1(tp_m, fp_m, fn_m))
def main():
print("precomputing modes (train, scene-leave-out) ...", flush=True)
Vs = {v: Venue(v) for v in VENUES}
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['Z']):,} train queries", flush=True)
print("precomputing modes (val) ...", flush=True)
VAL = {v: precompute(Vs[v], Vs[v].va, False) for v in VENUES}
print("training shared disambiguation head ...", flush=True)
head = train_head(TR)
torch.save(head.state_dict(), AN / 'method1_head.pt')
R = evaluate(head, VAL)
print("\n=== METHOD 1 (val, macro over windows) ===")
print(f" localization: soft={R['soft']:.2f}m hard(argmax mode)={R['hard']:.2f}m oracle(min@5)={R['oracle']:.2f}m")
print(f" [FCx baseline: top-1 3.31m, oracle min@3 1.60 / min@5 1.01]")
print(f" semantics: room sem-head={R['room_sem']:.3f} room mode-read={R['room_mode']:.3f} "
f"actF1 sem-head={R['actF1_sem']:.3f} actF1 mode-read={R['actF1_mode']:.3f}")
print(f" [FCx probe: room 0.462 actF1 0.282 | CamFormer: room 0.453 actF1 0.353]")
json.dump(R, open(AN / 'method1_metrics.json', 'w'), indent=1)
if __name__ == '__main__':
main()