Submetry-Lifted Relational Alignment / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()