Knieper Rollout Stability Metric / bench_run.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, sweep_baseline, evaluate, make_report
7
8SEEDS = tuple(range(8)); EPOCHS = 5; BATCH = 128; H = 8
9LRS = (0.0015, 0.003, 0.006); LAMBDAS = (0.01, 0.03, 0.10)
10CACHE = {}
11
12class DynamicsGRU(nn.Module):
13 def __init__(self, hidden=16):
14 super().__init__(); self.rnn=nn.GRU(3,hidden,batch_first=True); self.head=nn.Linear(hidden,1)
15 def forward(self,x,h=None,return_seq=False):
16 out,hn=self.rnn(x.view(x.shape[0],-1,3),h); pred=self.head(hn[-1])
17 return (pred,out,hn) if return_seq else pred
18
19def seed_all(s):
20 random.seed(s); np.random.seed(s); torch.manual_seed(s)
21
22def train_one(seed,lr,lam):
23 key=(int(seed),float(lr),float(lam))
24 if key in CACHE: return CACHE[key]
25 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
26 requested='cuda' if torch.cuda.is_available() else 'cpu'
27 for dev in ([requested,'cpu'] if requested=='cuda' else ['cpu']):
28 try:
29 model=DynamicsGRU().to(dev); xtr=ds['xtr'].to(dev); ytr=ds['ytr'].to(dev).view(-1,1)
30 opt=torch.optim.Adam(model.parameters(),lr=lr); mse=nn.MSELoss()
31 gen=torch.Generator(device=dev); gen.manual_seed(seed+10000)
32 for _ in range(EPOCHS):
33 order=torch.randperm(len(xtr),generator=gen,device=dev)
34 for st in range(0,len(xtr),BATCH):
35 ix=order[st:st+BATCH]; xb=xtr[ix]; task=mse(model(xb),ytr[ix])
36 if lam:
37 b=len(ix); h0=torch.zeros(1,b,16,device=dev)
38 eps=torch.randn(b,16,generator=gen,device=dev)*.02
39 _,a,_=model(xb,h0,True); _,bseq,_=model(xb,h0+eps.unsqueeze(0),True)
40 roll=(a-bseq).pow(2).sum(-1).sqrt().amax(1).div(eps.norm(dim=1)+1e-8).mean()
41 loss=task+lam*roll
42 else: loss=task
43 opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
44 model.eval()
45 with torch.no_grad():
46 xt=ds['xte'].to(dev); yt=ds['yte'].to(dev).view(-1,1)
47 metric=float(((model(xt)-yt)**2).mean())
48 b=len(xt); h0=torch.zeros(1,b,16,device=dev)
49 eps=torch.randn(b,16,generator=gen,device=dev)*.02
50 _,a,_=model(xt,h0,True); _,bb,_=model(xt,h0+eps.unsqueeze(0),True)
51 gain=float(((a-bb).pow(2).sum(-1).sqrt().amax(1)/(eps.norm(dim=1)+1e-8)).median())
52 CACHE[key]=(metric,gain); return CACHE[key]
53 except RuntimeError:
54 if dev=='cuda': continue
55 raise
56 raise RuntimeError('training failed')
57
58def fn(cfg): return lambda s: train_one(s,cfg['lr'],cfg.get('lambda',0.0))[0]
59
60def main():
61 # Baseline covers every LR used by the idea-side shared architecture.
62 base=sweep_baseline(fn,[{'lr':lr,'lambda':0.0} for lr in LRS])
63 best_lr=float(base['best_cfg']['lr'])
64 idea_runs=[]
65 for lam in LAMBDAS:
66 cfg={'lr':best_lr,'lambda':lam}; idea_runs.append((evaluate(fn(cfg),SEEDS),cfg))
67 idea,idea_cfg=min(idea_runs,key=lambda z:z[0]['mean'])
68 base_cfg=base['best_cfg']
69 sb=[train_one(s,base_cfg['lr'],0.0)[1] for s in SEEDS]
70 si=[train_one(s,idea_cfg['lr'],idea_cfg['lambda'])[1] for s in SEEDS]
71 red=1-float(np.mean(si))/float(np.mean(sb))
72 sig={'quantity':'median test-set G_H=max hidden separation / initial perturbation norm','H':H,
73 'predicted':'rollout penalty reduces finite-horizon gain','baseline_mean_GH':float(np.mean(sb)),
74 'idea_mean_GH':float(np.mean(si)),'observed_relative_reduction':float(red),
75 'predicted_vs_observed':{'predicted_direction':'decrease','observed_direction':'decrease' if red>0 else 'increase'},
76 'confirmed':bool(red>0.10),'paired_seed_GH_baseline':sb,'paired_seed_GH_idea':si,
77 'idea_config':idea_cfg,'track_choice_justification':'dynamics is the built-in control/stability track with multi-step pendulum rollouts.'}
78 rep=make_report('dynamics','rnn_small',base,idea,extra=sig)
79 rep['idea_sweep']=[{'cfg':c,'mean':r['mean'],'per_seed':r['per_seed']} for r,c in idea_runs]
80 rep['protocol_notes']={'epochs':EPOCHS,'n_train':400,'n_test':200,'baseline_lr_union':list(LRS),'idea_lambdas':list(LAMBDAS),'cache_used':True}
81 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
82 print(json.dumps(rep,indent=2))
83if __name__=='__main__': main()