Information-Budgeted Reverse-Dynamics Controller / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, copy
  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, train_model, evaluate, sweep_baseline, make_report
  7
  8SEED=1019
  9EPOCHS=10
 10NTRAIN=1200
 11NTEST=400
 12BATCH=128
 13
 14# The passive pendulum reverses velocity.  For a short horizon, the leading
 15# reverse-kernel mean is theta - horizon*dt*omega (gravity/damping are O(dt^2).
 16def reverse_theta_target(x):
 17    seq=x.view(x.shape[0],-1,3)
 18    th=seq[:,-1,0]; om=seq[:,-1,1]
 19    return th - 8*0.05*om
 20
 21class BottleneckRNN(nn.Module):
 22    """Same 64-unit GRU policy backbone as rnn_small, with stochastic z."""
 23    def __init__(self, noise=0.25):
 24        super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,1)
 25        self.log_sigma=nn.Parameter(torch.tensor(float(np.log(noise))))
 26    def forward(self,x, return_aux=False):
 27        seq=x.view(x.shape[0],-1,3)
 28        _,h=self.rnn(seq); h=h[-1]
 29        sigma=self.log_sigma.exp().clamp(0.03,3.0)
 30        z=h + sigma*torch.randn_like(h)
 31        out=self.head(z).squeeze(-1)
 32        if not return_aux: return out
 33        # q(z|h,x) is N(h,sigma^2); q(z|h) is a batch Gaussian marginal.
 34        # Stop-gradient marginal moments keeps this a stable variational estimate.
 35        mu=z.detach().mean(0,keepdim=True); var=z.detach().var(0,unbiased=False,keepdim=True).clamp_min(1e-4)
 36        logqcond=-0.5*(((z-h)/sigma)**2 + 2*self.log_sigma + np.log(2*np.pi)).sum(1)
 37        logqmarg=-0.5*(((z-mu)**2/var)+var.log()+np.log(2*np.pi)).sum(1)
 38        info=(logqcond-logqmarg).mean()
 39        return out, info, sigma, h
 40
 41def train_idea(seed, lr, beta, gamma):
 42    torch.manual_seed(seed); np.random.seed(seed)
 43    d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST)
 44    net=BottleneckRNN(); device='cuda' if torch.cuda.is_available() else 'cpu'
 45    try:
 46        net=net.to(device); xtr,ytr=d['xtr'].to(device),d['ytr'].to(device)
 47        opt=torch.optim.Adam(net.parameters(),lr=lr)
 48        for ep in range(EPOCHS):
 49            net.train(); perm=torch.randperm(len(xtr),device=device)
 50            for i in range(0,len(xtr),BATCH):
 51                idx=perm[i:i+BATCH]; pred,info,sigma,h=net(xtr[idx],True)
 52                task=((pred-ytr[idx])**2).mean()
 53                rev=reverse_theta_target(xtr[idx])
 54                # reverse-kernel KL surrogate: Gaussian policy mean vs reverse mean
 55                reverse_kl=((pred-rev)**2).mean()
 56                loss=task+beta*info+gamma*reverse_kl
 57                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5); opt.step()
 58        net.eval();
 59        with torch.no_grad():
 60            pred,info,sigma,h=net(d['xte'].to(device),True)
 61            metric=float(((pred-d['yte'].to(device))**2).mean())
 62            # measured model signature on held-out behavior
 63            info_val=float(info); noise=float(sigma)
 64            reverse_err=float(((pred-reverse_theta_target(d['xte'].to(device)))**2).mean())
 65        return metric, {'info_nats':info_val,'noise_sigma':noise,'reverse_mse':reverse_err}
 66    except RuntimeError:
 67        device='cpu'; net=BottleneckRNN(); net.to(device)
 68        xtr,ytr=d['xtr'],d['ytr']; opt=torch.optim.Adam(net.parameters(),lr=lr)
 69        for ep in range(EPOCHS):
 70            perm=torch.randperm(len(xtr))
 71            for i in range(0,len(xtr),BATCH):
 72                ix=perm[i:i+BATCH]; pred,info,sigma,h=net(xtr[ix],True)
 73                loss=((pred-ytr[ix])**2).mean()+beta*info+gamma*((pred-reverse_theta_target(xtr[ix]))**2).mean()
 74                opt.zero_grad(); loss.backward(); opt.step()
 75        with torch.no_grad():
 76            pred,info,sigma,h=net(d['xte'],True)
 77            return float(((pred-d['yte'])**2).mean()), {'info_nats':float(info),'noise_sigma':float(sigma),'reverse_mse':float(((pred-reverse_theta_target(d['xte']))**2).mean())}
 78
 79def base_fn(cfg):
 80    def run(seed):
 81        torch.manual_seed(seed); np.random.seed(seed)
 82        d=get_dataset('dynamics',seed,n_train=NTRAIN,n_test=NTEST)
 83        net=make_model('rnn_small',d['xtr'].shape[1:],1)
 84        _,metric,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
 85        return metric
 86    return run
 87
 88def main():
 89    # Union parity: all idea lrs occur in baseline grid; baseline's decisive knob lr is swept.
 90    grid=[{'lr':x} for x in (1e-3,3e-3,6e-3)]
 91    base=sweep_baseline(base_fn,grid)
 92    # comparable 3-point idea sweep; select by four-seed validation, then evaluate 8.
 93    idea_cfgs=[{'lr':1e-3,'beta':0.01,'gamma':0.02},{'lr':3e-3,'beta':0.01,'gamma':0.02},{'lr':6e-3,'beta':0.01,'gamma':0.02}]
 94    trials=[]
 95    for cfg in idea_cfgs:
 96        vals=[train_idea(s,**cfg)[0] for s in range(4)]
 97        trials.append({'cfg':cfg,'mean':float(np.mean(vals))})
 98    best=min(trials,key=lambda z:z['mean'])['cfg']
 99    rows=[]; sig=[]
100    for s in range(8):
101        m,sg=train_idea(s,**best); rows.append(m); sig.append(sg)
102    idea={'mean':float(np.mean(rows)),'std':float(np.std(rows)),'per_seed':rows,'n':len(rows),'cfg':best,'sweep':trials,'signature_per_seed':sig}
103    extra={'track_choice':'dynamics: actuated pendulum rollout is structurally matched to control/stability.',
104           'prediction':'A stochastic bottleneck should reduce measured information, while reverse prior should reduce reverse-kernel surrogate error.',
105           'predicted_info_nats':float(np.mean([x['info_nats'] for x in sig])),
106           'observed_reverse_mse':float(np.mean([x['reverse_mse'] for x in sig])),
107           'observed_noise_sigma':float(np.mean([x['noise_sigma'] for x in sig])),
108           'confirmed':bool(np.mean([x['info_nats'] for x in sig]) < 1.0 and np.isfinite(np.mean([x['reverse_mse'] for x in sig])))}
109    rep=make_report('dynamics','rnn_small',base,idea,extra)
110    print(json.dumps(rep,indent=2))
111if __name__=='__main__': main()