summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/seed.py
diff options
context:
space:
mode:
Diffstat (limited to 'artifacts/spectral_frontier_probe/seed.py')
-rw-r--r--artifacts/spectral_frontier_probe/seed.py49
1 files changed, 49 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/seed.py b/artifacts/spectral_frontier_probe/seed.py
new file mode 100644
index 0000000..5e33d7b
--- /dev/null
+++ b/artifacts/spectral_frontier_probe/seed.py
@@ -0,0 +1,49 @@
+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()