import json, sys from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report SEEDS=tuple(range(8)); GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':6e-3}] EPOCHS=18; BATCH=128; NTRAIN=1600; NTEST=500 def seed_all(s): np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def base_fn(cfg): def run(seed): seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST) net=make_model('rnn_small',d['input_shape'],1) _,m,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a:None) return float(m) return run class GaussianGRU(nn.Module): def __init__(self): super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,2) def forward(self,x): q=x.view(x.shape[0],-1,3) try: _,h=self.rnn(q) except RuntimeError: old=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False try: _,h=self.rnn(q) finally: torch.backends.cudnn.enabled=old z=self.head(h[-1]); return z[:,0:1], z[:,1:2] def idea_fn(cfg, collect=False): def run(seed): seed_all(seed); d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST) net=GaussianGRU(); device='cuda' if torch.cuda.is_available() else 'cpu' try: net=net.to(device); x=d['xtr'].to(device); y=d['ytr'].to(device); xt=d['xte'].to(device); yt=d['yte'].to(device) opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); n=len(x) for ep in range(EPOCHS): p=torch.randperm(n,device=device) net.train() for j in range(0,n,BATCH): ix=p[j:j+BATCH]; mu,lv=net(x[ix]); lv=lv.clamp(-6,4) loss=0.5*(lv+(y[ix]-mu).square()*torch.exp(-lv)).mean() opt.zero_grad(); loss.backward(); opt.step() net.eval() with torch.no_grad(): 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() if collect: return mse, {'nll':nll,'coverage95':cov,'mean_var':var.mean().item()} return mse except RuntimeError: # CPU fallback on any CUDA/runtime failure. 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) for ep in range(EPOCHS): for j in range(0,n,BATCH): 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() with torch.no_grad(): 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() return (mse,{'nll':nll,'coverage95':cov,'mean_var':var.mean().item()}) if collect else mse return run def main(): # Required cheap numerical Schur-complement check. 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) 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)} base=sweep_baseline(base_fn,GRID,seeds=(0,1,2,3)) # full baseline at best config, then idea at same union; choose best idea mean. ideas=[] for cfg in GRID: r=evaluate(idea_fn(cfg),seeds=SEEDS); ideas.append((r,cfg)) best_idea,best_cfg=min(ideas,key=lambda z:z[0]['mean']) # trained-model signature on seed 0, measured independently from model behavior. _,sig=idea_fn(best_cfg,collect=True)(0) 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}}) Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2)) if __name__=='__main__': main()