Knieper Rollout Stability Metric / bench_run.py

Failed on benchmark

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