A standard transformer spends the same compute on every token. The word 'the' and the digit that decides the answer to an arithmetic problem each pass through every attention layer and every MLP. Mixture of Depths (MoD), introduced by Raposo and colleagues at Google DeepMind in 2024 (arXiv 2404.02258), lets the model learn which tokens deserve a given block's compute. At each routed block a small router scores every token, the top k tokens go through the block, and the rest skip it along the residual stream.
The design choice that makes MoD interesting is that k is fixed in advance. The amount of compute per block is known before the data arrives, so tensor shapes stay static and the hardware stays efficient, while the choice of which tokens get the compute is dynamic. This article explains the block, works through the FLOP arithmetic, gives an illustrative implementation, and spends most of its time on the parts that decide whether MoD is usable: the top-k operation is not causal, and skipping tokens changes the KV cache and batching at inference. The paper does not publish code, so the code here is a teaching sketch, not a reference implementation.
The block, from first principles
Take a block l with input x for each of n tokens in a sequence. A router, here a single linear projection, gives each token a scalar score r = w · x. The block chooses a capacity, the number of tokens k it will process, and lets through the tokens whose scores are in the top k for that sequence. For a selected token the output is x + r · f(x), where f is the usual self-attention plus MLP computation applied only among the selected tokens. For any other token the output is simply x: the residual stream carries it past the block untouched.
Multiplying by r is not decoration. Top-k selection is a discrete choice and has no gradient. The paper multiplies the block output by the router weight, which, in its words, puts the router weights along the gradient path. If routing a token helps the loss, the gradient increases the score of that token's representation, and the router learns which tokens benefit from this block. Multiply only f(x), never the residual: scaling x itself would let the router attenuate the stream for selected tokens and would make skipped and selected tokens follow different scales.
The paper's best configurations used a capacity of 12.5 percent, so 87.5 percent of tokens route around each routed block, and routed every other block, keeping full-capacity blocks in between. The authors report that routing every other block was crucial for strong performance. Intuitively, the full blocks give every token regular updates and keep the representation of skipped tokens fresh, while the routed blocks concentrate extra compute where the router thinks it helps.
The FLOP arithmetic
Why does a low capacity save so much? A block has two kinds of cost. The linear layers (Q, K, V and output projections plus a 4x MLP) cost roughly 24 · n · d² FLOPs for n tokens of width d, linear in n. Attention scores and the weighted sum cost roughly 4 · n² · d, quadratic in n. Take d = 2,048 and a sequence of n = 4,096 tokens:
| Block | Linear FLOPs | Attention FLOPs | Total |
|---|---|---|---|
| Full (n = 4,096) | 4.12 × 10¹¹ | 1.37 × 10¹¹ | 5.50 × 10¹¹ |
| Routed (k = 512, 12.5%) | 5.15 × 10¹⁰ | 2.1 × 10⁹ | 5.37 × 10¹⁰ (incl. router) |
The routed block costs 9.8 percent of a full one: the linear part shrinks eightfold and the attention part sixty-four-fold, because attention is quadratic in the number of participating tokens. The router itself, one d-dimensional dot product per token, is negligible. With routing on every other block, a pair of blocks costs 54.9 percent of a dense pair, so a forward pass needs about 45 percent fewer FLOPs.
The paper uses those savings in two ways. In an isoFLOP analysis at training budgets of 6 × 10¹⁸, 2 × 10¹⁹ and 1 × 10²⁰ FLOPs, with models from 60M to 3B parameters, MoD models trained with the same total compute can be larger or see more data than the baseline, and the authors report they match or beat the isoFLOP-optimal baseline. Alternatively, a MoD model matching the baseline's quality needs fewer FLOPs per forward pass, and the authors report models that are upwards of 50 percent faster to step during post-training sampling. These are research-scale results; treat them as evidence the idea works, not as a guarantee for your model size.
An implementation sketch
The core of a routed block is gather, compute, scatter. The sketch wraps an existing pre-norm transformer block; it is written for clarity and assumes the wrapped block applies causal attention among the tokens it is given.
import torch
import torch.nn as nn
class MoDBlock(nn.Module):
# Illustrative Mixture-of-Depths wrapper around a standard pre-norm block.
# block(x, position_ids) must apply causal self-attention among the tokens it receives.
def __init__(self, block, d_model, capacity=0.125):
super().__init__()
self.block, self.capacity = block, capacity
self.router = nn.Linear(d_model, 1, bias=False)
def forward(self, x, position_ids): # x: [B, n, d]
B, n, d = x.shape
k = max(1, int(n * self.capacity))
logits = self.router(x).squeeze(-1) # [B, n]
top = torch.topk(logits, k, dim=-1).indices # [B, k], NOT in time order
idx, _ = torch.sort(top, dim=-1) # restore causal order
gather = idx.unsqueeze(-1).expand(-1, -1, d)
xs = torch.gather(x, 1, gather) # [B, k, d]
pos = torch.gather(position_ids, 1, idx) # ORIGINAL positions for RoPE
f = self.block(xs, pos) - xs # the block's update, without residual
w = torch.gather(logits, 1, idx).unsqueeze(-1) # router weight on the gradient path
return x.scatter_add(1, gather, w * f), logits, idxThree lines carry most of the risk. First, torch.topk returns indices ordered by score, not by time. Gathering in that order and then applying a causal mask lets later tokens leak into earlier ones, a bug that does not crash and quietly improves training loss. Sort the indices first. Second, positional information must use the original positions: with rotary embeddings, pass the gathered position_ids so that two selected tokens 3,000 positions apart are not treated as neighbours. Third, the update added to the residual is w * f, where f excludes the block's own residual connection, so the stream is not counted twice.
Because k is fixed, the gathered tensor has a static shape of [B, k, d]. That is the property that separates MoD from most adaptive-computation methods: the compiler and kernels see the same shapes every step, and there is no load-balancing problem of the kind mixture-of-experts needs auxiliary losses for, because each block simply takes exactly k tokens.
The causality problem
Top-k is a function of the whole sequence. Whether token 10 is in the top 12.5 percent depends on the scores of tokens 11 to 4,096. During training that is harmless, because the whole sequence is available and the routed block's attention is still causal among the selected tokens. During autoregressive sampling it is fatal: when generating token 10, tokens 11 onward do not exist, so the model cannot know whether token 10 would have made the cut.
The paper proposes two fixes. The first adds an auxiliary binary cross-entropy loss in which the router's outputs are the logits and the top-k selections are the targets. This pushes router scores for selected tokens above zero and for the rest below zero, so at sampling time the sign of a token's own score decides routing. The paper reports this costs about 0.2 to 0.3 percent in the main language-modelling objective. The second trains a small auxiliary MLP predictor, on a stop-gradient copy of the router's input, to predict whether a token will be in the top k; it does not affect the main model's training, and the paper reports no significant impact on step speed. The paper reports the sampling-time routing decision reaching upwards of 97 percent accuracy soon into training, and up to about 99 percent.
import torch.nn.functional as F
def router_aux_loss(logits, idx):
# Paper's first option: BCE with router logits as logits, top-k membership as targets,
# so that the sign of a logit alone predicts selection at sampling time.
target = torch.zeros_like(logits).scatter_(1, idx, 1.0)
return F.binary_cross_entropy_with_logits(logits, target)
class RoutePredictor(nn.Module):
# Paper's second option: a small MLP on a stop-gradient copy of the input
# predicts whether the token would have been in the top-k.
def __init__(self, d_model, hidden=256):
super().__init__()
self.net = nn.Sequential(nn.Linear(d_model, hidden), nn.SiLU(), nn.Linear(hidden, 1))
def loss(self, x, idx):
target = torch.zeros(x.shape[:2], device=x.device).scatter_(1, idx, 1.0)
return F.binary_cross_entropy_with_logits(self.net(x.detach()).squeeze(-1), target)
# At sampling time, per new token t and routed layer l:
# route = predictor(x_t) > 0 (or router_logit(x_t) > 0 with the aux-loss variant)
# if route: x_t = x_t + r_t * f(x_t), and K/V for t are written to layer l's cache
# else: x_t passes through unchanged, and layer l stores nothing for tThe consequence: at sampling time routing is a per-token threshold, so the routed fraction only approximates the training capacity. Static shapes are a training property.
What it does to serving
A transformer's KV cache stores keys and values for every past token at every layer. In a MoD model, a token that skipped routed layer l produced no keys or values at l, and future tokens at l attend only to the earlier tokens that were routed there. That has three consequences an inference engine must handle.
- Per-layer caches have different lengths. At 12.5 percent routing, a routed layer's cache holds roughly an eighth of the tokens of a full layer, which is a real memory saving, but block tables and paged allocation must be tracked per layer rather than per sequence. Engines built around one block table per request, as described in vLLM continuous batching, need changes.
- Batched decode becomes ragged. At each step and routed layer, a different subset of sequences in the batch routes. The engine must gather that subset, run the block, and scatter back, or run the whole batch and mask, which spends the compute MoD was meant to save. Savings in wall-clock time depend on doing the gather efficiently at small batch sizes, where decoding is memory-bound anyway.
- Prefill and decode must agree. During prefill the whole prompt is available, so top-k over the prompt is possible. But tokens generated later are routed by threshold. Using one rule for the prompt and another for generated tokens is a source of train-serve skew; pick a policy, apply it consistently, and evaluate the model under exactly that policy.
This is why a FLOP reduction on paper does not automatically become lower latency: a general-purpose serving stack needs explicit support for per-layer sparse caches first.
MoD next to its relatives
| Method | What varies | Compute per token | Static shapes | Main cost |
|---|---|---|---|---|
| Dense transformer | Nothing | Fixed | Yes | Wasted compute on easy tokens |
| Mixture of experts | Which MLP runs | Fixed (top-k experts) | Mostly; load balancing needed | Memory for all experts, routing imbalance |
| Early exit | How many layers run | Variable | No | Missing deep KV for exited tokens, ragged batches |
| Mixture of Depths | Whether a block runs | Variable, capacity-bounded | In training | Non-causal top-k, sparse per-layer KV |
MoD and mixture of experts are orthogonal, and the paper combines them as MoDE in two ways. Staged MoDE routes tokens around or towards a block before the self-attention step, then applies expert routing in the MLP. Integrated MoDE adds no-op experts among the conventional MLP experts, so choosing a no-op expert is the same as skipping. Early exit, covered in SLM early exit, decides depth per token at inference from confidence; MoD learns per-block routing during training with a known compute budget. For the expert side, see mixture of experts and the arithmetic in MoE math.
Failure modes
| Symptom | Cause | Fix |
|---|---|---|
| Training loss suspiciously good, generation poor | Unsorted top-k indices leak future tokens through the causal mask | Sort indices before gather; test causality by perturbing future tokens |
| Long-range quality drops | Routed tokens given compacted positions | Pass original position ids to RoPE |
| Quality drops when every block is routed | No full-capacity blocks to refresh skipped tokens | Route every other block, as the paper found |
| Generation quality worse than validation loss suggests | Sampling-time routing rule differs from training top-k | Train the aux loss or predictor; monitor its accuracy and routed fraction |
| No wall-clock speedup at inference | Engine runs the full batch and masks, or lacks per-layer caches | Implement gather/scatter decode and sparse per-layer KV, then measure |
| Routing collapses onto a fixed position pattern | Router learns position shortcuts, for example always routing the first tokens | Inspect routing maps per layer; compare with a random-routing control |
When to use it
MoD is a training-time decision; you cannot add it to a pretrained dense model without further training. It suits pretraining under a fixed compute budget, where saved FLOPs buy a larger model or cheaper forward passes. It suits serving on an engine without sparse per-layer caches least. Public evidence at large scale is limited, so run your own isoFLOP comparison against a carefully tuned dense baseline.
What to do next
- Implement the MoD wrapper on a small model and write a causality test: changing token t+1 must not change any output at position t or before.
- Train dense and MoD variants at the same FLOP budget, with capacity 12.5 percent on every other block as the starting point.
- Add the auxiliary predictor and log its accuracy and the routed fraction per layer during training.
- Evaluate generation with the sampling-time routing rule, not only teacher-forced validation loss.
- Before promising a latency win, prototype decode with per-layer sparse KV and measure wall-clock time at your target batch sizes.