Mixture of Experts
Be able to explain routing, expert capacity and load balancing in MoE models.
Prerequisites
- EThe MLP block: GELU, SwiGLUrequired
Intuition
An ordinary transformer runs the whole MLP block for every token. A mixture of experts replaces the MLP with many parallel experts and lets a small router pick out a handful per token.
┌──────────┐
token ──→ │ router │──→ picks expert 3 and 17 of 64
└──────────┘
│
┌──────────┴─────────┐
expert 3 expert 17
└──────────┬─────────┘
a weighted sum
The gain: the model can have far more parameters without every token costing more to compute.
| A dense model | MoE | |
|---|---|---|
| Total parameters | 8 B | 47 B |
| Active per token | 8 B | 13 B |
| Computation per token | 1× | ~1.6× |
| Memory | 16 GB | 94 GB |
The catch shows in the last row: every expert has to be in memory even if only two are used. MoE trades memory for capacity — a completely different bargain from quantisation.
Formal
The router is a linear layer followed by a softmax over the experts. The top (usually ) are chosen, and the output is the weighted sum:
Load balancing is the central problem. Without a counterweight a self-reinforcing spiral arises: an expert that happens to be chosen more often is trained more, gets better, and is chosen even more often. In the end a handful of experts are used and the rest are dead weights.
The solution is an auxiliary loss that punishes an uneven distribution:
where is the share of tokens routed to expert and is the average router probability for it. The term is minimised when both are evenly distributed. Typically .
Expert capacity. Every expert has a limit on how many tokens it accepts per batch:
With a capacity factor of 1.25 every expert tolerates 25 % above its proportional share. Tokens above the limit are dropped — they pass through unchanged via the residual connection. A low factor saves memory but drops more; a high one wastes.
Three practical problems:
| Problem | Why |
|---|---|
| Communication | the experts sit on different GPUs → an all-to-all exchange per layer |
| Unstable training | the router can oscillate; a z-loss and a lower learning rate on the router help |
| Fine-tuning | MoE models overfit more easily; fewer epochs, more regularisation |
Variants: shared experts that are always active (DeepSeek), expert choice instead of token choice (each expert picks its tokens, which gives perfect balance by construction), and very fine-grained experts (hundreds of small ones instead of eight large).
When MoE pays off: when you have plenty of memory but want more parameters per FLOP. When memory is limited, a dense model is nearly always the right choice.
Code
import torch, torch.nn as nn, torch.nn.functional as F
class MoE(nn.Module):
def __init__(self, d=512, hidden=2048, n_experts=8, k=2, capacity_factor=1.25):
super().__init__()
self.k, self.n = k, n_experts
self.capfactor = capacity_factor
self.router = nn.Linear(d, n_experts, bias=False)
self.experts = nn.ModuleList([
nn.Sequential(nn.Linear(d, hidden), nn.GELU(), nn.Linear(hidden, d))
for _ in range(n_experts)])
def forward(self, x):
B, T, d = x.shape
xf = x.reshape(-1, d) # (N, d)
N = xf.size(0)
logits = self.router(xf)
probability = F.softmax(logits, dim=-1) # (N, n)
weight, choice = probability.topk(self.k, dim=-1) # (N, k)
weight = weight / weight.sum(dim=-1, keepdim=True) # normalise the chosen ones
capacity = int(self.capfactor * N * self.k / self.n)
out = torch.zeros_like(xf)
dropped = 0
for e in range(self.n):
mask = (choice == e)
index = mask.any(dim=-1).nonzero(as_tuple=True)[0]
if len(index) > capacity: # over capacity → drop
dropped += len(index) - capacity
index = index[:capacity]
if len(index) == 0:
continue
w = (weight * mask.float()).sum(dim=-1)[index].unsqueeze(-1)
out[index] += w * self.experts[e](xf[index])
# The load-balancing loss: an even distribution of both the tokens and the probability
f = torch.zeros(self.n, device=x.device)
for e in range(self.n):
f[e] = (choice == e).any(dim=-1).float().mean()
P = probability.mean(dim=0)
aux = self.n * (f * P).sum()
self.latest = {"aux_loss": float(aux),
"dropped_share": round(dropped / max(N * self.k, 1), 4),
"distribution": [round(float(v), 3) for v in f]}
return out.view(B, T, d), aux
# Counting parameters: total against active
def parameters(d=4096, hidden=14336, n_experts=8, k=2, layers=32):
per_expert = 3 * d * hidden
total = layers * n_experts * per_expert
active = layers * k * per_expert
return {"total_bn": round(total / 1e9, 1), "active_bn": round(active / 1e9, 1),
"ratio": round(total / active, 1)}
print(parameters())
# {'total_bn': 45.1, 'active_bn': 11.3, 'ratio': 4.0}
# ↑ four times more parameters than are computed per token
# Collapsed routing looks like this
collapsed = [0.71, 0.24, 0.03, 0.01, 0.01, 0.0, 0.0, 0.0]
balanced = [0.13, 0.12, 0.13, 0.12, 0.13, 0.12, 0.12, 0.13]
for name, f in (("collapsed", collapsed), ("balanced", balanced)):
import math
entropy = -sum(p * math.log(p) for p in f if p > 0)
print(f"{name:<10} entropy {entropy:.3f} of a maximum of {math.log(8):.3f}")
# collapsed entropy 0.783 of a maximum of 2.079
# balanced entropy 2.079 of a maximum of 2.079
Logging the routing entropy is the simplest collapse detector: if it falls below about half the maximum, only a fraction of the experts are being used, and the weight of the auxiliary loss needs raising.
Mastery means
- Explains routing and the top-k choice
- Describes the load-balancing problem
- Knows what MoE saves and what it costs
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity — arXiv (open access; licence per article)
- arXiv — Mixtral of Experts — arXiv (open access; licence per article)
- arXiv — Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer — arXiv (open access; licence per article)