Finite-group relative message passing / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6
  7SEED = 563
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(4)
 10
 11
 12def torus(n, m):
 13    # vertices are (x,y), directed edges use displacement (+x,-x,+y,-y)
 14    N=n*m; edges=[]; rel=[]
 15    def ix(x,y): return (x%n)*m+(y%m)
 16    for x in range(n):
 17        for y in range(m):
 18            u=ix(x,y)
 19            for r,(dx,dy) in enumerate(((1,0),(-1,0),(0,1),(0,-1))):
 20                edges.append((u,ix(x+dx,y+dy))); rel.append(r)
 21    return N, np.asarray(edges,np.int64), np.asarray(rel,np.int64)
 22
 23
 24def quotient_check(n=5,m=7):
 25    # Signed cycle incidence vectors in Z^2. Every elementary torus cycle sums to 0.
 26    # The wrap cycles demonstrate finite quotient relations n*a_x=0 and m*a_y=0.
 27    cycles=[]
 28    for y in range(m):
 29        cycles.append(np.array([n,0],dtype=np.int64))
 30    for x in range(n):
 31        cycles.append(np.array([0,m],dtype=np.int64))
 32    for x in range(n):
 33        for y in range(m):
 34            cycles.append(np.array([0,0],dtype=np.int64)) # square commutator
 35    C=np.stack(cycles)
 36    local_max=int(np.max(np.abs(C[:n+m])))
 37    assert np.all(C[n+m:] == 0)
 38    # Labels in A=Z_n x Z_m; every edge displacement agrees modulo (n,m).
 39    N,E,R=torus(n,m); labels=np.array([(x,y) for x in range(n) for y in range(m)],dtype=np.int64)
 40    gens=np.array([[1,0],[-1,0],[0,1],[0,-1]],dtype=np.int64)
 41    residual=(labels[E[:,1]]-labels[E[:,0]]-gens[R])
 42    ok=np.all((residual[:,0]%n==0)&(residual[:,1]%m==0))
 43    # path independence: two paths' displacement difference is a cycle relation
 44    assert ok
 45    return {"group":f"Z_{n} x Z_{m}","cycle_rows":len(C),"max_wrap_incidence":local_max,
 46            "all_edge_label_differences_zero_in_quotient":bool(ok),
 47            "relation_classes":4,"directed_edges":int(len(E))}
 48
 49
 50class RelationMP(nn.Module):
 51    def __init__(self, d=16):
 52        super().__init__(); self.W=nn.Parameter(torch.randn(4,d,d)*.08)
 53        self.out=nn.Sequential(nn.Linear(d,d),nn.Tanh(),nn.Linear(d,1))
 54    def forward(self,x,edge,rel):
 55        z=torch.zeros_like(x)
 56        for r in range(4):
 57            mask=(rel==r); src=edge[mask,0]; dst=edge[mask,1]
 58            z.index_add_(0,dst,x[src]@self.W[r].T)
 59        return self.out(z)
 60
 61class GCN(nn.Module):
 62    def __init__(self,d=16):
 63        super().__init__(); self.W=nn.Parameter(torch.randn(d,d)*.08)
 64        self.out=nn.Sequential(nn.Linear(d,d),nn.Tanh(),nn.Linear(d,1))
 65    def forward(self,x,edge,rel):
 66        z=torch.zeros_like(x); z.index_add_(0,edge[:,1],x[edge[:,0]]@self.W.T)
 67        deg=torch.bincount(edge[:,1],minlength=x.shape[0]).float().unsqueeze(1)
 68        return self.out(z/deg)
 69
 70def data(n,m,seed):
 71    rng=np.random.default_rng(seed); N,E,R=torus(n,m)
 72    x=torch.tensor(rng.normal(size=(N,16)),dtype=torch.float32)
 73    # A relation-dependent one-hop convolution, with independent inputs per graph.
 74    coeff=torch.tensor([1.0,-0.7,0.45,-1.15])
 75    y=torch.zeros(N)
 76    for r in range(4): y.index_add_(0,torch.tensor(E[R==r,1]),coeff[r]*x[torch.tensor(E[R==r,0]),0])
 77    return x,torch.tensor(E),torch.tensor(R),y[:,None]
 78
 79def fit(model, n=6,m=6,steps=700):
 80    opt=torch.optim.Adam(model.parameters(),lr=.025)
 81    model.train()
 82    for t in range(steps):
 83        x,e,r,y=data(n,m,1000+t)
 84        pred=model(x,e,r); loss=((pred-y)**2).mean()
 85        opt.zero_grad(); loss.backward(); opt.step()
 86    return model
 87
 88def evaluate(model,n,m,count=20):
 89    model.eval(); vals=[]
 90    with torch.no_grad():
 91        for j in range(count):
 92            x,e,r,y=data(n,m,9000+j); vals.append(float(((model(x,e,r)-y)**2).mean()))
 93    return float(np.mean(vals))
 94
 95def main():
 96    math_result=quotient_check()
 97    # Same architecture scale, source training and larger-geometry transfer.
 98    rel=fit(RelationMP()); gcn=fit(GCN())
 99    results={
100      "source_6x6": {"GCN":evaluate(gcn,6,6),"quotient_tied":evaluate(rel,6,6)},
101      "transfer_8x8": {"GCN":evaluate(gcn,8,8),"quotient_tied":evaluate(rel,8,8)},
102      "parameters":{"GCN":sum(p.numel() for p in gcn.parameters()),"quotient_tied":sum(p.numel() for p in rel.parameters()),"untied_relation_bank_equivalent":4*16*16},
103      "compression_edges_over_generators":(2*6*6*4)/4,
104      "math":math_result}
105    Path('results.json').write_text(json.dumps(results,indent=2))
106    print(json.dumps(results,indent=2))
107
108if __name__=='__main__': main()