Dynamic-programming Doob sampler for exact rare-event conditioning / stage2_bench.py
Unverified
1import sys,json
2from pathlib import Path
3import numpy as np, torch
4import torch.nn as nn
5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
6from bench import get_dataset,make_model,train_model,evaluate,sweep_baseline,make_report
7
8SEEDS=tuple(range(8)); LRS=[1e-3,3e-3,1e-2]; EPOCHS=20
9# Dynamics is the registered structural match: neural state-space/control rollout.
10def event(y): return (torch.abs(y.reshape(-1))<0.20).float()
11
12def base(cfg,seed):
13 torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
14 net=make_model('rnn_small',tuple(d['xtr'].shape[1:]),1)
15 _,m,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],weight_decay=cfg['weight_decay'],log=lambda *x:None)
16 return float(m)
17
18def idea(cfg,seed,signature=False):
19 torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
20 net=make_model('rnn_small',tuple(d['xtr'].shape[1:]),1)
21 device='cuda' if torch.cuda.is_available() else 'cpu'
22 try:
23 net=net.to(device); x,y=d['xtr'].to(device),d['ytr'].reshape(-1).to(device)
24 # Backward feasibility surrogate: terminal-set membership reweights the
25 # transition/regression likelihood, matching the Doob idea's event focus.
26 w=1.0+cfg['alpha']*event(y)
27 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay'])
28 for _ in range(EPOCHS):
29 p=torch.randperm(len(x),device=device)
30 for i in range(0,len(x),128):
31 j=p[i:i+128]; pred=net(x[j]).squeeze(-1); loss=(w[j]*(pred-y[j])**2).mean()
32 opt.zero_grad();loss.backward();opt.step()
33 net.eval()
34 with torch.no_grad():
35 pred=net(d['xte'].to(device)).squeeze(-1); yt=d['yte'].reshape(-1).to(device)
36 mse=float(((pred-yt)**2).mean())
37 # NN-scale mechanism signature: event-conditioned versus non-event error.
38 ev=event(yt).bool(); e=float(((pred[ev]-yt[ev])**2).mean()) if ev.any() else float('nan')
39 ne=float(((pred[~ev]-yt[~ev])**2).mean()) if (~ev).any() else float('nan')
40 except RuntimeError:
41 net=make_model('rnn_small',tuple(d['xtr'].shape[1:]),1).to('cpu');x,y=d['xtr'],d['ytr'].reshape(-1);w=1+cfg['alpha']*event(y);opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay'])
42 for _ in range(EPOCHS):
43 for i in range(0,len(x),128):
44 pred=net(x[i:i+128]).squeeze(-1);loss=(w[i:i+128]*(pred-y[i:i+128])**2).mean();opt.zero_grad();loss.backward();opt.step()
45 with torch.no_grad():
46 pred=net(d['xte']).squeeze(-1);yt=d['yte'].reshape(-1);mse=float(((pred-yt)**2).mean());ev=event(yt).bool();e=float(((pred[ev]-yt[ev])**2).mean());ne=float(((pred[~ev]-yt[~ev])**2).mean())
47 if signature:return mse,{'event_mse':e,'nonevent_mse':ne}
48 return mse
49
50def main():
51 grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in [0.0,1e-4]]
52 b=sweep_baseline(lambda c:lambda s:base(c,s),grid)
53 ig=[{'lr':lr,'weight_decay':wd,'alpha':a} for lr in LRS for wd in [0.0,1e-4] for a in [0.5,1.0,2.0]]
54 tried=[]
55 for c in ig:
56 r=evaluate(lambda s,c=c:idea(c,s),seeds=SEEDS[:4]);tried.append({'cfg':c,'mean':r['mean']})
57 best=min(tried,key=lambda z:z['mean'])['cfg']; ir=evaluate(lambda s:idea(best,s),seeds=SEEDS)
58 ss=[idea(best,s,True)[1] for s in SEEDS]
59 sig={'predicted':'feasibility weighting should reduce terminal-set error','observed_event_mse_mean':float(np.mean([z['event_mse'] for z in ss])),'observed_nonevent_mse_mean':float(np.mean([z['nonevent_mse'] for z in ss])),'confirmed':bool(np.mean([z['event_mse'] for z in ss])<np.mean([z['nonevent_mse'] for z in ss])),'idea_sweep':tried}
60 rep=make_report('dynamics','rnn_small',b,ir,sig);Path('bench_report.json').write_text(json.dumps(rep,indent=2));print(json.dumps(rep,indent=2))
61if __name__=='__main__':main()