Lyapunov-Calibrated Multiplicative Noise / bench_run.py
Mechanism confirmed, baseline not beaten
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