How sparse gating unlocks trillion-parameter models at fraction-of-parameter compute: Switch Transformer roots, top-K routing, Mixtral 8x7B, DeepSeek-V3, load-balance losses, and expert parallelism.
The core intuition is simple: the parameter count of a model determines its capacity, but it is the active parameter count at each forward pass that determines its compute cost. Dense transformers conflate the two — every token visits every weight. Mixture-of-Experts decouples them.
175B total params. Every token activates all 175B. Training FLOPs scale linearly with both token count and parameter count. FLOP budget is fully spent every forward step.
46.7B total params. Every token activates 12.9B — two of eight expert FFNs. Same capacity as a ~46B dense model; compute of a ~13B dense model.
The MoE layer replaces the single feed-forward block in each transformer layer with E expert FFNs and a lightweight router. Each token is routed to K of the E experts; the outputs are weighted and summed before rejoining the residual stream. Attention layers are unchanged.
Kaplan et al. (2020) showed that loss scales as a power law with compute. MoE attacks the same loss target at lower FLOPs by growing the total parameter count (capacity) without growing active params (cost). A 2022 Google paper (Scaling Laws for Expert Specialisation) confirmed the MoE Pareto front sits consistently above dense models when matched on FLOPs rather than total params.
| Model | Total params | Active params/token | Ratio |
|---|---|---|---|
| GPT-3 | 175B | 175B | 1.0× |
| Switch-C (2021) | 1.6T | ~7B | ~0.004× |
| Mixtral 8x7B (2023) | 46.7B | 12.9B | 0.28× |
| Mixtral 8x22B (2024) | 141B | 39B | 0.28× |
| DeepSeek-V3 (2024) | 671B | 37B | 0.055× |
Fedus et al. (2022, NeurIPS) published the landmark Switch Transformer, scaling a T5-style encoder-decoder to 1.6 trillion parameters on a TPU v3 pod while using only a fraction of those params per token. The key architectural choice: route each token to exactly K=1 expert (hard top-1), reducing the routing from a soft mixture to a switch.
Each expert has a capacity of C = (tokens/experts) × capacity_factor. Overflow tokens skip the expert and pass through unchanged — this is the token drop mechanism. Capacity factor 1.25 is the default: 25% headroom above uniform load.
Switch used bfloat16 for most ops but float32 for the routing softmax and the weight update. This was an early observation that mixed-precision MoE training is brittle without careful casting at the router.
Top-1 routing makes the expert assignment a discrete decision, removing the soft-mixture gradient through the router. This forced Fedus et al. to rely entirely on the auxiliary load-balance loss to prevent all tokens collapsing to one expert. It worked surprisingly well — Switch-Base matched T5-Base quality in one-sixth the pre-training time on identical hardware.
import torch, torch.nn as nn, torch.nn.functional as F
class SwitchRouter(nn.Module):
def __init__(self, d_model: int, n_experts: int):
super().__init__()
self.w = nn.Linear(d_model, n_experts, bias=False)
def forward(self, x): # x: (B, T, D)
logits = self.w(x) # (B, T, E)
gates = F.softmax(logits, dim=-1)
idx = gates.argmax(dim=-1) # top-1 index
g = gates.gather(-1, idx.unsqueeze(-1)).squeeze(-1)
return idx, g # (B,T) and (B,T)
Two families of routing dominate the literature: token-choice (each token picks its K favourite experts) and expert-choice (each expert picks its favourite T tokens). They have complementary failure modes.
The original sparsely-gated MoE from Outrageously Large Neural Networks used top-2 routing with additive Gaussian noise to encourage exploration:
# H(x): noisy logits before softmax
H(x)i = (x · Wg)i + StandardNormal() · Softplus((x · Wnoise)i)
# Keep only top-K, set rest to -inf, then softmax
gates = Softmax(KeepTopK(H(x), k))
output = sum(gatesi · Experti(x) for i in top-K)
Expert-choice flips the direction: each of the E experts selects its top-C tokens from the batch. This guarantees perfect load balance by construction — no capacity-overflow drops. The downside: tokens are no longer guaranteed to be processed by any expert. For short sequences this can hurt quality.
Some recent work (e.g. DeepSeek-V2, Qwen-MoE) uses shared experts that every token visits unconditionally, plus a pool of routed experts chosen by top-K. This hybrid ensures that common, general-purpose computations are always applied, while rare specialised knowledge is routed sparsely.
Mistral AI released Mixtral 8x7B in December 2023 (Apache 2.0), the first widely-deployed open-weights MoE that outperformed LLaMA-2 70B on most benchmarks while activating fewer parameters per token than LLaMA-2 13B.
| Hyperparameter | Value | Notes |
|---|---|---|
| Total parameters | 46.7B | 8 experts × 7B per-expert FFN + shared attention |
| Active params / token | 12.9B | 2 of 8 experts selected each token |
| Experts per layer | 8 | K=2, so 2 experts selected |
| Hidden dim | 4096 | Same as LLaMA-2 13B |
| Intermediate dim | 14336 (per expert) | SwiGLU activated |
| Layers | 32 | MoE in all layers |
| Context length | 32768 | Sliding window attention + rope |
| Vocabulary | 32000 | SentencePiece BPE |
Mixtral uses a standard top-2 linear router with no jitter noise at inference (noise is added during training only, following Shazeer 2017). The gating weights are computed as:
g_raw = softmax(W_r · x) # (E,) = (8,)
top2 = topk(g_raw, k=2) # indices and values
g_norm = top2.values / top2.values.sum() # re-normalise to 1
y = sum(g_norm[i] * FFN_i(x) for i in top2.indices)
Mixtral layers alternate between full attention layers and sliding-window attention layers (window W=4096). This predates the dedicated long-context techniques in Arch 03 but already addresses the quadratic bottleneck for long sequences. The MoE FFN is orthogonal to the attention choice — any attention variant can be paired with MoE.
DeepSeek-V3 (December 2024) pushed MoE design further with two key innovations: fine-grained expert decomposition and multi-token prediction (MTP) as an auxiliary training objective. The model totals 671B parameters and is competitive with GPT-4o class models at a fraction of reported training cost.
Instead of 8 large experts, DeepSeek-V3 uses 256 fine-grained routed experts per layer (each with ~1/8 the FFN dimension of a coarse expert) plus 1 shared expert that every token visits. Top-8 out of 256 are selected. This gives:
8 experts, intermediate dim ~14K. Each expert is a full-size FFN. Top-2 selected. Expert overlap between different tokens is high — specialisation is limited.
256 routed experts, intermediate dim ~1.5K each. Plus 1 shared expert (full dim). Top-8 selected. Tokens can activate any combination of 256 fine-grained slots → far higher effective specialisation.
MoE architecture in DeepSeek-V3 is accompanied by an MTP auxiliary head that predicts the next 2 tokens simultaneously during training (in addition to the standard next-token loss). At inference the extra heads are discarded. MTP functions as a speculative decoding bootstrap: the auxiliary predictions seed the speculative drafts used by DeepSeek's inference system.
DeepSeek report training 671B total / 37B active on 2.048M H800 GPU-hours for the 14.8T token pre-train — approximately $5.6M at 2024 cloud rates. For comparison, Llama-3 405B (dense, ~6M A100-hours) cost roughly 3× more per token processed despite being 60% of the active parameter count.
Without any regularisation, MoE training collapses quickly: a small number of experts receive the majority of tokens (rich-get-richer routing), leaving the others undertrained. Multiple auxiliary loss terms have been proposed to combat this.
# f_i = fraction of tokens dispatched to expert i
# P_i = mean gate probability for expert i over the batch
# N = number of experts, n = tokens in batch
L_aux = N × sum( f_i × P_i for i in range(N) )
# f_i is non-differentiable; P_i provides the gradient signal.
# Minimised when all f_i = 1/N (uniform load).
L_total = L_CE + α × L_aux # α ~ 1e-2 to 1e-1
A second stabilisation term penalises large logit magnitudes entering the router softmax. Large logits cause numerical issues in bfloat16 and encourage winner-takes-all behaviour:
L_z = (1/B) × sum( log(sum(exp(x_i)))2 for x in batch )
# Penalises the log-sum-exp magnitude. Coefficient ~1e-3.
| Loss term | What it targets | Coefficient α |
|---|---|---|
| L_aux (importance) | Uniform token distribution across experts | 10−2 – 10−1 |
| L_z (router z-loss) | Router logit magnitude stability | 10−3 |
| L_expert_entropy | Maximise entropy of expert utilisation | varies |
| DeepSeek expert-bias | Learnable per-expert score bias for global balance | greedy update |
DeepSeek-V3 drops the differentiable auxiliary loss in favour of expert bias terms: each expert has a learnable scalar offset added to its gate score. After each training step, biases for overloaded experts are decremented and underloaded experts incremented by a fixed δ = 0.001. This avoids the gradient-vs-CE loss tension that arises when L_aux is too large.
At inference time, an MoE model with 671B total params cannot fit on a single GPU. The dominant parallelism strategy is Expert Parallelism (EP): shard the experts across devices so each device hosts a subset, then route tokens across the network to the correct device.
Two all_to_all collectives per MoE layer: dispatch (tokens to experts) and combine (results back). With NVLink (600 GB/s) this is fast within a node; across nodes via InfiniBand (400 Gb/s) it becomes the bottleneck. EP degree is chosen to keep experts on one node when possible.
If a request batch has uneven expert affinity, some GPUs finish early and wait. DeepSeek-V3's inference engine duplicates hot experts onto idle GPUs dynamically — a runtime load-balance step not present during training.
MoE training surfaces several failure modes not seen in dense transformers. Understanding them is essential before attempting a MoE run.
An expert that receives near-zero routing probability early in training never gets gradient signal, and so remains initialised-noise weights — which makes it even less likely to be selected. The feedback loop is self-reinforcing. Mitigation: jitter noise, z-loss, and careful warm-up with high α for L_aux before annealing.
The router can learn to route all tokens of a given linguistic class (e.g. punctuation) to the same expert, causing that expert to become highly specialised on a low-frequency pattern. Downstream experts then see a biased token distribution. Monitoring expert utilisation entropy during training catches this early.
An auxiliary loss coefficient α that is too large forces the model to sacrifice language modelling quality for perfect load balance. The result is periodic loss spikes as the router fights the CE gradient. Typical safe range: α ∈ [0.01, 0.1].
Router logits in bfloat16 have only 8 exponent bits. If expert embedding norms grow during training, softmax inputs saturate. The z-loss term (slide 06) and gradient clipping (max norm 1.0) are the standard mitigations.
1. Initialise all expert FFN weights identically (or with very small variance difference) to break symmetry gently. 2. Monitor router_entropy and expert_utilisation_cv (coefficient of variation) per layer every 100 steps. 3. Start α = 0.01 for the first 1% of tokens, then anneal to 0.001. 4. Log dead-expert count (experts with <1% of uniform load) — any persistent dead expert after 5% of training suggests a hyperparameter problem, not just noise.
Deck 02 moves to a completely different angle on the sequence modelling problem: state-space models (Mamba, S4, RWKV) that replace attention with recurrence — and ask whether the quadratic attention bottleneck can be avoided entirely rather than masked by sparsity.