Routh-Hurwitz Gain-Capped Optimizer / routh_bench.py
Mechanism confirmed, baseline not beaten
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()