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}")