Shell-Wise Balanced MoE Routing / shell_moe_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 1234
  6
  7def shell_ids(conf, K):
  8    order = np.argsort(conf, kind='stable')
  9    out = np.empty(len(conf), dtype=np.int64)
 10    for j, i in enumerate(order):
 11        out[i] = min(K - 1, (j * K) // len(conf))
 12    return out
 13
 14def shell_balance_route(logits, shells, K):
 15    """Top-1 exact per-shell balanced routing via greedy maximum-weight assignment."""
 16    n, e = logits.shape
 17    route = np.empty(n, dtype=np.int64)
 18    reassigned = 0
 19    score_loss = 0.0
 20    for k in range(K):
 21        ids = np.flatnonzero(shells == k)
 22        b = len(ids)
 23        if b == 0:
 24            continue
 25        target = np.full(e, b // e, dtype=int)
 26        rem = b - target.sum()
 27        demand = np.bincount(np.argmax(logits[ids], axis=1), minlength=e)
 28        for x in np.argsort(-demand, kind='stable')[:rem]:
 29            target[x] += 1
 30        pairs = [(float(logits[i, x]), int(i), x) for i in ids for x in range(e)]
 31        pairs.sort(key=lambda t: (-t[0], t[1], t[2]))
 32        used = np.zeros(e, dtype=int)
 33        assigned = set()
 34        for val, i, x in pairs:
 35            if i not in assigned and used[x] < target[x]:
 36                route[i] = x
 37                used[x] += 1
 38                assigned.add(i)
 39        assert len(assigned) == b and np.array_equal(used, target)
 40        pref = np.argmax(logits[ids], axis=1)
 41        reassigned += int(np.sum(route[ids] != pref))
 42        score_loss += float(np.sum(logits[ids, pref] - logits[ids, route[ids]]))
 43    return route, reassigned / n, score_loss / n
 44
 45def cv(route, E):
 46    loads = np.bincount(route, minlength=E).astype(float)
 47    return float(loads.std() / loads.mean()), loads.tolist()
 48
 49def toy_predictions():
 50    rng = np.random.default_rng(SEED)
 51    rows = []
 52    # Eq. (66) prediction: quota counts differ by at most one in every shell,
 53    # for all tested expert counts, shell counts, and batch sizes.
 54    for E in [3, 5, 8]:
 55        for K in [2, 4, 7]:
 56            n = 997
 57            logits = rng.normal(size=(n, E))
 58            sh = shell_ids(np.max(logits, axis=1), K)
 59            route, _, _ = shell_balance_route(logits, sh, K)
 60            worst = 0
 61            for k in range(K):
 62                counts = np.bincount(route[sh == k], minlength=E)
 63                worst = max(worst, int(counts.max() - counts.min()))
 64            rows.append({'prediction':'per-shell count spread <= 1', 'E':E, 'K':K,
 65                         'observed':worst, 'predicted_max':1})
 66
 67    # Eq. (66) also predicts identical polynomial marginals.  For each shell,
 68    # evaluate p_e(P)=sum_k count[e,k] P^k(1-P)^(K-k) on a probability grid.
 69    for E, K in [(4, 3), (8, 5), (8, 8)]:
 70        n = E * K * 32
 71        logits = rng.normal(size=(n, E))
 72        sh = shell_ids(np.max(logits, axis=1), K)
 73        route, _, _ = shell_balance_route(logits, sh, K)
 74        coeff = np.zeros((E, K))
 75        base_coeff = np.zeros((E, K))
 76        base = np.argmax(logits, axis=1)
 77        for k in range(K):
 78            for e in range(E):
 79                coeff[e, k] = np.sum((route == e) & (sh == k))
 80                base_coeff[e, k] = np.sum((base == e) & (sh == k))
 81        ps = np.linspace(0.01, 0.99, 99)
 82        idea_spread = []
 83        base_spread = []
 84        for P in ps:
 85            q = np.array([P**k * (1-P)**(K-k) for k in range(K)])
 86            idea_spread.append(float(np.ptp(coeff @ q)))
 87            base_spread.append(float(np.ptp(base_coeff @ q)))
 88        rows.append({'prediction':'Eq.66 marginal spread is zero for all P', 'E':E, 'K':K,
 89                     'observed_max_spread':max(idea_spread), 'predicted':0.0,
 90                     'baseline_max_spread':max(base_spread)})
 91    return rows
 92
 93def comparison():
 94    rng = np.random.default_rng(SEED + 1)
 95    E, K, n = 8, 8, 4096
 96    # Correlated logits intentionally create expert monopoly in confidence shells.
 97    common = rng.normal(size=(n, 1))
 98    logits = rng.normal(scale=.45, size=(n, E))
 99    logits[:, 0] += 1.8 * np.maximum(common[:, 0], 0)
100    logits[:, 1] += 1.1 * np.maximum(-common[:, 0], 0)
101    conf = np.max(logits, axis=1)
102    shells = shell_ids(conf, K)
103    base = np.argmax(logits, axis=1)
104    idea, frac, loss = shell_balance_route(logits, shells, K)
105    bcv, bloads = cv(base, E)
106    icv, iloads = cv(idea, E)
107    shell_cvs = []
108    for k in range(K):
109        ids = np.flatnonzero(shells == k)
110        shell_cvs.append({'shell':k, 'baseline_cv':cv(base[ids], E)[0], 'idea_cv':cv(idea[ids], E)[0], 'size':len(ids)})
111    return {'baseline_cv':bcv, 'idea_cv':icv, 'baseline_loads':bloads, 'idea_loads':iloads,
112            'reassigned_fraction':frac, 'router_score_loss_per_token':loss, 'shells':shell_cvs}
113
114def main():
115    out = {'seed':SEED, 'toy_predictions':toy_predictions(), 'comparison':comparison()}
116    Path('results.json').write_text(json.dumps(out, indent=2))
117    print(json.dumps(out, indent=2))
118
119if __name__ == '__main__':
120    main()