summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/diag4.py
diff options
context:
space:
mode:
Diffstat (limited to 'artifacts/spectral_frontier_probe/diag4.py')
-rw-r--r--artifacts/spectral_frontier_probe/diag4.py29
1 files changed, 29 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/diag4.py b/artifacts/spectral_frontier_probe/diag4.py
new file mode 100644
index 0000000..0f7a75a
--- /dev/null
+++ b/artifacts/spectral_frontier_probe/diag4.py
@@ -0,0 +1,29 @@
+import torch, numpy as np
+from scipy.optimize import linear_sum_assignment
+np.set_printoptions(precision=3, suppress=True, linewidth=200)
+rng=np.random.default_rng(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('artifacts/synth_v1/omit_size.pt',map_location='cpu')
+V=standardise(d['visual_field']); T=standardise(d['text_field']); N=len(V)
+wV,UV=np.linalg.eigh(V); wV=wV[::-1]; UV=UV[:,::-1]
+wT,UT=np.linalg.eigh(T); wT=wT[::-1]; UT=UT[:,::-1]
+truth=np.arange(N)
+def acc(p): return 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
+
+print("### 1. ORACLE orthogonal mixing O between top-r spectral embeddings")
+for r in (6,8,10,12,16,20,24,32,42):
+ for scale in ('none','sqrt','lam'):
+ f=lambda w: np.ones_like(w) if scale=='none' else (np.sqrt(np.abs(w)) if scale=='sqrt' else np.abs(w))
+ XV=UV[:,:r]*f(wV[:r]); XT=UT[:,:r]*f(wT[:r])
+ # oracle Procrustes using truth
+ M=XT.T@XV; u,s,vt=np.linalg.svd(M); O=u@vt
+ p=hung(XV,XT@O)
+ # also cosine-normalised rows
+ nV=XV/np.linalg.norm(XV,axis=1,keepdims=True); nT=(XT@O); nT=nT/np.linalg.norm(nT,axis=1,keepdims=True)
+ p2=hung(nV,nT)
+ print(f" r={r:3d} scale={scale:5s} oracle-O acc={acc(p):.3f} rownorm acc={acc(p2):.3f}")