summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/nmf_adv_prune_gw.py
blob: 2ecdfed6a67ecbaead22e9ae6849dce0fa8b4e9b (plain)
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
"""Second solver on the ORACLE-pruned field, so the refutation is not GRAMPA-specific."""
import numpy as np, torch, warnings, ot
warnings.filterwarnings('ignore')
from scipy.optimize import linear_sum_assignment
dev = torch.device('cuda:0'); DT = torch.float64
d = torch.load('/home/yurenh2/emm/artifacts/synth_v1/omit_size.pt', map_location='cpu')
V0 = d['visual_field'].double().numpy(); T0 = d['text_field'].double().numpy(); N = len(V0)
rng = np.random.default_rng(777); sigma = rng.permutation(N)
T0s = T0[np.ix_(sigma, sigma)]; truth = np.argsort(sigma)
acc = lambda p: float((p == truth).mean())
def standardise(M):
    M = np.asarray(M, float); mask = ~np.eye(len(M), dtype=bool); v = M[mask]
    out = (M - v.mean()) / v.std(); np.fill_diagonal(out, 0.0); return out
def sym_nmf_gpu(M, r, iters=3000, seed=0):
    Mg = torch.tensor(np.clip(M, 0, None), device=dev, dtype=DT)
    g = torch.Generator(device='cpu').manual_seed(seed)
    W = torch.abs(torch.randn(len(M), r, generator=g, dtype=DT)).to(dev) * float(np.sqrt(M.mean() / r))
    for _ in range(iters):
        W = W * (0.5 + 0.5 * (Mg @ W) / (W @ (W.T @ W) + 1e-12))
    return W.cpu().numpy()
Vs = standardise(V0); Ts = standardise(T0s)
CONST = float((Vs*Vs).sum() + (Ts*Ts).sum()); DEN = N*(N-1)
Ag = torch.tensor(Vs, device=dev, dtype=DT); Tg = torch.tensor(Ts, device=dev, dtype=DT)
energy = lambda p: (CONST - 2.0*float((Ts[np.ix_(p, p)]*Vs).sum()))/DEN
def descend(p0, A=None, max_steps=4000):
    A = Ag if A is None else A
    p = torch.tensor(np.asarray(p0), device=dev, dtype=torch.long)
    iu = torch.triu_indices(N, N, offset=1, device=dev)
    for _ in range(max_steps):
        B = Tg[p][:, p]; C = A @ B; dg = torch.diagonal(C)
        G = C + C.T - dg[:, None] - dg[None, :] + 2*A*B
        vals = G[iu[0], iu[1]]; k = int(vals.argmax())
        if float(vals[k]) <= 1e-12: break
        u, v = int(iu[0][k]), int(iu[1][k]); p[u], p[v] = p[v].clone(), p[u].clone()
    return p.cpu().numpy()
def gw(A, B, seed=0):
    q = np.ones(N)/N
    G, _ = ot.gromov.gromov_wasserstein(A, B, q, q, 'square_loss', log=True, max_iter=200)
    r_, c_ = linear_sum_assignment(-G); return c_
r = 42
WV = sym_nmf_gpu(V0, r, seed=1); WT = sym_nmf_gpu(T0s, r, seed=1)
nrm = lambda X: X/(np.linalg.norm(X, axis=0, keepdims=True)+1e-12)
Sx = nrm(WV).T @ nrm(WT)[truth]; rr, cc = linear_sum_assignment(-Sx); ocs = Sx[rr, cc]
print("baseline GW on raw:", end=" ")
p = gw(Vs, Ts); pd = descend(p); print(f"raw {acc(p):.4f} refined {acc(pd):.4f} E={energy(pd):.4f}")
for thr in (0.5, 0.65, 0.75):
    keepV = rr[ocs >= thr]
    Vc = standardise(WV[:, keepV] @ WV[:, keepV].T)
    p = gw(Vc, Ts); pd = descend(p, A=torch.tensor(Vc, device=dev, dtype=DT)); pd2 = descend(pd)
    print(f"prune thr={thr} kept {len(keepV)}/{r}: GW raw {acc(p):.4f} -> clean-desc {acc(pd):.4f} -> true-desc {acc(pd2):.4f} E={energy(pd2):.4f}")