Mixture of Experts

Load balancing and expert collapse

Nothing in a model's objective wants the experts used evenly, so a few take almost everything and the rest go dead unless you add a rule that forces the traffic apart.

On this page 9
  1. The short answer
  2. The two vegetable stalls
  3. The same thing happens inside the model
  4. Why it costs you
  5. The first fix: add a penalty
  6. The second fix, which is cleverer
  7. Where you have already seen the same idea
  8. Remember this
  9. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

Left alone, a mixture of experts sends nearly all its work to a few experts.

The two vegetable stalls

Two stalls sit side by side in the market, with the same vegetables at the same prices.

One morning a few more people happen to stop at the left stall. It sells out fast, so its stock stays fresh, so more people stop there. The right stall's stock sits and wilts, so fewer people stop, so it wilts further.

Within a month the left stall is packed and the right one is empty. Nobody decided this. A small early difference fed itself.

The same thing happens inside the model

The router sends a few extra words to expert three. Expert three gets more practice, so it improves faster.

Being better, it wins more of the router's choices. Being chosen more, it improves faster still.

Meanwhile expert eleven receives almost nothing. It never practises, so it never improves, so it is never chosen. It is dead weight in the file you downloaded.

This is called expert collapse, and it is the standard failure of this whole design.

Why it costs you

You paid for a hundred experts and are running about twelve. The memory bill is for a hundred.

Worse, experts live on different machines. If one machine holds the popular expert, everyone waits for that machine while the others idle.

The first fix: add a penalty

Add a rule to training that punishes uneven traffic. If the split is lopsided, the model is charged extra, so it learns to spread the work.

This works and it has a cost worth saying out loud. The model now has two aims: do the job well, and spread the work. They pull against each other.

The second fix, which is cleverer

Keep a small handicap number for each expert, used only when choosing.

An expert that took too much work this step gets its handicap nudged down. It is then picked slightly less often. An idle expert gets nudged up.

The important detail: the handicap changes who is picked, and never changes how much their answer counts. So the model's quality goal is not disturbed at all. Balance is arranged outside the learning, like a traffic policeman rather than a fine.

The measurements below show the difference. The penalty version balances perfectly and does the task worse. The handicap version balances well and keeps almost all of the quality.

Where you have already seen the same idea

  • A booking system that stops showing you a counter once its queue is long.
  • A delivery app that holds back orders from an overloaded restaurant.
  • Lane changes on a highway when one lane crawls.

Remember this

  • Nothing in the model's goal wants even usage, so a few experts take over.
  • One fix charges the model a penalty for imbalance, which pulls against quality.
  • A better fix nudges a per-expert handicap that changes who is picked, not what their answer is worth.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written against PyTorch 2.5.1, Python 3.10. CPU only. The script trains three tiny models and took about 22 seconds on this machine.

Three balancing policies, measured

balancing.py
import torch, torch.nn as nn, torch.nn.functional as F

D, E, K, N, V = 8, 16, 1, 256, 200      # dim, experts, top-k, tokens/step, vocabulary

def train(mode, steps=1500, alpha=0.01, gamma=0.05, seed=0):
    torch.manual_seed(seed)
    emb = torch.randn(V, D)                          # a fixed vocabulary of token vectors
    tgt = torch.tanh(emb @ torch.randn(D, D))        # what this layer must learn per token
    zipf = 1.0 / torch.arange(1, V + 1).float()      # real text is Zipfian, so make it so
    zipf = zipf / zipf.sum()

    router = nn.Linear(D, E, bias=False)
    experts = nn.ModuleList(nn.Linear(D, D, bias=False) for _ in range(E))
    opt = torch.optim.Adam([{"params": router.parameters(), "lr": 0.15},
                            {"params": experts.parameters(), "lr": 0.01}])
    bias = torch.zeros(E)                            # used only by the aux-loss-free mode
    gen = torch.Generator().manual_seed(seed + 1)
    hist = []

    for step in range(steps):
        ids = torch.multinomial(zipf, N, replacement=True, generator=gen)
        x, y = emb[ids], tgt[ids]

        probs = F.softmax(router(x), dim=-1)
        score = probs + bias if mode == "bias" else probs
        _, idx = score.topk(K, dim=-1)               # selection may use the bias
        w = probs.gather(-1, idx)                    # the WEIGHT never uses the bias
        w = w / w.sum(-1, keepdim=True)

        out = torch.zeros_like(x)
        counts = torch.zeros(E)
        for s in range(K):
            for e in range(E):
                hit = idx[:, s] == e
                counts[e] += hit.sum()
                if hit.any():
                    out[hit] += w[hit, s, None] * torch.tanh(experts[e](x[hit]))

        loss = F.mse_loss(out, y)
        if mode == "aux":                            # Switch Transformer, equations 4-6
            f = counts / (N * K)                     # fraction of tokens sent to each expert
            P = probs.mean(0)                        # mean router probability per expert
            loss = loss + alpha * E * (f * P).sum()

        loss.backward(); opt.step(); opt.zero_grad()

        if mode == "bias":                           # DeepSeek-V3, auxiliary-loss-free
            with torch.no_grad():
                over = counts > counts.mean()
                bias[over] -= gamma                  # overloaded: make it less attractive
                bias[~over] += gamma                 # idle: make it more attractive
        hist.append(counts.clone())
    return torch.stack(hist), loss.item()

labels = {"none": "no balancing", "aux": "auxiliary loss", "bias": "bias update"}
for mode in ("none", "aux", "bias"):
    hist, final = train(mode)
    share = hist[-100:].sum(0)
    share = share / share.sum()
    print(f"{labels[mode]:<16} task loss {final:.4f} | busiest {share.max():5.1%} "
          f"| quietest {share.min():.2%}")
    print("   ", " ".join(f"{v:5.1%}" for v in share.tolist()))
print(f"\nperfectly even would be {1/E:.1%} for all {E} experts")
Output
no balancing     task loss 0.0008 | busiest 26.4% | quietest 0.05%
     5.5% 26.4%  4.8%  8.7%  7.5%  2.5% 10.5%  1.4% 17.7%  5.1%  0.1%  2.0%  1.5%  1.2%  1.3%  3.9%
auxiliary loss   task loss 0.0144 | busiest  6.8% | quietest 5.88%
     6.6%  6.6%  6.5%  6.2%  6.4%  6.0%  5.9%  6.2%  6.0%  6.4%  6.8%  6.0%  6.1%  6.2%  5.9%  6.1%
bias update      task loss 0.0016 | busiest  8.7% | quietest 4.19%
     5.2%  8.4%  8.1%  5.0%  7.0%  4.2%  4.6%  4.7%  8.6%  6.3%  7.0%  6.3%  4.9%  8.7%  5.1%  5.8%

perfectly even would be 6.2% for all 16 experts

Reading the output

Without balancing, two experts take 44% of the traffic. Expert 1 gets 26.4%, expert 8 gets 17.7%, and expert 10 gets 0.1%. The spread between busiest and quietest is over five hundred times. That is collapse, arrived at from a symmetric initialisation, with no adversarial setup.

The unbalanced run has the best task loss: 0.0008. This is the single most important number here. Collapse is not a training failure. It is the optimiser doing its job. Concentrating traffic on a few well-practised experts genuinely fits the data better in the short run, and nothing in the objective objects.

The auxiliary loss balances perfectly and costs 18x the loss. Every expert lands between 5.9% and 6.8%. Task loss rises from 0.0008 to 0.0144. That is a real tax, and it is why the auxiliary coefficient is a tuning headache: too small and it does nothing, too large and it dominates.

The bias update gets most of both. Balance from 4.2% to 8.7%, task loss 0.0016 — twice the unbalanced loss instead of eighteen times. The bias steers selection without ever entering the gradient.

The caveat you must keep. These are three runs of a toy at one seed with one alpha. The ordering is the finding; the ratios are not a measurement of any real model. Lower alpha and the auxiliary run's loss improves while its balance degrades. That trade is the point.

The auxiliary loss, precisely

The Switch Transformer formulation, implemented in transformers as load_balancing_loss_func:

$$ \mathcal{L}{\text{aux}} = \alpha \cdot N \cdot \sum{i=1}^{N} f_i P_i $$

f_i is the fraction of tokens routed to expert i, and P_i is the mean router probability assigned to expert i across the batch. Only P_i carries gradient, since f_i comes from a hard topk. The product is minimised when both are uniform, and multiplying by N keeps its scale independent of the expert count.

Coefficients in the wild: Mixtral ships router_aux_loss_coef: 0.02; Qwen3-Next ships 0.001. Two orders of magnitude apart, which tells you how much this depends on the rest of the recipe.

The bias update, precisely

DeepSeek-V3's auxiliary-loss-free strategy. Their description: "At the end of each step, we will decrease the bias term by γ if its corresponding expert is overloaded, and increase it by γ if its corresponding expert is underloaded, where γ is a hyper-parameter called bias update speed."

Their settings: γ = 0.001 for the first 14.3T tokens, then 0.0; plus a complementary sequence-wise balance loss at α = 0.0001, deliberately tiny, to prevent extreme imbalance inside a single sequence.

The structural point is in this line of the script:

python
score = probs + bias           # bias affects WHO is selected
w = probs.gather(-1, idx)      # but NOT what the selection is worth

Add the bias to the gate weight too, and you are back to distorting the objective.

Diagnosing it in your own run

Log these every few hundred steps:

  • Max-to-mean load ratio. Should sit near 1. Above 2 you are heading for trouble.
  • Number of experts receiving under 1% of tokens. Should be zero.
  • Routing entropy. A sharp fall in the first few thousand steps is the early warning.
  • Fraction of dropped tokens, if you use a capacity limit — the subject of the next lesson.

Common mistakes

Averaging the auxiliary loss over layers incorrectly. It is defined per MoE layer and summed. Compute it once over a concatenation of all layers' router logits, as the reference implementation does, or you will silently scale it by the layer count.

Adding the balancing bias to the gate weight. Covered above. Symptom: balance looks fine, quality is worse than an aux-loss run.

Turning balancing off after warm-up and assuming it holds. DeepSeek sets γ = 0 only after 14.3T tokens, by which point routing has stabilised. Turning it off early re-starts the collapse.

Ignoring padding tokens in the load statistics. Padding routes somewhere, and counting it distorts f_i. The reference implementation takes an attention_mask for exactly this reason.

Measuring balance over a single batch. Per-batch imbalance is normal and mostly harmless. Persistent imbalance over thousands of steps is the failure.

Try it yourself

Set alpha = 0.001 and rerun the auxiliary mode. Task loss should improve and balance should worsen. Then set gamma = 0.005 in the bias mode and watch balance take much longer to establish. Those two sweeps are the entire tuning story.

What to learn next

Researcher — Mathematics and papers.

Why collapse is the optimum, not a bug

Let expert $i$ have quality $q_i$ on the tokens it receives, improving with the number of tokens it has trained on. Router preference is monotone in $q_i$. Then $\dot{n}_i \propto q_i$ and $\dot{q}_i \propto n_i$, a positive feedback loop with an unstable symmetric fixed point. Any perturbation grows.

Nothing in $\mathcal{L}_{\text{task}}$ opposes it. The experiment above confirms the direction: the collapsed run attains the lowest task loss of the three. Balance is a constraint imposed from outside for reasons of capacity utilisation and hardware efficiency, not a property the objective seeks.

The long-run cost is capacity: a model using 12 of 100 experts has the parameter budget of 12. But that cost is realised over a long horizon, and gradient descent optimises the next step.

Auxiliary losses

Shazeer et al. (2017) proposed two: a coefficient-of-variation penalty on the importance (summed gate values per expert), and a separate load loss requiring a smooth differentiable estimate of the discrete load, obtained through the noisy gate.

Fedus, Zoph and Shazeer (2021), Switch Transformer (arXiv:2101.03961) simplified this to a single differentiable-by-one-factor loss, equations 4–6:

$$ \mathcal{L}{\text{aux}} = \alpha N \sum{i=1}^{N} f_i P_i, \quad f_i = \frac{1}{T}\sum_{x} \mathbb{1}[\arg\max p(x) = i], \quad P_i = \frac{1}{T}\sum_{x} p_i(x) $$

$T$ is the token count and $p(x)$ the router distribution. Under uniform $f$ and $P$ the value is $\alpha$, independent of $N$. The gradient flows only through $P_i$.

Zoph et al. (2022), ST-MoE (arXiv:2202.08906) add a router z-loss penalising $\left(\log \sum_i e^{z_i}\right)^2$ on router logits, which bounds logit magnitude and removes a class of bf16 instabilities. Their paper is the standard reference for MoE training stability and fine-tuning.

Auxiliary-loss-free balancing

DeepSeek-V3 (arXiv:2412.19437) maintains a per-expert bias $b_i$ added to affinity scores for the top-$k$ selection only:

$$ \text{selected} = \operatorname{Top-}k_i\big(s_i + b_i\big), \qquad b_i \leftarrow b_i - \gamma\,\operatorname{sign}(n_i - \bar{n}) $$

$s_i$ is the sigmoid affinity, $n_i$ the token count for expert $i$ in the step and $\bar{n}$ the mean. The gate weight applied to the expert output uses $s_i$ alone.

This is a control loop, not a regulariser. It has no gradient, adds no term to the loss surface, and interacts with the objective only through which experts are exercised. DeepSeek reports better model quality than an equivalent auxiliary-loss run at comparable balance, which the toy above reproduces in direction if not in magnitude.

They retain a sequence-wise auxiliary loss at $\alpha = 10^{-4}$ to bound imbalance within a single sequence. Batch-level balance does not imply sequence-level balance, and severe within-sequence imbalance causes token dropping at inference where batches are small.

Expert choice as a structural answer

Zhou et al. (2022), Expert Choice Routing (arXiv:2202.09368) make balance exact by construction: each expert selects a fixed number of tokens. No auxiliary loss, no bias, no dropped tokens.

It is not usable for autoregressive decoding, since expert-side selection over a batch is not causal. For encoders and vision models it remains the cleanest solution available, and its absence from decoder-only language models is a constraint of the task rather than a judgement on the method.

Inference-time balance is a different problem

Training balance is averaged over enormous batches. Serving balance is per request, per step, with small batches, and the traffic distribution is your users' rather than your corpus's.

The serving-side answer is replication, not regularisation. DeepSeek's EPLB (Expert Parallelism Load Balancer) takes measured expert-load statistics and computes a placement, duplicating hot experts across devices with a configurable number of redundant slots. Their LPLB follow-on formulates the assignment as a linear program solved per batch. vLLM and SGLang both integrate EPLB.

That splits the problem cleanly: architecture-time balancing keeps experts alive during training, and deployment-time balancing keeps GPUs busy in production. Solving one does not solve the other — see serving a MoE across GPUs.

What to learn next