✨ Generative AI & LLMs · Lecture 22 of 30

Mixture of Experts: Scaling Parameters Without Scaling Compute

Mixture-of-experts layers route each token to a few of many expert networks, so models gain parameters without proportional compute. We cover gating, top-k routing, load balancing, capacity, training and serving challenges.

In a standard (dense) transformer, every token passes through every parameter. Doubling parameters doubles compute. Mixture of Experts (MoE) breaks this link: a layer contains many "expert" sub-networks, and a router sends each token to only a few of them. A model can have hundreds of billions of parameters while each token uses only a fraction — more knowledge capacity at a similar cost per token. Several prominent open models (e.g. Mixtral, DeepSeek-V3, Qwen MoE variants) use MoE, and it is widely reported to be used in frontier models.

History#

The idea dates to Jacobs, Jordan, Nowlan and Hinton (1991): several expert networks plus a gating network that decides which experts handle each input. Shazeer et al. (2017) scaled it to deep learning with the sparsely-gated MoE layer, placing thousands of experts between LSTM layers. GShard (2020) and the Switch Transformer (Fedus, Zoph & Shazeer, 2021) brought MoE into transformers at scale.

The MoE layer#

In a transformer, the MoE layer typically replaces the feed-forward network (FFN) in some or all blocks. There are $E$ experts $f_1, \dots, f_E$ (each an FFN). For a token representation $\mathbf{x}$:

  1. The router computes logits $\mathbf{g} = \mathbf{W}_r\mathbf{x}$ and probabilities $\mathbf{p} = \text{softmax}(\mathbf{g})$.
  2. Select the top-$k$ experts (commonly $k = 1$ or 2).
  3. Output the weighted combination:
$$ \mathbf{y} = \sum_{i \in \text{TopK}(\mathbf{p})}\tilde{p}_i\,f_i(\mathbf{x}) $$

where $\tilde{p}_i$ are the selected probabilities (often renormalised). Only $k$ experts run for that token. For example, a model with 8 experts and top-2 routing has roughly 8× the FFN parameters of a dense model but only about 2× the FFN compute per token. Mixtral 8x7B, for instance, has about 47B total parameters but uses about 13B per token.

python
import torch
import torch.nn as nn
import torch.nn.functional as F

class MoE(nn.Module):
    def __init__(self, d=256, d_ff=1024, n_experts=8, k=2):
        super().__init__()
        self.k = k
        self.router = nn.Linear(d, n_experts, bias=False)
        self.experts = nn.ModuleList([nn.Sequential(nn.Linear(d, d_ff), nn.GELU(), nn.Linear(d_ff, d))
                                      for _ in range(n_experts)])
    def forward(self, x):                                   # x: (tokens, d)
        probs = F.softmax(self.router(x), dim=-1)           # (tokens, E)
        topv, topi = probs.topk(self.k, dim=-1)
        topv = topv / topv.sum(-1, keepdim=True)            # renormalise over chosen experts
        out = torch.zeros_like(x)
        for e, expert in enumerate(self.experts):           # dispatch tokens to each expert
            mask = (topi == e)
            rows = mask.any(-1).nonzero(as_tuple=True)[0]
            if rows.numel():
                w = (topv * mask)[rows].sum(-1, keepdim=True)
                out[rows] += w * expert(x[rows])
        # auxiliary load-balancing loss (Switch Transformer style)
        frac_tokens = F.one_hot(topi[:, 0], len(self.experts)).float().mean(0)
        aux = len(self.experts) * (frac_tokens * probs.mean(0)).sum()
        return out, aux

y, aux = MoE()(torch.randn(64, 256))
print(y.shape, round(aux.item(), 3))

The central problem: load balancing#

Routers tend to collapse: a few experts receive most tokens, become better, and attract even more — while others are starved and never learn. Solutions:

  • Auxiliary load-balancing loss (Switch Transformer): encourage the fraction of tokens routed to each expert, $f_i$, and the average router probability, $P_i$, to be uniform:
$$ \mathcal{L}_{\text{aux}} = \alpha\,E\sum_{i=1}^{E}f_i\,P_i $$
  • Noisy gating: add noise to router logits during training to encourage exploration.
  • Expert capacity: each expert processes at most a fixed number of tokens per batch (capacity factor × tokens / experts); overflow tokens are dropped (passed through the residual connection) or rerouted.
  • Auxiliary-loss-free balancing: adjust per-expert bias terms dynamically based on load (used in DeepSeek-V3).
  • Router z-loss: penalise large router logits for numerical stability.

Design variations#

  • Top-1 routing (Switch) is simplest and cheapest; top-2 is common for quality.
  • Fine-grained experts: many smaller experts with higher $k$ allow more flexible combinations (DeepSeekMoE).
  • Shared experts: one or more experts always active, capturing common knowledge, while routed experts specialise.
  • Expert-choice routing: experts choose their top tokens, guaranteeing balance.

What do experts specialise in? Analyses often find specialisation by token type or syntax (punctuation, numbers, particular languages or code) more than by high-level topic.

Training and serving challenges#

  • Communication: experts are spread across devices (expert parallelism); tokens must be sent to their experts and back (all-to-all communication), which can dominate time.
  • Memory: all experts' parameters must be stored even though few are used per token — MoE saves compute, not memory.
  • Instability: routing is discrete; training can be less stable than dense models; fine-tuning MoE models can overfit more easily.
  • Batching at inference: different tokens use different experts, complicating efficient kernels; small-batch latency benefits are smaller than FLOP counts suggest.

When MoE makes sense#

MoE shines when you can afford the memory for many parameters and want higher quality per unit of compute — typically large-scale training and high-throughput serving. For small deployments on limited hardware, a dense model of similar active size is often simpler.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

✨ Generative AI & LLMs

Quantising Large Language Models for Efficient Inference

Quantisation shrinks LLM weights to 8, 4 or fewer bits so models run on smaller GPUs, laptops and phones. We cover outlier features, weight-only vs activation quantisation, GPTQ, AWQ, formats like GGUF, and how to evaluate quality loss.

Advanced⏱ 5 min#211
✨ Generative AI & LLMs

LLM Inference: KV Caching, Batching and Speculative Decoding

Serving LLMs efficiently is a systems problem. We analyse prefill vs decode phases, the KV cache and PagedAttention, continuous batching, speculative decoding, and the latency and throughput metrics that matter.

Advanced⏱ 6 min#213
✨ Generative AI & LLMs

LoRA and Parameter-Efficient Fine-Tuning (PEFT)

Full fine-tuning of billion-parameter models is expensive. PEFT methods train a tiny fraction of parameters. We derive LoRA's low-rank updates, QLoRA's 4-bit training, compare adapters and prompt tuning, and give practical recipes.

Advanced⏱ 5 min#210