summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/diag6.py
blob: 0158f9dd09d7cc58283d0212ad0875c062497293 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
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(1)

def icp(XV,XT,O,iters=40):
    for _ in range(iters):
        p=hung(XV,XT@O)
        u,s,vt=np.linalg.svd(XT[p].T@XV); Onew=u@vt
        if np.allclose(Onew,O,atol=1e-10): O=Onew; break
        O=Onew
    return hung(XV,XT@O),O

print("### BLIND: random-restart ICP over O(r), scored by QAP energy")
for r in (6,8,10,12,16):
    XV=UV[:,:r]*np.sqrt(np.abs(wV[:r])); XT=UT[:,:r]*np.sqrt(np.abs(wT[:r]))
    t0=time.time(); best=(1e9,None)
    R=200
    for t in range(R):
        O=ortho_group.rvs(r,random_state=int(rng.integers(1<<30)))
        p,_=icp(XV,XT,O)
        e=energy(p)
        if e<best[0]: best=(e,p)
    p=best[1]; pd=descend(p)
    print(f"  r={r:3d} R={R}: best-E {best[0]:.4f} acc {acc(p):.3f} | after descent E {energy(pd):.4f} acc {acc(pd):.3f}  [{time.time()-t0:.0f}s]")