A working MoE layer in PyTorch
Build the whole thing yourself, with a shared expert, a balancing loss, and the sorting trick that turns a slow loop over experts into a handful of contiguous matrix multiplies.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 7
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
A mixture-of-experts layer is about sixty lines of code, and the interesting part is the grouping.
The post office sorting table
Letters arrive at a post office in a jumble. Bangalore, Delhi, Bangalore, Chennai, Delhi, and so on.
The sorter does not carry one letter at a time to one van. He sweeps the pile into pigeonholes, one per city. Then each full pigeonhole goes out in a single load.
Sorting first costs a moment. Everything after it is done in bulk, which is where the time is saved.
Why a model needs the same trick
Each word has been assigned to a couple of experts. The assignments are jumbled, exactly like the letters.
The straightforward way is to go expert by expert. Scan the whole list, and pick out the words belonging to that expert. Eight experts means eight scans of everything.
The better way is to sort the assignments once. All the words for expert one sit together, then all for expert two, and so on. Each expert then works on one solid block.
Computers are dramatically faster on one solid block than on scattered items. In the measurement below, the sorted version runs in a little over half the time. The gap grows with size.
What else goes into the layer
A router, which scores the experts for each word and keeps the best few.
A shared expert, switched on for every word, holding what everything needs.
A balancing term, added to the training loss so the traffic does not pile onto a few experts.
A weighted recombine, which multiplies each expert's answer by its score before adding it back.
That is the whole layer. Every large sparse model is this, made bigger and spread across machines.
Why building one is worth an afternoon
Reading about routing tells you what happens. Writing it tells you where the sharp edges are.
You find out that the sorting step is most of the code. You find out that forgetting the score multiplication stops the router learning. You find out that the balancing term has to be carried out of the layer to reach the loss.
None of those are visible from a diagram.
Remember this
- Router, experts, a shared expert, a balancing term, and a weighted recombine.
- Sort the assignments by expert first, so each expert gets one solid block of work.
- Building it once teaches you things reading about it will not.
What to learn next
- Mixture of depths — the same routing idea applied to whole layers.
- Serving a MoE across GPUs — what changes when the experts are not local.
- ModuleList vs Sequential — why the experts are held in a
ModuleList.
Developer — Code and libraries.
Setup
pip install torchWritten against PyTorch 2.5.1, Python 3.10. CPU, a few seconds.
The complete layer
import time, torch, torch.nn as nn, torch.nn.functional as F
class SwiGLU(nn.Module):
def __init__(self, d, ff):
super().__init__()
self.gate = nn.Linear(d, ff, bias=False)
self.up = nn.Linear(d, ff, bias=False)
self.down = nn.Linear(ff, d, bias=False)
def forward(self, x):
return self.down(F.silu(self.gate(x)) * self.up(x))
class MoELayer(nn.Module):
"""Routed experts + one always-on shared expert, with sorted dispatch."""
def __init__(self, d=64, ff=128, n_exp=8, top_k=2, n_shared=1, aux_alpha=0.01):
super().__init__()
self.k, self.n, self.alpha = top_k, n_exp, aux_alpha
self.router = nn.Linear(d, n_exp, bias=False)
self.experts = nn.ModuleList(SwiGLU(d, ff) for _ in range(n_exp))
self.shared = nn.ModuleList(SwiGLU(d, ff) for _ in range(n_shared))
self.aux_loss = torch.tensor(0.0)
def forward(self, x):
shape = x.shape
x = x.reshape(-1, shape[-1])
T = x.shape[0]
logits = self.router(x)
probs = F.softmax(logits.float(), dim=-1).to(x.dtype) # float32 router, as in practice
w, idx = probs.topk(self.k, dim=-1)
w = w / w.sum(-1, keepdim=True)
flat = idx.reshape(-1) # (T*k,)
order = flat.argsort() # group assignments by expert
counts = torch.bincount(flat, minlength=self.n)
token_of = order // self.k # which token each slot belongs to
weight_of = w.reshape(-1)[order]
xs = x[token_of] # tokens laid out expert by expert
ys = torch.empty_like(xs)
start = 0
for e in range(self.n): # contiguous slice, no masking
end = start + int(counts[e])
if end > start:
ys[start:end] = self.experts[e](xs[start:end])
start = end
ys = ys * weight_of[:, None]
out = torch.zeros_like(x).index_add_(0, token_of, ys)
for s in self.shared: # the shared expert sees every token
out = out + s(x)
f = counts.float() / (T * self.k) # Switch load-balancing loss
P = probs.mean(0)
self.aux_loss = self.alpha * self.n * (f * P).sum()
return out.reshape(shape), counts
def naive_forward(self, x):
"""Reference implementation. Same maths, easy to check by eye, slow."""
shape = x.shape; x = x.reshape(-1, shape[-1])
probs = F.softmax(self.router(x).float(), dim=-1).to(x.dtype)
w, idx = probs.topk(self.k, dim=-1); w = w / w.sum(-1, keepdim=True)
out = torch.zeros_like(x)
for s in range(self.k):
for e in range(self.n):
hit = idx[:, s] == e
if hit.any():
out[hit] += w[hit, s, None] * self.experts[e](x[hit])
for s in self.shared:
out = out + s(x)
return out.reshape(shape)
torch.manual_seed(0)
moe = MoELayer()
x = torch.randn(4, 32, 64) # batch 4, 32 tokens, dim 64
fast, counts = moe(x)
slow = moe.naive_forward(x)
print("sorted dispatch matches the reference:", torch.allclose(fast, slow, atol=1e-5))
print("max difference:", (fast - slow).abs().max().item())
print("tokens per expert:", counts.tolist(), " (128 tokens x top-2 = 256 assignments)")
print("aux loss:", round(float(moe.aux_loss), 6))
per_expert = sum(p.numel() for p in moe.experts[0].parameters())
total_p = sum(p.numel() for p in moe.parameters())
active_p = (moe.k + len(moe.shared)) * per_expert + moe.router.weight.numel()
print(f"parameters: {total_p:,} total | {active_p:,} active per token "
f"({active_p/total_p:.1%})")
big = torch.randn(2048, 64)
for name, fn in (("naive loop", moe.naive_forward), ("sorted dispatch", lambda z: moe(z)[0])):
with torch.no_grad():
fn(big); t0 = time.perf_counter()
for _ in range(20): fn(big)
print(f" {name:<18}{(time.perf_counter()-t0)/20*1000:7.2f} ms per call (CPU)")
print("\ntraining it for 300 steps on a toy task")
torch.manual_seed(1)
model = MoELayer(d=32, ff=64, n_exp=8, top_k=2)
opt = torch.optim.Adam(model.parameters(), lr=3e-3)
gen = torch.Generator().manual_seed(2)
W = torch.randn(32, 32)
for step in range(301):
xb = torch.randn(256, 32, generator=gen)
yb = torch.tanh(xb @ W)
out, counts = model(xb)
loss = F.mse_loss(out, yb) + model.aux_loss
loss.backward(); opt.step(); opt.zero_grad()
if step % 100 == 0:
share = counts.float() / counts.sum()
print(f" step {step:>3} loss {loss.item():.4f} "
f"busiest expert {share.max():.1%} quietest {share.min():.1%}")sorted dispatch matches the reference: True max difference: 1.1920928955078125e-07 tokens per expert: [30, 33, 34, 28, 31, 38, 36, 26] (128 tokens x top-2 = 256 assignments) aux loss: 0.010024 parameters: 221,696 total | 74,240 active per token (33.5%) naive loop 8.59 ms per call (CPU) sorted dispatch 5.35 ms per call (CPU) training it for 300 steps on a toy task step 0 loss 0.8860 busiest expert 15.8% quietest 8.4% step 100 loss 0.3929 busiest expert 14.8% quietest 10.7% step 200 loss 0.3511 busiest expert 16.2% quietest 10.4% step 300 loss 0.3096 busiest expert 17.0% quietest 8.4%
Reading the output
The two implementations agree to 1.2e-07. Always keep a slow reference implementation next to a fast one and assert equality in a test. Dispatch bugs produce plausible-looking output and broken training, and this assertion is how you catch them in a minute rather than a week.
The two timing lines are the only numbers here that will not reproduce. Everything else in this output is deterministic. Three runs of this script on the same machine gave 9.11 / 5.08 ms, 8.59 / 5.35 ms and 30.08 / 14.73 ms — the third while the machine was busy with other work. The ratio stayed between 1.6x and 2.0x across all three.
Read the ratio, never the milliseconds. The gap widens with expert count, because the naive version scans the whole batch once per expert while the sorted version pays one sort regardless.
Balance holds through training. Busiest 17.0%, quietest 8.4%, against 12.5% for perfectly even. The auxiliary loss is doing its job without perfect uniformity, which is the normal healthy state.
Loss falls from 0.886 to 0.310. Note the aux term contributes about 0.010 at balance, so the task loss is around 0.30.
Line by line, for the parts that matter
F.softmax(logits.float(), ...) computes the router in float32 even if the model is bf16. Near-ties in the top-k flip under bf16 rounding, and routing instability is a real training failure. transformers does the same.
order = flat.argsort() is the sorting table. After it, all assignments for expert 0 come first, then expert 1, and so on.
token_of = order // self.k recovers which token each sorted slot came from, because slot i of the flattened (T, k) index tensor belongs to token i // k.
ys[start:end] = self.experts[e](xs[start:end]) is one contiguous matrix multiply per expert. No boolean masks, no gather inside the loop. In production this whole loop becomes a single grouped GEMM call taking the counts vector.
index_add_(0, token_of, ys) scatters results back and sums the contributions where one token used several experts. It handles duplicate indices correctly, which plain indexed assignment does not.
self.aux_loss is stored on the module because the loss is computed outside. Every MoE layer contributes one, and they are summed with the task loss. Forgetting to add them is a common and silent bug: training proceeds, balance degrades.
Wiring it into a transformer block
class Block(nn.Module):
def __init__(self, d, n_heads, **moe_kwargs):
super().__init__()
self.n1, self.n2 = nn.RMSNorm(d), nn.RMSNorm(d)
self.attn = nn.MultiheadAttention(d, n_heads, batch_first=True)
self.moe = MoELayer(d=d, **moe_kwargs)
def forward(self, x):
h = self.n1(x)
x = x + self.attn(h, h, h, need_weights=False)[0]
y, _ = self.moe(self.n2(x))
return x + yAttention is untouched. Only the feed-forward sublayer changes, which is the entire architectural delta between a dense transformer and a sparse one. nn.RMSNorm requires PyTorch 2.4 or newer.
Collect the auxiliary losses at the top level:
aux = sum(b.moe.aux_loss for b in model.blocks)
loss = task_loss + auxCommon mistakes
Using out[token_of] = ys instead of index_add_. With top-k above 1, several slots share a token, and indexed assignment keeps only the last one. Silent, and it halves your model.
Computing counts from a tensor and slicing with it without int(). Slicing with a 0-d tensor works but forces a synchronisation on GPU. Real implementations compute the offsets once on host.
Forgetting .reshape(shape) at the end. The layer flattens (B, T, d) to (B*T, d) and must restore it, or the next layer receives the wrong rank.
Dropping the aux loss. Nothing errors. Balance quietly degrades over thousands of steps.
Benchmarking this against a fused kernel. The loop over experts is Python. At 256 experts the loop overhead alone dominates. Use grouped GEMM, MegaBlocks or your framework's fused MoE path for anything real.
Try it yourself
Set n_shared=0 and retrain. Watch the task loss at step 300 and the balance figures. Then set aux_alpha=0 and run for 1000 steps, logging counts every 100 steps, to watch collapse begin.
What to learn next
- Mixture of depths — the same routing idea applied to whole layers.
- Serving a MoE across GPUs — what changes when the experts are not local.
- ModuleList vs Sequential — why the experts are held in a
ModuleList.
Researcher — Mathematics and papers.
Dispatch as a permutation
The layer computes
$$ y_t = \sum_{s \in S_{\text{shared}}} E_s(x_t) + \sum_{i \in \operatorname{Top-}k(x_t)} g_i(x_t)\, E_i(x_t) $$
The implementation problem is that $\operatorname{Top-}k(x_t)$ varies per token, so the work is ragged. Three formulations, in increasing order of what production uses:
Masked. For each expert, build a boolean mask over all $T$ tokens and gather. Cost $\Theta(NT)$ in mask work, independent of how many tokens each expert receives. This is naive_forward.
Sorted / permuted. Sort the $Tk$ assignments by expert, giving contiguous per-expert segments. Cost $\Theta(Tk \log(Tk))$ for the sort, then $N$ dense GEMMs of ragged sizes. This is forward.
Block-sparse. Express the whole layer as one block-sparse matmul over a block-diagonal expert weight matrix, so no per-expert kernel launch occurs at all. This is MegaBlocks (Gale, Narayanan, Young and Zaharia, 2022, arXiv:2211.15841), reported to never drop tokens with up to 40% end-to-end speed-up over Tutel.
The modern middle ground is grouped GEMM: one kernel invocation taking a vector of per-group row counts, avoiding both the $N$ launches of the sorted approach and the block-sparse machinery. Available in CUTLASS, Triton and vendor libraries, and it is what most current MoE stacks call.
Numerical and autograd notes
The permutation is differentiable as a gather, and index_add_ is its transpose, so autograd handles the round trip without custom code. Gradient correctness of a hand-written dispatch is worth testing directly with torch.autograd.gradcheck on a float64 copy — see verifying gradients with gradcheck.
Router precision is not optional. In bf16, adjacent router logits differing by less than $2^{-8}$ relative are indistinguishable, so top-$k$ selection becomes a function of rounding. Computing the router in float32 costs nothing (it is a $d \times N$ matmul) and removes a whole class of irreproducibility.
The auxiliary loss should be computed from the pre-selection probabilities, not the renormalised top-$k$ weights. Using the renormalised weights makes $P_i$ conditional on selection and destroys the gradient signal that pushes probability toward unused experts.
What this implementation omits
- Capacity and dropping. This layer is dropless by construction, since segment sizes are whatever the router produced. Add a cap if you are matching a fixed-buffer implementation — see expert capacity.
- Expert parallelism. All experts are local. Distributing them replaces the loop with all-to-all dispatch and combine, which is the subject of serving a MoE across GPUs.
- The z-loss. ST-MoE's router logit penalty, worth adding for any run past toy scale.
- Bias-based balancing. DeepSeek-V3's auxiliary-loss-free scheme, which replaces
aux_losswith a non-gradient control loop. - Fine-grained scaling. At $N = 256$ the Python loop is the bottleneck, not the maths.
Each is a small addition to this file, and together they are the distance between a teaching implementation and a training-grade one. The value of writing the sixty lines is that the additions then read as modifications to something you understand rather than as configuration flags.
What to learn next
- Mixture of depths — the same routing idea applied to whole layers.
- Serving a MoE across GPUs — what changes when the experts are not local.
- ModuleList vs Sequential — why the experts are held in a
ModuleList.