summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/nmf_adv_prune_control.py
blob: a3b23efcc9c1bae0a249c6ebcdc4cb8b0894a715 (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
51
52
53
54
55
56
57
58
59
60
61
"""CONTROL: is the 0.86 from PRUNING (needs oracle) or from symNMF RECONSTRUCTION alone (blind)?"""
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):
    q = np.ones(N)/N
    G, _ = ot.gromov.gromov_wasserstein(A, B, q, q, 'square_loss', log=True, max_iter=200)
    _, c = linear_sum_assignment(-G); return c
def run(tag, Vm, Tm=Ts):
    p = gw(Vm, Tm); pd = descend(p, A=torch.tensor(Vm, device=dev, dtype=DT)); pd2 = descend(pd)
    print(f"  {tag:46s} GW {acc(p):.4f} -> clean-desc {acc(pd):.4f} -> true-desc {acc(pd2):.4f}  E={energy(pd2):.4f}", flush=True)
    return acc(pd2)

r = 42
for seed in (1, 2, 3):
    print(f"--- symNMF seed {seed}, r={r}")
    WV = sym_nmf_gpu(V0, r, seed=seed); WT = sym_nmf_gpu(T0s, r, seed=seed)
    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]
    worst = rr[int(np.argmin(ocs))]
    Vrec = standardise(WV @ WV.T)                                  # BLIND: no pruning at all
    run("A. symNMF reconstruction, NO pruning (blind)", Vrec)
    keepV = rr[ocs >= 0.5]
    run(f"B. ORACLE prune worst {r-len(keepV)} (thr .5)", standardise(WV[:, keepV] @ WV[:, keepV].T))
    for t in range(3):                                              # BLIND: drop a random column
        j = np.random.default_rng(100+t).integers(r)
        kk = np.setdiff1d(np.arange(r), [j])
        run(f"C{t}. RANDOM single-column drop (col {j}, blind)", standardise(WV[:, kk] @ WV[:, kk].T))
    # also: text side reconstructed the same way (symmetric treatment)
    run("D. both sides symNMF-reconstructed (blind)", Vrec, standardise(WT @ WT.T))