Skip to content

Mixture of Experts: routing, sorted dispatch, Switch aux loss, aux-free bias

ModuleL7.8 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/modern/moe.py: MoE, topk_ids, dispatch, combine, load_balance_loss; and your own oracle tests in python/tests/l7-8-moe/
Contractcourse/contracts/py/tinyllm/modern/moe.pyi
Testscourse/tests/L7.8/test_moe.py (what they check: section 4), golden values from transformers 5.19.0 MixtralSparseMoeBlock (and load_balancing_loss_func), Qwen3MoeSparseMoeBlock, and DeepseekV3MoE in course/fixtures/L7.8/moe_hf.npz (course/oracle/L7.8/moe_hf.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL7.2 GatedMLP (every expert) · L0.4 Linear, ModuleList, Module · L0.2 the op library · L0.1 Tensor · M06.3 PCG32 (or --ref-deps)
Used byL7.9 builds MoE layers when tl_num_experts > 0 · later: C1’s MoE-vs-dense ablation at equal active parameters
MilestoneMS-L7 (your decoder loads and matches Hugging Face checkpoints)
Optional depthShazeer et al., “Outrageously Large Neural Networks” (2017); Fedus et al., “Switch Transformers” (2021), sections 2.1 and 2.2; Wang et al., “Auxiliary-Loss-Free Load Balancing” (2024); DeepSeek-V3 report, section 2.1.2
  • A router picks the top kk of EE experts per token and weighs their outputs; parameters grow with EE, compute with kk (test_hand_example).
  • Sorted dispatch gives each expert one contiguous slice of tokens and puts results back with the inverse permutation, and equals the per-token loop exactly (test_sorted_dispatch_equals_dense_loop).
  • The Switch loss E∑efePeE \sum_e f_e P_e is kk at perfect balance; ff is a count with no gradient, so the router learns through PP (test_load_balance_loss_bounds_and_gradient).
  • DeepSeek-V3’s bias chooses experts but never weighs them, and a sign rule balances the load with no loss term (test_aux_free_bias, test_bias_updates_balance_the_load).
Terminal window
ol start L7.8 # stubs moe.py; prints your test path and rung (R5)
ol tests L7.8 # the course tests
# write your oracle tests in python/tests/l7-8-moe/, then:
ol check L7.8 # course tests and the mutation grade of your tests
ol diff L7.8 # after passing: your code against the reference

Your gated MLP (L7.2) holds two thirds of a Llama block’s parameters, and every token pays for all of them. Mixtral, Qwen-MoE, and DeepSeek-V3 replace it with many smaller MLPs and a router: each token uses two (or eight of 256), so a model can hold far more knowledge at the same cost per token. C1 asks a concrete question with it: at equal active parameters, does an MoE beat the dense MLP on TinyStories? Answering it needs a router that matches the published ones, a dispatch that does not loop over tokens in Python, and a way to stop every token from crowding onto the same expert.

SymbolMeaningType / shape
NNtokens (every leading axis flattened)int
EE, kkexperts, experts per tokenint
zn∈REz_n \in \mathbb{R}^Erouter logits of token nn: WgxnW_g x_nfloat32[E]
pn=softmax(zn)p_n = \mathrm{softmax}(z_n)router probabilitiesfloat32[E]
Tn\mathcal{T}_nthe kk experts chosen for token nnids
wn,ew_{n,e}the weight of expert ee for token nnfloat32
fef_eassignments to expert ee, divided by NNfloat
PeP_emean of pn,ep_{n,e} over tokensfloat
beb_ethe aux-free correction biasfloat32[E]

yn=∑e∈Tnwn,e Experte(xn)  +  Shared(xn).y_n = \sum_{e \in \mathcal{T}_n} w_{n,e}\, \mathrm{Expert}_e(x_n) \;+\; \mathrm{Shared}(x_n).

routerchooses Tn\mathcal{T}_n byweights wn,ew_{n,e}used by
softmax_topktop kk of pnp_npn,ep_{n,e}, renormalized over Tn\mathcal{T}_n when norm_topkMixtral (norm), Qwen-MoE (configurable)
topk_softmaxtop kk of znz_nsoftmax over the chosen logitsSwitch-style
sigmoidtop kk of σ(zn)+b\sigma(z_n) + bσ(zn,e)\sigma(z_{n,e}), renormalized when norm_topkDeepSeek-V3

Then every weight is multiplied by routed_scaling (DeepSeek-V3 uses 2.5). Ties go to the lowest id, the course’s one tie rule. Shared experts (DeepSeek) see every token; SS shared experts of width ff are one gated MLP of width SfS f.

Flatten the NkN k assignments as a=nk+ja = n k + j. A stable sort by expert id gives perm; x_sorted = x[perm // k] puts each expert’s tokens in one slice, offsets (a prefix sum of the counts) bounds the slices, each expert runs once on its slice, and inv_perm (the inverse permutation, inv_perm[perm[i]] = i) puts each output back at its assignment before the weighted sum over the kk slots. This is the shape of every fast MoE kernel: one matrix product per expert instead of one per token.

A router left alone learns to favor a few experts: they get more gradient, improve, and attract more tokens. Switch Transformer’s loss, as Hugging Face computes it:

Laux=E∑e=1EfePe.L_{aux} = E \sum_{e=1}^{E} f_e P_e .

At perfect balance fe=k/Ef_e = k / E and Pe=1/EP_e = 1 / E, so Laux=kL_{aux} = k; concentrating raises it. ff comes from an argmax and has no gradient; the router learns through PP: ∂Laux/∂pn,e=Efe/N\partial L_{aux} / \partial p_{n,e} = E f_e / N, so the experts that got the most tokens get their probabilities pushed down hardest. MoE.route returns aux_loss_coef * L_aux.

The choice of Tn\mathcal{T}_n is piecewise constant: no gradient flows through it. The router still learns, through the weights wn,ew_{n,e} that multiply the chosen experts’ outputs and through PP in the loss. Away from ties, central differences agree with the analytic gradient.

An auxiliary loss also pulls the main objective. DeepSeek-V3 adds a bias only to the scores used for choosing, and after each step moves it by a fixed rate γ\gamma:

be←be+γ sign(load‾−loade).b_e \leftarrow b_e + \gamma\, \mathrm{sign}\big(\overline{\mathrm{load}} - \mathrm{load}_e\big).

An overloaded expert’s bias falls until it is chosen less. The bias never enters the weights, so the output is still a mixture of the experts’ own scores. It is plain state, not a trained parameter: no optimizer touches it and it is not in state_dict.

Four experts, top 2, two tokens with logits z0=(ln⁡2,ln⁡4,0,0)z_0 = (\ln 2, \ln 4, 0, 0) and z1=(ln⁡3,0,ln⁡3,ln⁡2)z_1 = (\ln 3, 0, \ln 3, \ln 2).

  • Softmax: p0=(2,4,1,1)/8=(1/4,1/2,1/8,1/8)p_0 = (2, 4, 1, 1)/8 = (1/4, 1/2, 1/8, 1/8) and p1=(3,1,3,2)/9p_1 = (3, 1, 3, 2)/9.
  • Token 0 takes experts 1 and 0 (largest first), weights (1/2,1/4)(1/2, 1/4) renormalized to (2/3,1/3)(2/3, 1/3). Token 1 has a tie between experts 0 and 2 at 1/31/3; both are in the top 2, lowest id first: (0,2)(0, 2), weights (1/2,1/2)(1/2, 1/2).
  • Assignments a=nk+ja = n k + j: (t0,e1),(t0,e0),(t1,e0),(t1,e2)(t_0, e_1), (t_0, e_0), (t_1, e_0), (t_1, e_2), experts [1,0,0,2][1, 0, 0, 2]. Stable sort: perm =(1,2,0,3)= (1, 2, 0, 3), tokens perm // 2 =(0,1,0,1)= (0, 1, 0, 1), counts (2,1,1,0)(2, 1, 1, 0), offsets =(0,2,3,4,4)= (0, 2, 3, 4, 4), inv_perm =(2,0,1,3)= (2, 0, 1, 3).
  • If the experts output (10,20,30,40)(10, 20, 30, 40) in sorted order, inv_perm puts back (30,10,20,40)(30, 10, 20, 40) in assignment order: token 0 gets 2330+1310=70/3\frac{2}{3} 30 + \frac{1}{3} 10 = 70/3, token 1 gets 1220+1240=30\frac{1}{2} 20 + \frac{1}{2} 40 = 30.
  • Loss: f=(2,1,1,0)/2f = (2, 1, 1, 0)/2, P=(7/24,11/36,11/48,25/144)P = (7/24, 11/36, 11/48, 25/144), Laux=4 (1⋅724+12⋅1136+12⋅1148)=161/72≈2.236L_{aux} = 4\,(1 \cdot \frac{7}{24} + \frac{1}{2} \cdot \frac{11}{36} + \frac{1}{2} \cdot \frac{11}{48}) = 161/72 \approx 2.236, above k=2k = 2.

This is test_hand_example.

def topk_ids(scores, k) -> NDArray # [N, k], ties to the lowest id
def load_balance_loss(router_probs: Tensor, topk_idx, n_experts) -> Tensor
def dispatch(x: Tensor, topk_idx, n_experts) -> tuple[Tensor, NDArray, NDArray] # x_sorted, offsets, inv_perm
def combine(y_sorted: Tensor, topk_w: Tensor, inv_perm) -> Tensor
class MoE(Module):
def __init__(self, d, d_ff_expert, n_experts, top_k, n_shared=0, router="softmax_topk", norm_topk=True,
aux_loss_coef=0.01, bias_update_rate=0.0, routed_scaling=1.0, act="silu", rng=None)
def route(self, x) -> tuple[NDArray, Tensor, Tensor] # idx, weights, aux
def forward(self, x) -> Tensor; aux_loss, expert_load (properties of the last forward)
def update_bias(self, expert_load) -> None
TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3: routing with a tie, dispatch arrays, combine, the lossyou and the test agree on every step
test_moe_goldengoldenMixtral, Qwen3-MoE without renormalizing, DeepSeek-V3 with bias, shared experts, scalingMoE checkpoints and configs in L7.9
test_sorted_dispatch_equals_dense_loopdifferentialbatched path vs per-token loop for three routersthe speed path is the correct path
test_dispatch_is_a_stable_groupingpropertycontiguous slices, stable order, exact inverse, gradients back to xkernels that read one slice per expert
test_topk_ties_go_to_lowest_idboundarythe tie rulereproducible routing in Python and C
test_load_balance_loss_bounds_and_gradientpropertykk at balance, larger when concentrated, Efe/NE f_e / N gradientthe router learns to spread load
test_gradcheck_router_and_expertsgradcheckfloat64 central differences to x and the routerC1 trains the MoE
test_aux_free_biasunitthe sign rule by hand; bias chooses, never weighsDeepSeek-V3’s balancing
test_bias_updates_balance_the_loadproperty60 updates spread a skewed router’s 64 tokensbalance without a loss term
test_state_and_idle_expertsunitHub key names; an idle expert; stats not parameterscheckpoints load; training loops read aux_loss
test_validationboundarybad top_k, router, rates, expert idsconfig bugs fail loudly

Your oracle is the per-token loop in numpy float64, with each router’s weights computed from the logits by the table in section 2.1 (sorting with an explicit (score, id) key for the tie rule), run for every router; add a hand-sized dispatch with known offsets and roundtrip, the tie rule, the loss and its gradient on a two-token case, a central difference for the router weight, the bias rule, and the validation errors. Import only tinyllm.modern.moe and tinyllm.autograd. ol check L7.8 requires 0.80 with every pitfall fault killed.

PitfallSymptomCaught by
1. renormalizing when the config says norm_topk_prob = falseQwen-MoE outputs scaled wrongtest_moe_golden (mutant s01)
2. weighing by the biased scorethe bias leaks into the output; DeepSeek disagreestest_moe_golden, test_aux_free_bias (mutant s02)
3. scattering back with perm instead of its inverseoutputs land on the wrong tokenstest_hand_example, test_dispatch_is_a_stable_grouping (mutant s03)
4. reading tokens in token order, not grouped by experteach expert runs on another expert’s tokenstest_dispatch_is_a_stable_grouping (mutant s04)
5. counting ff from probabilities instead of assignmentsthe loss no longer measures the routingtest_hand_example, test_load_balance_loss_bounds_and_gradient (mutant s05)
a bias rule with the sign flippedoverloaded experts get more tokenstest_bias_updates_balance_the_load (mutant s06)
ties to the highest idrouting differs between implementationstest_topk_ties_go_to_lowest_id (mutant s07)
shared experts skippedDeepSeek outputs miss a termtest_moe_golden (mutant s08)
routed_scaling ignoredDeepSeek outputs 2.5 times too small in the routed parttest_moe_golden (mutant s09)
router weights detachedthe router never learnstest_gradcheck_router_and_experts (mutant s10)
offsets that hold each expert’s endslices shifted by one experttest_dispatch_is_a_stable_grouping (mutant s11)
topk_softmax weighing by the full softmaxweights do not sum to 1test_sorted_dispatch_equals_dense_loop (mutant s12)
DirectionModuleHow it uses this
BackL7.2every expert and the shared experts are a GatedMLP
BackL0.4the router Linear, the experts’ ModuleList
BackL0.2softmax, sigmoid, gather, reshape, concat with their gradients
BackL0.1Tensor indexing for dispatch
BackM06.3the default initialization stream
ForwardL7.9tl_num_experts, tl_top_k_experts, first_k_dense_replace build MoE layers
ForwardC1MoE vs dense at equal active parameters (M05.1’s active count)
Your pieceProduction equivalentWhat it addsWhere to look
dispatch, combineMegaBlocks, vLLM fused_moegrouped GEMMs over sorted tokens, no padding to a capacityvllm/model_executor/layers/fused_moe/, MegaBlocks paper
routingDeepSeek-V3 group-limited routingtop groups first (n_group, topk_group), then experts inside themHF DeepseekV3TopkRouter
capacitySwitch Transformer capacity factordrop tokens over an expert’s capacitySwitch paper, section 2.2
expert parallelismDeepEP, all-to-all dispatchexperts on different GPUs, tokens moved by all-to-allDeepSeek DeepEP