Universal Trust-Region Neural Optimizer / bench_experiment.py

Failed on benchmark

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, 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()