import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, evaluate, sweep_baseline, make_report TRACK = 'correlated_token_moe_regression' NTR, NTE = 400, 200 EPOCHS, BATCH = 18, 128 DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' class MoE(nn.Module): def __init__(self, d=4, experts=3, hidden=32): super().__init__() self.experts = experts self.gate = nn.Sequential(nn.Linear(d, 32), nn.Tanh(), nn.Linear(32, experts)) self.ex = nn.ModuleList([nn.Sequential(nn.Linear(d, hidden), nn.Tanh(), nn.Linear(hidden, 1)) for _ in range(experts)]) def forward(self, x, prior=None): logits = self.gate(x) if prior is not None: logits = logits + torch.log(prior.clamp_min(1e-8)).view(1, 1, -1) a = F.softmax(logits, dim=-1) vals = torch.stack([e(x).squeeze(-1) for e in self.ex], dim=-1) return (a * vals).sum(-1).mean(-1), a def set_seed(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def prox(p, c, eta): return torch.softmax(torch.log(p.clamp_min(1e-12)) - eta * c, dim=-1) def run(seedno, mode, lr, eta=0.1, collect=False): set_seed(seedno) d = get_dataset(TRACK, seedno, NTR, NTE) xtr, ytr = d['xtr'], d['ytr'].squeeze(-1) xte, yte = d['xte'], d['yte'].squeeze(-1) dev = torch.device(DEVICE) try: net = MoE().to(dev) opt = torch.optim.Adam(net.parameters(), lr=lr) prior = torch.ones(3, device=dev) / 3 for ep in range(EPOCHS): perm = torch.randperm(len(xtr)) for st in range(0, len(xtr), BATCH): ix = perm[st:st+BATCH] xb, yb = xtr[ix].to(dev), ytr[ix].to(dev) pred, a = net(xb, prior if mode == 'mp' else None) loss = F.mse_loss(pred, yb) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): load = a.detach().mean((0, 1)) cost = (load - 1/3) + 0.25 * torch.relu(load - 0.45) if mode == 'mp': z = prox(prior, cost, eta) zload = load + (z - prior) zcost = (zload - 1/3) + 0.25 * torch.relu(zload - 0.45) prior = prox(prior, zcost, eta) else: prior = prox(prior, cost, eta) with torch.no_grad(): pred, a = net(xte.to(dev), prior if mode == 'mp' else None) metric = float(F.mse_loss(pred, yte.to(dev)).cpu()) load = a.mean((0,1)).cpu().numpy() train_pred, train_a = net(xtr.to(dev), prior if mode == 'mp' else None) train_load = train_a.mean((0,1)).cpu().numpy() out = {'metric': metric, 'load_var': float(np.var(load)), 'load': load.tolist(), 'train_load': train_load.tolist(), 'prior': prior.cpu().numpy().tolist()} return out if collect else metric except Exception: if str(dev) != 'cpu': return run_cpu(seedno, mode, lr, eta, collect) return float('nan') if not collect else {'metric': float('nan')} def run_cpu(seedno, mode, lr, eta=0.1, collect=False): global DEVICE old = DEVICE DEVICE = 'cpu' try: return run(seedno, mode, lr, eta, collect) finally: DEVICE = old def make_fn(mode): return lambda cfg: (lambda s: run(s, mode, cfg['lr'], cfg.get('eta', 0.1))) def main(): # Shared union: every idea lr is also evaluated by baseline. lrs = [0.001, 0.003, 0.006] base_grid = [{'lr': x, 'eta': e} for x in lrs for e in [0.03, 0.1, 0.3]] idea_grid = [{'lr': 0.001, 'eta': 0.03}, {'lr': 0.003, 'eta': 0.1}, {'lr': 0.006, 'eta': 0.3}] base = sweep_baseline(make_fn('md'), base_grid) best = min(idea_grid, key=lambda c: evaluate(make_fn('mp')(c), seeds=(0,1,2,3))['mean']) idea = evaluate(make_fn('mp')(best)) base_full = base['full'] sig = [] for s in range(8): b = run(s, 'md', base['best_cfg']['lr'], base['best_cfg'].get('eta',0.1), True) i = run(s, 'mp', best['lr'], best['eta'], True) sig.append({'seed':s, 'baseline':b, 'idea':i}) pred = np.array([z['idea']['prior'] for z in sig]) obs = np.array([z['idea']['load'] for z in sig]) 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.'} 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}) Path('bench_report.json').write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()