Lattice Error-Feedback Residual Blocks / lattice_bench.py
Beats tuned baseline
1import json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEEDS=list(range(8)); DEVICE='cuda' if torch.cuda.is_available() else 'cpu'
8class ResidualRNN(nn.Module):
9 def __init__(self,hidden=24,depth=8):
10 super().__init__(); self.inp=nn.Linear(2,hidden); self.blocks=nn.ModuleList([nn.Sequential(nn.Linear(hidden,hidden),nn.Tanh(),nn.Linear(hidden,hidden)) for _ in range(depth)]); self.out=nn.Linear(hidden,1)
11 def forward(self,x,mode='baseline',h=.05,diagnose=False):
12 z=torch.tanh(self.inp(x)); carry=torch.zeros_like(z); carries=[]; deltas=[]; qs=[]; sat=0
13 for block in self.blocks:
14 d=.15*block(z); deltas.append(d)
15 if mode=='feedback':
16 u=d+carry; q=torch.round(u/h)*h; carry=u-q; z=z+q; carries.append(carry); qs.append(q); sat+=int((u.abs()>6.35*h).sum().item())
17 else:
18 q=torch.round((z+d)/h)*h-z; z=z+q
19 return self.out(z),carries,sat,(deltas,qs) if diagnose else None
20
21def data(seed,n=400):
22 rng=np.random.default_rng(seed); x=rng.uniform(-1,1,(n,2)).astype('float32'); y=(.92*x[:,0]+.18*np.sin(x[:,1])+.10*x[:,1]).astype('float32')[:,None]; return x[:320],y[:320],x[320:],y[320:]
23def run(seed,mode,lr,epochs=35,h=.05):
24 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed); model=ResidualRNN().to(DEVICE); opt=torch.optim.Adam(model.parameters(),lr=lr); xt,yt,xe,ye=data(seed); xt,yt,xe,ye=[torch.from_numpy(a).to(DEVICE) for a in (xt,yt,xe,ye)]
25 for _ in range(epochs):
26 opt.zero_grad(); pred,_,_,_=model(xt,mode,h); ((pred-yt)**2).mean().backward(); opt.step()
27 with torch.no_grad():
28 pred,cs,sat,_=model(xe,mode,h); mse=((pred-ye)**2).mean().item(); cr=float(torch.stack(cs).abs().mean().item()) if cs else 0.; mx=float(torch.stack(cs).abs().max().item()) if cs else 0.
29 return {'seed':seed,'mse':mse,'mean_abs_carry':cr,'max_abs_carry':mx,'saturation':sat}
30def perm_p(a,b):
31 d=np.asarray(a)-np.asarray(b); obs=float(d.mean()); rng=np.random.default_rng(991); cnt=1
32 for _ in range(9999): cnt+=int(float((d*rng.choice([-1,1],len(d))).mean())<=obs)
33 return cnt/10000
34
35def main():
36 grid=[1e-3,3e-3,1e-2]; base={}; idea={}
37 for lr in grid: base[lr]=[run(s,'baseline',lr) for s in SEEDS]; idea[lr]=[run(s,'feedback',lr) for s in SEEDS]
38 bmeans={str(k):float(np.mean([r['mse'] for r in v])) for k,v in base.items()}; imeans={str(k):float(np.mean([r['mse'] for r in v])) for k,v in idea.items()}; bestb=min(bmeans,key=bmeans.get); besti=min(imeans,key=imeans.get); br=base[float(bestb)]; ir=idea[float(besti)]; bm=np.array([r['mse'] for r in br]); im=np.array([r['mse'] for r in ir])
39 # Re-test the conservation relation on trained idea models and actual validation inputs.
40 residuals=[]; bounds=[]
41 for s in SEEDS:
42 random.seed(s); np.random.seed(s); torch.manual_seed(s); model=ResidualRNN().to(DEVICE); opt=torch.optim.Adam(model.parameters(),lr=float(besti)); xt,yt,xe,ye=data(s); xt,xe=[torch.from_numpy(a).to(DEVICE) for a in (xt,xe)]
43 for _ in range(35): opt.zero_grad(); p,_,_,_=model(xt,'feedback',.05); ((p-torch.from_numpy(data(s)[1]).to(DEVICE))**2).mean().backward(); opt.step()
44 with torch.no_grad(): _,cs,_,dq=model(xe,'feedback',.05,True); ds=torch.stack(dq[0]); qs=torch.stack(dq[1]); c=cs[-1]; residuals.append(float((qs.sum(0)-ds.sum(0)+c).abs().mean().item())); bounds.append(float(c.abs().max().item()/.05))
45 sig={'prediction':'bounded carry and telescoping discrepancy at NN scale','observed_conservation_residual_mean':float(np.mean(residuals)),'observed_max_carry_over_h':float(np.max(bounds)),'observed_saturation_total':int(sum(r['saturation'] for r in ir)),'confirmed':bool(np.mean(residuals)<1e-5 and np.max(bounds)<=.5+1e-5)}
46 report={'track':'dynamics','official_bench_available':False,'baseline_sweep':bmeans,'idea_sweep':imeans,'baseline_best_lr':float(bestb),'idea_best_lr':float(besti),'baseline_per_seed':br,'idea_per_seed':ir,'paired_delta_mean':float((im-bm).mean()),'permutation_p_value':float(perm_p(im,bm)),'mechanism_signature':sig,'custom_track':None,'device':DEVICE}; Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
47if __name__=='__main__': main()