Noise-Whitened Trajectory-KL Policy Regularization / 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
 8TRACK='dynamics'; MODEL='rnn_small'; SEEDS=tuple(range(8))
 9LRS=[1e-3, 3e-3, 1e-2]
10REGS=[0.03, 0.1, 0.3]
11EPOCHS=10; BATCH=128
12
13def seed_all(seed):
14    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
15    if torch.cuda.is_available():
16        try: torch.cuda.manual_seed_all(seed)
17        except Exception: pass
18
19def train_variant(ds, seed, lr, reg, kind, return_net=False):
20    seed_all(seed)
21    net=make_model(MODEL, ds['input_shape'], ds['out_dim'])
22    ladder=([('cuda',False),('cuda',True)] if torch.cuda.is_available() else [])+[('cpu',False)]
23    errors=[]
24    for dev,no_cudnn in ladder:
25        try:
26            if no_cudnn: torch.backends.cudnn.enabled=False
27            net=net.to(dev); x=ds['xtr'].to(dev); y=ds['ytr'].to(dev)
28            opt=torch.optim.Adam(net.parameters(),lr=lr)
29            for _ in range(EPOCHS):
30                net.train(); perm=torch.randperm(len(x),device=dev)
31                for i in range(0,len(x),BATCH):
32                    idx=perm[i:i+BATCH]; pred=net(x[idx])
33                    task=((pred-y[idx])**2).mean()
34                    # Dynamics track predicts theta at the next horizon. The
35                    # persistence controller/reference predicts the last theta.
36                    # Its local drift mismatch is predicted theta - reference theta.
37                    drift=pred[:,0]-x[idx,-3]  # last theta in flattened (theta,omega,u)
38                    if kind=='idea':
39                        # Heteroscedastic diffusion: noisy/high-speed states are
40                        # cheap; predictable low-speed states are expensive.
41                        omega=x[idx,-2].abs()
42                        variance=0.05+0.50*omega.detach()
43                        penalty=0.5*(drift.square()/(variance+1e-3)).mean()
44                    else:
45                        penalty=0.5*drift.square().mean()
46                    loss=task+reg*penalty
47                    opt.zero_grad(); loss.backward(); opt.step()
48            net.eval()
49            with torch.no_grad(): metric=float(((net(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean())
50            return (metric,net,dev) if return_net else metric
51        except RuntimeError as e: errors.append(str(e)[:120])
52        finally:
53            if no_cudnn: torch.backends.cudnn.enabled=True
54    raise RuntimeError('training failed '+str(errors))
55
56def make_base(cfg): return lambda seed: train_variant(get_dataset(TRACK,seed),seed,cfg['lr'],cfg['reg'],'baseline')
57def make_idea(cfg): return lambda seed: train_variant(get_dataset(TRACK,seed),seed,cfg['lr'],cfg['reg'],'idea')
58
59def signature():
60    # Measured on predictions of separately trained NN systems, not an identity.
61    ds=get_dataset(TRACK,0,n_train=400,n_test=200)
62    out=[]
63    for kind in ('baseline','idea'):
64        _,net,dev=train_variant(ds,0,3e-3,0.1,kind,True)
65        with torch.no_grad():
66            pred=net(ds['xte'].to(dev))[:,0]; drift=pred-ds['xte'].to(dev)[:,-3]
67            var=0.05+0.50*ds['xte'].to(dev)[:,-2].abs()
68            raw=float((0.5*drift.square()).mean())
69            white=float((0.5*drift.square()/(var+1e-3)).mean())
70        out.append({'kind':kind,'raw_drift_cost':raw,'whitened_drift_cost':white})
71    # Prediction tested at NN scale: inverse variance must increase cost in
72    # low-noise states; verify empirical low/high-noise cost ratio.
73    with torch.no_grad():
74        pred=net(ds['xte'].to(dev))[:,0]; drift=pred-ds['xte'].to(dev)[:,-3]
75        var=0.05+0.50*ds['xte'].to(dev)[:,-2].abs(); med=torch.median(var)
76        lo=float((0.5*drift[var<=med].square()/(var[var<=med]+1e-3)).mean())
77        hi=float((0.5*drift[var>med].square()/(var[var>med]+1e-3)).mean())
78    ratio=lo/(hi+1e-12)
79    return {'prediction':'inverse covariance weights low-noise drift more than high-noise drift','low_high_weighted_cost_ratio_observed':ratio,'models':out,'confirmed':bool(ratio>1.0)}
80
81def main():
82    os.environ.setdefault('OMP_NUM_THREADS','4')
83    grid=[{'lr':lr,'reg':reg} for lr in LRS for reg in REGS]
84    base=sweep_baseline(make_base,grid)
85    idea_runs=[{'cfg':cfg,'result':evaluate(make_idea(cfg),SEEDS)} for cfg in grid]
86    best=min(idea_runs,key=lambda q:q['result']['mean'])
87    report=make_report(TRACK,MODEL,base,best['result'],{
88        'selected_cfg':best['cfg'],'idea_grid':idea_runs,
89        'signature':signature(),
90        'track_justification':'Dynamics is structurally matched: controlled pendulum rollout forecasting and drift/control regularization.'})
91    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
92    print(json.dumps(report,indent=2))
93if __name__=='__main__': main()