Mean-Square Proximal Relaxation Optimizer / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, random, sys
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report, count_params
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = tuple(range(4))
 11EPOCHS = 15
 12BATCH = 128
 13LR_GRID = [1e-3, 3e-3, 6e-3]
 14ALPHA_GRID = [0.25, 0.5, 0.75]
 15
 16
 17def seed_all(seed):
 18    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 19    if torch.cuda.is_available():
 20        try: torch.cuda.manual_seed_all(seed)
 21        except Exception: pass
 22
 23
 24def device():
 25    return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 26
 27
 28def proximal_train(model, ds, lr, alpha, epochs=EPOCHS, batch=BATCH, return_stats=False):
 29    """Blockwise stochastic proximal response T=w-lr*g, then w<-w+alpha(T-w).
 30    Blocks are parameter tensors; gradients are computed jointly but applied blockwise,
 31    making the intervention the optimizer update rather than the network architecture.
 32    """
 33    try:
 34        dev = device(); model = model.to(dev)
 35        xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
 36    except Exception:
 37        dev = torch.device('cpu'); model = model.to(dev)
 38        xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
 39    loss_fn = nn.MSELoss()
 40    params = [p for p in model.parameters() if p.requires_grad]
 41    rng = np.random.default_rng(12345)
 42    update_norms = []; grad_snapshots = []
 43    n = len(xtr)
 44    model.train()
 45    for ep in range(epochs):
 46        order = rng.permutation(n)
 47        for start in range(0, n, batch):
 48            ix = torch.as_tensor(order[start:start+batch], device=dev)
 49            model.zero_grad(set_to_none=True)
 50            pred = model(xtr[ix])
 51            loss = loss_fn(pred, ytr[ix])
 52            loss.backward()
 53            norms = []
 54            with torch.no_grad():
 55                for p in params:  # each tensor is one proximal block
 56                    if p.grad is None: continue
 57                    # T_hat = p - lr*g; relaxed response is p - alpha*lr*g.
 58                    step = alpha * lr * p.grad
 59                    p.sub_(step)
 60                    norms.append(float(step.norm().detach().cpu()))
 61            if norms:
 62                update_norms.append(float(np.sqrt(np.sum(np.square(norms)))))
 63                grad_snapshots.append(float(np.mean(np.square(norms))))
 64    model.eval()
 65    with torch.no_grad():
 66        metric = float(loss_fn(model(ds['xte'].to(dev)), ds['yte'].to(dev)).cpu())
 67    stats = {'update_rms': float(np.sqrt(np.mean(np.square(update_norms)))) if update_norms else 0.0,
 68             'update_std': float(np.std(update_norms)) if update_norms else 0.0,
 69             'n_updates': len(update_norms)}
 70    return model, metric, stats
 71
 72
 73def baseline_fn(cfg):
 74    def run(seed):
 75        seed_all(seed)
 76        ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
 77        model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 78        _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH, log=lambda *_: None)
 79        return float(metric)
 80    return run
 81
 82
 83def idea_fn(cfg, retain=None):
 84    def run(seed):
 85        seed_all(seed)
 86        ds = get_dataset('dynamics', seed, n_train=4000, n_test=1000)
 87        model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
 88        net, metric, stats = proximal_train(model, ds, cfg['lr'], cfg['alpha'])
 89        if retain is not None: retain[seed] = (net, ds, stats)
 90        return float(metric)
 91    return run
 92
 93
 94def main():
 95    # Baseline includes every lr used by the idea; alpha is the idea-only relaxation knob.
 96    base = sweep_baseline(baseline_fn, [{'lr': x} for x in LR_GRID], seeds=SWEEP_SEEDS)
 97    idea_grid = [{'lr': lr, 'alpha': a} for lr, a in zip(LR_GRID, ALPHA_GRID)]
 98    idea_sweep = []
 99    for cfg in idea_grid:
100        vals = [idea_fn(cfg)(s) for s in SWEEP_SEEDS]
101        idea_sweep.append({'cfg': cfg, 'mean': float(np.mean(vals)), 'per_seed': vals})
102    best_cfg = min(idea_sweep, key=lambda z: z['mean'])['cfg']
103
104    retained = {}
105    idea_vals = [idea_fn(best_cfg, retained)(s) for s in SEEDS]
106    idea_res = {'mean': float(np.mean(idea_vals)), 'std': float(np.std(idea_vals)),
107                'per_seed': idea_vals, 'n': 8, 'chosen_cfg': best_cfg,
108                'sweep': idea_sweep}
109
110    base_vals = []
111    for s in SEEDS:
112        seed_all(s); ds = get_dataset('dynamics', s, n_train=4000, n_test=1000)
113        m = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
114        _, metric, _ = train_model(m, ds, epochs=EPOCHS, lr=base['best_cfg']['lr'], batch=BATCH, log=lambda *_: None)
115        base_vals.append(float(metric))
116    base['full'] = {'mean': float(np.mean(base_vals)), 'std': float(np.std(base_vals)),
117                    'per_seed': base_vals, 'n': 8}
118
119    # Signature is measured on trained benchmark models: response/update noise and damping.
120    rms = [v[2]['update_rms'] for v in retained.values()]
121    std = [v[2]['update_std'] for v in retained.values()]
122    alpha = best_cfg['alpha']
123    # For the relaxed recursion, observed update magnitude should scale approximately alpha.
124    sig = {'prediction': 'relaxation damps stochastic block-response updates approximately linearly in alpha',
125           'predicted_alpha': alpha, 'observed_update_rms_mean': float(np.mean(rms)),
126           'observed_update_std_mean': float(np.mean(std)),
127           'predicted_vs_observed_relation': 'measured from trained rnn_small updates; alpha scaling was not quantitatively tested',
128           'confirmed': False}
129    report = make_report('dynamics', 'rnn_small', base, idea_res, {'mechanism_signature': sig})
130    report['mechanism_signature'] = sig
131    report['custom_track'] = None
132    with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2)
133    print(json.dumps(report, indent=2))
134
135if __name__ == '__main__': main()