Finite-Horizon Lyapunov Risk Monitor / bench_run.py
Mechanism confirmed, baseline not beaten
1import sys, json, math, random
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, sweep_baseline, make_report
7
8SEEDS=(0,1,2,3,4,5,6,7)
9SWEEP_SEEDS=(0,1,2,3)
10Z=1.645
11
12class MatchedRNN(nn.Module):
13 # Matched recurrent architecture for both systems; only the loss differs.
14 def __init__(self, input_dim, out_dim, hidden=64):
15 super().__init__()
16 self.cell=nn.GRUCell(3, hidden)
17 self.head=nn.Linear(hidden, out_dim)
18 self.hidden=hidden
19 def forward(self, x, monitor=False, sigma=0.04, K=3):
20 seq=x.view(x.shape[0],-1,3); B,T,_=seq.shape
21 h=torch.zeros(B,self.hidden,device=x.device)
22 if monitor:
23 qs=torch.randn(K,B,self.hidden,device=x.device)
24 qs=qs/(qs.norm(dim=2,keepdim=True)+1e-8); sums=[]
25 sums=torch.zeros(K,B,device=x.device)
26 for t in range(T):
27 old=h.detach().requires_grad_(monitor)
28 h=self.cell(seq[:,t],old)
29 if monitor:
30 # Directional finite-difference JVP (cheap monitor; detached update).
31 newq=[]
32 eps=1e-3
33 for k in range(K):
34 vp=(self.cell(seq[:,t],old + eps*qs[k])-h)/eps
35 if sigma:
36 vp=vp + sigma*torch.randn_like(vp)*vp.detach().std().clamp_min(1e-4)
37 n=vp.norm(dim=1)+1e-8
38 newq.append(vp/(n[:,None])); sums[k]=sums[k]+n.log()
39 qs=torch.stack(newq)
40 out=self.head(h)
41 return (out,sums/T) if monitor else out
42
43def seed_all(s):
44 random.seed(s); np.random.seed(s); torch.manual_seed(s)
45 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
46
47def run_one(seed, cfg, risk):
48 seed_all(seed)
49 ds=get_dataset('dynamics', seed, n_train=2000, n_test=500)
50 dev='cuda' if torch.cuda.is_available() else 'cpu'
51 try:
52 torch.zeros(1,device=dev)
53 except Exception: dev='cpu'
54 net=MatchedRNN(np.prod(ds['input_shape']),ds['out_dim']).to(dev)
55 x=ds['xtr'].to(dev); y=ds['ytr'].to(dev)
56 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg.get('weight_decay',0.0))
57 bs=128
58 for ep in range(4):
59 net.train(); perm=torch.randperm(len(x),device=dev)
60 for j in range(0,len(x),bs):
61 ix=perm[j:j+bs]; result=net(x[ix],monitor=(risk=='ucb'),sigma=cfg.get('sigma',.04),K=3)
62 pred=result[0] if risk=='ucb' else result
63 loss=((pred-y[ix])**2).mean()
64 if risk=='ucb':
65 lam=result[1]; mu=lam.mean(); sd=lam.std(unbiased=True)
66 loss=loss+cfg['rho']*torch.relu(mu+Z*sd).square()
67 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
68 net.eval()
69 with torch.no_grad():
70 pred=net(ds['xte'].to(dev)); metric=float(((pred-ds['yte'].to(dev))**2).mean())
71 # Re-test behavior on trained models with fresh perturbations.
72 vals=[]
73 net.train()
74 with torch.enable_grad():
75 for j in range(0,len(ds['xte']),128):
76 _,lam=net(ds['xte'][j:j+128].to(dev),monitor=True,sigma=cfg.get('sigma',.04),K=8)
77 vals.append(lam.detach().cpu().numpy().ravel())
78 a=np.concatenate(vals); mu=float(a.mean()); sd=float(a.std(ddof=1))
79 return metric, {'metric':metric,'ftle_mean':mu,'ftle_sd':sd,'positive_fraction':float((a>0).mean()),'ucb':mu+Z*sd}
80
81def main():
82 # Search-space parity: every idea lr is included in the baseline sweep.
83 lrs=[0.001,0.003]
84 base_grid=[{'lr':lr,'weight_decay':0.0} for lr in lrs]
85 def base_fn(cfg): return lambda s: run_one(s,cfg,'base')[0]
86 base=sweep_baseline(base_fn,base_grid,seeds=SWEEP_SEEDS)
87 idea_grid=[{'lr':base['best_cfg']['lr'],'weight_decay':0.0,'rho':r,'sigma':.04} for r in (.05,.15,.30)]
88 idea_blocks=[]
89 for cfg in idea_grid:
90 per=[]; sig=[]
91 for s in SEEDS:
92 m,z=run_one(s,cfg,'ucb'); per.append(m); sig.append(z)
93 idea_blocks.append({'cfg':cfg,'res':{'mean':float(np.mean(per)),'std':float(np.std(per)),'per_seed':per,'n':len(per)},'sig':sig})
94 best=min(idea_blocks,key=lambda q:q['res']['mean'])
95 # Official baseline best re-evaluation is already produced by sweep_baseline full.
96 signature=[]
97 for z in best['sig']: signature.append(z)
98 mu=float(np.mean([z['ftle_mean'] for z in signature])); sd=float(np.mean([z['ftle_sd'] for z in signature]))
99 pred=.5*math.erfc(-mu/(math.sqrt(2)*sd)); obs=float(np.mean([z['positive_fraction'] for z in signature]))
100 extra={'prediction':'Gaussian p+=Phi(mu/sd) for finite-horizon FTLE','predicted_positive_fraction':pred,'observed_positive_fraction':obs,'absolute_error':abs(pred-obs),'confirmed':bool(abs(pred-obs)<=.10),'idea_sweep':[{k:v for k,v in b.items() if k!='sig'} for b in idea_blocks]}
101 rep=make_report('dynamics','rnn_small',base,best['res'],extra)
102 rep['official_architecture_note']='Matched GRUCell implementation exposes recurrent state for differentiable FTLE; baseline and idea share all parameters/optimizer/data.'
103 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
104 print(json.dumps(rep,indent=2))
105if __name__=='__main__': main()