Fermionic circuit message passing / fermionic_mvp.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6SEED=1481
  7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  8
  9def j_sign(j):
 10    return (-1)**(j*(j-1)//2)
 11
 12def wedge2(x,y):
 13    return torch.outer(x,y)-torch.outer(y,x)
 14
 15def fermionic_local(messages, require_even=True):
 16    # degree 0 and degree 2 truncation; odd incident degree is rejected
 17    n=len(messages)
 18    if require_even and n % 2: return torch.tensor(0., dtype=messages[0].dtype), torch.zeros((messages[0].numel(),messages[0].numel()))
 19    if n < 2: return torch.tensor(1., dtype=messages[0].dtype), torch.zeros((messages[0].numel(),messages[0].numel()))
 20    A=torch.zeros((messages[0].numel(),messages[0].numel()),dtype=messages[0].dtype)
 21    for i in range(n):
 22        for k in range(i+1,n): A += wedge2(messages[i],messages[k])
 23    return torch.tensor(1., dtype=messages[0].dtype), A
 24
 25def math_checks():
 26    out={}
 27    # Prediction 1: every odd number of incident odd half-edges is exactly killed.
 28    odd_res=[]
 29    d=4
 30    for n in range(1,10,2):
 31        xs=[torch.randn(d) for _ in range(n)]
 32        s,A=fermionic_local(xs)
 33        odd_res.append(float(abs(s)+torch.linalg.norm(A)))
 34    out['odd_degree_predicted_zero']= {'predicted':0.0,'observed_max_residual':max(odd_res),'all_residuals':odd_res}
 35    # Prediction 2: parity sign is (-1)^|F|, swept over all masks and graph sizes.
 36    parity=[]
 37    for E in range(1,9):
 38        maxerr=0
 39        for mask in range(1<<E):
 40            m=mask.bit_count(); lhs=(-1)**sum(((2 if False else 0) for _ in []))
 41            # construct a cycle: each selected edge contributes to two endpoint degrees
 42            deg=[0]*E
 43            for e in range(E):
 44                if mask>>e&1: deg[e]+=1; deg[(e+1)%E]+=1
 45            # Odd local degrees vanish; compare signs only on surviving masks.
 46            if any(q % 2 for q in deg):
 47                continue
 48            rhs=(-1)**m
 49            local=1
 50            for q in deg: local*=(-1)**(q//2)
 51            maxerr=max(maxerr,abs(local-rhs))
 52        parity.append({'edges':E,'max_abs_error':maxerr})
 53    out['parity_predicted_exact']= {'predicted_error':0.0,'sweep':parity}
 54    # Prediction 3: random degree-2 wedge energy scales as d(d-1); normalized energy is stable.
 55    scaling=[]
 56    for d in [2,3,4,6,8,12]:
 57        vals=[]
 58        for _ in range(400):
 59            x=torch.randn(d); y=torch.randn(d); vals.append(float(torch.linalg.norm(wedge2(x,y))**2))
 60        mean=np.mean(vals); predicted=2*d*(d-1)
 61        scaling.append({'ell':d//2,'dimension':d,'predicted_mean':predicted,'observed_mean':float(mean),'normalized':float(mean/predicted)})
 62    out['wedge_energy_scaling']=scaling
 63    return out
 64
 65class GraphSet:
 66    def __init__(self,ntrain=240, ntest=120, min_n=6,max_n=10):
 67        self.train=self.make(ntrain,min_n,max_n); self.test=self.make(ntest, max_n+1, max_n+5)
 68    def make(self,N,lo,hi):
 69        arr=[]
 70        for z in range(N):
 71            n=random.randint(lo,hi); y=z%2
 72            edges=[]
 73            if y==1: # cycle
 74                edges=[(i,(i+1)%n) for i in range(n)]
 75            else: # tree/path, same node count and degree statistics mostly
 76                edges=[(i,i+1) for i in range(n-1)]
 77                random.shuffle(edges)
 78            arr.append((n,edges,y))
 79        return arr
 80
 81def batch_graph(g, max_n):
 82    n,edges,y=g
 83    # node scalar is constant; challenge is topology
 84    X=torch.ones(n,1)
 85    return X,edges,y
 86
 87class GIN(nn.Module):
 88    def __init__(self,h=16, fermionic=False,d=4):
 89        super().__init__(); self.fermionic=fermionic; self.d=d
 90        self.inp=nn.Linear(1,h)
 91        self.e=nn.ModuleList([nn.Sequential(nn.Linear(h,h),nn.ReLU(),nn.Linear(h,h)) for _ in range(3)])
 92        if fermionic:
 93            self.odd=nn.ModuleList([nn.Linear(h,d) for _ in range(3)])
 94            self.mix=nn.ModuleList([nn.Sequential(nn.Linear(h+d*d,h),nn.ReLU(),nn.Linear(h,h)) for _ in range(3)])
 95        self.cls=nn.Sequential(nn.Linear(h, h),nn.ReLU(),nn.Linear(h,2))
 96    def forward(self,X,edges):
 97        h=self.inp(X); n=h.shape[0]
 98        adj=[[] for _ in range(n)]
 99        for u,v in edges: adj[u].append(v); adj[v].append(u)
100        for t in range(3):
101            agg=torch.zeros_like(h); odd=[[] for _ in range(n)]
102            for v in range(n):
103                for u in adj[v]:
104                    agg[v]+=self.e[t](h[u])
105                    if self.fermionic: odd[v].append(self.odd[t](h[u]))
106            if self.fermionic:
107                feats=[]
108                for v in range(n):
109                    A=fermionic_local(odd[v])[1].reshape(-1)
110                    feats.append(torch.cat([agg[v],A]))
111                h=h+self.mix[t](torch.stack(feats))
112            else: h=h+agg
113        return self.cls(h.mean(0))
114
115def train_eval(fermionic, data):
116    torch.manual_seed(SEED+int(fermionic)); model=GIN(fermionic=fermionic)
117    opt=torch.optim.Adam(model.parameters(),lr=2e-3)
118    for ep in range(12):
119        random.shuffle(data.train); model.train()
120        for g in data.train:
121            X,e,y=batch_graph(g,16); loss=nn.functional.cross_entropy(model(X,e)[None,:],torch.tensor([y]))
122            opt.zero_grad(); loss.backward(); opt.step()
123    model.eval(); good=0; total=0
124    with torch.no_grad():
125        for g in data.test:
126            X,e,y=batch_graph(g,16); good += int(model(X,e).argmax().item()==y); total+=1
127    return good/total
128
129def main():
130    checks=math_checks(); data=GraphSet(); b=train_eval(False,data); f=train_eval(True,data)
131    result={'seed':SEED,'checks':checks,'classification':{'baseline_gin_accuracy':b,'fermionic_accuracy':f,'test_graphs':len(data.test),'note':'test cycles have longer node counts than training'}}
132    with open('results.json','w') as h: json.dump(result,h,indent=2)
133    print(json.dumps(result,indent=2))
134if __name__=='__main__': main()