Lattice Error-Feedback Residual Blocks / lattice_bench.py

✓✓ Beats tuned baseline

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