Shell-Wise Balanced MoE Routing / shell_moe_experiment.py
Mechanism confirmed, baseline not beaten
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()