KL Mirror-Prox for coupled routing / run_bench.py
Failed on benchmark
1import sys, json, math, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
9from bench import get_dataset, evaluate, sweep_baseline, make_report
10
11TRACK = 'correlated_token_moe_regression'
12NTR, NTE = 400, 200
13EPOCHS, BATCH = 18, 128
14DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
15
16class MoE(nn.Module):
17 def __init__(self, d=4, experts=3, hidden=32):
18 super().__init__()
19 self.experts = experts
20 self.gate = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, experts))
21 self.ex = nn.ModuleList([nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1)) for _ in range(experts)])
22 def forward(self, x, prior=None):
23 logits = self.gate(x)
24 if prior is not None:
25 logits = logits + torch.log(prior.clamp_min(1e-8)).view(1, 1, -1)
26 a = F.softmax(logits, dim=-1)
27 vals = torch.stack([e(x).squeeze(-1) for e in self.ex], dim=-1)
28 return (a * vals).sum(-1).mean(-1), a
29
30def set_seed(s):
31 random.seed(s); np.random.seed(s); torch.manual_seed(s)
32 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
33
34def prox(p, c, eta):
35 return torch.softmax(torch.log(p.clamp_min(1e-12)) - eta * c, dim=-1)
36
37def run(seedno, mode, lr, eta=0.1, collect=False):
38 set_seed(seedno)
39 d = get_dataset(TRACK, seedno, NTR, NTE)
40 xtr, ytr = d['xtr'], d['ytr'].squeeze(-1)
41 xte, yte = d['xte'], d['yte'].squeeze(-1)
42 dev = torch.device(DEVICE)
43 try:
44 net = MoE().to(dev)
45 opt = torch.optim.Adam(net.parameters(), lr=lr)
46 prior = torch.ones(3, device=dev) / 3
47 for ep in range(EPOCHS):
48 perm = torch.randperm(len(xtr))
49 for st in range(0, len(xtr), BATCH):
50 ix = perm[st:st+BATCH]
51 xb, yb = xtr[ix].to(dev), ytr[ix].to(dev)
52 pred, a = net(xb, prior if mode == 'mp' else None)
53 loss = F.mse_loss(pred, yb)
54 opt.zero_grad(); loss.backward(); opt.step()
55 with torch.no_grad():
56 load = a.detach().mean((0, 1))
57 cost = (load - 1/3) + 0.25 * torch.relu(load - 0.45)
58 if mode == 'mp':
59 z = prox(prior, cost, eta)
60 zload = load + (z - prior)
61 zcost = (zload - 1/3) + 0.25 * torch.relu(zload - 0.45)
62 prior = prox(prior, zcost, eta)
63 else:
64 prior = prox(prior, cost, eta)
65 with torch.no_grad():
66 pred, a = net(xte.to(dev), prior if mode == 'mp' else None)
67 metric = float(F.mse_loss(pred, yte.to(dev)).cpu())
68 load = a.mean((0,1)).cpu().numpy()
69 train_pred, train_a = net(xtr.to(dev), prior if mode == 'mp' else None)
70 train_load = train_a.mean((0,1)).cpu().numpy()
71 out = {'metric': metric, 'load_var': float(np.var(load)), 'load': load.tolist(), 'train_load': train_load.tolist(), 'prior': prior.cpu().numpy().tolist()}
72 return out if collect else metric
73 except Exception:
74 if str(dev) != 'cpu':
75 return run_cpu(seedno, mode, lr, eta, collect)
76 return float('nan') if not collect else {'metric': float('nan')}
77
78def run_cpu(seedno, mode, lr, eta=0.1, collect=False):
79 global DEVICE
80 old = DEVICE
81 DEVICE = 'cpu'
82 try:
83 return run(seedno, mode, lr, eta, collect)
84 finally:
85 DEVICE = old
86
87def make_fn(mode):
88 return lambda cfg: (lambda s: run(s, mode, cfg['lr'], cfg.get('eta', 0.1)))
89
90def main():
91 # Shared union: every idea lr is also evaluated by baseline.
92 lrs = [0.001, 0.003, 0.006]
93 base_grid = [{'lr': x, 'eta': e} for x in lrs for e in [0.03, 0.1, 0.3]]
94 idea_grid = [{'lr': 0.001, 'eta': 0.03}, {'lr': 0.003, 'eta': 0.1}, {'lr': 0.006, 'eta': 0.3}]
95 base = sweep_baseline(make_fn('md'), base_grid)
96 best = min(idea_grid, key=lambda c: evaluate(make_fn('mp')(c), seeds=(0,1,2,3))['mean'])
97 idea = evaluate(make_fn('mp')(best))
98 base_full = base['full']
99 sig = []
100 for s in range(8):
101 b = run(s, 'md', base['best_cfg']['lr'], base['best_cfg'].get('eta',0.1), True)
102 i = run(s, 'mp', best['lr'], best['eta'], True)
103 sig.append({'seed':s, 'baseline':b, 'idea':i})
104 pred = np.array([z['idea']['prior'] for z in sig])
105 obs = np.array([z['idea']['load'] for z in sig])
106 signature = {'predicted_prior_vs_observed_load_mae': float(np.mean(np.abs(pred-obs))), 'predicted_prior_mean': pred.mean(0).tolist(), 'observed_load_mean': obs.mean(0).tolist(), 'confirmed': bool(np.mean(np.abs(pred-obs)) < 0.12), 'note':'Measured on trained MoE systems; lower MAE means the coupled prior tracked realized routing load.'}
107 report = make_report(TRACK, 'custom_moe_shared', base, idea, {'mechanism_signature': signature, 'selection': {'baseline_union_grid':base_grid, 'idea_grid':idea_grid, 'idea_best':best}, 'per_seed_behavior':sig})
108 Path('bench_report.json').write_text(json.dumps(report, indent=2))
109 print(json.dumps(report, indent=2))
110
111if __name__ == '__main__': main()