Recursive Bellman Variance Targets / stage2_bench.py
Failed on benchmark
1import os, sys, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
7
8SEEDS=tuple(range(8))
9# Same union of learning rates is swept for both systems.
10GRID=[{'lr':1e-3,'epochs':8},{'lr':3e-3,'epochs':8},{'lr':1e-2,'epochs':8}]
11NTR,NTE=400,200
12BATCH=128
13
14class MeanVarianceRNN(nn.Module):
15 """The bench rnn_small backbone with the required mean and variance heads."""
16 def __init__(self, input_shape):
17 super().__init__()
18 self.rnn=nn.GRU(3,64,batch_first=True)
19 self.mean=nn.Linear(64,1)
20 self.rawvar=nn.Linear(64,1)
21 def encode(self,x):
22 _,h=self.rnn(x.view(x.shape[0],-1,3)); return h[-1]
23 def forward(self,x):
24 h=self.encode(x)
25 return self.mean(h)
26 def both(self,x):
27 h=self.encode(x)
28 return self.mean(h), torch.nn.functional.softplus(self.rawvar(h))+1e-4
29
30def seed_all(seed):
31 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
32
33def device_run(fn):
34 # Small batches/models, with explicit GPU fallback as required by the harness.
35 dev='cuda' if torch.cuda.is_available() else 'cpu'
36 try: return fn(dev)
37 except RuntimeError:
38 return fn('cpu')
39
40def train_base(seed,cfg, capture=False):
41 seed_all(seed); d=get_dataset('dynamics',seed,NTR,NTE)
42 def run(dev):
43 net=make_model('rnn_small',d['input_shape'],1).to(dev)
44 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
45 x,y=d['xtr'].to(dev),d['ytr'].to(dev)
46 for _ in range(cfg['epochs']):
47 net.train(); p=torch.randperm(len(x),device=dev)
48 for i in range(0,len(x),BATCH):
49 z=net(x[p[i:i+BATCH]]); loss=((z-y[p[i:i+BATCH]])**2).mean()
50 opt.zero_grad(); loss.backward(); opt.step()
51 net.eval()
52 with torch.no_grad():
53 pred=net(d['xte'].to(dev)); err=(pred-d['yte'].to(dev))**2
54 metric=float(err.mean())
55 if capture: return metric, {'pred':pred.cpu().numpy().ravel(),'err':err.cpu().numpy().ravel()}
56 return metric
57 return device_run(run)
58
59def recursive_q(net,x):
60 """One-step Bellman variance estimate from four stochastic child branches.
61 Branch noise represents rollout/transition uncertainty in the final observed
62 (theta, omega, action); q_rec=E[q_child]+Var(m_child), with stop-gradient."""
63 b=x.shape[0]; k=4
64 scales=x.new_tensor([.01,.03,.15]).view(1,1,3)
65 eps=torch.randn((b,k,3),device=x.device)*scales
66 branches=x[:,None,:].repeat(1,k,1).view(b*k,-1)
67 branches[:,-3:]=branches[:,-3:]+eps.view(b*k,3)
68 mu,q=net.both(branches); mu=mu.view(b,k); q=q.view(b,k)
69 return (q.mean(1)+(mu-mu.mean(1,keepdim=True)).pow(2).mean(1)).detach()
70
71def train_idea(seed,cfg,capture=False):
72 seed_all(seed); d=get_dataset('dynamics',seed,NTR,NTE)
73 def run(dev):
74 net=MeanVarianceRNN(d['input_shape']).to(dev)
75 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
76 x,y=d['xtr'].to(dev),d['ytr'].to(dev)
77 for _ in range(cfg['epochs']):
78 net.train(); p=torch.randperm(len(x),device=dev)
79 for i in range(0,len(x),BATCH):
80 ix=p[i:i+BATCH]; mu,q=net.both(x[ix]); qr=recursive_q(net,x[ix])
81 # recursive target is detached when weighting mean loss; q head is
82 # separately regressed to the recursive target.
83 meanloss=.5*((y[ix]-mu).pow(2)/(qr+1e-3)+torch.log(qr+1e-3)).mean()
84 varloss=.05*(torch.log(q+1e-4)-torch.log(qr+1e-4)).pow(2).mean()
85 loss=meanloss+varloss
86 opt.zero_grad(); loss.backward(); opt.step()
87 net.eval()
88 with torch.no_grad():
89 mu,q=net.both(d['xte'].to(dev)); err=(mu-d['yte'].to(dev))**2
90 # Signature is measured on this trained model, not an analytical toy.
91 qr=recursive_q(net,d['xte'].to(dev))
92 metric=float(err.mean())
93 if capture: return metric, {'pred':mu.cpu().numpy().ravel(),'err':err.cpu().numpy().ravel(),'q':q.cpu().numpy().ravel(),'qr':qr.cpu().numpy().ravel()}
94 return metric
95 return device_run(run)
96
97def base_factory(cfg): return lambda seed: train_base(seed,cfg)
98def idea_factory(cfg): return lambda seed: train_idea(seed,cfg)
99
100def main():
101 # Baseline tuning uses the harness sweep (4 seeds), then final paired 8-seed result.
102 base=sweep_baseline(base_factory,GRID)
103 idea_runs=[]
104 for cfg in GRID:
105 idea_runs.append({'cfg':cfg,'result':evaluate(idea_factory(cfg),SEEDS)})
106 best=min(idea_runs,key=lambda z:z['result']['mean'])
107 idea=best['result']; idea['selected_cfg']=best['cfg']
108 # Capture the same selected systems once for a trained-model signature.
109 bvals=[]; ivals=[]; ratios=[]; correlations=[]
110 for s in SEEDS:
111 bm,bd=train_base(s,base['best_cfg'],True); im,idat=train_idea(s,best['cfg'],True)
112 bvals.append(bm); ivals.append(im)
113 ratios.append(float(np.mean(idat['err'])/(np.mean(idat['qr'])+1e-12)))
114 correlations.append(float(np.corrcoef(idat['err'],idat['qr'])[0,1]))
115 base['full']={'mean':float(np.mean(bvals)),'std':float(np.std(bvals)),'per_seed':bvals,'n':8}
116 idea={'mean':float(np.mean(ivals)),'std':float(np.std(ivals)),'per_seed':ivals,'n':8,'selected_cfg':best['cfg']}
117 sig={'prediction':'recursive predicted variance should calibrate observed squared error near 1 and track heteroscedasticity',
118 'predicted_vs_observed':{'mean_q_recursive':float(np.mean([np.mean(idat['qr'])])),
119 'mean_observed_squared_error':float(np.mean([np.mean(idat['err'])])),
120 'calibration_ratio_mean':float(np.mean(ratios)),
121 'calibration_ratio_per_seed':ratios,'q_error_correlation_mean':float(np.mean(correlations))},
122 'confirmed':bool(.8 <= np.mean(ratios) <= 1.2)}
123 rep=make_report('dynamics','rnn_small',base,idea,sig)
124 rep['protocol_notes']={'structural_match':'multi-step stochastic actuated pendulum rollout; recursive uncertainty applies to transition branches',
125 'baseline':'standard MSE training', 'idea':'mean/variance RNN with detached recursive child q plus between-child variance',
126 'n_train':NTR,'n_test':NTE,'branch_count':4,'grid_union':GRID}
127 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
128 print(json.dumps(rep,indent=2))
129if __name__=='__main__': main()