Submetry-Lifted Relational Alignment / experiment.py
Mechanism confirmed, baseline not beaten
1import math, json, time
2import numpy as np
3from scipy.optimize import linear_sum_assignment
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8SEED=2140
9np.random.seed(SEED); torch.manual_seed(SEED)
10
11def sinkhorn(C, a=None, b=None, eps=0.1, iters=60):
12 C=np.asarray(C,float); n,m=C.shape
13 if a is None: a=np.ones(n)/n
14 if b is None: b=np.ones(m)/m
15 a=np.asarray(a); b=np.asarray(b)
16 eps=max(float(eps),1e-10)
17 logK=-C/eps
18 la=np.log(a); lb=np.log(b)
19 u=np.zeros(n); v=np.zeros(m)
20 for _ in range(iters):
21 u=la-np.logaddexp.reduce(logK+v[None,:],axis=1)
22 v=lb-np.logaddexp.reduce(logK+u[:,None],axis=0)
23 P=np.exp(logK+u[:,None]+v[None,:])
24 return P
25
26def lifted_cost(W,V,P):
27 C=(W[:,:,None,None]-V[None,None,:,:])**2
28 return float(np.einsum('ik,jl,ijkl->',P,P,C,optimize=True))
29
30def quotient_sinkhorn(W,V,eps_scale=.03):
31 n=W.shape[0]; # node-pair costs
32 C=np.zeros((n,n))
33 # quadratic assignment cost used to obtain a node coupling; this is the
34 # standard lifted objective evaluated approximately by alternating updates.
35 P=np.ones((n,n))/n
36 for _ in range(5):
37 C=np.einsum('jl,ikjl->ik',P,(W[:,:,None,None]-V[None,None,:,:])**2)*n
38 med=np.median(C[C>1e-12]) if np.any(C>1e-12) else 1.
39 P=sinkhorn(C,eps=max(eps_scale*med,1e-5),iters=80)
40 return P, lifted_cost(W,V,P)
41
42def make_graph(n=6):
43 x=np.random.rand(n,n); return (x+x.T)/2
44
45def toy_checks():
46 W=make_graph(); perm=np.random.permutation(len(W)); V=W[np.ix_(perm,perm)]
47 # identity and optimized lifted costs
48 I=np.eye(len(W))/len(W)
49 fixed=lifted_cost(W,V,I)
50 inv=np.argsort(perm); Prelabel=np.zeros_like(I); Prelabel[np.arange(len(W)),inv]=1/len(W)
51 exact_relabel=lifted_cost(W,V,Prelabel)
52 P,q=quotient_sinkhorn(W,V)
53 marg=max(np.max(abs(P.sum(1)-1/len(W))),np.max(abs(P.sum(0)-1/len(W))))
54 # noise scaling prediction: d_Q is proportional to sigma for p=2
55 sigmas=np.array([0,.01,.03,.06,.12,.24])
56 vals=[]
57 for s in sigmas:
58 reps=[]
59 for r in range(4):
60 N=np.random.randn(*W.shape); N=(N+N.T)/2
61 Vn=W+float(s)*N
62 # Compare the unpermuted graph to its noisy version under identity.
63 reps.append(math.sqrt(lifted_cost(W,Vn,I)))
64 vals.append(np.mean(reps))
65 slope=float(np.dot(sigmas,vals)/max(np.dot(sigmas,sigmas),1e-12))
66 pred_intercept=vals[0]
67 # regularization prediction: entropy rises with epsilon
68 base=(W[:,:,None,None]-V[None,None,:,:])**2
69 med=np.median(base)
70 epses=med*np.array([.003,.01,.03,.1,.3,1.0])
71 ents=[]; costs=[]
72 for e in epses:
73 # pairwise node cost from uniform lifted current iterate
74 P=np.ones((len(W),len(W)))/len(W)
75 for _ in range(5):
76 C=np.einsum('jl,ikjl->ik',P,base)*len(W)
77 P=sinkhorn(C,eps=e,iters=80)
78 ent=float(-(P[P>0]*np.log(P[P>0])).sum())
79 ents.append(ent); costs.append(lifted_cost(W,V,P))
80 assert lifted_cost(W,W,I) < 1e-12
81 return dict(identity_zero=math.sqrt(lifted_cost(W,W,I)), exact_relabel_root=math.sqrt(exact_relabel), fixed_root=math.sqrt(fixed), quotient_root=math.sqrt(q), marginal_error=marg,
82 sigmas=sigmas.tolist(), noise_roots=np.array(vals).tolist(), noise_slope=slope,
83 entropy_eps=epses.tolist(), entropies=ents, reg_costs=costs)
84
85class OrderEquivariantNet(nn.Module):
86 def __init__(self,n):
87 super().__init__(); self.node=nn.Sequential(nn.Linear(n,24),nn.ReLU(),nn.Linear(24,12),nn.ReLU()); self.cls=nn.Linear(12,2)
88 def forward(self,A):
89 one = A.ndim == 2
90 if one: A = A.unsqueeze(0)
91 H=self.node(A) # each row sees its relational neighborhood
92 out=self.cls(H.mean(1))
93 return (out[0] if one else out), H
94
95def data(n=8, count=100):
96 As=[]; ys=[]
97 for y in range(2):
98 for _ in range(count//2):
99 p=.25 if y==0 else .65
100 X=(np.random.rand(n,n)<p).astype('float32'); X=np.triu(X,1); X=X+X.T
101 np.fill_diagonal(X,0); As.append(X); ys.append(y)
102 ix=np.random.permutation(len(ys)); return np.array(As)[ix],np.array(ys)[ix]
103
104def train_compare():
105 np.random.seed(SEED); torch.manual_seed(SEED)
106 A,y=data(); n=A.shape[1]; cut=int(.7*len(A))
107 At=torch.tensor(A[:cut]); yt=torch.tensor(y[:cut]); Av=torch.tensor(A[cut:]); yv=torch.tensor(y[cut:])
108 def run(align):
109 torch.manual_seed(SEED+int(align)); model=OrderEquivariantNet(n); opt=torch.optim.Adam(model.parameters(),lr=.01)
110 t0=time.time()
111 for step in range(35):
112 ix=torch.randperm(cut)[:16]; X=At[ix]; target=yt[ix]
113 if align:
114 X2=[]; Ps=[]
115 for x in X.numpy():
116 p=np.random.permutation(n); X2.append(x[np.ix_(p,p)])
117 # coupling is computed between relational rows, with detached Sinkhorn
118 C=((x[:,None,:]-X2[-1][None,:,:])**2).mean(2)
119 Ps.append(sinkhorn(C,eps=.1*max(np.median(C),1e-3),iters=35))
120 X2=torch.tensor(np.array(X2),dtype=torch.float32)
121 logits,H=model(X); logits2,H2=model(X2)
122 align_loss=0.0
123 for h1,h2,P in zip(H,H2,Ps):
124 Pt=torch.tensor(P,dtype=h1.dtype)
125 align_loss=align_loss+((h1[:,None,:]-h2[None,:,:])**2*Pt[:,:,None]).sum()
126 align_loss=align_loss/len(H)
127 loss=F.cross_entropy(logits,target)+.15*F.mse_loss(logits,logits2)+.05*align_loss
128 else:
129 logits,H=model(X); loss=F.cross_entropy(logits,target)
130 opt.zero_grad(); loss.backward(); opt.step()
131 with torch.no_grad():
132 pred=model(Av)[0].argmax(1); acc=float((pred==yv).float().mean())
133 # prediction variance over 12 arbitrary node permutations
134 vars=[]
135 for x in Av[:20]:
136 ls=[]
137 for _ in range(5):
138 p=torch.randperm(n); ls.append(torch.softmax(model(x[p][:,p])[0],0)[1].item())
139 vars.append(np.var(ls))
140 return dict(acc=acc,perm_prob_variance=float(np.mean(vars)),seconds=time.time()-t0)
141 return run(False),run(True)
142
143def main():
144 toy=toy_checks(); base,idea=train_compare()
145 out={'toy':toy,'baseline':base,'idea':idea}
146 with open('results.json','w') as f: json.dump(out,f,indent=2)
147 print(json.dumps(out,indent=2))
148if __name__=='__main__': main()