DL S077 · Mixture of Experts (MoE) Intuition
Why Mixtral has 47B parameters but only ~13B active per token — the router, the experts, the load-balancing loss. Part of the 'Deep Learning & LLMs From Scratch' 80-session self-study series.
🎯 Understand exactly why Mixtral-8x7B has 47B total parameters but only ~13B active per token, how the router decides, and why the load-balancing loss is the whole ballgame.
Series: Deep Learning & LLMs From Scratch — 80 sessions · Session 77 / 80 · Module M12 · ~2 hours
The story we're starting with
Here's an economic argument that will unlock MoE for you. A dense 70B model requires 70B parameters worth of memory and 70B parameters worth of compute per token. Memory and compute are locked together. Now suppose you noticed that on any given token, most of the model's specialized capacity is wasted — the "Python code" circuits aren't firing for a Shakespeare token, the "French grammar" circuits aren't firing for a math problem. What if you could have all the specialists on standby, pay memory for all of them, but only activate the two most relevant per token? You'd get the capacity of a 70B model at the compute cost of a 13B model. Same wall-clock speed, same VRAM (roughly, per weight), but you can memorize twice the world.
That's Mixture of Experts. Mixtral-8x7B has 8 "experts" (each is a full MLP block) in place of the usual single MLP. A tiny router (one linear layer) looks at each token and picks the top-2 experts to run. The other 6 sit idle for that token. Aggregate parameter count: 47B (7B backbone + 8 experts × ~5B each). Active per token: ~13B. Quality: matches Llama 2 70B. Inference cost: like a 13B model. That's the trick, and it's been in the literature since Shazeer 2017 — it just took someone (Mistral) building it at scale and shipping the weights for the field to actually notice.
Today we go under the hood: what the router computes, why it needs a load-balancing loss to not collapse, what "auxiliary loss" and "expert parallelism" actually mean, and why MoE is a nightmare to serve at inference despite being cheap on paper.
- Compute the total-vs-active parameter count for any MoE spec (given #experts, top-k, and expert size).
- Write the router forward pass and the top-k gating in ~20 lines of PyTorch.
- Explain why an MoE without a load-balancing loss will collapse to a single expert, and what the auxiliary loss looks like.
- Reason about MoE inference bottlenecks: memory bandwidth, expert imbalance, batch-size effects.
- Decide whether MoE makes sense for your capstone (spoiler: probably not, and you'll know why).
Prerequisites
- S036–S040 · Transformers. MoE is a swap-in for the MLP block; you need to know what an MLP block is doing.
- S057 · Distributed training basics. Expert parallelism is a distinct sharding pattern you should have some vocabulary for.
- S015 · Softmax and cross-entropy. The router's decision is a masked softmax.
1 · What changes in the transformer block
Standard decoder block:
MoE block:
Nothing else changes. Attention is unaffected. The router+experts substitution happens per-layer (Mixtral does it every layer; some designs do it every other layer to save memory).
1.1 · What an "expert" is
Just an MLP. Same architecture as the dense MLP block:
class Expert(nn.Module):
def __init__(self, d_model, d_ff):
super().__init__()
self.w1 = nn.Linear(d_model, d_ff, bias=False) # gate
self.w3 = nn.Linear(d_model, d_ff, bias=False) # up
self.w2 = nn.Linear(d_ff, d_model, bias=False) # down
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x)) # SwiGLUThat's a SwiGLU MLP — same as Llama's. You have N of them in parallel, only k get run per token.
1.2 · What the router computes
class Router(nn.Module):
def __init__(self, d_model, n_experts):
super().__init__()
self.gate = nn.Linear(d_model, n_experts, bias=False)
def forward(self, x): # x: (B, T, d_model)
logits = self.gate(x) # (B, T, n_experts)
return logitsOne linear layer. That's the entire router. In Mixtral, it maps a 4096-d hidden state to 8 scores. The top-2 experts by score are selected; the softmax over just those two gives the mixture weights.
2 · Top-k gating, step by step
Say n_experts=8, top_k=2, and a token's router logits are:
logits = [0.2, -1.3, 2.5, 0.1, 1.8, -0.5, 0.0, 0.7]Top-2: experts 2 (score 2.5) and 4 (score 1.8). Everything else contributes nothing to this token.
Compute the softmax over just those two:
top2_logits = [2.5, 1.8]
weights = softmax([2.5, 1.8]) = [0.668, 0.332]Now run expert 2 and expert 4 on x, and combine:
y = 0.668 * Expert_2(x) + 0.332 * Expert_4(x)Experts 0, 1, 3, 5, 6, 7 never run. We saved 75% of the MLP compute for this token.
2.1 · In code
def moe_forward(x, router, experts, top_k=2):
# x: (B, T, d), router: Router, experts: list of Expert
logits = router(x) # (B, T, E)
top_vals, top_idx = logits.topk(top_k, dim=-1) # both (B, T, k)
weights = F.softmax(top_vals, dim=-1) # (B, T, k) — normalized over the k picked
y = torch.zeros_like(x)
for slot in range(top_k):
expert_idx = top_idx[..., slot] # (B, T)
w = weights[..., slot].unsqueeze(-1) # (B, T, 1)
# Naive: run each expert on the tokens routed to it (production is more subtle)
for e in range(len(experts)):
mask = (expert_idx == e) # (B, T)
if mask.any():
y[mask] += w[mask] * experts[e](x[mask])
return ySlow, but shows what's happening. Production implementations dispatch tokens to experts in batched form (grouped_gemm) to keep GPUs busy.
Build a toy MoE with 8 experts, top-2 routing, and a tiny 2-layer transformer on WikiText-2 or similar. Train two variants for 500 steps each: (a) main cross-entropy only, (b) main + 0.01 × auxiliary load-balancing loss. Every 50 steps, log usage[e] = fraction of tokens routed to expert e for e in 0..7. In (a) you'll see the entropy of the usage distribution crater within 100 steps — one or two experts eat everything, the rest go to zero. In (b) the distribution stays close to uniform (usage entropy ≈ log 8). That plot is why every MoE paper has section 3.1 titled "load balancing."
3 · The load-balancing problem
Here's the puzzle. If nothing constrains the router, it can decide "expert 3 is great" and route every token to expert 3. Now you have a 47B-parameter model whose compute cost equals a 13B model — but the 13B is the same 13B every time. The other 7 experts are dead weight. Their gradients are near-zero because they're never called. This is called expert collapse and it's the #1 failure mode of MoE.
3.1 · The auxiliary loss
Switch Transformer's fix (later adopted by Mixtral): add a small auxiliary loss that penalizes unbalanced routing.
Define, for a batch:
f_e= fraction of tokens routed to expert e (a proportion, sums to k across experts).P_e= average router probability assigned to expert e.
Auxiliary loss:
L_aux = n_experts * Σ_e (f_e * P_e)Minimized when f_e = 1/n_experts and P_e = 1/n_experts for all e — i.e., uniform routing. Multiplied by a small coefficient (Mixtral uses 0.001–0.01) and added to the main next-token loss:
L_total = L_next_token + α * L_auxThat's it. Two lines of extra code. Without them, the model collapses within a few thousand steps. With them, load stays roughly balanced.
3.2 · Why not just enforce balance?
You could hard-cap tokens-per-expert with a "capacity factor" (Google's approach): once expert e has been assigned C tokens in a batch, further tokens routed to e get dropped (or routed to their second choice). This is called token dropping. It works but wastes some tokens. Mixtral uses the softer auxiliary-loss approach and skips capacity capping at training time.
4 · The parameter-count arithmetic
Mixtral-8x7B specs:
- 32 decoder layers
- d_model = 4096, d_ff = 14336
- 8 experts per MoE layer, top-2 routed
- MoE layer replaces the MLP in every decoder block
Per-layer parameters:
- Attention (Q, K, V, O projections + GQA): ~42M params (with GQA, KV heads are 8 vs 32 Q heads).
- One expert MLP (SwiGLU, w1/w2/w3):
3 × 4096 × 14336 ≈ 176Mparams. - Router:
4096 × 8 = 33kparams (negligible). - Total per layer:
42M + 8 × 176M ≈ 42M + 1.41B ≈ 1.45B - Active per layer (top-2):
42M + 2 × 176M ≈ 394M
× 32 layers + embeddings (~130M):
- Total: ~46.7B parameters (matches "47B" label).
- Active per token: ~12.9B (matches "13B" quality proxy).
That's the whole trick. You memorized 8× the MLP capacity, you compute 2/8 = 25% of it per token.
5 · Inference is the hard part
MoE is elegant on paper and painful in production. Here's why.
5.1 · Memory bandwidth
Every expert must be loaded into VRAM even though only 2 run per token. So a 47B MoE occupies as much VRAM as a 47B dense model, but at compute cost of 13B. Great for throughput per FLOP, terrible for what you can fit on one GPU.
5.2 · Expert imbalance at inference
At training, batch-averaged load balancing works because you have lots of tokens. At inference — especially single-batch, low-latency chat — one prompt might route almost everything to expert 3 for a stretch, leaving expert 3's GPU pegged and 7 others idle. vLLM and TensorRT-LLM have specific optimizations for this (batched expert-grouped GEMMs, expert parallelism across GPUs).
5.3 · Expert parallelism
The natural sharding for a big MoE: put each expert on a different GPU. Now every token that routes to expert 3 has to move its activation to GPU 3, run the MLP, and move back. This is called all-to-all communication and its cost dominates on slow interconnects. It's fine on NVLink; a disaster on PCIe.
Naive setup: expert-parallel across 4 GPUs, PCIe interconnect. All-to-all was 40% of step time. Mixtral was slower than a dense 13B on one GPU. Fix that worked: replicate the whole model on each GPU (needs 90+ GB per GPU with FP16), use tensor parallelism instead of expert parallelism. Faster, uses more memory. There's no free lunch here.
Tried to LoRA-fine-tune Mixtral on ~50k in-domain examples. After a few hundred steps, benchmarks tanked. Turned out the router had drifted: 60% of tokens were routing to just 2 experts. Fine-tuning had bypassed the aux-loss (I'd only added LoRA to the attention). Fix: freeze router entirely during LoRA, or explicitly include the aux loss. Prefer freezing.
INT4 quantization is usually ~1 point of perplexity loss on dense models. On Mixtral it was ~5 points. Reason: expert-specific weight distributions are more diverse than a single MLP's, and uniform per-tensor scales fit them poorly. Fix: per-expert quantization scales (AWQ / GPTQ variants that treat each expert independently).
6 · When MoE makes sense for you
Yes, use MoE when:
- You have lots of GPU memory but limited FLOP budget per token (e.g., high-QPS serving with big VRAM boxes).
- Your data is genuinely diverse (multi-language, code + prose, multiple domains) — different experts can genuinely specialize.
- You're training from scratch at 30B+ scale and can afford the routing complexity.
No, dense is better when:
- You're at hobby scale (< 10B). MoE overhead isn't worth it below this — the router adds complexity without buying much.
- You're memory-constrained (one 24GB GPU). Dense fits easier.
- Your workload is homogeneous (one domain, one language). Experts won't specialize meaningfully.
- You're doing your capstone. Ship dense. Add MoE in v2 if you find a use case.
7 · Diagram — the MoE block
8 · The 2024–2025 MoE landscape — from Mixtral to DeepSeek-V3
The field went from "MoE is a research curiosity" (2022) to "the largest open-weights model on Earth is MoE" (Dec 2024) in about 20 months. The vocabulary you need:
8.1 · Mixtral 8×7B and 8×22B (Mistral, Dec 2023 / Apr 2024)
Mixtral 8×7B (arXiv:2401.04088) was the first widely-used open-weights MoE. 8 experts per layer, top-2 routing, 47B total / 13B active. Apache-licensed. Every downstream MoE paper compares to it. Simple recipe, no tricks — which made it a great teaching model.
8.2 · DeepSeek-V2 → V3 → R1 (2024–2025)
DeepSeek-V2 (May 2024, arXiv:2405.04434) introduced two things that changed the field:
- MLA — Multi-head Latent Attention. Instead of caching K and V per head, project to a low-rank latent (d_c ≈ 512), cache that, and reconstruct K/V during attention. KV-cache shrinks 5–10×. See S037 for the details you already covered.
- Fine-grained expert segmentation + shared experts. Instead of 8 experts of size 7B, use 160 experts of size ~350M plus 2 always-on "shared" experts. Fine-grained specialization gives each expert a narrower niche; shared experts absorb the common cross-domain capabilities so specialists don't have to relearn them.
DeepSeek-V3 (Dec 2024, arXiv:2412.19437) scaled this to 256 routed experts + 1 shared per layer, 61 layers, 671B total / 37B active per token. Highlights:
- Auxiliary-loss-free load balancing — they replaced the aux loss with a per-expert bias that is nudged up when an expert is under-used and down when over-used, applied only during routing decisions (not gradient path). Cleaner gradients, better quality.
- Multi-token prediction (MTP) — a second output head trained to predict the next-next token in parallel. Provides an auxiliary training signal (better representations) and enables speculative-decoding-style inference speedups at deploy time.
- FP8 training end-to-end. Training cost: ~2,788K H800-hours, ~$5.5M in rented compute. Compare Llama 3 405B: ~30M H100-hours. ~5× cheaper for arguably better quality on many benchmarks.
DeepSeek-R1 (Jan 2025, arXiv:2501.12948) then took V3 and did large-scale RL for reasoning (GRPO, no PPO critic), producing an open o1-class reasoner. The MoE architecture was unchanged; the alignment recipe was the news.
8.3 · Qwen2/3-MoE, DBRX, Llama 4
- Qwen1.5/2/3-MoE (Alibaba, 2024–2025) — similar fine-grained expert scheme as DeepSeek-V2. Qwen3-235B-A22B is currently one of the strongest open MoEs.
- DBRX (Databricks, Mar 2024, 132B/36B active) — 16 experts, top-4 routing. Notable for training entirely on Databricks' own infra as a demo.
- Llama 4 Scout / Maverick / Behemoth (Meta, Apr 2025) — Meta finally shipped MoE, with 16–16 experts (Scout) up to 128 experts (Behemoth), native multimodal, 10M context on Scout via iRoPE (interleaved layers with no positional encoding on some layers).
8.4 · The "active parameters" mental model (Muennighoff / Chinchilla for MoE)
Hoffmann's Chinchilla laws were for dense models. For MoE, the compute-optimal scaling behaves closer to active parameters, but data-hungriness scales with total parameters (Krajewski et al., 2024, arXiv:2402.07871). Practical implication: if a Mixtral-scale MoE (~13B active, ~47B total) trained on Chinchilla-optimal-for-13B tokens (~260B), it's under-trained for its total-parameter memory footprint. DeepSeek trained V3 on 14.8T tokens — way past dense-Chinchilla-optimal, but appropriate for the total-parameter capacity. Rule of thumb: budget tokens as if training a dense model of size √(active × total).
8.5 · Serving MoE in 2025 — vLLM and SGLang
Both vLLM (v1, 2025) and SGLang added first-class expert parallelism + expert-batched GEMMs (Grouped GEMM via CUTLASS / TE). The key optimization: batch tokens across a sequence that route to the same expert, run one big matmul, unshuffle. Combined with speculative decoding (S069) using a small dense draft model against a big MoE target, you can get 2–4× tokens/sec vs naive routing.
Deployment reality: DeepSeek-V3 needs ~8×H100-80G in FP8 to serve; ~16× in BF16. Not a laptop model. But Mixtral 8×7 fits in ~90GB FP16 or ~45GB INT4 (2× 4090s). Choose the scale that matches your hardware.
8.6 · Further reading
- Sparsely-Gated MoE (the original): https://arxiv.org/abs/1701.06538 · Shazeer et al. 2017
- Switch Transformer: https://arxiv.org/abs/2101.03961 · Fedus et al. 2021
- Mixtral: https://arxiv.org/abs/2401.04088
- DeepSeek-V2 (MLA + fine-grained MoE): https://arxiv.org/abs/2405.04434
- DeepSeek-V3 tech report: https://arxiv.org/abs/2412.19437
- DeepSeek-R1: https://arxiv.org/abs/2501.12948
- Scaling laws for MoE: https://arxiv.org/abs/2402.07871 · Krajewski et al. 2024
- Llama 4 blog: https://ai.meta.com/blog/llama-4-multimodal-intelligence/
- MegaBlocks (block-sparse MoE kernels): https://arxiv.org/abs/2211.15841
9 · Try it yourself
One evening of MoE from scratch:
- Take your 3-layer nano-transformer from S042.
- Swap one MLP block for a 4-expert top-2 MoE (the 20-line block from §1).
- Train on TinyStories (~500MB). Log per-expert token counts every 100 steps.
- Plot: does the aux loss keep counts balanced? Set the aux coefficient to 0 and re-run. Watch one or two experts eat everything.
- Bonus: add DeepSeek's aux-loss-free bias nudging (a scalar bias per expert, updated by an exponential moving average of relative load). Compare.
You'll come away with a felt sense for why routing is the whole game.
10 · Recap
MoE swaps one MLP for N experts and a router that picks k. You pay N× the memory, spend k/N× the compute. Everything hard about MoE is about (a) keeping the router balanced (aux loss, or DeepSeek's bias nudging), (b) making all-to-all communication cheap (expert parallelism on fast interconnects), and (c) quantizing per-expert distributions correctly. The 2024–2025 frontier is fine-grained experts (100s of small experts + shared), FP8 training, and treating MTP as both a training regularizer and an inference speedup. If you're serving MoE below hyperscaler scale, you're mostly picking pre-trained weights and tuning vLLM flags — which is a fine place to be.
"An MoE model with 8 experts and top-2 routing uses a quarter of the compute, so it's a quarter of the cost to serve. Sparse activation means sparse resource usage."
Sparse in FLOPs, dense in memory. Every expert's weights must be resident, because any token in the batch may route to any of them — you hold all parameters and use a fraction per token. Since decode is memory-bandwidth-bound rather than compute-bound, the resource that MoE saves is the one you were not short of, and the resource it does not save is the one that limits you.
The belief is seductive because it is genuinely true for training, which is where MoE's headline results come from and where the story you heard was written. Training is compute-bound and runs at large batch, so cutting FLOPs per token cuts wall-clock almost proportionally — the claim is not marketing, it is correct in its original context. The failure is transporting a compute-regime conclusion into a bandwidth-regime deployment. It is also reinforced by parameter counts being quoted two ways, active and total, in a way that invites you to use the flattering number for whichever property you are discussing.
Compute both budgets from a config instead of from the marketing line:
E, K = 8, 2 # experts, top-k
d_model, d_ff, L = 4096, 14336, 32
attn_per_layer = 4 * d_model * d_model
expert_per_layer = 3 * d_model * d_ff # gated MLP
total_params = L * (attn_per_layer + E * expert_per_layer)
active_params = L * (attn_per_layer + K * expert_per_layer)
print('total GB bf16:', round(total_params * 2 / 1e9, 1)) # what you must hold
print('active GB bf16:', round(active_params * 2 / 1e9, 1)) # what you compute with
print('ratio:', round(total_params / active_params, 2))The first number sizes your GPU. The second sizes your FLOPs. Quoting one when you mean the other is the entire misconception.
Why does an MoE layer need an explicit load-balancing loss? If some experts are genuinely more useful, letting the router prefer them seems like the correct outcome.
- 1The router is trained end-to-end on the task loss, and its gradient rewards routing a token to whichever expert currently handles it best.forced by · there is no term in the task objective that cares which expert was used, only that the output was good
- 2An expert that receives more tokens receives more gradient updates and therefore improves faster than its peers.forced by · parameters only move when tokens flow through them, so token share directly sets learning rate in effect
- 3Improving faster makes it the best choice for still more tokens, which is a positive feedback loop with no opposing force in the task loss.forced by · the router's preference and the expert's competence reinforce each other, and nothing in the objective penalises concentration
- 4The fixed point of that loop is collapse: a few experts absorb nearly all traffic and the rest are never selected, so they never improve and never will.forced by · an expert receiving no tokens receives no gradient, which is an absorbing state rather than a slow one
- 5Collapse destroys the entire premise — you are now holding the memory of E experts while getting the capacity of a handful, which is strictly worse than a dense model of the same footprint.forced by · unused parameters cost memory and bandwidth budget without contributing capacity
- 6So an auxiliary term must penalise imbalance directly, adding the opposing force the task loss structurally cannot provide.forced by · a feedback loop with no negative term can only be stopped from outside the loop
Therefore load balancing is not a tuning refinement — it is what prevents an MoE from degenerating into a dense model wearing an expensive costume. The instability is structural, not incidental.
And note the prediction. First: since the mechanism is a runaway feedback loop, imbalance should appear early in training and be hard to reverse later — so monitor per-expert token counts from the first steps rather than checking at the end, and expect that a collapse caught late cannot be fixed by simply raising the balance coefficient. Second: the balance coefficient must have a genuine sweet spot, because too large a value forces routing toward uniformity, which means routing carries no information and the experts cannot specialise at all — sweep it and expect quality to degrade at both extremes. Third: at inference there is no balance loss and no training signal, so routing is determined entirely by input distribution; a production traffic mix unlike the training mix should produce measurably more skewed expert utilisation, which is a real and often overlooked source of inference-time load imbalance across devices.
Total parameters are the library: every book must be on the shelves, physically present, taking up floor space whether or not anyone opens it. Active parameters are the desk: how many books you actually read per token. MoE builds a much bigger library without enlarging the desk.
That is a real win when your constraint is reading speed — training, and any regime where you are compute-bound. It is no win at all when your constraint is floor space, which is exactly the situation at inference, where holding all the weights in HBM and streaming them is what sets your limit. The router is the librarian, and if it sends everyone to the same three shelves you have paid for a library and built a bookshelf.
- Sparse in FLOPs, dense in memory. Quote both parameter counts or you are misleading someone, possibly yourself.
- Routing collapse is the default outcome without an explicit counter-force. Monitor per-expert token share from step one.
- Expert parallelism turns routing into network traffic: tokens must travel to the device holding their expert, so an all-to-all sits in the critical path of every layer.
- Inference-time load balance depends on the input distribution, which you do not control and which will not match training.
- MoE buys capacity per FLOP. If your bottleneck is memory or bandwidth rather than compute, it buys you nothing.
Fire this the moment you see: a parameter count quoted without saying total or active · an MoE deployment that OOMs on hardware sized from the active count · expert utilisation histograms nobody is looking at · training loss that plateaus while a few experts dominate · MoE proposed to reduce serving cost on a memory-bound deployment · multi-GPU MoE latency dominated by communication rather than compute · quality that varies by input domain more than a dense model would.
You have a fixed memory budget for serving. Dense model, or a sparse model with the same footprint?
Decide by measuring which regime you are in before comparing architectures at all. MoE's advantage is quality per FLOP, so it pays exactly when FLOPs are what you are short of. At single-stream decode you are short of bandwidth, and a sparse model gives you a large memory bill for a benefit you cannot collect.
If you do serve MoE, instrument expert utilisation in production as a first-class metric, not as a training-time curiosity. Routing at inference is driven entirely by input distribution, and a shift in traffic shows up as skew — which manifests as load imbalance across devices and as quality that quietly varies by customer or domain. It is also one of the few metrics that gives you early warning of distribution shift for free, which makes it worth collecting even when everything is fine.
🧠 Retention scaffold
One-line summary (write it in your own words): _______________________________
Spaced review: re-read §1 (the block) + §3 (load balancing) in 24 hours. Revisit alongside the DeepSeek-V3 tech report on day 7.
Next session (S078): your model has a fixed context window from pretraining. What if you want to feed it a whole book? Position Interpolation, NTK-aware RoPE, YaRN, LongRoPE — how to stretch 4k to 128k+ without a full retrain.
Sticky note (keep on your desk): MoE = N experts, top-k routing, k/N compute at N× memory. Everything else is routing tricks.
Previous: ← DL S076 · Next: DL S078 → Long Context