27. Mixture of experts
In this chapter
- How a mixture-of-experts layer decouples a model's size from the compute per token.
- The router: softmax, top-k, renormalization, and a gated shared expert.
- Dispatch: from a readable per-expert loop to the sort-by-expert layout that grouped matrix multiplications consume.
- The economics of MoE inference: why batch-1 decode is fast, why memory is the price, and what happens as the batch grows.
- Loading a real Qwen3-MoE checkpoint and matching Hugging Face's output.
You will build
TopKRouter.forward and Experts.forward_grouped in engine/moe.py. Your Qwen3 grows into Qwen3-MoE, and the same block becomes Flash-Next's MoE in Chapter 30.
Time: 5-6 hours. GPU: optional.
Size without compute
In a dense transformer, every parameter takes part in every token. Doubling the parameters doubles the compute per token and the bytes read per decode step. A mixture of experts (MoE) breaks that link. Each MLP is replaced by $E$ smaller MLPs, the experts, and a router that sends each token to only $k$ of them.
dense layer: x ──► MLP ──► y every token uses all of it
MoE layer: x ──► router ──► top-k of E experts ──► weighted sum ──► y
Qwen3-30B-A3B has 128 experts per layer and uses 8 per token. It stores 30.5B parameters but each token is processed by about 3.3B active ones. The model knows roughly what a 30B model knows, and computes like a 3B one.
python run.py moe
{"model": "Qwen3-30B-A3B", "expert_params_stored_B": 29.0, "expert_params_active_B": 1.82, "active_fraction": 0.063}
{"model": "Qwen3-235B-A22B", "expert_params_stored_B": 227.15, "expert_params_active_B": 14.24, "active_fraction": 0.063}
{"model": "Qwen3.8-Flash-Next", "expert_params_stored_B": 121.09, "expert_params_active_B": 2.66, "active_fraction": 0.022}
(Expert parameters only; attention and embeddings add the rest.) Flash-Next, the capstone’s target, routes each token to 10 of 512 experts plus a shared expert: it stores 121B parameters of experts and uses 2.2% of them per token.
The router
The router is one linear layer from the hidden state to $E$ scores. Qwen-family models then use:
- Softmax over all experts, in FP32: $\pi = \operatorname{softmax}(W_r x)$.
- Top-k: keep the $k$ largest probabilities and their expert indices.
- Renormalize (when
norm_topk_probis set): divide the kept weights by their sum so they add to 1.
The layer’s output is the weighted sum of the chosen experts’ outputs: $y = \sum_{i \in \text{top-}k} w_i, \text{expert}_i(x)$.
Example. Four experts, $k = 2$, router probabilities $[0.1, 0.5, 0.3, 0.1]$. Experts 1 and 2 are chosen with weights $0.5$ and $0.3$; renormalized, they become $0.625$ and $0.375$.
class TopKRouter(nn.Module):
def __init__(self, hidden, experts, top_k, normalize=True):
super().__init__()
self.weight = nn.Parameter(torch.zeros(experts, hidden))
self.top_k, self.normalize = top_k, normalize
def forward(self, x):
"""x [N, D] -> (logits [N, E], weights [N, k], experts [N, k]). (Your engine: Chapter 27)
Softmax over all experts in FP32, keep the k largest, and (if normalize) rescale the
kept probabilities to sum to 1.
"""
logits = F.linear(x, self.weight)
probs = torch.softmax(logits, dim=-1, dtype=torch.float32)
weights, experts = probs.topk(self.top_k, dim=-1)
if self.normalize:
weights = weights / weights.sum(-1, keepdim=True)
return logits, weights.to(x.dtype), experts
Two details matter for exact parity. The softmax must be in FP32: the router’s decision is discrete, so a BF16 rounding difference can swap two near-tied experts and change the output completely. And ties in topk are broken by index, which matters when comparing against another implementation (Chapter 29 meets this problem in earnest).
Other routers exist: DeepSeek-V3 uses a sigmoid score per expert with a learned bias for load balancing, and some models route groups of experts first. Each checkpoint documents its own; from_hf refuses configurations it doesn’t implement.
The shared expert
Some models add one shared expert that every token uses, alongside the routed ones. In Qwen3-Next and Flash-Next, its output is scaled by a learned gate, $\sigma(w_s^\top x)$, so each token decides how much of the shared knowledge to mix in:
$$ y = \sum_{i \in \text{top-}k} w_i, \text{expert}_i(x) + \sigma(w_s^\top x), \text{shared}(x). $$
class SparseMoeBlock(nn.Module):
"""Router + routed experts (+ an optional shared expert whose output is scaled by
sigmoid(shared_expert_gate(x)), as in Qwen3-Next and Flash-Next)."""
def __init__(self, hidden, experts, top_k, intermediate, shared_intermediate=0, normalize=True):
super().__init__()
self.gate = TopKRouter(hidden, experts, top_k, normalize)
self.experts = Experts(experts, hidden, intermediate)
self.shared_expert = SwiGLU(hidden, shared_intermediate) if shared_intermediate else None
self.shared_expert_gate = nn.Linear(hidden, 1, bias=False) if shared_intermediate else None
self.dispatch = "grouped"
def forward(self, x):
shape = x.shape
flat = x.reshape(-1, shape[-1])
_, weights, experts = self.gate(flat)
run = self.experts.forward_grouped if self.dispatch == "grouped" else self.experts.forward_loop
out = run(flat, weights, experts)
if self.shared_expert is not None:
out = out + torch.sigmoid(self.shared_expert_gate(flat)) * self.shared_expert(flat)
return out.reshape(shape)
Dispatch: getting tokens to their experts
Experts are stored stacked: one tensor gate_up_proj [E, 2I, D] holding every expert’s gate and up projections, and one down_proj [E, D, I]. Checkpoints usually store each expert’s matrices separately (mlp.experts.17.gate_proj.weight), and the loader stacks them.
class Experts(nn.Module):
"""E SwiGLU experts in two stacked tensors: gate_up_proj [E, 2I, D] (gate rows, then up rows)
and down_proj [E, D, I]. One tensor per kind keeps the layout friendly to grouped kernels."""
def __init__(self, experts, hidden, intermediate):
super().__init__()
self.gate_up_proj = nn.Parameter(torch.empty(experts, 2 * intermediate, hidden))
self.down_proj = nn.Parameter(torch.empty(experts, hidden, intermediate))
nn.init.normal_(self.gate_up_proj, std=0.02)
nn.init.normal_(self.down_proj, std=0.02)
@property
def num_experts(self):
return self.gate_up_proj.shape[0]
def expert(self, e, x):
gate, up = F.linear(x, self.gate_up_proj[e]).chunk(2, dim=-1)
return F.linear(F.silu(gate) * up, self.down_proj[e])
def forward_loop(self, x, weights, experts):
"""Readable reference: for each expert, gather its tokens, run it, scatter-add back."""
out = torch.zeros_like(x)
for e in experts.unique().tolist():
token, slot = torch.where(experts == e)
out.index_add_(0, token, self.expert(e, x[token]) * weights[token, slot, None])
return out
def forward_grouped(self, x, weights, experts):
"""Sort the N*k assignments by expert so each expert's rows are contiguous: this is the
layout a grouped GEMM consumes (one launch, many independent matmuls). (Your engine: Chapter 27)"""
n, k = experts.shape
flat_expert = experts.reshape(-1)
order = flat_expert.argsort(stable=True)
token_of = order // k # which token each sorted assignment came from
counts = torch.bincount(flat_expert, minlength=self.num_experts).tolist()
rows = x[token_of] # [N*k, D] grouped by expert
results, start = [], 0
for e, count in enumerate(counts): # a grouped GEMM does these as one kernel
if count:
results.append(self.expert(e, rows[start:start + count]))
start += count
expert_out = torch.cat(results) * weights.reshape(-1)[order, None]
return torch.zeros_like(x).index_add_(0, token_of, expert_out)
The loop dispatch is the readable reference: for each expert used in the batch, find its tokens, run the expert on them, and add the weighted result back to those tokens’ rows. Each expert is a separate small matmul, and the loop body runs once per active expert, with a host sync from unique().
The grouped dispatch prepares the layout a fast kernel needs:
- Flatten the $N \times k$ assignments and sort them by expert (
argsort, stable). Now each expert’s rows are contiguous. - Permute the token rows into that order (
x[token_of]), duplicating each token $k$ times. - Run each expert on its contiguous slice. A grouped GEMM kernel does all of these as one launch: many independent matmuls of different sizes.
- Multiply by the routing weights and unpermute with
index_add_, summing each token’s $k$ contributions.
On a CPU, with 512 tokens, 32 experts and $k = 4$:
{"tokens": 512, "top_k": 4, "assignments": 2048, "busiest_expert": 80, "idlest_expert": 52, "ideal_per_expert": 64}
{"dispatch": "loop", "ms": 16.1, "checksum": 1027.975}
{"dispatch": "grouped", "ms": 9.3, "checksum": 1027.975}
Same result, and even the Python-level grouped version is faster. On a GPU, production MoE layers fuse the permutation, the grouped GEMMs and the unpermutation into a few kernels (vLLM’s fused_moe Triton kernel, MegaBlocks, DeepGEMM).
The first line shows load imbalance: with a random router, one expert got 80 tokens and another 52, against an ideal of 64. Trained routers are pushed toward balance during training, with an auxiliary loss (Switch Transformer: $E \sum_i f_i P_i$, the fraction of tokens routed to expert $i$ times its mean probability) or a bias adjusted on the fly (DeepSeek-V3’s auxiliary-loss-free balancing). At inference, imbalance means the busiest expert sets the step time.
The economics of MoE inference
Batch-1 decode is fast. A decode step reads only the experts its token uses. For Qwen3-30B-A3B in BF16, that’s about 6.6 GB per token (3.3B active parameters × 2 bytes), not 61 GB: decode runs at the speed of a 3B dense model.
Memory is the price. All experts must be resident, because the next token may need any of them. A 30B MoE needs 30B parameters of memory, the same as a 30B dense model. That’s why MoE models pair so well with quantization (experts are most of the bytes, Chapter 20) and with large-memory machines: the 128 GB of unified memory in a DGX Spark holds Flash-Next with 4-bit experts.
Larger batches touch more experts. With $B$ tokens each choosing $k$ of $E$ experts roughly uniformly, a layer touches about $E,\big(1 - (1 - k/E)^B\big)$ distinct experts:
| batch | Qwen3-30B-A3B ($E = 128$, $k = 8$) | Flash-Next ($E = 512$, $k = 10$) |
|---|---|---|
| 1 | 8 | 10 |
| 8 | 52 | 75 |
| 32 | 112 | 240 |
| 128 | 128 (all) | 471 |
By batch 32, a decode step of Qwen3-30B-A3B reads almost every expert, the full 61 GB, while each expert does only a little work (about $Bk/E = 2$ tokens each). MoE decode at moderate batch is therefore more memory-bound per token than a dense model with the same active parameters. Throughput still grows with batch until each expert gets enough tokens to be compute-bound, which needs batches in the hundreds or thousands. That’s why large MoE deployments use expert parallelism: experts spread over many GPUs, tokens exchanged with all-to-all communication, so that each GPU holds few experts and gets many tokens for each (Chapter 41).
Prefill is easy. Thousands of prompt tokens give every expert plenty of rows; grouped GEMMs run near peak.
Qwen3-MoE
Qwen3-MoE is Qwen3 with every MLP replaced by a SparseMoeBlock, and nothing else changed. Your Chapter 17 code gives you everything but the block:
@dataclass
class Qwen3MoeConfig(Qwen3Config):
num_experts: int = 8
num_experts_per_tok: int = 2
moe_intermediate_size: int = 64
norm_topk_prob: bool = True
@classmethod
def from_hf(cls, raw):
if raw.get("model_type") != "qwen3_moe":
raise ValueError("Expected model_type 'qwen3_moe'")
if raw.get("mlp_only_layers") or raw.get("decoder_sparse_step", 1) != 1:
raise ValueError("Dense layers inside a Qwen3-MoE stack are not implemented")
raw = dict(raw)
if "num_experts" not in raw and "num_local_experts" in raw: # the name Transformers 5 writes
raw["num_experts"] = raw["num_local_experts"]
dense = Qwen3Config.from_hf({**raw, "model_type": "qwen3"})
own = [f.name for f in fields(cls) if not hasattr(Qwen3Config, f.name)]
missing = [name for name in own if name not in raw]
if missing: # never fall back to defaults for the model's shape
raise ValueError(f"config.json lacks {missing}")
return cls(**vars(dense), **{name: raw[name] for name in own})
class Qwen3MoeLayer(Qwen3Layer):
def __init__(self, cfg):
super().__init__(cfg)
self.mlp = SparseMoeBlock(cfg.hidden_size, cfg.num_experts, cfg.num_experts_per_tok,
cfg.moe_intermediate_size, 0, cfg.norm_topk_prob)
class Qwen3Moe(Qwen3):
"""Qwen3 with every MLP replaced by a routed-expert block."""
def __init__(self, cfg):
super().__init__(cfg)
self.model.layers = nn.ModuleList(Qwen3MoeLayer(cfg) for _ in range(cfg.num_hidden_layers))
@torch.no_grad()
def load_qwen3_moe(directory, device="cpu", dtype=torch.bfloat16):
cfg = Qwen3MoeConfig.from_hf(read_config(directory))
model = Qwen3Moe(cfg).to(device=device, dtype=dtype)
params = dict(model.named_parameters())
pending = {}
for name, value in snapshot_tensors(directory):
parts = name.split(".")
if ".mlp.experts." in name: # model.layers.L.mlp.experts.E.kind.weight
layer, e, kind = int(parts[2]), int(parts[5]), parts[6]
pending.setdefault(layer, {})[(e, kind)] = value
continue
if name == "lm_head.weight" and cfg.tie_word_embeddings:
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
for layer, tensors in pending.items():
gate_up, down = stack_expert_tensors(tensors, cfg.num_experts, cfg.moe_intermediate_size)
block = model.model.layers[layer].mlp.experts
assign(block.gate_up_proj, gate_up, f"layer {layer} gate_up_proj")
assign(block.down_proj, down, f"layer {layer} down_proj")
return model.eval()
The loader collects the per-expert tensors as it streams the safetensors shards, then stacks them per layer. Qwen3-30B-A3B is 61 GB in BF16, so this chapter’s tests build tiny random Qwen3-MoE checkpoints, save them in Hugging Face’s format, and compare your model’s logits with transformers’ Qwen3MoeForCausalLM, with and without renormalization. The largest logit difference is about $10^{-7}$ in FP32, on logits of magnitude 0.5.
Note
Writing that test exposed a real bug. Transformers 5 saves the expert count as
num_local_experts, while Qwen’s published configs saynum_experts. The first version offrom_hfsilently fell back to its default of 8 experts and passed every test that happened to use 8. Now it accepts both names and refuses a config that lacks the field. Two lessons: never let a model’s shape come from a default, and make test models use non-default sizes everywhere.
def moe_parameter_split(hidden, layers, experts, top_k, intermediate, shared_intermediate=0):
"""(stored expert parameters, active expert parameters per token) for the MLP part only."""
per_expert = 3 * hidden * intermediate
shared = 3 * hidden * shared_intermediate
stored = layers * (experts * per_expert + shared + experts * hidden)
active = layers * (top_k * per_expert + shared + experts * hidden)
return stored, active
Build it
Engine milestone 27: mixture of experts. Implement TopKRouter.forward and Experts.forward_grouped in engine/moe.py (the loop dispatch, the block, the model and the loader are provided).
pytest tests/test_ch27_moe.py
python run.py moe --impl engine
The tests check the router’s selection and renormalization, that grouped dispatch equals the loop for random routings (including experts that receive no tokens), the parameter split, and logit parity with Hugging Face’s Qwen3-MoE on saved tiny checkpoints. If you have the memory, load Qwen/Qwen3-30B-A3B with load_model and chat with it through your Chapter 18 engine.
Stretch exercises
- ★ Run a tiny Qwen3-MoE on 1,000 tokens of text and record every layer’s expert counts. Is the load more balanced in early or late layers? (Use a real checkpoint if you can; random routers are uninformative.) Where:
experiments/ch27.py(create it), registering hooks onTopKRouterinengine/moe.py. - ★★ Write a Triton grouped GEMM: one program per (expert, output tile), with a prefix-sum array of row offsets telling each program where its expert’s rows start. Where: new
engine/kernels/triton_grouped.py, called byExperts.forward_groupedinengine/moe.py. - ★★ Quantize only the experts to 4 bits with Chapter 20’s
QuantLinearlogic adapted to stacked tensors, leaving attention and the router in BF16. Measure memory and logit error. Where: add quantized stacked-expert storage/dispatch inengine/moe.py, usingengine.quant. - ★★★ Expert offloading: keep experts in CPU memory, keep a GPU cache of the most recently used ones, and copy missing experts on demand. Measure cache hit rates over a conversation. Where: add an expert-cache variant in
engine/moe.py; Chapter 40’sengine/offload.pyprovides the later integration.
Check your understanding
- Why does an MoE’s decode speed at batch 1 depend on active parameters, while its memory depends on total parameters?
- Why must the router’s softmax run in FP32?
- What does sorting the assignments by expert buy you?
- Why does a moderate batch make MoE decode read nearly all expert weights?
- What problem does a load-balancing loss solve, and why does imbalance matter at inference too?
Going deeper
- Shazeer et al., Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer (2017); Fedus, Zoph and Shazeer, Switch Transformers (2021); Jiang et al., Mixtral of Experts (2024); DeepSeek-AI, DeepSeekMoE (2024) and the DeepSeek-V3 technical report (fine-grained and shared experts, auxiliary-loss-free balancing).
- Gale et al., MegaBlocks: Efficient Sparse Training with Mixture-of-Experts (2022), on grouped and block-sparse kernels.
- The Qwen3 technical report (2025) for Qwen3-MoE, and the Qwen3-Next model card for the gated shared expert.
- vLLM’s
fused_moeTriton kernel and SGLang’s MoE runner; GPU Mode L11 (Sparsity) for the broader picture of sparse computation on GPUs.