Wasserstein Speed-Limit Controller / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, math, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, make_model, sweep_baseline, evaluate, make_report
  8
  9SEEDS = tuple(range(8))
 10SWEEP_SEEDS = tuple(range(4))
 11BATCH = 128
 12EPOCHS = 12
 13# Union of all rates tried by both methods. Baseline also sweeps its central Adam knob.
 14LRS = [1e-3, 3e-3, 1e-2]
 15WEIGHT_DECAYS = [0.0, 1e-4]
 16
 17
 18def seed_all(seed):
 19    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 20    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 21
 22
 23def device():
 24    return 'cuda' if torch.cuda.is_available() else 'cpu'
 25
 26
 27def flat_params(model):
 28    return torch.cat([p.detach().reshape(-1).cpu() for p in model.parameters()])
 29
 30
 31def sliced_w2(a, b, n_proj=32, rng=None):
 32    # Dimension-corrected sliced estimator, matching stage-1 implementation.
 33    rng = np.random.default_rng(0) if rng is None else rng
 34    d = a.shape[1]
 35    q = rng.normal(size=(n_proj, d)); q /= np.linalg.norm(q, axis=1, keepdims=True)
 36    pa = np.sort(a @ q.T, axis=0); pb = np.sort(b @ q.T, axis=0)
 37    return float(d * np.mean((pa - pb) ** 2))
 38
 39
 40def batches(n, batch, rng):
 41    ix = rng.permutation(n)
 42    return [ix[i:i+batch] for i in range(0, n, batch)]
 43
 44
 45def run(seed, lr, controlled, weight_decay=0.0, collect=False):
 46    seed_all(seed)
 47    ds = get_dataset('dynamics', seed, n_train=400, n_test=400)
 48    dev = device()
 49    try:
 50        model = make_model('rnn_small', ds['input_shape'], ds['out_dim']).to(dev)
 51        xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
 52        xte, yte = ds['xte'].to(dev), ds['yte'].to(dev)
 53        # rnn_small consumes [N, 8, 3] despite flattened dataset storage.
 54        xtr, xte = xtr.reshape(-1, 8, 3), xte.reshape(-1, 8, 3)
 55        opt = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)
 56        loss_fn = nn.MSELoss()
 57        rng = np.random.default_rng(seed + 91)
 58        old = flat_params(model).numpy()[None, :]
 59        ratios, displacements, etas = [], [], []
 60        eta = lr; eta_min, eta_max = lr * .25, lr * 1.5
 61        step = 0
 62        for ep in range(EPOCHS):
 63            for ib in batches(len(xtr), BATCH, rng):
 64                xb, yb = xtr[ib], ytr[ib]
 65                opt.zero_grad(set_to_none=True)
 66                pred = model(xb)
 67                loss = loss_fn(pred, yb)
 68                loss.backward()
 69                if controlled:
 70                    # The intervention is a speed-limited SGD-like update. Adam's
 71                    # moments are deliberately not used on this side.
 72                    with torch.no_grad():
 73                        for p in model.parameters():
 74                            if p.grad is not None:
 75                                noise = torch.randn_like(p) * math.sqrt(2.0 * eta * 1e-4)
 76                                p.add_(p.grad, alpha=-eta)
 77                                p.add_(noise)
 78                else:
 79                    opt.step()
 80                step += 1
 81                if controlled and step % 4 == 0:
 82                    now = flat_params(model).numpy()[None, :]
 83                    delta = now - old; dt = 4.0 * max(eta, 1e-9); D = 1e-4
 84                    sigma = float(np.mean((delta / dt) ** 2) / D)
 85                    w2 = sliced_w2(old, now, n_proj=32, rng=np.random.default_rng(seed + step))
 86                    rhs = max(D * dt * sigma * dt, 1e-12)
 87                    r = w2 / rhs
 88                    ratios.append(r); displacements.append(w2); etas.append(eta)
 89                    # cautious feedback: high empirical motion/action ratio lowers rate.
 90                    target = .72
 91                    factor = float(np.clip((target / max(r, 1e-8)) ** .25, .65, 1.18))
 92                    eta = float(np.clip(eta * factor, eta_min, eta_max))
 93                    old = now
 94        with torch.no_grad():
 95            metric = float(loss_fn(model(xte), yte).cpu())
 96        if collect:
 97            return metric, {'mean_ratio': float(np.mean(ratios)) if ratios else float('nan'),
 98                            'max_ratio': float(np.max(ratios)) if ratios else float('nan'),
 99                            'final_eta': eta, 'mean_displacement': float(np.mean(displacements)) if displacements else float('nan'),
100                            'n_intervals': len(ratios)}
101        return metric
102    except Exception:
103        # Robust CPU fallback required by the bench environment.
104        if dev == 'cuda':
105            torch.cuda.empty_cache()
106            torch.cuda.is_available = lambda: False
107            return run(seed, lr, controlled, weight_decay, collect)
108        raise
109
110
111def main():
112    baseline_grid = [{'lr': lr, 'weight_decay': wd} for lr in LRS for wd in WEIGHT_DECAYS]
113    base = sweep_baseline(lambda cfg: lambda s: run(s, cfg['lr'], False, cfg['weight_decay']), baseline_grid, seeds=SWEEP_SEEDS)
114    # Idea uses the baseline's selected lr plus two nearby/shared settings.
115    idea_grid = sorted(set(LRS + [base['best_cfg']['lr']]))
116    idea_cfg = min(idea_grid, key=lambda lr: np.mean([run(s, lr, True, 0.0) for s in SWEEP_SEEDS]))
117    idea = evaluate(lambda s: run(s, idea_cfg, True, 0.0), seeds=SEEDS)
118    # trained-model signature, measured independently on all paired models at selected settings
119    sig = [run(s, idea_cfg, True, 0.0, True)[1] for s in SEEDS]
120    signature = {'quantity': 'W2^2/(D*dt*Sigma)', 'predicted': '<= 1 in stable intervals',
121                 'observed_mean_ratio': float(np.nanmean([x['mean_ratio'] for x in sig])),
122                 'observed_max_ratio_mean': float(np.nanmean([x['max_ratio'] for x in sig])),
123                 'observed_final_eta_mean': float(np.mean([x['final_eta'] for x in sig])),
124                 'confirmed': bool(np.nanmean([x['mean_ratio'] for x in sig]) <= 1.25)}
125    report = make_report('dynamics', 'rnn_small', base, idea, {'signature': signature, 'idea_grid': idea_grid,
126        'baseline_grid': baseline_grid, 'protocol': '8 paired seeds; 4-seed baseline/idea selection'})
127    report['idea_configs_considered'] = [{'lr': x} for x in idea_grid]
128    Path('bench_report.json').write_text(json.dumps(report, indent=2))
129    print(json.dumps(report, indent=2))
130
131if __name__ == '__main__': main()