Mixture of Experts

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.

On this page 9
  1. The short answer
  2. The toll plaza
  3. Why not give every word the same effort
  4. The design that avoids the wall
  5. How it relates to experts
  6. The problem that has held it back
  7. Where this stands
  8. Remember this
  9. What to learn next

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 smaller

How 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

Developer — Code and libraries.

Setup

bash
pip install torch

Written against PyTorch 2.5.1, Python 3.10. CPU, instant.

A mixture-of-depths block, and its central problem

mod.py
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.")
Output
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

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

ApproachWhat variesCausal at decode
Adaptive Computation Time (Graves, 2016)number of recurrent stepsyes
Early exit / CALMhow many layers before stoppingyes
Mixture of Depthswhich tokens enter each layerneeds a predictor
Mixture of Recursionshow many times a shared block is reappliedyes

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