Recursive Bellman Variance Targets / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  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()