diff options
Diffstat (limited to 'artifacts/spectral_frontier_probe/diag11.py')
| -rw-r--r-- | artifacts/spectral_frontier_probe/diag11.py | 52 |
1 files changed, 52 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/diag11.py b/artifacts/spectral_frontier_probe/diag11.py new file mode 100644 index 0000000..4977c3c --- /dev/null +++ b/artifacts/spectral_frontier_probe/diag11.py @@ -0,0 +1,52 @@ +import numpy as np, torch, warnings, time +warnings.filterwarnings('ignore') +from scipy.optimize import linear_sum_assignment +from sklearn.decomposition import NMF +np.set_printoptions(precision=3,suppress=True,linewidth=200) +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 hung(A,B): + C=((A**2).sum(1)[:,None]+(B**2).sum(1)[None,:]-2*A@B.T); r,c=linear_sum_assignment(C); return c +def sym_nmf(M,r,iters=3000,seed=0): + """symmetric NMF M ~ W W^T by multiplicative updates (Ding et al.)""" + M=np.clip(M,0,None).copy(); np.fill_diagonal(M,np.clip(np.diag(M),0,None)) + rg=np.random.default_rng(seed); W=np.abs(rg.standard_normal((len(M),r)))*np.sqrt(M.mean()/r) + for _ in range(iters): + num=M@W; den=W@(W.T@W)+1e-12 + W=W*(0.5+0.5*num/den) + return W +for r in (24,42): + t0=time.time() + WV=sym_nmf(V0,r,seed=1); WT=sym_nmf(T0s,r,seed=1) + print(f"r={r} symNMF resid V {np.linalg.norm(V0-WV@WV.T)/np.linalg.norm(V0):.3f} T {np.linalg.norm(T0s-WT@WT.T)/np.linalg.norm(T0s):.3f} [{time.time()-t0:.0f}s]") + A=WV.copy(); B=WT.copy() + # ORACLE column match for reference + An=A/np.linalg.norm(A,axis=0,keepdims=True); Bn=B/np.linalg.norm(B,axis=0,keepdims=True) + rr,cc=linear_sum_assignment(-(An[np.arange(N)].T@Bn[truth])) + print(" oracle col cos:",np.sort((An.T@Bn[truth])[rr,cc])[::-1][:10]) + # BLIND column match by permutation-invariant column signatures + def sig(X): + Xn=X/ (np.linalg.norm(X,axis=0,keepdims=True)+1e-12) + q=np.quantile(Xn,np.linspace(0.5,1.0,16),axis=0).T + return np.hstack([q,(Xn>0.02).mean(0)[:,None],np.linalg.norm(X,axis=0)[:,None]/np.linalg.norm(X)]) + sA=sig(A); sB=sig(B); m=hung(sA,sB) + agree=float((m==cc[np.argsort(rr)]).mean()) + def scene_match(colmap): + AV=A; BT=B[:,colmap] + AV=AV/np.linalg.norm(AV,axis=1,keepdims=True).clip(1e-9); BT=BT/np.linalg.norm(BT,axis=1,keepdims=True).clip(1e-9) + return hung(AV,BT) + p=scene_match(m); print(f" BLIND stat col-match agrees w/ oracle {agree:.2f} -> scene acc {acc(p):.3f}") + # alternate: scene Hungarian <-> column Hungarian + colmap=m.copy() + for it in range(25): + p=scene_match(colmap) + Ar=A/np.linalg.norm(A,axis=1,keepdims=True).clip(1e-9) + Br=B/np.linalg.norm(B,axis=1,keepdims=True).clip(1e-9) + rr2,cc2=linear_sum_assignment(-(Ar.T@Br[p])); newmap=cc2[np.argsort(rr2)] + if (newmap==colmap).all(): break + colmap=newmap + p=scene_match(colmap) + print(f" after co-alternation ({it+1} iters): col agree {float((colmap==cc[np.argsort(rr)]).mean()):.2f} -> scene acc {acc(p):.3f}") |
