Skip to content
AI-grafen
FAI engineeringTransformer architecture· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

Mixture of Experts

Be able to explain routing, expert capacity and load balancing in MoE models.

Prerequisites

Intuition

An ordinary transformer runs the whole MLP block for every token. A mixture of experts replaces the MLP with many parallel experts and lets a small router pick out a handful per token.

          ┌──────────┐
token ──→ │  router  │──→ picks expert 3 and 17 of 64
          └──────────┘
                │
     ┌──────────┴─────────┐
  expert 3            expert 17
     └──────────┬─────────┘
          a weighted sum

The gain: the model can have far more parameters without every token costing more to compute.

A dense modelMoE
Total parameters8 B47 B
Active per token8 B13 B
Computation per token1×~1.6×
Memory16 GB94 GB

The catch shows in the last row: every expert has to be in memory even if only two are used. MoE trades memory for capacity — a completely different bargain from quantisation.

Formal

The router is a linear layer followed by a softmax over the experts. The top kk (usually k=2k = 2) are chosen, and the output is the weighted sum:

y=∑i∈top-kgi⋅Ei(x),g=softmax(Wrx)y = \sum_{i \in \text{top-}k} g_i \cdot E_i(x), \qquad g = \mathrm{softmax}(W_r x)

Load balancing is the central problem. Without a counterweight a self-reinforcing spiral arises: an expert that happens to be chosen more often is trained more, gets better, and is chosen even more often. In the end a handful of experts are used and the rest are dead weights.

The solution is an auxiliary loss that punishes an uneven distribution:

Laux=α⋅N∑i=1Nfi⋅Pi\mathcal{L}_{\text{aux}} = \alpha \cdot N \sum_{i=1}^{N} f_i \cdot P_i

where fif_i is the share of tokens routed to expert ii and PiP_i is the average router probability for it. The term is minimised when both are evenly distributed. Typically α≈0.01\alpha \approx 0.01.

Expert capacity. Every expert has a limit on how many tokens it accepts per batch:

C=tokens per batchthe number of experts×the capacity factorC = \frac{\text{tokens per batch}}{\text{the number of experts}} \times \text{the capacity factor}

With a capacity factor of 1.25 every expert tolerates 25 % above its proportional share. Tokens above the limit are dropped — they pass through unchanged via the residual connection. A low factor saves memory but drops more; a high one wastes.

Three practical problems:

ProblemWhy
Communicationthe experts sit on different GPUs → an all-to-all exchange per layer
Unstable trainingthe router can oscillate; a z-loss and a lower learning rate on the router help
Fine-tuningMoE models overfit more easily; fewer epochs, more regularisation

Variants: shared experts that are always active (DeepSeek), expert choice instead of token choice (each expert picks its tokens, which gives perfect balance by construction), and very fine-grained experts (hundreds of small ones instead of eight large).

When MoE pays off: when you have plenty of memory but want more parameters per FLOP. When memory is limited, a dense model is nearly always the right choice.

Code

import torch, torch.nn as nn, torch.nn.functional as F

class MoE(nn.Module):
    def __init__(self, d=512, hidden=2048, n_experts=8, k=2, capacity_factor=1.25):
        super().__init__()
        self.k, self.n = k, n_experts
        self.capfactor = capacity_factor
        self.router = nn.Linear(d, n_experts, bias=False)
        self.experts = nn.ModuleList([
            nn.Sequential(nn.Linear(d, hidden), nn.GELU(), nn.Linear(hidden, d))
            for _ in range(n_experts)])

    def forward(self, x):
        B, T, d = x.shape
        xf = x.reshape(-1, d)                          # (N, d)
        N = xf.size(0)
        logits = self.router(xf)
        probability = F.softmax(logits, dim=-1)        # (N, n)
        weight, choice = probability.topk(self.k, dim=-1)   # (N, k)
        weight = weight / weight.sum(dim=-1, keepdim=True)  # normalise the chosen ones

        capacity = int(self.capfactor * N * self.k / self.n)
        out = torch.zeros_like(xf)
        dropped = 0
        for e in range(self.n):
            mask = (choice == e)
            index = mask.any(dim=-1).nonzero(as_tuple=True)[0]
            if len(index) > capacity:                  # over capacity → drop
                dropped += len(index) - capacity
                index = index[:capacity]
            if len(index) == 0:
                continue
            w = (weight * mask.float()).sum(dim=-1)[index].unsqueeze(-1)
            out[index] += w * self.experts[e](xf[index])

        # The load-balancing loss: an even distribution of both the tokens and the probability
        f = torch.zeros(self.n, device=x.device)
        for e in range(self.n):
            f[e] = (choice == e).any(dim=-1).float().mean()
        P = probability.mean(dim=0)
        aux = self.n * (f * P).sum()

        self.latest = {"aux_loss": float(aux),
                       "dropped_share": round(dropped / max(N * self.k, 1), 4),
                       "distribution": [round(float(v), 3) for v in f]}
        return out.view(B, T, d), aux

# Counting parameters: total against active
def parameters(d=4096, hidden=14336, n_experts=8, k=2, layers=32):
    per_expert = 3 * d * hidden
    total = layers * n_experts * per_expert
    active = layers * k * per_expert
    return {"total_bn": round(total / 1e9, 1), "active_bn": round(active / 1e9, 1),
            "ratio": round(total / active, 1)}

print(parameters())
# {'total_bn': 45.1, 'active_bn': 11.3, 'ratio': 4.0}
#  ↑ four times more parameters than are computed per token

# Collapsed routing looks like this
collapsed = [0.71, 0.24, 0.03, 0.01, 0.01, 0.0, 0.0, 0.0]
balanced = [0.13, 0.12, 0.13, 0.12, 0.13, 0.12, 0.12, 0.13]
for name, f in (("collapsed", collapsed), ("balanced", balanced)):
    import math
    entropy = -sum(p * math.log(p) for p in f if p > 0)
    print(f"{name:<10} entropy {entropy:.3f} of a maximum of {math.log(8):.3f}")
# collapsed  entropy 0.783 of a maximum of 2.079
# balanced   entropy 2.079 of a maximum of 2.079

Logging the routing entropy is the simplest collapse detector: if it falls below about half the maximum, only a fraction of the experts are being used, and the weight of the auxiliary loss needs raising.

Mastery means

  • Explains routing and the top-k choice
  • Describes the load-balancing problem
  • Knows what MoE saves and what it costs

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences