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
47
48
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()
|