Lyapunov-Calibrated Multiplicative Noise / bench_run.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys,json,random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 7from bench import get_dataset,make_report
 8from bench.protocol import evaluate,sweep_baseline
 9DEVICE=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10
11class ResidualRNN(nn.Module):
12    def __init__(self,mode='fixed',fixed_q=.03,kappa=.7,dt=.1,target=0.,hidden=48):
13        super().__init__(); self.mode=mode; self.fixed_q=fixed_q; self.kappa=kappa; self.dt=dt; self.target=target
14        self.inp=nn.Linear(3,hidden); self.rec=nn.Linear(hidden,hidden); self.head=nn.Linear(hidden,1); self.last=[]
15    def transition(self,h,x):
16        return torch.tanh(self.inp(x)+self.rec(h))
17    def estimate_r(self,h,x):
18        v=torch.ones_like(h); v=v/(torch.linalg.vector_norm(v,dim=1,keepdim=True)+1e-8)
19        h0=h.detach().requires_grad_(True); y=self.transition(h0,x.detach())
20        jv=torch.autograd.grad(y,h0,v,retain_graph=False,create_graph=False)[0]
21        gain=torch.linalg.vector_norm(v+jv,dim=1).clamp(1e-6,10.)
22        return torch.log(gain).mean().detach().clamp(-5.,2.)
23    def forward(self,x,collect=False):
24        seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.rec.out_features,device=x.device); obs=[]
25        for t in range(seq.shape[1]):
26            h=h.nan_to_num(0.).clamp(-5.,5.); f=self.transition(h,seq[:,t].clamp(-5.,5.))
27            r=self.estimate_r(h,seq[:,t]) if (self.mode=='adaptive' or collect) else torch.tensor(0.,device=x.device)
28            q=(2*self.kappa*torch.relu(r-self.target)*self.dt).clamp(0,.12) if self.mode=='adaptive' else torch.tensor(self.fixed_q,device=x.device)
29            u=torch.sigmoid(h).clamp(.001,.999); pv=(q*u*(1-u)).clamp(0,.12)
30            z=1.+torch.sqrt(pv)*torch.randn_like(h); h=(z*f).clamp(-5.,5.)
31            if collect: obs.append((float(r.cpu()),float(q.cpu()),float(pv.mean().detach().cpu()),float(z.mean().detach().cpu()),float((z*z).mean().detach().cpu())))
32        if collect:self.last=obs
33        return self.head(h.clamp(-5.,5.))
34
35def train(seed,mode,lr,knob,device,return_net=False):
36    random.seed(seed);np.random.seed(seed);torch.manual_seed(seed)
37    ds=get_dataset('dynamics',seed,n_train=400,n_test=200); net=ResidualRNN(mode=mode,fixed_q=knob if mode=='fixed' else 0.,kappa=knob if mode=='adaptive' else .7).to(device)
38    opt=torch.optim.Adam(net.parameters(),lr=lr); xtr,ytr=ds['xtr'].to(device),ds['ytr'].to(device)
39    for ep in range(10):
40        net.train(); perm=torch.randperm(len(xtr),device=device)
41        for i in range(0,len(perm),64):
42            ix=perm[i:i+64]; loss=((net(xtr[ix])-ytr[ix])**2).mean(); opt.zero_grad(set_to_none=True); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.); opt.step()
43    net.eval(); xte,yte=ds['xte'].to(device),ds['yte'].to(device); oldm,oldq=net.mode,net.fixed_q; net.mode='none';net.fixed_q=0.
44    with torch.no_grad(): metric=float(((net(xte)-yte)**2).mean().cpu())
45    net.mode,net.fixed_q=oldm,oldq
46    return (metric,net,xte[:64]) if return_net else metric
47
48def mk(c,mode,device):return lambda s:train(s,mode,float(c['lr']),float(c['knob']),device)
49def main(device):
50    lrs=[.001,.003,.006]; bg=[{'lr':l,'knob':q} for l in lrs for q in [.01,.03,.06]]; ig=[{'lr':l,'knob':k} for l in lrs for k in [.35,.7,1.4]]
51    base=sweep_baseline(lambda c:mk(c,'fixed',device),bg); tried=[{'cfg':c,'mean':evaluate(mk(c,'adaptive',device),seeds=(0,1,2,3))['mean']} for c in ig]; best=min(tried,key=lambda z:z['mean'])['cfg']; idea=evaluate(mk(best,'adaptive',device))
52    # Re-test the predicted variance law on trained models using repeated stochastic passes.
53    rows=[]
54    for s in range(8):
55        _,net,x=train(s,'adaptive',best['lr'],best['knob'],device,True); moments=[]
56        with torch.enable_grad():
57            for _ in range(16): net(x,collect=True); moments.append(np.asarray(net.last,dtype=float))
58        a=np.asarray(moments); pred=np.nanmean(a[:,:,2],axis=0); ez=np.nanmean(a[:,:,3],axis=0); ez2=np.nanmean(a[:,:,4],axis=0); emp=ez2-ez*ez
59        for p,e in zip(pred,emp):
60            if np.isfinite(p) and np.isfinite(e) and p>1e-10: rows.append((p,max(e,0.)))
61    ar=np.asarray(rows); slope=float(np.dot(ar[:,0],ar[:,1])/np.dot(ar[:,0],ar[:,0])) if len(ar) else float('nan'); sig={'trained_model':'ResidualRNN dynamics forecaster','n_observations':int(len(ar)),'predicted_conditional_gate_variance':'q*u*(1-u), q=2*kappa*dt*[r-r_target]+','observed_variance_fit_slope':slope,'expected_slope':1.0,'relative_error':abs(slope-1.) if np.isfinite(slope) else None,'confirmed':bool(np.isfinite(slope) and abs(slope-1.)<.25)}
62    rep=make_report('dynamics','rnn_small',base,idea,{'idea_sweep':tried,'mechanism_signature':sig}); rep['custom_track']=None; Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps({'device':str(device),'best_baseline':base['best_cfg'],'best_idea':best,'comparison':rep['comparison'],'signature':sig},indent=2))
63if __name__=='__main__':
64    try:main(DEVICE)
65    except RuntimeError as e:
66        if DEVICE.type=='cuda': print('CUDA fallback:',str(e)[:160]);torch.cuda.empty_cache();main(torch.device('cpu'))
67        else:raise