KL Mirror-Prox for coupled routing / run_bench.py

Failed on benchmark

Raw ⬇ ZIP
  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()