Mixture of Experts

The router

The router is a tiny layer that picks which experts each token visits, and it has to learn that job through a hard selection step that gives it almost no feedback.

On this page 9
  1. The short answer
  2. The enquiry counter
  3. How it decides
  4. The hard part, which is not obvious
  5. The trick that makes it learn at all
  6. The two ways to arrange the choice
  7. A practical shortcut used at scale
  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

The router is a small piece that reads each word and decides which experts should handle it.

The enquiry counter

You have stood in front of an enquiry counter. You hand over a form. The clerk glances at it for two seconds, and writes a counter number on top.

The clerk cannot solve your problem. He does not know your case. His entire skill is knowing where to send you.

The router is that clerk. It is a tiny fraction of the model's size, and the whole system depends on it getting the direction right.

How it decides

The router gives every expert a score for this particular word. Then it keeps the best few and ignores the rest.

The kept scores are also used as weights. An expert scored twice as strongly contributes twice as much to the answer.

   word: "chapati"

   router scores:   E0  E1  E2  E3  E4  E5  E6  E7
                    0.1 0.6 0.0 0.2 0.9 0.1 0.0 0.3
                                        ^       ^
   keep the top two:            E4 (0.9) and E1 (0.6)
   answer = 0.6 x E4's output  +  0.4 x E1's output

The hard part, which is not obvious

The router has to learn its job, and its job is a choice. Choices have no in-between.

Think of a new librarian who only ever sends readers to shelves A and B. Nobody comes back from shelf C, so no one ever tells him whether shelf C would have been better. He has no way to find out, so he keeps sending everyone to A and B.

That is the router's problem exactly. It learns from the experts it picked, and hears nothing from the ones it skipped.

The trick that makes it learn at all

Multiply each expert's output by its score before adding it in.

Now the score is part of the answer, not only part of the decision. If an expert did well, raising its score improves the result, and the router gets told so.

Take that multiplication away and the router receives no feedback at all. Not a small amount. None. The code below proves it in one line of output.

This part is genuinely subtle and worth reading twice.

The two ways to arrange the choice

Each word picks its experts. The natural way, and the one nearly everything uses. Its weakness is that nothing stops every word picking the same expert.

Each expert picks its words. Every expert takes a fixed number of the words that suit it best. Balance is perfect by construction.

The second sounds better and has a fatal flaw for chat models. To pick its favourites, an expert must see all the words first. A model writing one word at a time does not have them.

So it works for training and for reading a document, and not for writing. That is why the first method won.

A practical shortcut used at scale

With hundreds of experts spread over many machines, a word could be sent anywhere, and the network becomes the bottleneck.

So large models restrict the choice. First pick a few machines, then pick experts only from those. Fewer places for the word to travel, and it costs very little in quality.

Remember this

  • The router scores every expert for every word and keeps the best few.
  • Those scores also weight the experts' answers, and that is the only reason the router can learn.
  • Words picking experts is standard; experts picking words balances perfectly but cannot write text.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written against PyTorch 2.5.1, Python 3.10. CPU, instant.

Where the router's gradient comes from

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

torch.manual_seed(0)
d, n_exp, top_k = 4, 6, 2

router = nn.Linear(d, n_exp, bias=False)
experts = nn.ModuleList(nn.Linear(d, d, bias=False) for _ in range(n_exp))
x = torch.randn(3, d)                                   # 3 tokens

def run(use_gate_weight):
    router.zero_grad(set_to_none=True)
    for e in experts:
        e.zero_grad(set_to_none=True)
    probs = F.softmax(router(x), dim=-1)
    w, idx = probs.topk(top_k, dim=-1)
    w = w / w.sum(-1, keepdim=True)                     # renormalise the survivors
    out = torch.zeros_like(x)
    for s in range(top_k):
        for e in range(n_exp):
            hit = idx[:, s] == e
            if hit.any():
                y = experts[e](x[hit])
                out[hit] += (w[hit, s, None] * y) if use_gate_weight else y
    out.sum().backward()
    g = router.weight.grad
    return idx, (None if g is None else g.clone())

idx, grad_with = run(True)
print("expert chosen per token (top-2):")
for t in range(3):
    print(f"  token {t}: {idx[t].tolist()}")
chosen = sorted(set(idx.flatten().tolist()))
print("experts chosen by anyone:", chosen)

print("\nrouter gradient row norms, WITH the gating weight in the output")
for e in range(n_exp):
    tag = "chosen" if e in chosen else "never chosen"
    print(f"  expert {e}: {grad_with[e].norm():.6f}   ({tag})")

_, grad_without = run(False)
print("\nrouter gradient WITHOUT the gating weight (experts summed raw):")
print("  router.weight.grad is", grad_without)

# Scoring function: softmax couples the experts, sigmoid does not.
logits = torch.tensor([2.0, 1.0, 0.5, -1.0, -2.0, 0.0])
print("\nsame router logits, two scoring functions")
print("  softmax:", [round(v, 4) for v in F.softmax(logits, 0).tolist()], "sum =",
      round(float(F.softmax(logits, 0).sum()), 4))
print("  sigmoid:", [round(v, 4) for v in torch.sigmoid(logits).tolist()], "sum =",
      round(float(torch.sigmoid(logits).sum()), 4))
Output
expert chosen per token (top-2):
  token 0: [1, 5]
  token 1: [3, 1]
  token 2: [4, 0]
experts chosen by anyone: [0, 1, 3, 4, 5]

router gradient row norms, WITH the gating weight in the output
  expert 0: 2.140871   (chosen)
  expert 1: 3.234802   (chosen)
  expert 2: 0.000000   (never chosen)
  expert 3: 0.382025   (chosen)
  expert 4: 2.140871   (chosen)
  expert 5: 3.160457   (chosen)

router gradient WITHOUT the gating weight (experts summed raw):
  router.weight.grad is None

same router logits, two scoring functions
  softmax: [0.5573, 0.205, 0.1243, 0.0277, 0.0102, 0.0754] sum = 1.0
  sigmoid: [0.8808, 0.7311, 0.6225, 0.2689, 0.1192, 0.5] sum = 3.1225

Reading the output

router.weight.grad is None without the gating weight. Not small, not noisy: absent. Remove the multiplication and the router's output stops being connected to the loss at all, because topk returns indices and indices carry no gradient. The router would never train.

That single line is the reason every MoE implementation multiplies expert outputs by their routing weight, and why the weight is the normalised probability rather than a constant.

Expert 2 has gradient exactly 0.000000. It was never chosen by any of the three tokens. Note it is exactly zero rather than very small, and the reason is the renormalisation w / w.sum(): the softmax denominator over all experts cancels, leaving a softmax over the chosen ones alone. The unchosen logits drop out of the computation entirely.

This is the exploration problem in one number. Experts that are not chosen receive no signal about whether they should have been.

Softmax sums to 1; sigmoid sums to 3.12. Softmax couples the experts: raising one score lowers every other. Sigmoid scores each expert independently, so a token can find two experts genuinely good rather than being forced to rank them.

DeepSeek-V3's config declares scoring_func: "sigmoid" with norm_topk_prob: true and routed_scaling_factor: 2.5. Sigmoid scoring is the newer choice at large expert counts; Mixtral and Qwen3 use softmax.

The knobs on a real router

Read from configs on the Hub:

FieldMixtral-8x7BDeepSeek-V3Qwen3-Next-80B
num_experts_per_tok2810
experts8256 routed + 1 shared512 + 1 shared
scoringsoftmaxsigmoidsoftmax
norm_topk_prob—truetrue
groupingnonen_group: 8, topk_group: 4none
router_jitter_noise0.0——

n_group: 8, topk_group: 4 is node-limited routing. Experts are partitioned into 8 groups matched to physical nodes, each token first picks its best 4 groups, and only then picks 8 experts from within them. It bounds how many machines one token's activations must be sent to, which is the whole subject of expert parallelism.

Token choice and expert choice

Token choice (everything above): each token takes its top-$k$ experts. Simple, causal, and unbalanced by default.

Expert choice (Zhou et al., 2022): each expert takes its top-$c$ tokens from the batch. Load is perfectly even by construction, and no token is dropped for capacity reasons — but a token may be picked by many experts or by none.

Expert choice cannot be used for autoregressive decoding. Selecting the best tokens for an expert requires seeing the whole batch of tokens at once, which leaks information from later positions to earlier ones. It is usable for encoders and for training-time throughput, not for generation.

Common mistakes

Dropping the gating weight multiplication. Demonstrated above. The failure is silent in the forward pass and total in the backward pass.

Forgetting to renormalise after topk. Without it, output scale varies with router confidence, and the layer's contribution to the residual stream fluctuates. Set norm_topk_prob or divide by the sum.

Computing the router in the model dtype. Routers are usually computed in float32 even in a bf16 model, because bf16 has few mantissa bits and near-ties in the top-$k$ flip between runs. transformers does softmax(router_logits.float(), dim=-1) for this reason.

Adding jitter noise at inference. router_jitter_noise multiplies the input by uniform noise during training to encourage exploration. Guard it with self.training; Mixtral ships it set to 0.0 anyway.

Assuming routing is stable across a rerun. Change the batch, the padding or the dtype and near-tied routing decisions flip. This is one of several reasons an MoE model is less reproducible token-for-token than a dense one.

Try it yourself

Set top_k = 1 and rerun. Fewer experts receive gradient, and the collapse risk rises. Then replace softmax with sigmoid in the scoring and check whether the same experts are chosen for each token.

What to learn next

Researcher — Mathematics and papers.

The gating function

Shazeer et al. (2017) define noisy top-$k$ gating:

$$ H(x)_i = (x W_g)_i + \varepsilon_i \cdot \operatorname{softplus}\big((x W_{\text{noise}})_i\big), \quad \varepsilon_i \sim \mathcal{N}(0,1) $$

$$ G(x) = \operatorname{softmax}\big(\operatorname{keep_top_k}(H(x), k)\big) $$

with non-selected entries set to $-\infty$ before the softmax. The learned noise scale $W_{\text{noise}}$ provides exploration and, in the original paper, makes the load a differentiable function of the parameters so a balancing loss can be defined on it.

Modern implementations mostly drop the learned noise. Mixtral's router_jitter_noise is a simpler multiplicative input perturbation and ships at 0.0.

Why the gradient path is what it is

$\operatorname{topk}$ produces indices, and $\partial \text{idx} / \partial \text{logits} = 0$ almost everywhere. Gradient reaches the router only through the magnitude of the retained gate values:

$$ \frac{\partial \mathcal{L}}{\partial W_g} = \sum_{i \in \text{top-}k} \frac{\partial \mathcal{L}}{\partial y} \cdot E_i(x) \cdot \frac{\partial G(x)_i}{\partial W_g} $$

This has three consequences worth stating plainly.

The router learns only about experts it already picks. A rich-get-richer dynamic that is the direct cause of expert collapse — see load balancing.

With renormalised top-$k$ weights, unselected logits have exactly zero gradient. Renormalisation cancels the global softmax denominator, reducing the effective gate to a softmax over the selected set. The zero in the output above is exact, not numerical.

The router optimises the loss, not balance. Any balance the model achieves comes from an added term or an added rule, never from the objective.

Straight-through estimators, Gumbel-softmax relaxations and reinforcement-learning formulations of routing have all been tried. None has displaced the plain gate-weight path in production models, which is a mildly surprising empirical fact given how crude that path is.

Softmax against sigmoid

Softmax over experts imposes $\sum_i p_i = 1$, so scores are comparative. At $N = 256$ or $512$ the average probability is tiny, gradients through the gate become small, and the coupling makes every expert's score depend on every other's.

Sigmoid scores each expert independently. DeepSeek-V3 uses scoring_func: sigmoid with topk_method: noaux_tc and a routed_scaling_factor: 2.5 to restore output magnitude after selection. The pairing of sigmoid scoring with a bias-based balancing rule instead of an auxiliary loss is deliberate: a bias added to a sigmoid score shifts selection without distorting the gate weight used in the output, which a softmax makes awkward.

Expert choice

Zhou et al. (2022), Mixture-of-Experts with Expert Choice Routing (arXiv:2202.09368) invert the assignment. With $T$ tokens, $N$ experts and capacity factor $c$, each expert selects its top $\lceil Tc/N \rceil$ tokens by affinity.

Properties: load is exactly balanced by construction, no auxiliary loss is required, no token is dropped for overflow, and tokens receive a variable number of experts — a form of adaptive compute the model gets for free.

The disqualifying property for decoders is that expert-side top-$k$ over a batch is not causal. Training a decoder with expert choice and then decoding with token choice creates a train/test mismatch. It remains the right choice for encoders and for vision models.

Grouped and hierarchical routing

DeepSeek-V3's node-limited routing (n_group: 8, topk_group: 4) selects at most 4 of 8 expert groups per token, then 8 experts within them. The all-to-all communication volume per token is bounded by the number of groups touched, not by the number of experts, which caps the network cost of scaling $N$.

This is the clearest example of an architectural choice made for the interconnect rather than for quality. Expect more of them: at 512 experts across many nodes, routing is a network problem wearing a modelling costume.

Stability

ST-MoE (Zoph et al., 2022) traces MoE training instabilities to large router logits and adds a router z-loss penalising the log-sum-exp of the router logits, which keeps them in a range where bf16 rounding does not flip selections. The paper is written as a design guide and its stability and fine-tuning sections remain the most useful practical reference on the topic.

Practical rules that follow: compute the router in float32, keep logits bounded, and treat any run where routing entropy collapses in the first few thousand steps as failed rather than recoverable.

What to learn next