Covariance-Conditioned Neural Rollouts / bench_runner.py
Mechanism confirmed, baseline not beaten
1import json, sys
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_model, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS=tuple(range(8)); GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':6e-3}]
10EPOCHS=18; BATCH=128; NTRAIN=1600; NTEST=500
11
12def seed_all(s):
13 np.random.seed(s); torch.manual_seed(s)
14 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
15
16def base_fn(cfg):
17 def run(seed):
18 seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST)
19 net=make_model('rnn_small',d['input_shape'],1)
20 _,m,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a:None)
21 return float(m)
22 return run
23
24class GaussianGRU(nn.Module):
25 def __init__(self):
26 super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,2)
27 def forward(self,x):
28 q=x.view(x.shape[0],-1,3)
29 try: _,h=self.rnn(q)
30 except RuntimeError:
31 old=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False
32 try: _,h=self.rnn(q)
33 finally: torch.backends.cudnn.enabled=old
34 z=self.head(h[-1]); return z[:,0:1], z[:,1:2]
35
36def idea_fn(cfg, collect=False):
37 def run(seed):
38 seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST)
39 net=GaussianGRU(); device='cuda' if torch.cuda.is_available() else 'cpu'
40 try:
41 net=net.to(device); x=d['xtr'].to(device); y=d['ytr'].to(device); xt=d['xte'].to(device); yt=d['yte'].to(device)
42 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); n=len(x)
43 for ep in range(EPOCHS):
44 p=torch.randperm(n,device=device)
45 net.train()
46 for j in range(0,n,BATCH):
47 ix=p[j:j+BATCH]; mu,lv=net(x[ix]); lv=lv.clamp(-6,4)
48 loss=0.5*(lv+(y[ix]-mu).square()*torch.exp(-lv)).mean()
49 opt.zero_grad(); loss.backward(); opt.step()
50 net.eval()
51 with torch.no_grad():
52 mu,lv=net(xt); var=torch.exp(lv).clamp_min(1e-6); mse=(mu-yt).square().mean().item()
53 nll=(0.5*(torch.log(2*torch.pi*var)+(yt-mu).square()/var)).mean().item()
54 cov=((yt-mu).square()<=5.9915*var).float().mean().item()
55 if collect: return mse, {'nll':nll,'coverage95':cov,'mean_var':var.mean().item()}
56 return mse
57 except RuntimeError:
58 # CPU fallback on any CUDA/runtime failure.
59 net=GaussianGRU().cpu(); x=d['xtr']; y=d['ytr']; xt=d['xte']; yt=d['yte']; opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); n=len(x)
60 for ep in range(EPOCHS):
61 for j in range(0,n,BATCH):
62 mu,lv=net(x[j:j+BATCH]); lv=lv.clamp(-6,4); loss=0.5*(lv+(y[j:j+BATCH]-mu).square()*torch.exp(-lv)).mean(); opt.zero_grad(); loss.backward(); opt.step()
63 with torch.no_grad():
64 mu,lv=net(xt); var=torch.exp(lv).clamp_min(1e-6); mse=(mu-yt).square().mean().item(); nll=(0.5*(torch.log(2*torch.pi*var)+(yt-mu).square()/var)).mean().item(); cov=((yt-mu).square()<=5.9915*var).float().mean().item()
65 return (mse,{'nll':nll,'coverage95':cov,'mean_var':var.mean().item()}) if collect else mse
66 return run
67
68def main():
69 # Required cheap numerical Schur-complement check.
70 rng=np.random.default_rng(988); a=rng.normal(size=(9,9)); S=a@a.T+.2*np.eye(9); A=S[:4,:4]; B=S[4:,:4]; C=S[4:,4:]; sch=C-B@np.linalg.solve(A,B.T)
71 math_check={'joint_min_eig':float(np.linalg.eigvalsh(S).min()),'schur_min_eig':float(np.linalg.eigvalsh(sch).min()),'psd':bool(np.linalg.eigvalsh(sch).min()>-1e-10)}
72 base=sweep_baseline(base_fn,GRID,seeds=(0,1,2,3))
73 # full baseline at best config, then idea at same union; choose best idea mean.
74 ideas=[]
75 for cfg in GRID:
76 r=evaluate(idea_fn(cfg),seeds=SEEDS); ideas.append((r,cfg))
77 best_idea,best_cfg=min(ideas,key=lambda z:z[0]['mean'])
78 # trained-model signature on seed 0, measured independently from model behavior.
79 _,sig=idea_fn(best_cfg,collect=True)(0)
80 rep=make_report('dynamics','rnn_small',base,best_idea,extra={'mechanism_signature':{'prediction':'joint Gaussian head yields calibrated predictive variance; 95% interval coverage near nominal 0.95','observed_nll':sig['nll'],'observed_coverage95':sig['coverage95'],'observed_mean_variance':sig['mean_var'],'confirmed':bool(abs(sig['coverage95']-.95)<.10)},'math_check':math_check,'idea_sweep':[{'cfg':c,'result':r} for r,c in ideas],'budget':{'epochs':EPOCHS,'batch':BATCH,'n_train':NTRAIN,'n_test':NTEST}})
81 Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
82if __name__=='__main__': main()