Shell-Wise Balanced MoE Routing / shell_moe_bench_track.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import numpy as np
 2
 3META = {
 4    "name": "correlated_token_moe_regression",
 5    "domain": "moe-routing",
 6    "description": "Correlated token groups with regime-dependent nonlinear targets for expert routing."
 7}
 8
 9
10def get_dataset(seed, n_train, n_test):
11    def make(n, s):
12        rng = np.random.RandomState(s)
13        x = rng.normal(size=(n, 16, 4)).astype(np.float32)
14        regime = (x[:, :, 0].mean(1) > 0).astype(np.float32)
15        token = np.where(
16            regime[:, None] > 0,
17            np.sin(x[:, :, 0]) + 0.45 * x[:, :, 1] ** 2,
18            np.cos(x[:, :, 1]) - 0.45 * x[:, :, 0] ** 2,
19        )
20        y = (
21            token.mean(1)
22            + 0.25 * x[:, :, 2].mean(1)
23            + 0.15 * x[:, :, 3].mean(1)
24            + rng.normal(0, 0.06, n)
25        ).astype(np.float32)
26        return x, y[:, None]
27
28    xtr, ytr = make(n_train, seed)
29    xte, yte = make(n_test, seed + 5000)
30    return {
31        "xtr": xtr,
32        "ytr": ytr,
33        "xte": xte,
34        "yte": yte,
35        "task": "regression",
36        "metric": "mse",
37        "out_dim": 1,
38    }