Cyclic Lie-Bracket Residual Block / train_toy.py
Beats tuned baseline
1import json, random, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED=19
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8try:
9 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10 if device.type=='cuda':
11 torch.zeros(1,device=device)
12except Exception:
13 device=torch.device('cpu')
14
15N=2048
16rng=np.random.default_rng(SEED)
17X=torch.tensor(rng.uniform(-1,1,(N,2)),dtype=torch.float32)
18Y=(X[:,0:1]*X[:,1:2]+.25*torch.sin(2*X[:,0:1]-X[:,1:2]))
19train_slice=slice(0,1536); test_slice=slice(1536,None)
20
21class Field(nn.Module):
22 def __init__(self,d,h=32):
23 super().__init__(); self.net=nn.Sequential(nn.LayerNorm(d),nn.Linear(d,h),nn.GELU(),nn.Linear(h,d))
24 def forward(self,x): return self.net(x)
25
26class Baseline(nn.Module):
27 def __init__(self,d=16):
28 super().__init__(); self.inp=nn.Linear(2,d); self.f=Field(d); self.out=nn.Linear(d,1); self.gate=nn.Parameter(torch.tensor(.1))
29 def forward(self,x):
30 z=torch.tanh(self.inp(x)); return self.out(z+self.gate*self.f(z))
31
32class Cyclic(nn.Module):
33 def __init__(self,d=16, random_order=True):
34 super().__init__(); self.inp=nn.Linear(2,d); self.f1=Field(d); self.f2=Field(d); self.out=nn.Linear(d,1)
35 self.gate=nn.Parameter(torch.tensor(.1)); self.random_order=random_order
36 def forward(self,x):
37 z=torch.tanh(self.inp(x)); base=z
38 # Center fields pointwise, preserving noncommutativity while canceling
39 # their first-order average. Random centered run times have variance 1.
40 u=self.f1(z); v=self.f2(z); mean=(u+v)/2
41 fields=(u-mean,v-mean)
42 a=torch.randn((),device=z.device); b=torch.randn((),device=z.device)
43 if self.training and self.random_order and torch.rand((),device=z.device)<.5:
44 z=z+self.gate*.35*a*fields[1]; z=z+self.gate*.35*b*(self.f1(z)-self.f2(z))/2
45 else:
46 z=z+self.gate*.35*a*fields[0]; z=z+self.gate*.35*b*(self.f2(z)-self.f1(z))/2
47 return self.out(z)
48
49def train(kind):
50 torch.manual_seed(SEED)
51 model=Baseline() if kind=='baseline' else Cyclic()
52 model.to(device); xx=X.to(device); yy=Y.to(device)
53 opt=torch.optim.Adam(model.parameters(),lr=2e-3)
54 losses=[]; t=time.time()
55 for step in range(350):
56 model.train(); opt.zero_grad(); pred=model(xx[train_slice]); loss=((pred-yy[train_slice])**2).mean(); loss.backward(); opt.step()
57 losses.append(float(loss.detach().cpu()))
58 model.eval()
59 with torch.no_grad():
60 tr=float(((model(xx[train_slice])-yy[train_slice])**2).mean().cpu())
61 te=float(((model(xx[test_slice])-yy[test_slice])**2).mean().cpu())
62 return {'train_mse':tr,'test_mse':te,'loss_step_50':losses[49],'loss_step_350':losses[-1], 'seconds':time.time()-t, 'parameters':sum(p.numel() for p in model.parameters())}
63
64# Explicit order signal using the learned fields and shared input, with fixed
65# unit run times; this is separate from stochastic training forward passes.
66def order_signal(model):
67 if not isinstance(model,Cyclic): return None
68 model.eval(); z=torch.tanh(model.inp(X[test_slice].to(device)))
69 with torch.no_grad():
70 u=model.f1(z); v=model.f2(z); m=(u+v)/2; u=u-m; v=v-m; e=.1
71 ab=z+e*u; ab=ab+e*(model.f2(ab)-model.f1(ab))/2
72 ba=z+e*v; ba=ba+e*(model.f1(ba)-model.f2(ba))/2
73 return float((ab-ba).norm(dim=1).mean().cpu())
74
75if __name__=='__main__':
76 b=train('baseline'); c=train('cyclic')
77 # Retrain one cyclic model only to measure order gap consistently.
78 torch.manual_seed(SEED); cm=Cyclic().to(device); opt=torch.optim.Adam(cm.parameters(),lr=2e-3)
79 for _ in range(350):
80 opt.zero_grad(); l=((cm(X.to(device))[train_slice]-Y.to(device)[train_slice])**2).mean(); l.backward(); opt.step()
81 c['order_signal']=order_signal(cm)
82 out={'device':str(device),'baseline':b,'cyclic':c}
83 print(json.dumps(out,indent=2)); open('toy_results.json','w').write(json.dumps(out,indent=2))