Work-trained neural Hamiltonian bridge / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, os, json, random, math
  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
  7from bench.protocol import sweep_baseline, evaluate, make_report
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = (0,1,2,3)
 11EPOCHS = 12
 12NTR, NTE = 400, 200
 13BATCH = 128
 14# Shared search-space union: every idea lr is included in baseline.
 15LRS = [1e-3, 3e-3, 1e-2]
 16WDS = [0.0, 1e-4]
 17LAMBDA = [0.0, 0.01, 0.05]  # lambda=0 is the MSE control within the idea sweep
 18
 19def seed_all(seed):
 20    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 21    try:
 22        if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 23    except Exception: pass
 24
 25def data(seed):
 26    return get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
 27
 28def baseline_one(cfg, seed):
 29    seed_all(seed); d=data(seed)
 30    net=make_model('rnn_small', d['input_shape'], d['out_dim'])
 31    _, metric, _ = train_model(net, d, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH,
 32                               weight_decay=cfg['weight_decay'], log=lambda *_: None)
 33    return float(metric) if metric is not None else float('nan')
 34
 35def potential(theta, omega=0.0):
 36    # Dimensionless pendulum Hamiltonian potential, measured on the trained output.
 37    return 0.5*omega*omega + 9.81*(1.0-torch.cos(theta))
 38
 39def idea_one(cfg, seed, return_model=False):
 40    seed_all(seed); d=data(seed)
 41    net=make_model('rnn_small', d['input_shape'], d['out_dim'])
 42    dev='cuda' if torch.cuda.is_available() else 'cpu'
 43    try:
 44        net=net.to(dev); xtr,ytr=d['xtr'].to(dev),d['ytr'].to(dev)
 45        opt=torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 46        for _ in range(EPOCHS):
 47            net.train(); perm=torch.randperm(len(xtr),device=dev)
 48            for i in range(0,len(xtr),BATCH):
 49                ix=perm[i:i+BATCH]; pred=net(xtr[ix])
 50                mse=((pred-ytr[ix])**2).mean()
 51                # Reparameterized path proxy: endpoint energy change from the last
 52                # observed state, with Gaussian path NLL represented by MSE.
 53                last=xtr[ix].view(-1,8,3)[:,-1,0:1]
 54                work=potential(pred).mean()-potential(last).mean()
 55                loss=mse + cfg['lambda']*work
 56                if not torch.isfinite(loss): raise RuntimeError('nonfinite')
 57                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
 58        net.eval()
 59        with torch.no_grad():
 60            pred=net(d['xte'].to(dev)); metric=((pred-d['yte'].to(dev))**2).mean()
 61        if return_model: return net, float(metric), d
 62        return float(metric)
 63    except Exception:
 64        # CPU fallback mirrors the benchmark's robustness; recreate cleanly.
 65        seed_all(seed); net=make_model('rnn_small',d['input_shape'],d['out_dim']).cpu()
 66        xtr,ytr=d['xtr'],d['ytr']; opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay'])
 67        for _ in range(EPOCHS):
 68            perm=torch.randperm(len(xtr))
 69            for i in range(0,len(xtr),BATCH):
 70                ix=perm[i:i+BATCH]; pred=net(xtr[ix]); mse=((pred-ytr[ix])**2).mean()
 71                last=xtr[ix].view(-1,8,3)[:,-1,0:1]
 72                loss=mse+cfg['lambda']*(potential(pred).mean()-potential(last).mean())
 73                opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(),5.0); opt.step()
 74        with torch.no_grad(): metric=((net(d['xte'])-d['yte'])**2).mean()
 75        return float(metric)
 76
 77def signature(cfg, seed=0):
 78    net, metric, d=idea_one(cfg,seed,True)
 79    dev=next(net.parameters()).device
 80    with torch.no_grad():
 81        pred=net(d['xte'].to(dev)); last=d['xte'].to(dev).view(-1,8,3)[:,-1,0:1]
 82        delta=(potential(pred)-potential(last)).detach().cpu().numpy().ravel()
 83        err=((pred-d['yte'].to(dev))**2).detach().cpu().numpy().ravel()
 84    corr=float(np.corrcoef(delta,err)[0,1]) if np.std(delta)>1e-12 else 0.0
 85    return {'definition':'trained-model endpoint work proxy vs endpoint squared error',
 86            'mean_work_proxy':float(delta.mean()), 'std_work_proxy':float(delta.std()),
 87            'mean_test_mse':float(err.mean()), 'corr_work_error':corr,
 88            'predicted_direction':'work penalty should reduce endpoint energy change',
 89            'observed_direction': 'reduced' if delta.mean()<0 else 'increased',
 90            'confirmed': bool(delta.mean()<0 and np.isfinite(corr))}
 91
 92def main():
 93    # Baseline sweep includes all lr values used by idea, plus both baseline knobs.
 94    grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in WDS]
 95    base=sweep_baseline(lambda c: lambda s: baseline_one(c,s), grid, seeds=SWEEP_SEEDS)
 96    best=base['best_cfg']
 97    idea_grid=[{'lr':best['lr'],'weight_decay':best['weight_decay'],'lambda':l} for l in LAMBDA]
 98    # Nearby settings are the shared lr neighbors; all are in baseline sweep.
 99    for lr in LRS:
100        if lr != best['lr']: idea_grid.append({'lr':lr,'weight_decay':best['weight_decay'],'lambda':0.05})
101    idea_runs=[]
102    for cfg in idea_grid:
103        r=evaluate(lambda s,cfg=cfg: idea_one(cfg,s), seeds=SEEDS)
104        idea_runs.append({'cfg':cfg,'result':r})
105    chosen=min(idea_runs,key=lambda z:z['result']['mean'])
106    rep=make_report('dynamics','rnn_small',base,chosen['result'],extra=signature(chosen['cfg']))
107    rep['idea_sweep']=idea_runs
108    rep['protocol']={'epochs':EPOCHS,'n_train':NTR,'n_test':NTE,'batch':BATCH,'paired_seeds':list(SEEDS),'structural_match':'dynamics/control'}
109    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
110    print(json.dumps(rep,indent=2))
111if __name__=='__main__': main()