Routh-Hurwitz Gain-Capped Optimizer / routh_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, sweep_baseline, make_report
  8
  9SEEDS = tuple(range(8))
 10LRS = [1e-3, 3e-3, 1e-2]
 11EPOCHS, BATCH, MOMENTUM, RHO = 8, 128, 0.9, 0.8
 12
 13def seed_all(seed):
 14    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 15    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 16
 17def cap_update(net, previous_grad, previous_param, lr):
 18    """Online secant estimate followed by cubic Routh-Hurwitz cap.
 19
 20    The NN supplies curvature through gradient/parameter secants.  The
 21    dimensionless damping and frequency are the stated conservative model
 22    coordinates; no target labels or oracle dynamics are used in the cap.
 23    """
 24    num = den = 0.0
 25    current = []
 26    for p, oldg, oldp in zip(net.parameters(), previous_grad, previous_param):
 27        if p.grad is None:
 28            current.append(None); continue
 29        g = p.grad.detach()
 30        dp = p.detach() - oldp
 31        dg = g - oldg
 32        num += float((dg * dp).sum())
 33        den += float((dp * dp).sum())
 34        current.append(g.clone())
 35    curvature = max(1e-4, num / max(den, 1e-12))
 36    # Paper coefficients with r/L=.2 and omega0=1.
 37    r, L, omega = 0.2, 1.0, 1.0
 38    a1 = 2*r/L; a2 = (r/L)**2 + omega**2; kappa = 1.5*omega/L
 39    g = lr * curvature
 40    gmax = RHO * a1*a2 / kappa
 41    effective_lr = min(lr, gmax/curvature)
 42    chi = kappa * (effective_lr*curvature) / (a1*a2)
 43    return effective_lr, chi, curvature, current
 44
 45def train_one(ds, seed, lr, capped):
 46    seed_all(seed)
 47    requested_device = 'cuda' if torch.cuda.is_available() else 'cpu'
 48    try:
 49        return _train(ds, seed, lr, capped, requested_device)
 50    except RuntimeError:
 51        if requested_device == 'cuda':
 52            torch.cuda.empty_cache()
 53            return _train({k:(v.cpu() if torch.is_tensor(v) else v) for k,v in ds.items()}, seed, lr, capped, 'cpu')
 54        raise
 55
 56def _train(ds, seed, lr, capped, device):
 57    # Same architecture and base optimizer hyperparameters on both sides.
 58    net = make_model('rnn_small', tuple(ds['input_shape']), ds['out_dim']).to(device)
 59    x, y = ds['xtr'].to(device), ds['ytr'].to(device)
 60    opt = torch.optim.SGD(net.parameters(), lr=lr, momentum=MOMENTUM)
 61    lossf = nn.MSELoss()
 62    prev_g = [torch.zeros_like(p) for p in net.parameters()]
 63    prev_p = [p.detach().clone() for p in net.parameters()]
 64    history, chis, gains, curvatures = [], [], [], []
 65    for _ in range(EPOCHS):
 66        net.train(); perm = torch.randperm(len(x), device=device); total = 0.0
 67        for start in range(0, len(x), BATCH):
 68            idx = perm[start:start+BATCH]
 69            loss = lossf(net(x[idx]), y[idx])
 70            opt.zero_grad(set_to_none=True); loss.backward()
 71            if capped:
 72                effective, chi, curvature, new_g = cap_update(net, prev_g, prev_p, lr)
 73                scale = effective / lr
 74                for p in net.parameters():
 75                    if p.grad is not None: p.grad.mul_(scale)
 76                gains.append(effective * curvature); chis.append(chi); curvatures.append(curvature)
 77            else:
 78                effective, chi, curvature, new_g = lr, float('nan'), float('nan'), [p.grad.detach().clone() if p.grad is not None else None for p in net.parameters()]
 79            opt.step()
 80            for j, p in enumerate(net.parameters()):
 81                if p.grad is not None:
 82                    prev_g[j] = new_g[j] if new_g[j] is not None else p.grad.detach().clone()
 83                    prev_p[j] = p.detach().clone()
 84            total += float(loss) * len(idx)
 85        history.append(total / len(x))
 86    net.eval()
 87    with torch.no_grad():
 88        metric = float(((net(ds['xte'].to(device)) - ds['yte'].to(device))**2).mean())
 89    return {'metric': metric, 'history': history,
 90            'max_chi': float(max(chis)) if chis else float('nan'),
 91            'mean_effective_lr': float(np.mean([lr if not capped else min(lr, RHO*.4*1.04/(1.5*max(c,1e-4))) for c in curvatures])) if capped and curvatures else lr,
 92            'mean_observed_gain': float(np.mean(gains)) if gains else float('nan'),
 93            'max_curvature': float(max(curvatures)) if curvatures else float('nan')}
 94
 95def dataset(seed):
 96    return get_dataset('dynamics', seed, n_train=400, n_test=200)
 97
 98def metric_fn(seed, lr, capped):
 99    return train_one(dataset(seed), seed, lr, capped)['metric']
100
101def factory(capped):
102    return lambda cfg: (lambda seed: metric_fn(seed, cfg['lr'], capped))
103
104def main():
105    # Independent math sanity check: cubic pole real part changes sign at chi=1.
106    r, L, w = .2, 1., 1.; a1=2*r/L; a2=(r/L)**2+w*w; k=1.5*w/L
107    root_check=[]
108    for frac in (.8, 1.0, 1.2):
109        roots=np.roots([1.,a1,a2,k*frac*a1*a2])
110        root_check.append({'chi':frac, 'max_real_root':float(np.max(roots.real))})
111    grid=[{'lr':v} for v in LRS]
112    sweep=sweep_baseline(factory(False), grid)
113    best_lr=float(sweep['best_cfg']['lr'])
114    base_full=sweep['full']
115    idea_full={'per_seed':[metric_fn(s,best_lr,True) for s in SEEDS]}
116    # Include the two nearby idea settings; these were all baseline-swept too.
117    idea_all={str(lr):[train_one(dataset(s),s,lr,True) for s in SEEDS] for lr in LRS}
118    sig_runs=[train_one(dataset(s),s,best_lr,True) for s in SEEDS]
119    base_sig=[train_one(dataset(s),s,best_lr,False) for s in SEEDS]
120    extra={'mechanism_signature':{
121        'prediction':'online RH cap enforces chi <= rho=0.8',
122        'predicted_max_chi':RHO,
123        'observed_max_chi_idea':float(max(x['max_chi'] for x in sig_runs)),
124        'observed_max_chi_baseline':float(max(x['max_chi'] for x in base_sig)),
125        'observed_mean_gain_idea':float(np.mean([x['mean_observed_gain'] for x in sig_runs])),
126        'confirmed':bool(max(x['max_chi'] for x in sig_runs) <= RHO+1e-6)}}
127    report=make_report('dynamics','rnn_small',{'best_cfg':sweep['best_cfg'],'sweep':sweep['sweep'],'full':base_full},idea_full,extra)
128    out={'root_check':root_check,'bench_report':report,'idea_settings':idea_all,'custom_track':None}
129    Path('bench_results.json').write_text(json.dumps(out,indent=2,default=float)); print(json.dumps(out,indent=2,default=float))
130
131if __name__ == '__main__': main()