summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/confirm.py
diff options
context:
space:
mode:
Diffstat (limited to 'artifacts/spectral_frontier_probe/confirm.py')
-rw-r--r--artifacts/spectral_frontier_probe/confirm.py49
1 files changed, 49 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/confirm.py b/artifacts/spectral_frontier_probe/confirm.py
new file mode 100644
index 0000000..d0fc003
--- /dev/null
+++ b/artifacts/spectral_frontier_probe/confirm.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
+from worldalign.synth_fast_gate import fast_pair_descent
+dev='cuda:3'
+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)
+Vt=torch.tensor(V,dtype=torch.float32,device=dev); Tt=torch.tensor(T,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))
+def descend(p,steps=4000):
+ P=torch.as_tensor(np.asarray(p),dtype=torch.long,device=dev)
+ return fast_pair_descent(Tt,Vt,P,steps).cpu().numpy()
+truth=np.arange(N); acc=lambda p: float((p==truth).mean())
+wV,UV=np.linalg.eigh(V); wV=wV[::-1]; UV=UV[:,::-1]
+wT,UT=np.linalg.eigh(T); 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
+rng=np.random.default_rng(20260801)
+def icp(XV,XT,O,iters=30):
+ 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)
+t0=time.time(); pool=[]
+for r in (4,6,8,10,12,14,16,20,24):
+ XV=UV[:,:r]*np.sqrt(np.abs(wV[:r])); XT=UT[:,:r]*np.sqrt(np.abs(wT[:r]))
+ cand=[]
+ for t in range(120):
+ p=icp(XV,XT,ortho_group.rvs(r,random_state=int(rng.integers(1<<30))))
+ cand.append((energy(p),p))
+ cand.sort(key=lambda z:z[0])
+ best=None
+ for e,p in cand[:5]:
+ pd=descend(p); ed=energy(pd)
+ if best is None or ed<best[0]: best=(ed,pd)
+ pool.append((best[0],best[1],r))
+ print(f" r={r:3d}: pre-descent best E {cand[0][0]:.4f} (acc {acc(cand[0][1]):.3f}) -> post-descent E {best[0]:.4f} acc {acc(best[1]):.3f}",flush=True)
+pool.sort(key=lambda z:z[0])
+print(f"\nBLIND PICK (lowest E over all r): r={pool[0][2]} E={pool[0][0]:.4f} ACC={acc(pool[0][1]):.3f}")
+print(f"E(truth)={energy(truth):.4f} total wall {time.time()-t0:.0f}s")