Contact-Splitting Momentum Optimizer / bench_contact.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random, sys
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8import bench
  9
 10SEEDS = tuple(range(8))
 11# Identical learning-rate union on both sides; baseline also sweeps its beta knob.
 12IDEA_GRID = [{'lr': lr, 'gamma': 0.10} for lr in (0.003, 0.006, 0.012)]
 13BASE_GRID = [{'lr': lr, 'beta': beta} for lr in (0.003, 0.006, 0.012) for beta in (0.85, 0.95)]
 14EPOCHS, BATCH = 18, 64
 15
 16
 17def seed_all(seed):
 18    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
 19    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 20
 21
 22def make_net(ds, seed):
 23    seed_all(seed)
 24    return bench.make_model('mlp_tiny', tuple(ds['input_shape']), int(ds['out_dim']))
 25
 26
 27def loss_fn(out, y, task):
 28    return nn.CrossEntropyLoss()(out, y.long().view(-1)) if task == 'classification' else nn.MSELoss()(out, y)
 29
 30
 31def run(seed, method, cfg, signature=False):
 32    seed_all(seed)
 33    ds = bench.get_dataset('tabular', seed, n_train=400, n_test=200)
 34    dev = 'cuda' if torch.cuda.is_available() else 'cpu'
 35    try:
 36        return _run(seed, ds, method, cfg, dev, signature)
 37    except Exception:
 38        if dev == 'cuda':
 39            return _run(seed, ds, method, cfg, 'cpu', signature)
 40        raise
 41
 42
 43def _run(seed, ds, method, cfg, dev, signature):
 44    net = make_net(ds, seed).to(dev)
 45    xtr, ytr = ds['xtr'].to(dev), ds['ytr'].to(dev)
 46    xte, yte = ds['xte'].to(dev), ds['yte'].to(dev)
 47    params = [p for p in net.parameters() if p.requires_grad]
 48    # M=I is the prescribed initial MVP; p is the contact momentum.
 49    p = [torch.zeros_like(q) for q in params]
 50    s = 0.0
 51    rng = np.random.default_rng(seed + 991)
 52    cert_rates, certs = [], []
 53    grad_evals = 0
 54
 55    def grad_batch(xb, yb):
 56        nonlocal grad_evals
 57        net.zero_grad(set_to_none=True)
 58        out = net(xb); loss = loss_fn(out, yb, ds['task']); loss.backward()
 59        grad_evals += 1
 60        return float(loss.detach().cpu()), [q.grad.detach().clone() for q in params]
 61
 62    def test_metric():
 63        net.eval()
 64        with torch.no_grad():
 65            val = loss_fn(net(xte), yte, ds['task'])
 66        net.train(); return float(val.cpu())
 67
 68    net.train()
 69    n = xtr.shape[0]
 70    for epoch in range(EPOCHS):
 71        order = rng.permutation(n)
 72        for st in range(0, n, BATCH):
 73            ix = torch.as_tensor(order[st:st+BATCH], device=dev)
 74            xb, yb = xtr[ix], ytr[ix]
 75            lr = float(cfg['lr'])
 76            if method == 'baseline':
 77                # Two ordinary momentum updates, matching the two contact gradients.
 78                for j in range(2):
 79                    f, g = grad_batch(xb, yb)
 80                    beta = float(cfg['beta'])
 81                    with torch.no_grad():
 82                        for k in range(len(params)):
 83                            p[k].mul_(beta).add_(g[k], alpha=-lr / 2.0)
 84                            params[k].add_(p[k])
 85            else:
 86                gamma = float(cfg['gamma'])
 87                # K(h/2): x += h p/2 and s += h K/2.
 88                with torch.no_grad():
 89                    kin = 0.5 * sum((q*q).sum() for q in p)
 90                    s += (lr / 2.0) * float(kin.cpu())
 91                    for q, v in zip(params, p): q.add_(v, alpha=lr / 2.0)
 92                f1, g1 = grad_batch(xb, yb)
 93                with torch.no_grad():
 94                    for v, gg in zip(p, g1): v.add_(gg, alpha=-lr / 2.0)
 95                    s -= lr / 2.0 * f1
 96                    damp = math.exp(-gamma * lr)
 97                    for v in p: v.mul_(damp)
 98                    s *= damp
 99                f2, g2 = grad_batch(xb, yb)
100                with torch.no_grad():
101                    for v, gg in zip(p, g2): v.add_(gg, alpha=-lr / 2.0)
102                    s -= lr / 2.0 * f2
103                    kin = 0.5 * sum((q*q).sum() for q in p)
104                    s += (lr / 2.0) * float(kin.cpu())
105                    for q, v in zip(params, p): q.add_(v, alpha=lr / 2.0)
106                H = float((0.5 * sum((q*q).sum() for q in p)).cpu()) + f2 + gamma*s
107                if certs and abs(certs[-1]) > 1e-8 and abs(H) > 1e-8 and certs[-1]*H > 0:
108                    cert_rates.append(-(math.log(abs(H))-math.log(abs(certs[-1]))) / lr)
109                certs.append(H)
110    metric = test_metric()
111    result = {'metric': metric, 'grad_evals': grad_evals}
112    if signature:
113        result['cert_rate_mean'] = float(np.mean(cert_rates)) if cert_rates else float('nan')
114        result['cert_rate_n'] = len(cert_rates)
115    return result
116
117
118def evaluate(method, cfg, seeds=SEEDS, signature=False):
119    vals, rows = [], []
120    for seed in seeds:
121        r = run(seed, method, cfg, signature)
122        vals.append(r['metric']); rows.append(r)
123    return {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals), 'details': rows}
124
125
126def main():
127    # Fair small baseline sweep on four seeds, then full paired evaluation.
128    sweep = []
129    for cfg in BASE_GRID:
130        r = evaluate('baseline', cfg, seeds=(0,1,2,3))
131        sweep.append({'cfg': cfg, 'mean': r['mean']})
132    best = min(sweep, key=lambda z: z['mean'])['cfg']
133    base_full = evaluate('baseline', best, SEEDS)
134    base_block = {'best_cfg': best, 'sweep': sweep, 'full': base_full}
135
136    idea_sweep = []
137    for cfg in IDEA_GRID:
138        # Union parity is satisfied because all idea lrs occur in BASE_GRID.
139        r = evaluate('idea', cfg, seeds=(0,1,2,3))
140        idea_sweep.append({'cfg': cfg, 'mean': r['mean']})
141    idea_best = min(idea_sweep, key=lambda z: z['mean'])['cfg']
142    idea_full = evaluate('idea', idea_best, SEEDS, signature=True)
143    # Re-test the stage-1 prediction on trained models: expected empirical rate ~ gamma.
144    rates = [d['cert_rate_mean'] for d in idea_full['details'] if np.isfinite(d['cert_rate_mean'])]
145    observed = float(np.mean(rates)) if rates else float('nan')
146    gamma = idea_best['gamma']
147    sig = {'prediction': 'certificate conformal rate ~= gamma', 'predicted': gamma,
148           'observed': observed, 'absolute_error': abs(observed-gamma) if np.isfinite(observed) else None,
149           'n_models': len(rates), 'confirmed': bool(np.isfinite(observed) and abs(observed-gamma) <= 0.08)}
150    report = bench.make_report('tabular', 'mlp_tiny', base_block, idea_full,
151                               {'idea_sweep': idea_sweep, 'mechanism_signature': sig,
152                                'budget': {'epochs': EPOCHS, 'batch': BATCH, 'baseline_gradient_evals': base_full['details'][0]['grad_evals'], 'idea_gradient_evals': idea_full['details'][0]['grad_evals']}})
153    Path('bench_report.json').write_text(json.dumps(report, indent=2, allow_nan=False))
154    print(json.dumps(report, indent=2, allow_nan=False))
155
156if __name__ == '__main__':
157    main()