Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Describe MoE and how Mixtral routes tokens at inference time?
What you need to know
The routing rule
The Mixtral paper writes the gate as a softmax over the top-K logits:
logits = x · W_router # one score per expert (8)G(x) = softmax(TopK(logits, K=2)) # the other 6 get weight 0y = Σ over the 2 chosen experts of G_i(x) · Expert_i(x)1import torch, torch.nn as nn, torch.nn.functional as F2torch.manual_seed(0)3d, n_exp, k = 16, 8, 24router = nn.Linear(d, n_exp, bias=False)5experts = nn.ModuleList([nn.Sequential(nn.Linear(d, 4*d), nn.SiLU(), nn.Linear(4*d, d))6 for _ in range(n_exp)])7x = torch.randn(5, d) # 5 tokens8w, idx = router(x).topk(k, dim=-1) # pick 2 experts per token9w = F.softmax(w, dim=-1) # weights over the 2 chosen10y = torch.zeros_like(x)11for e in range(n_exp):12 tok, slot = (idx == e).nonzero(as_tuple=True)13 if len(tok):14 y[tok] += w[tok, slot, None] * experts[e](x[tok])15print(idx.tolist()) # [[3, 2], [5, 0], [3, 6], [0, 5], [7, 5]]Five tokens, five different expert pairs. Real kernels group tokens by expert and run each expert once as a batched matmul, as the loop does here.
Where "46.7B total, 12.9B active" comes from
Mixtral uses d_model = 4096, FFN width 14,336, 32 layers and GQA with 8 KV heads.
one expert FFN per layer: 3 × 4096 × 14336 ≈ 176Mall experts: 176M × 8 × 32 ≈ 45.1Battention (shared) ≈ 1.34B, embeddings + head ≈ 0.26Btotal ≈ 46.7Bactive = 2 experts × 32 layers × 176M + 1.34B + 0.26B ≈ 12.9BSo "8x7B" is not eight 7B models. Only the FFNs are copied; attention and embeddings are shared.
What the router learns
Routing happens per token and per layer. The Mixtral paper looked for topic specialisation — for example, maths versus biology — and found no clear pattern; routing followed syntax and token identity more than subject.
Training details interviewers probe
- Load balancing. Without help, the router sends most tokens to a few experts and the rest stop learning. The Switch Transformer (Fedus et al., 2021) added an auxiliary loss that rewards even use. DeepSeek-V3 instead adjusts a per-expert bias during training, with no extra loss term.
- Capacity. If too many tokens pick one expert in a batch, some systems drop the overflow; newer kernels handle uneven loads without dropping.
- Newer designs. DeepSeek-V3 uses many small experts — 256 routed plus 1 shared expert, 8 routed chosen per token — so each token gets a finer mix.
A real-life example
An Indian-language news app wants a stronger model for summarising and translating articles, and considers Mixtral 8x7B. In fp16, all 46.7B parameters take about 93 GB — more than one 80 GB GPU. The team either shards it across two GPUs or quantises to 4-bit (about 26 GB) and uses one.
For one user at night (batch size 1), each token needs only its 2 experts per layer, so a decode step reads about 13B parameters — it runs like a 13B model. At peak traffic, the batch covers many tokens and almost all 8 experts are used in every layer, so each step reads nearly all 47B parameters; throughput stays good only because many tokens share that read. The compute is "13B", but the memory is "47B" — the team sizes the cluster for memory.
Follow-up questions to expect
- "Why top-2 rather than top-1?" — Top-1 (the Switch Transformer) is cheapest; top-2 gives each token a blend of two experts and a smoother training signal. Newer fine-grained MoEs choose 8 of many small experts.
- "Does the router pick experts per prompt?" — No, per token and per layer. A single sentence can pass through dozens of different expert pairs.
- "Why not make attention MoE too?" — Most parameters and FLOPs of a block sit in the FFN, so that is where sparsity pays off most; attention also holds the KV cache, which experts would complicate.