Residual-Scenario Safety Training / bench_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import os, 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, make_model, train_model, sweep_baseline, make_report
  7
  8SEEDS = tuple(range(8)); TUNE = (0,1,2,3)
  9LRS = [1e-3, 3e-3, 1e-2]
 10MUS = [0.05, 0.2, 0.5]
 11EPOCHS = 12; BATCH = 128
 12
 13def seed_all(s):
 14    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 15    try:
 16        if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 17    except Exception: pass
 18
 19def get_residuals(ds, seed):
 20    seed_all(seed+10000)
 21    m=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
 22    m,_,_=train_model(m,ds,epochs=3,lr=3e-3,batch=BATCH,log=lambda *a:None)
 23    if m is None: return None
 24    dev=next(m.parameters()).device
 25    with torch.no_grad(): p=m(ds['xtr'].to(dev)).cpu()
 26    r=(ds['ytr']-p).reshape(-1)
 27    scale=max(float(r.std()),1e-5)
 28    return r.float(), scale
 29
 30def scenario_train(ds, seed, lr, mu):
 31    seed_all(seed); m=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
 32    rb=get_residuals(ds,seed)
 33    if rb is None: return float('nan'), {}
 34    residuals,scale=rb
 35    try:
 36        dev=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 37        m=m.to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev)
 38        r=residuals.to(dev); lo=float(y.min())-0.05; hi=float(y.max())+0.05
 39        opt=torch.optim.Adam(m.parameters(),lr=lr)
 40        gen=torch.Generator(device=dev); gen.manual_seed(seed+20000)
 41        for _ in range(EPOCHS):
 42            m.train(); perm=torch.randperm(len(x),device=dev)
 43            for j in range(0,len(x),BATCH):
 44                ix=perm[j:j+BATCH]; pred=m(x[ix]); nom=((pred-y[ix])**2).mean()
 45                ind=torch.randint(len(r),(8,len(ix)),generator=gen,device=dev)
 46                ys=pred[:,None,:]+r[ind].T[:,:,None]*scale
 47                h=torch.relu(lo-ys)+torch.relu(ys-hi)
 48                loss=nom+mu*h.mean(); opt.zero_grad(); loss.backward(); opt.step()
 49        m.eval()
 50        with torch.no_grad():
 51            metric=float(((m(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean())
 52            pred=m(x).cpu(); obs=(pred-ds['ytr']).reshape(-1)
 53            sg=torch.Generator().manual_seed(seed+30000)
 54            draws=obs[torch.randint(len(obs),(8,len(obs)),generator=sg)]
 55            yy=pred[:,None,:]+draws.T[:,:,None]
 56            slack=torch.relu(lo-yy)+torch.relu(yy-hi)
 57        return metric, {'residual_std':float(obs.std()),'scenario_violation_rate':float((slack>0).float().mean()),'mean_slack':float(slack.mean())}
 58    except RuntimeError:
 59        return float('nan'), {}
 60
 61def baseline_train(ds,seed,lr):
 62    seed_all(seed); m=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
 63    try:
 64        m,metric,_=train_model(m,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *a:None)
 65        if m is None:return float('nan'),{}
 66        dev=next(m.parameters()).device
 67        with torch.no_grad():
 68            pred=m(ds['xtr'].to(dev)).cpu(); obs=(pred-ds['ytr']).reshape(-1)
 69            lo=float(ds['ytr'].min())-0.05; hi=float(ds['ytr'].max())+0.05
 70            sg=torch.Generator().manual_seed(seed+30000)
 71            draws=obs[torch.randint(len(obs),(8,len(obs)),generator=sg)]
 72            slack=torch.relu(lo-(pred[:,None,:]+draws.T[:,:,None]))+torch.relu((pred[:,None,:]+draws.T[:,:,None])-hi)
 73        return float(metric),{'residual_std':float(obs.std()),'scenario_violation_rate':float((slack>0).float().mean()),'mean_slack':float(slack.mean())}
 74    except RuntimeError:return float('nan'),{}
 75
 76def main():
 77    # Baseline grid includes every lr and every method-side mu value (mu is inert for MSE).
 78    grid=[{'lr':lr,'mu':mu} for lr in LRS for mu in [0.0]+MUS]
 79    def base_factory(cfg): return lambda s: baseline_train(get_dataset('dynamics',s,400,200),s,cfg['lr'])[0]
 80    base=sweep_baseline(base_factory,grid,seeds=TUNE)
 81    # Tune the idea on exactly the same four seeds, then only the selected config is scored full.
 82    idea_sweep=[]
 83    for lr in LRS:
 84        for mu in MUS:
 85            vals=[scenario_train(get_dataset('dynamics',s,400,200),s,lr,mu)[0] for s in TUNE]
 86            idea_sweep.append({'cfg':{'lr':lr,'mu':mu},'mean':float(np.nanmean(vals))})
 87    best=min(idea_sweep,key=lambda z:z['mean'])['cfg']
 88    rows=[]; vals=[]; sig=[]
 89    for s in SEEDS:
 90        ds=get_dataset('dynamics',s,400,200)
 91        b,bs=baseline_train(ds,s,base['best_cfg']['lr'])
 92        v,isig=scenario_train(ds,s,best['lr'],best['mu'])
 93        vals.append(v); sig.append((bs,isig))
 94    idea={'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':[float(v) for v in vals],'n':8,'best_cfg':best,'sweep':idea_sweep}
 95    br={'mean':float(np.mean([baseline_train(get_dataset('dynamics',s,400,200),s,base['best_cfg']['lr'])[0] for s in SEEDS])),'std':base['full']['std'],'per_seed':[float(baseline_train(get_dataset('dynamics',s,400,200),s,base['best_cfg']['lr'])[0]) for s in SEEDS],'n':8}
 96    # Use the already paired baseline values from rows, avoiding any synthetic signature.
 97    bvals=[]
 98    for s in SEEDS: bvals.append(baseline_train(get_dataset('dynamics',s,400,200),s,base['best_cfg']['lr'])[0])
 99    br={'mean':float(np.mean(bvals)),'std':float(np.std(bvals)),'per_seed':[float(v) for v in bvals],'n':8}
100    signature={'baseline_residual_std':float(np.mean([x[0]['residual_std'] for x in sig])),'idea_residual_std':float(np.mean([x[1]['residual_std'] for x in sig])),'baseline_mean_slack':float(np.mean([x[0]['mean_slack'] for x in sig])),'idea_mean_slack':float(np.mean([x[1]['mean_slack'] for x in sig])),'prediction':'scenario penalty reduces empirical residual-boundary slack','confirmed':bool(np.mean([x[1]['mean_slack'] for x in sig]) < np.mean([x[0]['mean_slack'] for x in sig]))}
101    rep=make_report('dynamics','rnn_small',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':br},idea,{'mechanism_signature':signature,'track_justification':'Dynamics is structurally matched: the task is controlled pendulum rollout and the intervention constrains predicted outputs under empirical error scenarios.','idea_sweep':idea_sweep,'custom_track':None})
102    os.makedirs('artifacts',exist_ok=True)
103    json.dump(rep,open('artifacts/bench_report.json','w'),indent=2); print(json.dumps(rep,indent=2))
104if __name__=='__main__': main()