ISS-Certified Sampled Optimizer Wrapper / stage2_iss_bench.py

Failed on benchmark

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