Universal Trust-Region Neural Optimizer / bench_experiment.py
Failed on benchmark
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, train_model, evaluate, sweep_baseline, make_report
8
9EPOCHS = 18
10BATCH = 128
11SEEDS = tuple(range(8))
12LRS = [1e-3, 3e-3, 1e-2]
13WEIGHT_DECAYS = [0.0, 1e-4]
14
15
16def seed_all(seed):
17 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
18 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
19
20
21def baseline_run(cfg, seed, keep=False):
22 seed_all(seed)
23 ds = get_dataset('tabular', seed, n_train=400, n_test=400)
24 net = make_model('mlp_tiny', tuple(ds['input_shape']), ds['out_dim'])
25 net, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'],
26 batch=BATCH, weight_decay=cfg['weight_decay'],
27 log=lambda *_: None)
28 out = {'metric': float(metric), 'history': hist}
29 return out if keep else float(metric)
30
31
32def params(model):
33 return [p for p in model.parameters() if p.requires_grad]
34
35
36def flatten(xs):
37 return torch.cat([x.detach().reshape(-1) for x in xs])
38
39
40def assign(model, vec):
41 pos = 0
42 with torch.no_grad():
43 for p in params(model):
44 n = p.numel(); p.copy_(vec[pos:pos+n].view_as(p)); pos += n
45
46
47def trust_run(cfg, seed, keep=False):
48 seed_all(seed)
49 ds = get_dataset('tabular', seed, n_train=400, n_test=400)
50 device = 'cuda' if torch.cuda.is_available() else 'cpu'
51 try:
52 net = make_model('mlp_tiny', tuple(ds['input_shape']), ds['out_dim']).to(device)
53 x = ds['xtr'].to(device); y = ds['ytr'].to(device)
54 lossf = nn.MSELoss()
55 delta = float(cfg['delta0']); delta_max = 10.0
56 hist, rhos, rejects, radii, pred_obs = [], [], 0, [], []
57 for _ in range(EPOCHS):
58 net.train(); ps = params(net)
59 net.zero_grad(set_to_none=True)
60 old = lossf(net(x), y); old.backward()
61 g = flatten([p.grad for p in ps])
62 # Diagonal empirical-Fisher curvature plus damping; this is the local B model.
63 b = flatten([p.grad * p.grad for p in ps]).clamp_min(1e-6) + cfg['damping']
64 gn = float(g.norm())
65 if gn < 1e-12:
66 break
67 # Exact trust-region solution for positive diagonal B via bisection on lambda.
68 def step_for(lam): return -g / (b + lam)
69 s = step_for(0.0)
70 if float(s.norm()) > delta:
71 lo, hi = 0.0, 1.0
72 while float(step_for(hi).norm()) > delta: hi *= 2.0
73 for _ in range(25):
74 mid = (lo + hi) / 2
75 if float(step_for(mid).norm()) > delta: lo = mid
76 else: hi = mid
77 s = step_for(hi)
78 cauchy_len = min(delta, gn / float(b.max()))
79 sc = -cauchy_len * g / (gn + 1e-12)
80 pred = -(torch.dot(g, s) + 0.5 * torch.dot(b * s, s))
81 cpred = -(torch.dot(g, sc) + 0.5 * torch.dot(b * sc, sc))
82 if float(pred) < 0.1 * float(cpred):
83 s, pred = sc, cpred
84 old_vec = flatten([p for p in ps])
85 assign(net, old_vec + s)
86 with torch.no_grad(): new = lossf(net(x), y)
87 ared = old.detach() - new
88 rho = float(ared / (pred + 1e-12))
89 boundary = float(s.norm()) >= 0.99 * delta
90 accepted = rho >= 0.1 and float(pred) > 0
91 if not accepted:
92 assign(net, old_vec); delta *= 0.25; rejects += 1; value = float(old)
93 else:
94 value = float(new)
95 if rho > 0.75 and boundary: delta = min(2.0 * delta, delta_max)
96 hist.append(value); rhos.append(rho); radii.append(delta)
97 pred_obs.append({'predicted': float(pred), 'observed': float(ared), 'rho': rho})
98 net.eval()
99 with torch.no_grad(): metric = float(((net(ds['xte'].to(device)) - ds['yte'].to(device)) ** 2).mean())
100 out = {'metric': metric, 'history': hist, 'reject_fraction': rejects / max(1, EPOCHS),
101 'median_rho': float(np.median(rhos)) if rhos else float('nan'),
102 'final_radius': delta, 'radii': radii, 'pred_obs': pred_obs}
103 return out if keep else metric
104 except RuntimeError:
105 # Robust CPU fallback for shared/unsupported CUDA environments.
106 torch.cuda.empty_cache()
107 old = torch.cuda.is_available
108 torch.cuda.is_available = lambda: False
109 try: return trust_run(cfg, seed, keep)
110 finally: torch.cuda.is_available = old
111
112
113def main():
114 base_grid = [{'lr': lr, 'weight_decay': wd} for lr in LRS for wd in WEIGHT_DECAYS]
115 base = sweep_baseline(lambda cfg: lambda seed: baseline_run(cfg, seed), base_grid)
116 best = base['best_cfg']
117 idea_grid = [{'lr': lr, 'weight_decay': best['weight_decay'], 'delta0': d, 'damping': 0.01}
118 for lr in LRS for d in [0.05, 0.2, 0.8]]
119 # Equal-budget idea sweep on seeds 0..3, then full paired run for its best config.
120 tried = []
121 for cfg in idea_grid:
122 vals = [trust_run(cfg, s) for s in (0, 1, 2, 3)]
123 tried.append({'cfg': cfg, 'mean': float(np.mean(vals)), 'per_seed': vals})
124 ibest = min(tried, key=lambda z: z['mean'])['cfg']
125 idea = evaluate(lambda seed: trust_run(ibest, seed), seeds=SEEDS)
126 extra_vals = [trust_run(ibest, s, keep=True) for s in SEEDS]
127 allro = [q for r in extra_vals for q in r['pred_obs']]
128 ratios = np.array([q['observed'] / q['predicted'] for q in allro if q['predicted'] > 1e-10 and np.isfinite(q['observed'])])
129 signature = {'predicted_decrease_mean': float(np.mean([q['predicted'] for q in allro])),
130 'observed_decrease_mean': float(np.mean([q['observed'] for q in allro])),
131 'rho_median': float(np.median(ratios)) if len(ratios) else float('nan'),
132 'rho_iqr': [float(np.quantile(ratios, .25)), float(np.quantile(ratios, .75))] if len(ratios) else [],
133 'reject_fraction_mean': float(np.mean([r['reject_fraction'] for r in extra_vals])),
134 'confirmed': bool(len(ratios) > 0 and 0.5 <= float(np.median(ratios)) <= 1.5)}
135 report = make_report('tabular', 'mlp_tiny', base, idea, signature)
136 report['idea']['sweep'] = tried
137 report['protocol_notes'] = 'Tabular is the built-in optimizer track; both systems use identical mlp_tiny, data, epochs, and paired seeds. Baseline is Adam via train_model; idea changes only the training optimizer.'
138 Path('bench_report.json').write_text(json.dumps(report, indent=2, allow_nan=False))
139 print(json.dumps(report, indent=2, allow_nan=False))
140
141if __name__ == '__main__': main()