Finite-Horizon Lyapunov Risk Monitor / bench_run.py

Mechanism confirmed, baseline not beaten

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