Composed Trusted Reachable Families for Recurrent Networks / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random, math
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  9
 10SEEDS=(0,1,2,3,4,5,6,7)
 11LR_GRID=[1e-3,3e-3,6e-3]
 12EPOCHS=12
 13BATCH=128
 14ALPHA=0.02
 15RADIUS=0.08
 16DELTA=1e-3
 17
 18def seed_all(s):
 19    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 20    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 21
 22def cell(net, h, x):
 23    """One GRU step, using the exact parameterization of bench.models.rnn_small."""
 24    r=net.rnn
 25    wi, wh = r.weight_ih_l0, r.weight_hh_l0
 26    bi = r.bias_ih_l0; bh = r.bias_hh_l0
 27    gi=F.linear(x,wi,bi); gh=F.linear(h,wh,bh)
 28    ir, iz, inn = gi.chunk(3,-1); hr, hz, hnn=gh.chunk(3,-1)
 29    rr=torch.sigmoid(ir+hr); zz=torch.sigmoid(iz+hz)
 30    nnv=torch.tanh(inn + rr*hnn)
 31    return (1-zz)*nnv + zz*h
 32
 33def nominal_and_penalty(net, x):
 34    """Propagate R_k=A_k R_k and penalize one-step nonlinear defect."""
 35    b=x.shape[0]; dev=x.device; hidden=net.rnn.hidden_size
 36    x=x.view(b,-1,3)
 37    h=torch.zeros(b,hidden,device=dev)
 38    # one fixed random direction per hidden dimension, normalized per sample
 39    d=torch.randn(b,hidden,device=dev)
 40    d=d/(d.norm(dim=1,keepdim=True)+1e-8)
 41    R=torch.ones(b,1,device=dev)
 42    total=0.0; maxv=0.0
 43    # scalar gamma with direction d; U=0, matching initial-state uncertainty
 44    for k in range(x.shape[1]):
 45        u=x[:,k,:]
 46        hn=cell(net,h,u)
 47        # directional Jacobian-vector estimate at nominal state; detached for stable monitor
 48        eps=DELTA
 49        ap=(cell(net,h+eps*d,u)-cell(net,h-eps*d,u))/(2*eps)
 50        ap=ap.detach()
 51        # R is scalar amplitude multiplying d; affine predicted next state
 52        pred=hn + R*ap
 53        pert=cell(net,h + RADIUS*R*d,u)
 54        defect=(pert-pred).pow(2).mean(dim=1)
 55        total=total+defect.mean()
 56        maxv=max(maxv,float(defect.sqrt().max().detach().cpu()))
 57        # propagate direction with local Jacobian-vector; keep nominal rollout graph
 58        R=(ap*d).norm(dim=1,keepdim=True).detach() + 1e-6
 59        h=hn
 60    return total/x.shape[1], maxv
 61
 62def idea_train(seed, lr):
 63    seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
 64    net=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
 65    device='cuda' if torch.cuda.is_available() else 'cpu'
 66    try:
 67        net=net.to(device); xtr,ytr=ds['xtr'].to(device),ds['ytr'].to(device)
 68        opt=torch.optim.Adam(net.parameters(),lr=lr)
 69        lossf=nn.MSELoss()
 70        for _ in range(EPOCHS):
 71            net.train(); perm=torch.randperm(len(xtr),device=device)
 72            for i in range(0,len(xtr),BATCH):
 73                ix=perm[i:i+BATCH]; xb=xtr[ix]; yb=ytr[ix]
 74                pred=net(xb); task=lossf(pred,yb)
 75                reach,_=nominal_and_penalty(net,xb)
 76                loss=task+ALPHA*reach
 77                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
 78        net.eval()
 79        with torch.no_grad(): metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean().cpu())
 80        return metric
 81    except RuntimeError:
 82        # CPU retry is deliberately independent and deterministic
 83        seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
 84        net=make_model('rnn_small',ds['input_shape'],ds['out_dim']).cpu(); xtr,ytr=ds['xtr'],ds['ytr']
 85        opt=torch.optim.Adam(net.parameters(),lr=lr)
 86        for _ in range(EPOCHS):
 87            perm=torch.randperm(len(xtr))
 88            for i in range(0,len(xtr),BATCH):
 89                ix=perm[i:i+BATCH]; task=((net(xtr[ix])-ytr[ix])**2).mean(); reach,_=nominal_and_penalty(net,xtr[ix]); loss=task+ALPHA*reach
 90                opt.zero_grad(); loss.backward(); opt.step()
 91        with torch.no_grad(): return float(((net(ds['xte'])-ds['yte'])**2).mean())
 92
 93def baseline_fn(cfg):
 94    def run(seed):
 95        seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
 96        net=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
 97        _,m,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a:None)
 98        return m
 99    return run
100
101def signature(seed, lr):
102    # Re-test v(r) on a trained model, not on the toy analytic recurrence.
103    seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
104    net=make_model('rnn_small',ds['input_shape'],ds['out_dim']); device='cuda' if torch.cuda.is_available() else 'cpu'
105    try: net=net.to(device)
106    except Exception: net=net.cpu(); device='cpu'
107    # short canonical training for signature
108    x,y=ds['xtr'][:64].to(device),ds['ytr'][:64].to(device); opt=torch.optim.Adam(net.parameters(),lr=lr)
109    for _ in range(EPOCHS):
110        task=((net(x)-y)**2).mean(); opt.zero_grad(); task.backward(); opt.step()
111    net.eval(); xx=ds['xte'][:16].to(device); b=xx.shape[0]; h=torch.zeros(b,64,device=device); d=torch.ones_like(h); d=d/d.norm(dim=1,keepdim=True); R=torch.ones(b,1,device=device)
112    vals=[]
113    with torch.no_grad():
114      for rad in [0.02,0.04,0.08]:
115        h=torch.zeros(b,64,device=device); R=torch.ones(b,1,device=device); vmax=0.
116        for k in range(xx.shape[1] if xx.dim()>2 else 8):
117          u=xx[:,k,:] if xx.dim()>2 else xx[:,3*k:3*k+3]; hn=cell(net,h,u); eps=DELTA
118          ap=(cell(net,h+eps*d,u)-cell(net,h-eps*d,u))/(2*eps); pert=cell(net,h+rad*R*d,u); vmax=max(vmax,float((pert-(hn+rad*R*ap)).norm(dim=1).max().cpu())); R=(ap*d).norm(dim=1,keepdim=True)+1e-6; h=hn
119        vals.append(vmax)
120    slope=float(np.polyfit(np.log([.02,.04,.08]),np.log(np.maximum(vals,1e-12)),1)[0])
121    return {'radii':[.02,.04,.08],'observed_violation':vals,'observed_log_slope':slope,'predicted_slope':2.0,'confirmed':bool(1.5<slope<2.5)}
122
123def main():
124    # Union parity: baseline evaluates every lr considered by idea.
125    grid=[{'lr':v} for v in LR_GRID]
126    base=sweep_baseline(baseline_fn,grid,seeds=(0,1,2,3))
127    idea_grid=LR_GRID
128    idea={}
129    for lr in idea_grid:
130        idea[lr]=evaluate(lambda s,lr=lr: idea_train(s,lr),seeds=SEEDS)
131    best_lr=min(idea,key=lambda z:idea[z]['mean']); idea_res=idea[best_lr]
132    rep=make_report('dynamics','rnn_small',base,idea_res,{'method':'multi_step_affine_reachability_penalty','alpha':ALPHA,'signature':signature(0,best_lr),'idea_lr_results':{str(k):v for k,v in idea.items()}})
133    rep['protocol_notes']={'structural_match':'dynamics recurrent/control track','same_architecture':True,'baseline_lr_union':LR_GRID,'idea_lr_union':LR_GRID,'epochs':EPOCHS,'train_samples':400,'test_samples':200}
134    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
135if __name__=='__main__': main()