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