In a mixture-of-experts layer the router is the smallest piece of compute and the one that decides everything else: which expert weights each token touches, how many rows each expert's matrix multiply gets, how many bytes cross the network in the all-to-all, and whether tokens are silently dropped. The router itself is one small matrix multiply and a top-k. The consequences are a page of arithmetic that is worth doing by hand once.
This article goes through that arithmetic as the GPU executes it: the three scoring recipes used by Switch Transformer, Mixtral and DeepSeek-V3, where gradients actually flow, capacity and drop probabilities with real numbers, the balance losses and the bias-based alternative, the sort and prefix-sum index math that turns routing decisions into contiguous expert batches, and the numerical rules that keep routing identical across ranks. Architecture choices such as expert placement and shared experts are covered in MoE transformer architecture; parallel layouts in expert parallelism.
What the router computes
For T tokens of hidden size d and E experts, the router computes logits z = xW_r with W_r of shape [d, E], so z has shape [T, E]. For DeepSeek-V3 (d = 7168, E = 256 routed experts) that is about 1.8 million router weights per MoE layer, tiny next to the experts, and about 2·d·E ≈ 3.7 MFLOP per token: noise compared with the expert GEMMs. The cost of routing is therefore not FLOPs. It is the launch overhead of the small kernels that follow and, far more, what the decisions do to the rest of the layer.
Compute the router in fp32 even when the model runs in bf16. bf16 has 7 explicit mantissa bits (8 bits of precision); two logits that differ by less than about 0.4 percent can round to the same value or swap order, and a swapped top-k changes which expert runs. Many implementations cast the input and W_r to fp32 for this one GEMM; the extra cost is negligible.
One layer end to end
Three scoring recipes
Three recipes cover most deployed models. Take one token with four experts, top-2, and logits z = [2.0, 1.0, 0.5, −1.0].
| Recipe | Used by | Computation | Gates for the example |
|---|---|---|---|
| Softmax, then top-k | Switch (k = 1) | p = softmax(z); keep the top-k p unchanged | softmax = [0.609, 0.224, 0.136, 0.030]; experts 0, 1 with 0.609 and 0.224 |
| Top-k, then softmax | Mixtral (8 experts, k = 2) | softmax over the k selected logits only | experts 0, 1 with 0.731 and 0.269 |
| Sigmoid, bias for selection, normalise | DeepSeek-V3 (256 routed + 1 shared, k = 8) | s = sigmoid(z); select top-k of s + b; gates = s / sum of selected s, times 2.5 | s = [0.881, 0.731, 0.622, 0.269]; with b = 0: 0.546, 0.454 before scaling |
Mixtral's gates always sum to 1, so the layer output is a convex combination of the chosen experts. Softmax-then-top-k gates sum to less than 1, which shrinks the output when the router is unsure. Sigmoid scores are independent per expert, so adding experts does not dilute every score; DeepSeek-V3 normalises the selected scores and multiplies by its routed_scaling_factor of 2.5 to restore magnitude. Its published config also sets n_group = 8 and topk_group = 4: experts are split into eight groups, each group is scored by the sum of its two best experts, and a token may only pick experts from the four best groups. With groups aligned to nodes, that caps each token's dispatch at four nodes, which is a direct bound on all-to-all traffic.
import torch
def route(x, w_router, k, recipe, bias=None):
logits = x.float() @ w_router.float() # [T, E] in fp32
if recipe == "softmax_topk": # Switch
probs = logits.softmax(-1)
gate, idx = probs.topk(k, dim=-1)
elif recipe == "topk_softmax": # Mixtral
top, idx = logits.topk(k, dim=-1)
gate = top.softmax(-1)
else: # DeepSeek-V3 style
s = logits.sigmoid()
choice = s + bias if bias is not None else s # bias affects selection only
idx = choice.topk(k, dim=-1).indices
gate = s.gather(-1, idx)
gate = gate / gate.sum(-1, keepdim=True) * 2.5
return idx, gate.to(x.dtype), logits
Where the gradient flows
Top-k returns integer indices, and indices have no gradient. The router learns only through the gate values: the layer output is y = Σ_i g_i · E_i(x) over the selected experts, so ∂y/∂g_i = E_i(x), and the loss gradient on each gate is the dot product of the upstream gradient with that expert's output. If expert 0's output helps the loss more than expert 1's, gradient raises z_0 relative to z_1. Unselected experts receive no signal at all from this token, which is exactly why routers collapse: an expert that is never chosen never gets the gradient that would make it worth choosing. With top-1 and a softmax gate, the selected gate is still a function of all logits through the normaliser, so other logits do receive gradient; with Mixtral's recipe only the selected logits do.
Capacity and drop arithmetic
Fixed-shape implementations give each expert a buffer of C slots: C = ⌈cf · T · k / E⌉, where cf is the capacity factor. Assignments beyond C are dropped: the token skips that expert and passes through on the residual path. Take T = 4096 tokens per device batch, E = 64, k = 2, so the mean load is T·k/E = 128 assignments per expert. Even with a perfectly uniform router, each expert's load is random, roughly binomial with mean 128 and standard deviation about 11. Computing the expected overflow exactly for that binomial gives the drop rates below.
| Capacity factor | C (slots per expert) | Assignments dropped, uniform router |
|---|---|---|
| 1.0 | 128 | 3.5% |
| 1.25 | 160 | 0.008% |
| 1.5 | 192 | below 1e-9 |
| 2.0 | 256 | effectively zero |
Real routers are not uniform. If one expert attracts twice its share, 256 assignments meet 160 slots at cf = 1.25 and 96 of them, 37.5 percent of that expert's traffic, are dropped. So the capacity factor is really a bet on skew, and its cost is linear: every expert's GEMM and every all-to-all buffer is padded to C rows, so cf = 2.0 doubles expert compute and dispatch bytes compared with a perfect fit. Dropless implementations remove the bet by sizing each expert's batch from the actual counts and running a grouped GEMM over variable row counts; skew then shows up as a slower expert, not a lost token.
Balancing: auxiliary loss, z-loss and bias
Three mechanisms push the router toward even load. The Switch Transformer auxiliary loss is L_aux = α · E · Σ_i f_i · P_i, where f_i is the fraction of tokens dispatched to expert i and P_i is the mean router probability for i over the batch. f_i comes from the argmax and has no gradient; the gradient flows through P_i, lowering the probability of overloaded experts. Its value is α at uniform load. Switch used α = 0.01; Mixtral's published Hugging Face config uses a coefficient of 0.02.
The router z-loss from ST-MoE, L_z = (1/T) Σ_t (log Σ_i exp z_ti)², penalises large logits. Large logits make the softmax saturate and amplify bf16 rounding; ST-MoE used a coefficient of 0.001. For the example token, log-sum-exp is 2.495 and its square 6.23.
DeepSeek-V3 replaced most of the auxiliary loss with a bias per expert, added to the scores only for top-k selection. After each step the bias of an overloaded expert decreases by γ and an underloaded one increases by γ; the paper used γ = 0.001 for most of pre-training, plus a very small sequence-wise balance loss with α = 0.0001. Because the bias never enters the gate values, balancing does not distort the output mixture. In the example, a bias of −0.3 on expert 0 changes the selection to experts 1 and 2, whose gates are still computed from the unbiased scores: 0.540 and 0.460 before scaling.
def update_bias(bias, idx, num_experts, gamma=1e-3):
load = torch.bincount(idx.flatten(), minlength=num_experts).float() # all-reduce across DP first
bias -= gamma * torch.sign(load - load.mean()) # no gradient, no optimizer
return bias
Dispatch index math
After routing, the T·k (token, expert) assignments are scattered in token order. Experts want contiguous rows. The conversion is a stable sort plus a prefix sum, and the same counts drive the all-to-all.
def dispatch(x, idx, num_experts):
T, k = idx.shape
flat = idx.flatten() # [T*k] expert id per slot
order = torch.argsort(flat, stable=True) # slots grouped by expert
token_of_slot = order // k # which token each sorted slot came from
counts = torch.bincount(flat, minlength=num_experts) # rows per expert
offsets = torch.cumsum(counts, 0) - counts # start row of each expert
x_sorted = x[token_of_slot] # [T*k, d], contiguous per expert
return x_sorted, order, counts, offsets
def combine(y_sorted, order, gate, T):
k = gate.shape[1]
y_slots = torch.empty_like(y_sorted)
y_slots[order] = y_sorted # undo the sort
y = y_slots.view(T, k, -1) * gate.unsqueeze(-1) # weight by gates
return y.sum(1) # [T, d]With expert parallelism, counts grouped by destination rank become the all-to-all send sizes, and each rank first exchanges counts so receivers can allocate buffers; the message-size effects are in MoE all-to-all communication. Production kernels fuse the sort, the gather and the gate multiply, because each unfused step reads and writes a [T·k, d] tensor. Keep order and counts for the backward pass; recomputing them under activation checkpointing must reproduce bit-identical routing, or the backward pass sends gradients to the wrong experts.
Numerics and determinism
Routing must be identical wherever it is computed. Under tensor parallelism each rank may compute the router on the same replicated input; if any rank uses a different reduction order or precision, two ranks can disagree on a near-tie and the all-to-all deadlocks or mixes tokens. Either compute the router once and broadcast the indices, or guarantee identical fp32 computation on every rank. Ties at exact equality are rare in fp32 but common after bf16 casts, which is another reason to select in fp32.
Inference adds a subtler issue: routing depends on batch composition only through capacity. Drop-based serving can change a token's output depending on which other requests share its batch, which breaks reproducibility and caching assumptions. Serve dropless; MoE serving covers the rest.
What to monitor
- Per-expert load per layer: max over mean of
counts; above about 1.5 at your capacity factor means drops or a straggler expert. - Drop rate per layer, if your implementation drops; it should be logged, never inferred.
- Router entropy averaged over tokens: a steady fall toward zero early in training signals collapse.
- Mean log-sum-exp of the logits: rising values precede bf16 instabilities and loss spikes.
- For bias balancing, the bias range: biases that keep growing mean the router is fighting the data.
Failure modes
| Failure | Cause | Fix |
|---|---|---|
| A few experts take most tokens | No or too weak balancing; collapse early in training | Aux loss or bias updates from step 0; watch entropy |
| Loss spikes, NaNs in router | Logits grow; bf16 saturation | fp32 router, z-loss |
| Hang in all-to-all | Ranks computed different top-k | Broadcast indices or identical fp32 router on every rank |
| Wrong gradients after checkpointing | Recomputed routing differs from forward | Save idx and order; do not recompute them |
| Quality drop at long context or large batch | Capacity overflow drops tokens | Raise cf, log drops, or go dropless |
| Output depends on batch neighbours | Drop-based inference | Dropless serving |
Trade-offs
Softmax gates are simple but make every score compete; sigmoid gates scale to hundreds of experts but need explicit normalisation. An auxiliary loss balances reliably but adds a gradient that competes with the language-modelling loss; bias balancing avoids that interference but is a controller with its own gain to tune. Capacity buffers give static shapes and predictable memory at the cost of drops or padding; dropless gives exact routing at the cost of variable shapes and a slowest-expert step time. Group-limited routing cuts network traffic at a small cost in routing freedom. The math and the hardware push in the same direction: fewer, more even, more local assignments are cheaper.
What to do next
- Check that your router GEMM and top-k run in fp32, and that routing is computed identically on every tensor-parallel rank.
- Write down your recipe (softmax-top-k, top-k-softmax or sigmoid with bias) and confirm gate sums behave as you expect.
- Compute C for your T, E, k and cf, and log the actual drop rate per layer.
- Add per-expert load, router entropy and mean log-sum-exp to your training dashboards.
- Save routing indices for the backward pass if you use activation checkpointing.
- Size all-to-all buffers from measured counts, and consider group-limited routing if dispatch crosses many nodes.
- Serve dropless, and test that a request's output does not change with batch composition.