Composed Trusted Reachable Families for Recurrent Networks / stage2_bench.py
Failed on benchmark
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()