ISS-Certified Sampled Optimizer Wrapper / stage2_iss_bench.py
Failed on benchmark
1import sys, json, copy, 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, evaluate, sweep_baseline, make_report
7
8SEEDS=tuple(range(8)); SWEEP=tuple(range(4)); EPOCHS=18; BATCH=128
9# Union of all learning rates: baseline and certified wrapper both see these.
10GRID=[{'lr':1e-3,'M':1},{'lr':3e-3,'M':1},{'lr':1e-2,'M':1}]
11IDEA_GRID=GRID
12
13def seed_all(s):
14 random.seed(s); np.random.seed(s); torch.manual_seed(s)
15
16def device_model(ds):
17 # train_model's fallback is not usable because this is a modified optimizer loop.
18 try:
19 dev=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
20 return make_model('rnn_small', ds['input_shape'], ds['out_dim']).to(dev),dev
21 except Exception:
22 return make_model('rnn_small', ds['input_shape'], ds['out_dim']).to('cpu'),torch.device('cpu')
23
24def hidden_energy(net,x):
25 got={}
26 def hook(mod, inp, out):
27 h=out[1]
28 got['v']=(h*h).mean()
29 h=net.rnn.register_forward_hook(hook)
30 try: net(x)
31 finally: h.remove()
32 return got.get('v', torch.tensor(0.,device=x.device))
33
34def train(seed,cfg,certified=False, collect=False):
35 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=160)
36 try: net,dev=device_model(ds)
37 except Exception: net=make_model('rnn_small',ds['input_shape'],ds['out_dim']); dev=torch.device('cpu')
38 xtr,ytr=[ds[k].to(dev) for k in ('xtr','ytr')]; xte,yte=[ds[k].to(dev) for k in ('xte','yte')]
39 lossf=nn.MSELoss(); opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
40 M=cfg.get('M',1); lam=0.02; c=1.0
41 accepted=rejected=checks=violations=0; cert_margins=[]; prev_batch=None
42 for ep in range(EPOCHS):
43 net.train(); perm=torch.randperm(len(xtr),device=dev); opt.zero_grad(set_to_none=True)
44 for bi,i in enumerate(range(0,len(xtr),BATCH)):
45 ix=perm[i:i+BATCH]; xb,yb=xtr[ix],ytr[ix]
46 loss=lossf(net(xb),yb); loss.backward()
47 if ((bi+1)%M and i+BATCH<len(xtr)): continue
48 # Save proposed state and old energy on a fixed current probe.
49 if certified:
50 probe=xb[:min(64,len(xb))]
51 net.eval()
52 with torch.no_grad(): vold=hidden_energy(net,probe).detach()
53 old={k:v.detach().clone() for k,v in net.state_dict().items()}
54 # Adam already has a proposal only after step; try decreasing step sizes.
55 accepted_this=False; base_lr=opt.param_groups[0]['lr']
56 for attempt in range(8):
57 opt.param_groups[0]['lr']=base_lr*(0.5**attempt)
58 opt.step(); opt.zero_grad(set_to_none=True)
59 with torch.no_grad(): vnew=hidden_energy(net,probe).detach()
60 # d is an observed input perturbation scale; certificate is ISS form.
61 d=(probe[:,3:]-probe[:,:-3]).pow(2).mean().sqrt().detach()
62 cert=vnew-vold+lam*vold-c*d*d
63 checks+=1; violations += int(float(cert)>0); cert_margins.append(float(cert))
64 if float(cert)<=0:
65 accepted+=1; accepted_this=True; break
66 net.load_state_dict(old)
67 opt.param_groups[0]['lr']=base_lr
68 if not accepted_this:
69 rejected+=1
70 # safe fallback: retain the previous parameters, i.e. zero control update.
71 opt.zero_grad(set_to_none=True)
72 else:
73 opt.step(); opt.zero_grad(set_to_none=True)
74 net.train()
75 net.eval()
76 with torch.no_grad(): metric=float(lossf(net(xte),yte))
77 if collect:
78 return metric, {'accepted':accepted,'rejected':rejected,'checks':checks,'violations':violations,
79 'violation_rate':violations/max(1,checks),'mean_certificate':float(np.mean(cert_margins)) if cert_margins else 0.0}
80 return metric
81
82def baseline_fn(cfg): return lambda seed: train(seed,cfg,False)
83def idea_fn(cfg): return lambda seed: train(seed,cfg,True)
84
85def main():
86 # Baseline sweep on four seeds, then full paired evaluation; idea 3-config sweep uses same union.
87 base=sweep_baseline(baseline_fn,GRID,seeds=SWEEP)
88 idea_sweep=[]
89 for cfg in IDEA_GRID:
90 r=evaluate(idea_fn(cfg),seeds=SWEEP); idea_sweep.append({'cfg':cfg,'mean':r['mean']})
91 best=min(idea_sweep,key=lambda z:z['mean'])['cfg']
92 idea=evaluate(idea_fn(best),seeds=SEEDS)
93 rep=make_report('dynamics','rnn_small',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':base['full']},idea,
94 {'prediction':'Lyapunov certificate filters sampled updates; rejected proposals should have positive certificate and accepted proposals nonpositive.',
95 'trained_model_measurements': [train(s,best,True,True)[1] for s in SEEDS],
96 'confirmed': False})
97 rep['idea']['sweep']=idea_sweep; rep['protocol_notes']='Baseline and idea share rnn_small, data, epochs, batch, Adam, and lr/M grid; only certificate rejection differs.'
98 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
99 print(json.dumps(rep,indent=2))
100if __name__=='__main__': main()