import sys, time, torch, numpy as np sys.path.insert(0,'/home/yurenh2/emm') from scipy.optimize import linear_sum_assignment from scipy.stats import ortho_group dev='cuda:0' def standardise(M): M=np.asarray(M,dtype=np.float64); mask=~np.eye(len(M),dtype=bool); v=M[mask] out=(M-v.mean())/v.std(); np.fill_diagonal(out,0.0); return out d=torch.load('/home/yurenh2/emm/artifacts/synth_v1/omit_size.pt',map_location='cpu') V=standardise(d['visual_field']); T=standardise(d['text_field']); N=len(V) rng=np.random.default_rng(31337); sigma=rng.permutation(N) Ts=T[np.ix_(sigma,sigma)]; truth=np.argsort(sigma) Vt=torch.tensor(V,dtype=torch.float32,device=dev); Tt=torch.tensor(Ts,dtype=torch.float32,device=dev) CONST=float((Tt*Tt).sum()+(Vt*Vt).sum()) def energy(p): P=torch.as_tensor(np.asarray(p),dtype=torch.long,device=dev) return (CONST-2.0*float((Tt[P[:,None],P[None,:]]*Vt).sum()))/(N*(N-1)) acc=lambda p: float((p==truth).mean()) wV,UV=np.linalg.eigh(V); wV=wV[::-1]; UV=UV[:,::-1] wT,UT=np.linalg.eigh(Ts); wT=wT[::-1]; UT=UT[:,::-1] 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 icp(XV,XT,O,iters=25): for _ in range(iters): p=hung(XV,XT@O); u,s,vt=np.linalg.svd(XT[p].T@XV); On=u@vt if np.allclose(On,O,atol=1e-10): O=On; break O=On return hung(XV,XT@O) sols=[] t0=time.time() for r in (10,12,14,16,18,20): XV=UV[:,:r]*np.sqrt(np.abs(wV[:r])); XT=UT[:,:r]*np.sqrt(np.abs(wT[:r])) for t in range(60): p=icp(XV,XT,ortho_group.rvs(r,random_state=int(rng.integers(1<<30)))) sols.append((energy(p),acc(p),p,r)) sols.sort(key=lambda z:z[0]) E=np.array([s[0] for s in sols]); A=np.array([s[1] for s in sols]) print(f"{len(sols)} ICP solutions in {time.time()-t0:.0f}s; corr(E,acc)={np.corrcoef(E,A)[0,1]:.3f}") print("lowest-10 (E,acc):", [(round(e,4),round(a,3)) for e,a,_,_ in sols[:10]]) print("E quantiles",np.quantile(E,[0,.1,.5,.9,1]).round(3)," acc of best-E:",A[0]) for m in (5,10,20,40): votes=np.zeros((N,N)) for e,a,p,r in sols[:m]: votes[np.arange(N),p]+=1 conf=votes.max(1); pick=np.argsort(-conf) pm=votes.argmax(1) for k in (12,25,50,100): sel=pick[:k]; prec=float((pm[sel]==truth[sel]).mean()) print(f" consensus over {m} lowest-E sols: precision@{k} = {prec:.3f}", end='') print()