Tiny Local Recurrence with Adaptive Computation / bench_experiment_corrected.py
Unverified
1import json, random, sys
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, train_model, sweep_baseline, make_report, count_params
7SEEDS=tuple(range(8)); SWEEP=tuple(range(4)); EPOCHS=15; BATCH=128; LRS=[1e-3,3e-3,6e-3]
8kept={}
9class FixedResidual(nn.Module):
10 def __init__(self,d=64,depth=8):
11 super().__init__(); self.enc=nn.Linear(24,d); self.blocks=nn.ModuleList([nn.Sequential(nn.LayerNorm(d),nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)) for _ in range(depth)]); self.head=nn.Linear(d,1)
12 def forward(self,x):
13 s=self.enc(x)
14 for b in self.blocks: s=s+.25*b(s)
15 return self.head(s)
16class AdaptiveResidual(nn.Module):
17 def __init__(self,d=64,tmax=8,alpha=.25):
18 super().__init__(); self.enc=nn.Linear(24,d); self.norm=nn.LayerNorm(d); self.rule=nn.Sequential(nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)); self.halt=nn.Linear(d,1); self.head=nn.Linear(d,1); self.tmax=tmax; self.alpha=alpha
19 def forward_stats(self,x):
20 s=self.enc(x); acc=torch.zeros_like(s); mass=torch.zeros(x.shape[0],1,device=x.device); steps=torch.zeros_like(mass)
21 for _ in range(self.tmax):
22 s=s+self.alpha*self.rule(self.norm(s)); h=torch.sigmoid(self.halt(s)); delta=torch.minimum(h,1-mass); acc=acc+delta*s; steps=steps+delta*(1+steps); mass=mass+delta
23 acc=acc+(1-mass)*s
24 return self.head(acc),steps.squeeze(1),mass.squeeze(1)
25 def forward(self,x): return self.forward_stats(x)[0]
26def seed(s):
27 random.seed(s); np.random.seed(s); torch.manual_seed(s)
28 try:
29 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
30 except Exception: pass
31def ds(s): return get_dataset('dynamics',s,n_train=4000,n_test=1000)
32def base_run(cfg,keep=False):
33 def f(s):
34 seed(s); d=ds(s); m=FixedResidual(); n,v,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
35 if keep: kept[('b',s)]=(n,d)
36 return float(v)
37 return f
38def idea_run(cfg,keep=False):
39 def f(s):
40 seed(s); d=ds(s); m=AdaptiveResidual(tmax=cfg['tmax']); n,v,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
41 if keep: kept[('i',s)]=(n,d)
42 return float(v)
43 return f
44def ev(vals): return {'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':[float(x) for x in vals],'n':len(vals)}
45def main():
46 grid=[{'lr':x} for x in LRS]
47 b=sweep_baseline(base_run,grid,seeds=SWEEP)
48 ig=[]
49 for lr in LRS:
50 v=[idea_run({'lr':lr,'tmax':8})(s) for s in SWEEP]; ig.append({'cfg':{'lr':lr,'tmax':8},'mean':float(np.mean(v)),'per_seed':v})
51 best=min(ig,key=lambda z:z['mean'])['cfg']
52 bv=[base_run(b['best_cfg'],True)(s) for s in SEEDS]; iv=[idea_run(best,True)(s) for s in SEEDS]
53 b['full']=ev(bv); ir=ev(iv); ir['chosen_cfg']=best; ir['sweep']=ig
54 with torch.no_grad():
55 steps=[]; masses=[]
56 for s in SEEDS:
57 n,d=kept[('i',s)]; dev=next(n.parameters()).device; _,st,ma=n.forward_stats(d['xte'].to(dev)); steps.append(float(st.mean().cpu())); masses.append(float(ma.mean().cpu()))
58 sig={'prediction':'shared adaptive recurrence performs fewer than Tmax weighted updates on trained inputs','predicted':{'weighted_steps':'< 8'},'observed':{'mean_weighted_steps':float(np.mean(steps)),'mean_mass':float(np.mean(masses)),'per_seed_steps':steps},'confirmed':bool(np.mean(steps)<8),'parameter_counts':{'baseline':count_params(kept[('b',0)][0]),'idea':count_params(kept[('i',0)][0])}}
59 r=make_report('dynamics','rnn_small',b,ir,sig); r['corrected_protocol']=True
60 open('bench_report_corrected.json','w').write(json.dumps(r,indent=2)); print(json.dumps(r,indent=2))
61if __name__=='__main__': main()