Mixture of depths
Mixture of depths routes tokens past whole layers instead of between experts, so easy words get less computation than hard ones while the total cost per batch stays fixed and predictable.
- 12 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Mixture of depths lets most words drive straight past a layer, and sends only a chosen few in for processing.
The toll plaza
You have driven through a toll plaza with a tag on the windscreen. Most vehicles roll through without stopping.
A few are waved aside into the inspection bays. There is a fixed number of bays, so a fixed number of vehicles gets inspected, however busy the road is.
The plaza never jams unpredictably, because the capacity is decided in advance. That fixed capacity is the clever part of this design.
Why not give every word the same effort
The word "the" and the final digit of a long calculation are not equally hard. A model spends exactly the same work on both.
That is wasteful, and every attempt to fix it has run into the same wall. If different words take different amounts of work, the work per batch becomes unpredictable. Unpredictable work is slow work on a graphics card.
The design that avoids the wall
Decide in advance that only, say, one word in eight goes through this layer. The rest slide past on the shortcut that every layer already has.
The choice of which words is made fresh for each batch. The number of words is fixed forever.
ordinary layer mixture of depths
-------------- -----------------
t0 -> [ layer ] -> t0 --------------->
t1 -> [ layer ] -> t1 -> [ layer ] ->
t2 -> [ layer ] -> t2 --------------->
t3 -> [ layer ] -> t3 --------------->
4 tokens processed 1 token processed
cost: fixed cost: fixed, and smallerHow it relates to experts
This is the same routing idea as a mixture of experts, with one change.
In a mixture of experts, the router chooses between several experts. Here it chooses between one expert and doing nothing at all.
Doing nothing costs nothing, which is where the saving comes from.
The problem that has held it back
To pick the busiest one word in eight, you have to look at all eight and compare them.
When a model writes text it produces one word at a time. It cannot compare a word against words it has not written yet.
The published fix is a second, tiny predictor. It guesses whether this word would have been picked, using only the past. It works, and it is another moving part that has to be trained and can be wrong.
Where this stands
The published results are encouraging. Equal quality for the same training budget, far fewer operations per pass, and meaningfully faster generation.
It has not become standard. The major open models covered in this section route between experts, not past layers. Research on the idea has continued steadily, so treat it as promising rather than settled.
Remember this
- Only a fixed fraction of words go through each layer. The rest take the shortcut.
- Fixed capacity keeps the cost predictable, which is why this design works where earlier ones did not.
- Choosing the busiest words needs to see the future, so generation needs an extra predictor.
What to learn next
- Early exit and layer skipping — stopping the stack instead of skipping parts of it.
- The router — the gradient path this design inherits.
- Knowledge distillation — the other way to make a model do less work per token.
Developer — Code and libraries.
Setup
pip install torchWritten against PyTorch 2.5.1, Python 3.10. CPU, instant.
A mixture-of-depths block, and its central problem
import torch, torch.nn as nn
torch.manual_seed(0)
T, d, CAP = 16, 8, 0.25 # 16 tokens, dim 8, only a quarter go through
class MoDBlock(nn.Module):
"""A transformer block that only processes the top-scoring fraction of tokens."""
def __init__(self, d, capacity):
super().__init__()
self.router = nn.Linear(d, 1, bias=False)
self.block = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
self.capacity = capacity
def forward(self, x):
n = x.shape[0]
k = max(1, int(n * self.capacity))
score = self.router(x).squeeze(-1) # one number per token
_, keep = score.topk(k) # the busiest k tokens in the sequence
keep, _ = keep.sort()
out = x.clone() # everyone else takes the residual path
w = torch.sigmoid(score[keep])[:, None] # gate weight, so the router gets gradient
out[keep] = x[keep] + w * self.block(x[keep])
return out, keep
mod = MoDBlock(d, CAP)
x = torch.randn(T, d)
out, keep = mod(x)
print(f"{T} tokens, capacity {CAP:.0%} -> {len(keep)} processed, {T-len(keep)} skipped")
print("processed token positions:", keep.tolist())
print("rows that came out unchanged:",
[t for t in range(T) if torch.allclose(out[t], x[t])])
flops_full = T * (2 * d * 4 * d * 2)
flops_mod = len(keep) * (2 * d * 4 * d * 2) + T * d
print(f"\nblock FLOPs: full {flops_full:,} mixture-of-depths {flops_mod:,} "
f"({flops_mod/flops_full:.1%})")
# The catch: top-k over the sequence lets a LATER token change an EARLIER token's fate.
print("\nnow change only the LAST token, and watch an earlier token lose its slot")
x2 = x.clone()
x2[15] = mod.router.weight[0] * 10.0 # make the final token score very high
_, keep2 = mod(x2)
print(" original selection:", keep.tolist())
print(" after editing only token 15:", keep2.tolist())
dropped = sorted(set(keep.tolist()) - set(keep2.tolist()))
print(f" token(s) {dropped} were processed before and are skipped now,")
print(" purely because of a token that comes AFTER them.")
print("\nA decoder does not have its future tokens yet, so this rule cannot be")
print("applied at generation time without a separate causal predictor.")16 tokens, capacity 25% -> 4 processed, 12 skipped processed token positions: [3, 5, 11, 14] rows that came out unchanged: [0, 1, 2, 4, 6, 7, 8, 9, 10, 12, 13, 15] block FLOPs: full 16,384 mixture-of-depths 4,224 (25.8%) now change only the LAST token, and watch an earlier token lose its slot original selection: [3, 5, 11, 14] after editing only token 15: [3, 11, 14, 15] token(s) [5] were processed before and are skipped now, purely because of a token that comes AFTER them. A decoder does not have its future tokens yet, so this rule cannot be applied at generation time without a separate causal predictor.
Reading the output
Twelve of sixteen rows came out bit-identical to the input. The residual path is not an approximation of the layer; for a skipped token the layer is the identity function.
FLOPs drop to 25.8% of the block. Not 25% exactly: the router itself costs T * d, and it runs on every token including the skipped ones. That overhead is small and it is not zero.
Token 5 lost its slot because of token 15. This is the demonstration that matters. Nothing about token 5 changed. A token nine positions later scored higher, took the last bay, and evicted it. Top-$k$ over a sequence is a comparison across positions, and comparison across positions is not causal.
That single fact is why mixture of depths is harder to deploy than mixture of experts. An MoE router looks at one token in isolation; an MoD router ranks tokens against each other.
The two ways round it
Train an auxiliary causal predictor. A small head learns to predict "would the top-$k$ have selected me", from the token's own representation alone. At generation time, use its output with a threshold instead of the ranking. This is the approach in the original paper.
Use a threshold rather than a rank. Process any token whose score exceeds a fixed value. Perfectly causal, and it gives up the fixed capacity that made the design attractive — batch cost is now data-dependent again.
The first keeps the static graph and adds a component that can be wrong. The second is simple and reintroduces the problem the design was solving. Neither is clean.
The relationship to mixture of experts
Write the MoD block as an MoE layer with two experts:
expert 0 = the transformer block
expert 1 = the identity function (free)with top-1 routing and a capacity constraint on expert 0. Everything from the earlier lessons transfers: capacity factors, load balancing, the gate-weight gradient path. The gate weight torch.sigmoid(score[keep]) in the code above exists for exactly the reason established in the router — without it, the router receives no gradient.
The two can be combined. Route past the block, and route among experts when you do not. The original paper calls this MoDE.
Common mistakes
Applying it to every layer. The paper interleaves MoD blocks with ordinary blocks. Skipping at every layer starves tokens that were never selected anywhere.
Forgetting the router runs on skipped tokens. You still pay T * d per layer. At high skip rates this becomes a visible fraction of what is left.
Assuming the FLOP saving is a wall-clock saving. Gathering and scattering the selected tokens costs memory traffic. The saving is real and smaller than the FLOP count suggests, especially at small capacity where the gather dominates.
Training with top-k and generating with a threshold. A train-test mismatch that will show up as a quality drop nobody can locate.
Try it yourself
Set CAP = 0.125, the paper's figure, and recompute the FLOP fraction. Then add a second MoDBlock and apply it only on alternate layers, checking that no token is skipped by every layer in the stack.
What to learn next
- Early exit and layer skipping — stopping the stack instead of skipping parts of it.
- The router — the gradient path this design inherits.
- Knowledge distillation — the other way to make a model do less work per token.
Researcher — Mathematics and papers.
The method
Raposo, Ritter, Richards, Lillicrap, Humphreys and Santoro (2024), Mixture-of-Depths: Dynamically allocating compute in transformer-based language models (arXiv:2404.02258).
Each MoD block holds a scalar router $r_\theta: \mathbb{R}^{d} \to \mathbb{R}$. For a sequence of $T$ tokens with capacity $c$, let $\mathcal{S} = \operatorname{Top-}k({r_\theta(x_t)}, k = \lfloor cT \rfloor)$. Then
$$ x'_t = \begin{cases} x_t + g_t \cdot f_\theta(x_t) & t \in \mathcal{S} \ x_t & \text{otherwise} \end{cases} $$
$f_\theta$ is the full block (attention plus MLP) and $g_t$ a gate derived from the router score, present so gradient reaches $r_\theta$.
The key property, in the authors' terms, is that this uses "a static computation graph with known tensor sizes" while remaining "dynamic and context-sensitive". The set of processed tokens varies; the count does not. Every previous adaptive-computation scheme surrendered one or the other.
Their reported result: models "match baseline performance for equivalent FLOPS and wall-clock times to train, but require a fraction of the FLOPs per forward pass, and can be upwards of 50% faster to step during post-training sampling." Typical settings are $c = 0.125$ with MoD blocks on alternating layers.
Why the capacity is expert-choice
The selection is over tokens for a fixed number of slots, which is exactly expert choice routing (Zhou et al., 2022, arXiv:2202.09368) applied to a two-expert layer where the second expert is the identity.
That inheritance brings the good property — perfect load balance without an auxiliary loss, since capacity is filled exactly — and the bad one: expert-choice selection is not causal. The demonstration above is a minimal instance.
The paper's remedy is an auxiliary predictor trained to reproduce the top-$k$ decision from a single token's representation, used autoregressively at sampling. It restores causality at the cost of an approximation error whose downstream effect is not straightforward to bound.
Position in the adaptive-compute literature
| Approach | What varies | Causal at decode |
|---|---|---|
| Adaptive Computation Time (Graves, 2016) | number of recurrent steps | yes |
| Early exit / CALM | how many layers before stopping | yes |
| Mixture of Depths | which tokens enter each layer | needs a predictor |
| Mixture of Recursions | how many times a shared block is reapplied | yes |
MoD skips over layers while preserving the rest of the stack, whereas early exit stops the stack entirely. That distinction matters: a token that exits early receives no processing in any later layer, while a token skipped by an MoD block can still be selected by the next one. Later comparisons argue MoD preserves quality better than early exit for this reason.
Follow-up work
- MoDification: Mixture of Depths Made Easy (NAACL 2025) converts existing dense checkpoints to MoD, reducing the cost of adoption.
- Mixture-of-Recursions (arXiv:2507.10524) learns per-token recursive depth over a shared block, giving parameter sharing and adaptive depth together.
- GateSkip (arXiv:2510.13876) gates attention and MLP branches independently with a small linear gate and sigmoid on the residual stream.
- LayerRoute (arXiv:2606.01838) adapts a pretrained model post hoc with sequence-level rather than token-level routing, which sidesteps the causality problem entirely at coarser granularity.
The movement from token-level to sequence-level routing in that last entry is worth noting. Sequence-level decisions are causal, cheap, and compatible with continuous batching, at the cost of the fine-grained allocation that motivated the idea.
Why it has not become standard
Read the configs of the large open sparse models in this section — Mixtral, DeepSeek-V3, Qwen3, gpt-oss, Qwen3-Next — and all of them route among feed-forward experts. None routes past layers.
Plausible reasons, none conclusive: the causality workaround adds a trained component that must be validated separately; the wall-clock gain is smaller than the FLOP gain because of gather and scatter overhead; and the same engineering effort spent on expert sparsity yields a larger, better-understood return. Adaptive depth remains one of the clearest unexploited efficiency axes in the architecture, and the steady stream of follow-up papers suggests the question is open rather than closed.
What to learn next
- Early exit and layer skipping — stopping the stack instead of skipping parts of it.
- The router — the gradient path this design inherits.
- Knowledge distillation — the other way to make a model do less work per token.