import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn torch.set_num_threads(2) sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import sweep_baseline, evaluate, make_report, get_dataset META = {"name": "congestion_moe_sequence", "domain": "moe-routing", "description": "Registered multi-token expert regression track"} DEVICE = "cuda" if torch.cuda.is_available() else "cpu" # CUDA initialization is known to be flaky on this host; train_model-style fallback. def safe_device(): if DEVICE != "cuda": return "cpu" try: torch.zeros(1, device="cuda").sum().item() return "cuda" except Exception: return "cpu" def project_simplex(v, z=1.0): # Batched Euclidean projection, differentiable almost everywhere. u, _ = torch.sort(v, dim=-1, descending=True) cssv = torch.cumsum(u, dim=-1) - z ind = torch.arange(1, v.shape[-1] + 1, device=v.device, dtype=v.dtype) cond = u - cssv / ind > 0 rho = cond.sum(dim=-1).clamp_min(1).long() - 1 theta = cssv.gather(-1, rho.unsqueeze(-1)).squeeze(-1) / (rho + 1).to(v.dtype) return (v - theta.unsqueeze(-1)).clamp_min(0) class TinyMoE(nn.Module): def __init__(self, d=4, h=8, experts=4, temperature=0.7, idea=False, steps=4, eta=0.8): super().__init__() self.experts = nn.ModuleList([nn.Sequential(nn.Linear(d, h), nn.Tanh(), nn.Linear(h, h), nn.Tanh()) for _ in range(experts)]) self.router = nn.Linear(d, experts) self.head = nn.Sequential(nn.Linear(h, h), nn.Tanh(), nn.Linear(h, 1)) self.temperature, self.idea, self.steps, self.eta = temperature, idea, steps, eta self.last_load = None self.last_soft_load = None def forward(self, x): # Sequence-level groups: each sequence is one routing player; token outputs # are mixed using the same group allocation, preserving multi-token structure. b, l, d = x.shape logits = self.router(x).mean(dim=1) / max(self.temperature, 1e-5) a = torch.exp(logits - logits.detach().amax(dim=-1, keepdim=True)) soft = torch.softmax(logits, dim=-1) if not self.idea: alloc = soft else: # b_i is a positive capacity prior; equal capacity is standard in this tiny bench. cap = torch.full((a.shape[-1],), 0.35, device=x.device, dtype=x.dtype) alloc = torch.full_like(a, 1.0 / a.shape[-1]) for _ in range(self.steps): D = cap + alloc.sum(dim=0) g = a * (D.unsqueeze(0) - alloc) / (D.unsqueeze(0).square() + 1e-8) old_pay = (a * alloc / D.unsqueeze(0)).sum() step = self.eta proposal = project_simplex(alloc + step * g) new_pay = (a * proposal / (cap + proposal.sum(dim=0)).unsqueeze(0)).sum() # conservative backtracking, while retaining gradients through accepted proposal for _ in range(5): if new_pay.detach() >= old_pay.detach() - 1e-7: break step *= 0.5 proposal = project_simplex(alloc + step * g) new_pay = (a * proposal / (cap + proposal.sum(dim=0)).unsqueeze(0)).sum() alloc = proposal tok = torch.stack([e(x.reshape(-1, d)).reshape(b, l, -1) for e in self.experts], dim=2) pooled = (tok * alloc[:, None, :, None]).sum(dim=2).mean(dim=1) self.last_load = alloc.detach().sum(0).cpu().numpy() self.last_soft_load = soft.detach().sum(0).cpu().numpy() return self.head(pooled) def run_one(seed, lr, temperature, idea, epochs=4, return_model=False): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) ds = get_dataset("congestion_moe_sequence", seed, 240, 100) dev = safe_device() model = TinyMoE(temperature=temperature, idea=idea).to(dev) xtr, ytr = torch.tensor(ds["xtr"], dtype=torch.float32, device=dev), torch.tensor(ds["ytr"], dtype=torch.float32, device=dev) xte, yte = torch.tensor(ds["xte"], dtype=torch.float32, device=dev), torch.tensor(ds["yte"], dtype=torch.float32, device=dev) opt = torch.optim.Adam(model.parameters(), lr=lr) model.train() for ep in range(epochs): g = torch.Generator(device="cpu"); g.manual_seed(seed * 100 + ep) ix = torch.randperm(len(xtr), generator=g, device="cpu").to(dev) for start in range(0, len(ix), 128): j = ix[start:start+128] loss = (model(xtr[j]) - ytr[j]).square().mean() opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): metric = float((model(xte) - yte).square().mean().cpu()) if return_model: return metric, model del model if dev == "cuda": torch.cuda.empty_cache() return metric def main(): # Union parity: all idea learning rates are also baseline candidates; baseline's # decisive temperature knob is swept as well. lrs = [0.0015, 0.003, 0.006] temps = [0.5, 0.7, 1.0] base_grid = [{"lr": lr, "temperature": t} for lr in lrs for t in temps] idea_grid = [{"lr": lr, "temperature": 0.7} for lr in lrs] def base_factory(cfg): return lambda seed: run_one(seed, cfg["lr"], cfg["temperature"], False) base = sweep_baseline(base_factory, base_grid) # Required three idea settings, evaluated on all eight paired seeds. idea_runs = [] for cfg in idea_grid: r = evaluate(lambda s, c=cfg: run_one(s, c["lr"], c["temperature"], True)) idea_runs.append({"cfg": cfg, "result": r}) best = min(idea_runs, key=lambda q: q["result"]["mean"]) idea = best["result"] # Signature is measured from trained systems on the same held-out task inputs. sig_rows = [] for s in range(8): bm, bmodel = run_one(s, base["best_cfg"]["lr"], base["best_cfg"]["temperature"], False, return_model=True) im, imodel = run_one(s, best["cfg"]["lr"], best["cfg"]["temperature"], True, return_model=True) sig_rows.append({"seed": s, "baseline_load_cv": float(np.std(bmodel.last_load)/(np.mean(bmodel.last_load)+1e-9)), "idea_load_cv": float(np.std(imodel.last_load)/(np.mean(imodel.last_load)+1e-9)), "baseline_mse": bm, "idea_mse": im}) base_c = base["full"]; cmp = __import__("bench").compare_results(base_c, idea) observed_cv_delta = float(np.mean([r["idea_load_cv"]-r["baseline_load_cv"] for r in sig_rows])) signature = {"prediction": "congestion-aware routing reduces expert-load CV while retaining a fixed row budget", "predicted_direction": "negative", "observed_load_cv_delta": observed_cv_delta, "per_seed": sig_rows, "confirmed": bool(observed_cv_delta < 0)} report = make_report("congestion_moe_sequence", "tiny_moe", base, idea, {"idea_sweep": idea_runs, **signature}) report["custom_track"] = {"name": META["name"], "file": "custom_moe_track.py", "domain": META["domain"]} report["device"] = safe_device(); report["epochs"] = 4; report["batch"] = 128 Path("bench_report.json").write_text(json.dumps(report, indent=2)) print(json.dumps(report, indent=2)) if __name__ == "__main__": main()