Fermionic circuit message passing / fast_mvp.py
Beats tuned baseline
1import json, random, numpy as np, torch
2from torch import nn
3SEED=1481
4random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
5
6def wedge(x,y): return x[..., :,None]*y[...,None,:]-y[..., :,None]*x[...,None,:]
7def j_sign(j): return (-1)**(j*(j-1)//2)
8def local_scalar(degree):
9 return 0.0 if degree % 2 else 1.0
10
11def checks():
12 # Prediction 1: odd incident degree is annihilated exactly.
13 residual=[]
14 for n in range(1,10,2):
15 xs=torch.randn(n,4); residual.append(abs(local_scalar(n)))
16 # Prediction 2: surviving closed edge masks obey product_v (-1)^(j_v/2)=(-1)^|F|.
17 rows=[]
18 for E in range(1,9):
19 errs=[]; surviving=0
20 for mask in range(1<<E):
21 deg=[0]*E
22 for e in range(E):
23 if mask>>e&1: deg[e]+=1; deg[(e+1)%E]+=1
24 if any(x%2 for x in deg): continue
25 surviving+=1; local=np.prod([(-1)**(x//2) for x in deg]); rhs=(-1)**mask.bit_count(); errs.append(abs(local-rhs))
26 rows.append({'edges':E,'surviving_masks':int(surviving),'max_error':float(max(errs) if errs else 0)})
27 # Prediction 3: E||x wedge y||^2=2d(d-1) for iid N(0,1).
28 scale=[]
29 for d in [2,4,8,12]:
30 x=torch.randn(20000,d); y=torch.randn(20000,d)
31 val=(wedge(x,y)**2).sum( dim=(1,2)).mean().item(); pred=2*d*(d-1)
32 scale.append({'dimension':d,'predicted':pred,'observed':val,'ratio':val/pred})
33 return {'odd_degree_predicted_zero':{'predicted_max_residual':0,'observed_max_residual':max(residual)},'parity_sign':{'predicted_max_error':0,'sweep':rows},'wedge_energy_scaling':scale}
34
35def graph(n,cycle):
36 A=np.zeros((n,n),np.float32)
37 es=[(i,(i+1)%n) for i in range(n)] if cycle else [(i,i+1) for i in range(n-1)]
38 for i,j in es:A[i,j]=A[j,i]=1
39 return torch.tensor(A)
40class Net(nn.Module):
41 def __init__(self,ferm=False,d=4,h=12):
42 super().__init__(); self.ferm=ferm; self.d=d
43 self.inp=nn.Linear(1,h); self.layers=nn.ModuleList(); self.odd=nn.ModuleList(); self.mix=nn.ModuleList()
44 for _ in range(3):
45 self.layers.append(nn.Sequential(nn.Linear(h,h),nn.ReLU(),nn.Linear(h,h)))
46 if ferm:
47 self.odd.append(nn.Linear(h,d)); self.mix.append(nn.Sequential(nn.Linear(h+d*d,h),nn.ReLU(),nn.Linear(h,h)))
48 self.out=nn.Linear(h,2)
49 def forward(self,A):
50 h=self.inp(torch.ones(A.shape[0],1));
51 for k in range(3):
52 z=self.layers[k](h); agg=A@z
53 if self.ferm:
54 o=self.odd[k](h); # pairwise degree-2 exterior coefficient, with incident neighbor messages
55 # Sums of pair wedges: (sum o)(sum o)^T - sum(o o^T), equivalent to sum_{u<w} wedge.
56 so=A@o; pair=so[:,:,None]*so[:,None,:]-(A[:,:,None,None]*o[None,:, :,None]*o[None,:,None,:]).sum(1)
57 h=h+self.mix[k](torch.cat([agg,pair.reshape(A.shape[0],-1)],1))
58 else: h=h+agg
59 return self.out(h.mean(0))
60def run(ferm):
61 torch.manual_seed(SEED+(1 if ferm else 0)); model=Net(ferm); opt=torch.optim.Adam(model.parameters(),lr=.004)
62 train=[(graph(n,c),c) for n in range(6,11) for c in [0,1]]
63 for _ in range(25):
64 random.shuffle(train)
65 for A,y in train:
66 loss=nn.functional.cross_entropy(model(A)[None],torch.tensor([y])); opt.zero_grad(); loss.backward(); opt.step()
67 with torch.no_grad():
68 test=[(graph(n,c),c) for n in range(11,16) for c in [0,1]]
69 acc=np.mean([model(A).argmax().item()==y for A,y in test])
70 return float(acc)
71def main():
72 result={'seed':SEED,'checks':checks(),'classification':{'baseline_gin_accuracy':run(False),'fermionic_accuracy':run(True),'train_nodes':'6-10','test_nodes':'11-15'}}
73 open('results.json','w').write(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))
74if __name__=='__main__':main()