diff options
Diffstat (limited to 'artifacts/spectral_frontier_probe/band.py')
| -rw-r--r-- | artifacts/spectral_frontier_probe/band.py | 26 |
1 files changed, 26 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/band.py b/artifacts/spectral_frontier_probe/band.py new file mode 100644 index 0000000..cca368a --- /dev/null +++ b/artifacts/spectral_frontier_probe/band.py @@ -0,0 +1,26 @@ +import numpy as np, torch +from scipy.optimize import linear_sum_assignment +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); truth=np.arange(N) +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 +acc=lambda p: float((p==truth).mean()) +CONST=float((T*T).sum()+(V*V).sum()) +en=lambda p:(CONST-2*(T[np.ix_(p,p)]*V).sum())/(N*(N-1)) +for r in (12,16,24): + XV=UV[:,:r]*np.sqrt(np.abs(wV[:r])); XT=UT[:,:r]*np.sqrt(np.abs(wT[:r])) + u,s,vt=np.linalg.svd(XT.T@XV); O=u@vt + idx=np.arange(r); band_mass=[] + for b in (1,2,3,4,6,r): + Mk=(np.abs(idx[:,None]-idx[None,:])<=b).astype(float) + frac=float((O**2*Mk).sum()/ (O**2).sum()) + Ob=O*Mk; uu,ss,vv=np.linalg.svd(Ob); Ob=uu@vv + p=hung(XV,XT@Ob) + band_mass.append((b,round(frac,3),round(acc(p),3),round(en(p),4))) + print(f"r={r}: (bandwidth, mass of oracle-O inside band, acc after re-orthogonalising, E) -> {band_mass}") + print(f" free params: full {r*(r-1)//2}, band3 {sum(min(3,r-1-i) for i in range(r))}") |
