Contact-Splitting Momentum Optimizer / bench_contact.py
Failed on benchmark
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()