Fermionic circuit message passing / fast_mvp.py

✓✓ Beats tuned baseline

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