Cyclic Lie-Bracket Residual Block / train_toy.py

✓✓ Beats tuned baseline

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