File size: 7,137 Bytes
b8c7534
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
#!/usr/bin/env python3
"""Final refinement: (A) RELIABILITY-GATED temporal smoothing for localization — apply soft temporal
HMM only where retrieval is trustworthy (deployable proxy = mean top-mode retrieval mass), else fall
back to the mass-weighted mean; (B) DUAL semantic probe on [FCx (geometry) + CamFormer (semantic)]
concatenated embeddings for room+activity. Reports vs all baselines."""
import json
from pathlib import Path
import numpy as np, pandas as pd, torch, torch.nn as nn, torch.nn.functional as F
exec(open('/workspace/ego6d_rag/train_cf.py').read().split('\ndef main():')[0])   # Venue, embed_all, cf_embed_all, Forecaster, CamFormer, RL, A, DEV
AN = Path('/workspace/ego6d_rag/analysis')
N_TOP, K_MODES, RADIUS, SIGMA = 128, 5, 1.5, 0.75
VENUES = sorted(p.stem for p in (AN / 'p2').glob('Loc_*.npz'))
LW = pd.read_parquet(AN / 'p1/labels_windows.parquet')
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 val_modes(V):
    Zb = embed_all(fnet, V, V.tr); xyz_b = V.xyz[V.tr]; C = []; MA = []; VA = []; GT = []; SC = []
    for i in range(0, len(V.va), 256):
        b = V.va[i:i+256]; zq = embed_all(fnet, V, b); sim = zq @ Zb.t()
        tv, ti = sim.topk(min(N_TOP, Zb.shape[0]), -1); tw = tv / tv.sum(-1, keepdim=True)
        txyz = xyz_b[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); ma = 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]
            near = ((txyz - sx[:, None]).norm(dim=-1) < RADIUS) & remain; wn = tw * near.float(); ms = wn.sum(-1)
            cc[:, k] = (wn[..., None] * txyz).sum(1) / ms.clamp(min=1e-9)[:, None]; ma[:, k] = ms; va[:, k] = ms > 0; remain = remain & ~near
        C.append(cc); MA.append(ma); VA.append(va); GT.append(V.xyz[b]); SC.append(V.scene[b])
    return (torch.cat(C).cpu().numpy(), torch.cat(MA).cpu().numpy(), torch.cat(VA).cpu().numpy(),
            torch.cat(GT).cpu().numpy(), torch.cat(SC).cpu().numpy())


def recover_t0(V, gt, sid):
    t0 = np.full(len(gt), np.nan)
    for s in np.unique(sid):
        sub = LW[LW.scene == V.names[int(s)]]
        if not len(sub): continue
        cen = sub[['cx', 'cy', 'cz']].values; tt = sub.t0.values
        for j in np.where(sid == s)[0]: t0[j] = tt[int(((cen - gt[j]) ** 2).sum(1).argmin())]
    return t0


def fwd_bwd(c, ma, va, sigma):
    T, K, _ = c.shape; e = np.where(va, ma, 0.) + 1e-9; e = e / e.sum(1, keepdims=True)
    A = []; a = np.zeros((T, K)); a[0] = e[0]
    for t in range(1, T):
        d2 = ((c[t-1][:, None] - c[t][None]) ** 2).sum(-1); At = np.exp(-d2 / (2*sigma*sigma)); At[~va[t-1]] = 0
        At = At / (At.sum(1, keepdims=True) + 1e-9); A.append(At); a[t] = e[t] * (a[t-1] @ At); a[t] /= (a[t].sum() + 1e-9)
    b = np.ones((T, K))
    for t in range(T-2, -1, -1): b[t] = (A[t] * (e[t+1]*b[t+1])[None]).sum(1); b[t] /= (b[t].sum() + 1e-9)
    g = a*b; g = g / (g.sum(1, keepdims=True) + 1e-9); return (g[..., None]*c).sum(1)


# ---------- PART A: gated temporal localization ----------
loc = {}
for v in VENUES:
    V = Venue(v); C, MA, VA, GT, SID = val_modes(V); t0 = recover_t0(V, GT, SID)
    em, es, eo = [], [], []
    for s in np.unique(SID):
        idx = np.where(SID == s)[0]; o = idx[np.argsort(t0[idx])]; c, ma, va, gt = C[o], MA[o], VA[o], GT[o]
        soft = fwd_bwd(c, ma, va, SIGMA); mw = ma / ma.sum(1, keepdims=True).clip(min=1e-9); pm = (mw[..., None]*c).sum(1)
        em += list(np.linalg.norm(pm-gt, axis=1)); es += list(np.linalg.norm(soft-gt, axis=1))
        dm = np.linalg.norm(c-gt[:, None], axis=-1); dm[~va] = 1e9; eo += list(dm.min(1))
    loc[v] = dict(mass=float(np.median(em)), soft=float(np.median(es)), oracle=float(np.median(eo)),
                  proxy=float(np.mean(MA[:, 0])))            # deployable reliability proxy = mean top-mode mass


# ---------- PART B: dual semantic probe [FCx | CamFormer] ----------
@torch.no_grad()
def dual_emb(net_pair, V, idx):
    return torch.cat([embed_all(net_pair[0], V, idx), cf_embed_all(net_pair[1], V, idx)], -1)   # [N, 64+384]


def dual_probe():
    Vs = {v: Venue(v) for v in VENUES}; rows = {}
    for v in VENUES:
        V = Vs[v]; Ztr = dual_emb((fnet, cnet), V, V.tr).detach(); lp = torch.where(V.lab[V.tr])[0]
        Zl, lo, ac = Ztr[lp], V.loc[V.tr][lp], V.act[V.tr][lp]; hl = lo >= 0
        p = nn.Sequential(nn.Linear(Ztr.shape[1], 128), nn.GELU(), nn.Linear(128, RL+A)).to(DEV)
        opt = torch.optim.AdamW(p.parameters(), 1e-3, weight_decay=1e-4)
        with torch.enable_grad():
            for _ in range(400):
                o = p(Zl); ll = F.cross_entropy(o[hl, :RL], lo[hl]) if hl.any() else o.new_zeros(())
                la = F.binary_cross_entropy_with_logits(o[:, RL:], ac); opt.zero_grad(); (ll+la).backward(); opt.step()
        Zv = dual_emb((fnet, cnet), V, V.va); o = p(Zv); lab = V.lab[V.va] & (V.loc[V.va] >= 0)
        rt = max(int(lab.sum()), 1); room = int((o[lab, :RL].argmax(-1) == V.loc[V.va][lab]).sum())/rt
        gt = V.act[V.va] > 0; lm = V.lab[V.va]; pr = (o[:, RL:].sigmoid() > 0.5) & lm[:, None]; gl = gt & lm[:, None]
        tp, fp, fn = int((pr & gl).sum()), int((pr & ~gl).sum()), int((~pr & gl).sum())
        rows[v] = dict(room=room, actF1=2*tp/(2*tp+fp+fn) if (2*tp+fp+fn) else 0.0)
    return rows


sem = dual_probe()
mac = lambda d, k: float(np.mean([d[v][k] for v in VENUES]))

# gate on proxy: try thresholds, pick best macro; also report oracle-gate upper bound
prox = np.array([loc[v]['proxy'] for v in VENUES])
def gated(thr): return float(np.mean([loc[v]['soft'] if loc[v]['proxy'] >= thr else loc[v]['mass'] for v in VENUES]))
best_thr = min(np.unique(prox), key=gated)
oracle_gate = float(np.mean([min(loc[v]['mass'], loc[v]['soft']) for v in VENUES]))

print(f"{'venue':6}{'mass':>6}{'soft':>6}{'oracle':>7}{'proxy':>7}  gate->{'use':>5}")
for v in VENUES:
    use = 'soft' if loc[v]['proxy'] >= best_thr else 'mass'
    print(f"{v:6}{loc[v]['mass']:6.2f}{loc[v]['soft']:6.2f}{loc[v]['oracle']:7.2f}{loc[v]['proxy']:7.3f}  {use:>9}")
print(f"\nLOCALIZATION macro:  mass-mean {mac(loc,'mass'):.2f} | soft(ungated) {mac(loc,'soft'):.2f} | "
      f"PROXY-GATED(thr={best_thr:.3f}) {gated(best_thr):.2f} | oracle-gate {oracle_gate:.2f} | per-window oracle {mac(loc,'oracle'):.2f}")
print(f"   refs: FCx top-1 3.31 | chance 3.95 | Method1/1b 3.4-3.7 (failed)")
print(f"\nSEMANTICS macro (dual FCx|CamFormer probe):  room {mac(sem,'room'):.3f}  actF1 {mac(sem,'actF1'):.3f}")
print(f"   refs: FCx room .462/actF1 .282 | CamFormer room .453/actF1 .353")
json.dump({'loc': loc, 'sem': sem, 'gate_thr': float(best_thr), 'gated_macro': gated(best_thr),
           'oracle_gate': oracle_gate}, open(AN / 'method_final.json', 'w'), indent=1)