Fermionic circuit message passing / fermionic_mvp.py
Beats tuned baseline
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()