Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Describe MoE and how Mixtral routes tokens at inference time?


One token through one Mixtral MoE layerHidden state xfor this tokenRouter scoresall 8 expertsKeep thetop 2 scoresSoftmax the 2into weightsWeighted sum of2 expert outputsThe next token, and the next layer, may choose a different pair.
Only two of eight FFNs run per token, so Mixtral computes like 12.9B parameters but must hold all 46.7B in memory.

What you need to know

The routing rule

The Mixtral paper writes the gate as a softmax over the top-K logits:

Text
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)
Python
import torch, torch.nn as nn, torch.nn.functional as Ftorch.manual_seed(0)d, n_exp, k = 16, 8, 2router = nn.Linear(d, n_exp, bias=False)experts = nn.ModuleList([nn.Sequential(nn.Linear(d, 4*d), nn.SiLU(), nn.Linear(4*d, d))                         for _ in range(n_exp)])x = torch.randn(5, d)                           # 5 tokensw, idx = router(x).topk(k, dim=-1)              # pick 2 experts per tokenw = F.softmax(w, dim=-1)                        # weights over the 2 choseny = torch.zeros_like(x)for e in range(n_exp):    tok, slot = (idx == e).nonzero(as_tuple=True)    if len(tok):        y[tok] += w[tok, slot, None] * experts[e](x[tok])print(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.

Text
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.9B

So "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.