How do you stop one expert from eating everything?
Part 2 ended on the sentence that makes this part necessary: the router’s only gradient is the gate values of the experts it already chose. An expert it stops choosing sends back nothing, and the router learns nothing about it.
That is a positive feedback loop with no damping term. An expert that wins a few extra tokens early gets a few extra gradient updates, becomes slightly better at the tokens it sees, and wins more of them. Left alone, a Mixture of Experts is a winner-take-all system.
The failure is quiet. The loss curve looks fine for a long time, because a bank in which four experts do all the work is still a perfectly good four-expert model. What is being destroyed is the capacity you paid the memory bill for, and nothing in the objective mentions it.
Three mechanisms exist to stop it, and a good answer distinguishes them by what they change: the loss, the selection, or the numerics.
The 30-second version
It is router collapse, and it is the default behaviour rather than a bug. The gradient only reaches experts that are selected, so an early lead compounds: a slightly favoured expert gets more tokens, improves faster, and gets more tokens. The classic fix is the auxiliary loss from GShard and Switch Transformer, $\alpha N \sum_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 it received. Both vectors sum to one, so the product is minimised when both are flat, and the $N$ in front makes that minimum exactly 1. The subtlety is that $f_i$ is a histogram of hard decisions and has no gradient, so all of the pressure goes through $P_i$. DeepSeek-V3 replaced it with a per-expert bias added to the score used for the top-$k$ comparison and nowhere else, nudged up for starved experts and down for crowded ones after every step, which balances the load without adding a term to the loss or distorting any gradient. Separately, ST-MoE's router z-loss penalises the squared log-sum-exp of the router logits, which is about numerical stability rather than balance. In practice you log the per-layer load histogram every step, because collapse usually starts in one layer.The figure runs the same sixteen-expert router three times on the same token stream: once with nothing, once with the auxiliary loss, once with the bias. Press play on the first step and watch it fall over.
Why does a router collapse?
Two rules are enough, and they are both true of any Mixture of Experts trained by gradient descent.
Tokens improve an expert. An expert that receives more tokens takes more gradient steps on the task, so it gets better at it.
A better expert scores higher. The router is trained to send tokens where the loss goes down, so an expert that is better at a token attracts more of them.
Compose those and you have compounding. The figure’s first step models exactly this and nothing else, with the two rules written at the top of its source, and it collapses to a single expert in about a hundred steps. Nothing exotic is required: no bad data, no bad initialisation, no bug.
What makes it dangerous is the shape of the curve. The load imbalance sits flat for the first forty steps, drifts for thirty more, then falls off a cliff. If your only instrument is the training loss, the first sign of trouble arrives long after the mechanism has started.
What does collapse actually cost you?
Nothing you can see on a memory profile, which is the point.
The parameters are all still there. The checkpoint is the same size, the HBM bill from Part 1 has not moved by a byte, and the FLOPs per token are unchanged because $k$ experts still run. What has moved is how much of the bank is learning anything. An expert with no tokens gets no gradient, so it stops improving, and a sixteen-expert bank in which one expert does everything is a one-expert model with sixteen experts’ worth of storage.
Scale that to a frontier model. If a quarter of DeepSeek-V3’s routed experts went quiet, about 163 billion parameters would stop learning while continuing to occupy HBM and continuing to be sharded, checkpointed and shipped.
That is why the load histogram, logged per layer per step, is the cheapest instrument in an MoE run and the first thing to add. It costs an $N$-element reduction and it is the only place collapse is visible early.
What does the auxiliary loss actually penalise?
The formulation from GShard, made standard by Switch Transformer:
\[L_{\text{bal}} = \alpha \, N \sum_{i=1}^{N} f_i \, P_i\]where $f_i$ is the fraction of tokens in the batch dispatched to expert $i$, and $P_i$ is the mean router probability assigned to expert $i$ over those tokens. Both vectors sum to one over the $N$ experts.
Their dot product is smallest when both are flat, and the $N$ in front normalises the minimum to exactly 1. Its maximum is $N$, when one expert takes everything. So the number is readable: 1 means perfectly even, 16 on a sixteen-expert bank means total collapse, and in practice it hovers a little above 1 because routing noise never lets it settle.
Switch used $\alpha = 10^{-2}$, having swept it from $10^{-1}$ down to $10^{-5}$ and found that value balanced the load quickly without interfering with the training loss. That sweep is worth remembering: too small and the mechanism does nothing, too large and you are optimising for uniformity rather than for the task.
Why the product of two vectors and not the variance of one?
Because you need something differentiable, and $f$ is not.
$f_i$ is a histogram. It counts how many tokens the hard top-$k$ actually sent to expert $i$, so it is piecewise constant in the router’s weights and its gradient is zero almost everywhere. Penalising the variance of $f$ alone would give you a number that describes the problem and a gradient that does nothing about it.
$P_i$, the mean router probability, is smooth. Multiplying the two gives a surrogate that reads as “wherever a lot of tokens actually went, push the probability down”, and its gradient is
\[\frac{\partial L_{\text{bal}}}{\partial s_j} = \alpha N P_j \Big( f_j - \sum_i f_i P_i \Big)\]which is positive for over-used experts and negative for under-used ones. Being able to say which factor carries the gradient is the standard follow-up question, and it separates people who have read the formula from people who have implemented it.
There is a second-order version of this worth knowing, because it is an implementation detail that changes the model rather than the speed. The batch you compute $f_i$ over is a choice.
Most frameworks compute it per micro-batch, which at frontier scale is a handful of sequences, so the term effectively demands that every individual sequence spread itself evenly over all the experts. A micro-batch of pure code is then forced to use the whole bank.
Qiu and colleagues showed in January 2025 that computing $f_i$ over the global batch instead, which costs one extra all-reduce of an $N$-element vector, improves both perplexity and the domain specialisation of the experts, at scales up to 42.8 billion parameters and 400 billion tokens.
What is DeepSeek’s auxiliary-loss-free balancing?
The observation behind it is that the auxiliary loss is a foreign term. Its gradient is not trying to make the model better at the task; it is trying to make the histogram flat, and those two gradients are added together and pull in different directions.
So remove it. Keep a per-expert bias $b_i$, add it to the score used for the top-$k$ comparison, and leave the gate value alone:
\[\text{selection: } \text{TopK}\big(s_i + b_i\big), \qquad \text{gate: } g_i \text{ from } s_i \text{ alone}\]After each step, decrease $b_i$ by a small $\gamma$ for experts that were over-subscribed and increase it for those that were under-subscribed. No term is added to the loss. No gradient is distorted. It is a controller bolted to the side of the router rather than an objective, and it works because it changes which experts are chosen without changing how much they count.
DeepSeek’s numbers for it, on models up to 3 billion parameters, are a perplexity of 9.50 against the auxiliary loss’s 9.56 at 1B parameters and 100B tokens, with a maximum load violation of 0.04 against 0.72. Better balance and slightly better perplexity, which is the claim.
Two footnotes, both worth having. In DeepSeek-V3 the bias update rate was $\gamma = 0.001$ for the first 14.3 trillion tokens and zero for the rest, and they kept a complementary sequence-wise auxiliary loss at $\alpha = 0.0001$ to stop any single sequence from piling onto one expert. So the “loss-free” method ships alongside a small loss, with its own controller mostly switched off.
And the paper was rejected at ICLR 2025 and shipped in V3 anyway. In the public reviews the authors conceded that the interference-gradient motivation was based on intuition and needed more rigorous validation. Knowing that is a better answer than either endorsing or dismissing the method.
What is the z-loss for, and why is it a different problem?
The router z-loss, from the ST-MoE paper, is not about balance at all.
\[L_z = \frac{1}{B} \sum_{b=1}^{B} \Big( \log \sum_{j=1}^{N} e^{x_j^{(b)}} \Big)^2\]It penalises the squared log-sum-exp of the router logits. Squaring means the penalty grows with the size of the logits and says nothing about which expert is winning, so it constrains magnitude and leaves the ranking alone.
The reason it exists is numerical. Router logits drift upward over a long run, the way attention logits do, and large numbers have large rounding errors in low precision, which an exponential then amplifies.
ST-MoE chose a coefficient of 0.001 by sweeping for the best model quality after pretraining, and they pair it with a second recommendation: cast the router’s input to float32 before the softmax and cast the dispatch tensors back to bfloat16 afterwards. The router tensor is tiny, so the fp32 cast costs nothing and removes the problem at the source.
This is the router’s version of the QK-norm story from the Transformer series. An exponential of a drifting logit is a stability hazard wherever it appears, and the fix is always to bound the logit rather than to fix the exponential.
Does a balanced router mean specialised experts?
No, and this is the question that separates a memorised answer from an understood one.
Every metric in this part counts how many tokens each expert received. None of them says anything about whether the assignment means anything. Two routings can have identical perfect balance, zero maximum violation and an auxiliary loss of exactly 1, with one assigning tokens by content and the other assigning them at random. The figure’s last step draws both.
What the evidence says is genuinely mixed, and the disagreement is informative. Mistral’s own routing analysis of Mixtral found no obvious pattern by topic: the expert assignment distribution for arXiv papers, PubMed abstracts and philosophy papers is very similar at every layer. What they did find is syntactic and positional structure. Python’s self and indentation tokens route consistently, and consecutive tokens pick the same first-choice expert about 28% of the time at layer 15 against a 12.5% random baseline.
ST-MoE found clear specialisation in their encoder experts, punctuation, verbs, proper nouns, and explicitly did not find it in the decoder. They also looked for language specialisation in a multilingual model and found none: experts handled English, Japanese, French and Chinese indiscriminately.
AI2’s OLMoE, trained from scratch rather than upcycled, does report domain and vocabulary specialisation, with one layer-0 expert nearly 100% specialised to arXiv, and they hypothesise that Mixtral shows less of it because it was upcycled from a dense model, so all its experts started from the same optimum.
So the honest position: specialisation at the level of token class and vocabulary is real and measurable; specialisation by subject matter is contested and may depend on how the model was initialised; specialisation by language does not appear. “Which expert handles medicine” is the wrong question. “Is any expert idle” is the right one.
What do you actually watch during a run?
Four things, and none of them is the training loss.
The per-layer load histogram, every step. Collapse starts in one layer and spreads, so a model-wide average hides it. The single number to derive from it is the maximum violation: the busiest expert’s load over the even load, minus one.
The number of experts receiving zero tokens. It is the metric that turns into lost capacity, and it is a step function: fine, fine, fine, then several.
The router logit magnitude, which is what the z-loss is defending, and which tells you whether the numerics are drifting before a spike appears.
And the auxiliary loss value itself, if you are using one, remembering that it never reaches its minimum and is not supposed to. What matters is the trend, not the level.
Rapid fire: can you do these from memory?
- Describe the feedback loop that makes a router collapse, in two sentences.
- Write the auxiliary load-balancing loss, define both vectors, and say what its minimum and maximum are.
- Which factor of the auxiliary loss carries the gradient, and why is the other one there at all?
- Why does the batch you compute the dispatch fractions over change the model, not just the speed?
- Where exactly does the per-expert bias enter, and where does it deliberately not enter?
- Write the router z-loss and say what problem it is solving.
- Give two routings with identical perfect load balance and completely different meaning.
- Name four things you would log every step in an MoE run.
Part 4 leaves the loss function and asks what it costs to actually run this: two all-to-alls a layer, and the reason a sparse model is harder to batch during decoding than a dense one.