Standard Transformer architectures allocate an identical amount of compute to every token in a sequence — spending the exact same FLOP budget on simple punctuation marks ("," or "the") as on complex reasoning tokens. Mixture of Depths (MoD) — introduced by DeepMind (Raposo et al., 2024) — allows Transformers to dynamically route tokens around self-attention and MLP blocks, focusing FLOPs exclusively on tokens that require deep contextual processing.
Standard Transformers spend uniform compute per token; Mixture of Depths routes tokens dynamically to save 50% FLOPs with zero loss in accuracy.
The inefficiency of static depth in Transformers
In a standard 32-layer Transformer, every token passes sequentially through all 32 self-attention and MLP blocks. Yet empirical probing demonstrates that over 60% of intermediate layer transformations on common words act as simple identity mappings.
MoD sets a strict token capacity budget K for each block (e.g. 50% of sequence length T). Only the top K tokens with the highest routing router weights pass through the attention/MLP block; the remaining tokens skip the block entirely via a residual shortcut!
PyTorch Mixture of Depths (MoD) layer implementation
import torch
import torch.nn as nn
import torch.nn.functional as F
class MixtureOfDepthsBlock(nn.Module):
"""
Mixture of Depths (MoD) block that dynamically routes
the top-K tokens through an expensive block while bypassing others.
"""
def __init__(self, d_model: int, block: nn.Module, capacity_factor: float = 0.5):
super().__init__()
self.block = block
self.capacity_factor = capacity_factor
self.router = nn.Linear(d_model, 1) # Router scalar score
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x shape: (batch_size, seq_len, d_model)
B, T, D = x.shape
K = int(T * self.capacity_factor) # Compute capacity budget
# 1. Compute router scalar weights for each token
router_logits = self.router(x).squeeze(-1) # (B, T)
# 2. Select top-K tokens per batch sequence
topk_weights, topk_indices = torch.topk(router_logits, K, dim=-1)
topk_weights = F.softmax(topk_weights, dim=-1) # (B, K)
# 3. Gather selected tokens for block execution
batch_idx = torch.arange(B).unsqueeze(-1).expand(-1, K)
selected_tokens = x[batch_idx, topk_indices] # (B, K, D)
# 4. Pass ONLY top-K tokens through expensive block
processed_tokens = self.block(selected_tokens) # (B, K, D)
# 5. Multiply by router weights and scatter back into residual stream
weighted_processed = processed_tokens * topk_weights.unsqueeze(-1)
out = x.clone() # Residual shortcut for non-selected tokens
out[batch_idx, topk_indices] += weighted_processed
return out
Performance & throughput benchmarks
- 50% Reduction in FLOPs: MoD cuts total forward pass compute in half while preserving pre-training loss parity with dense baselines.
- Increased Inference Step Speed: Fast-path token skipping reduces KV cache memory reads and increases generation speeds.
- MoE Synergies: Combining MoD (dynamic depth) with MoE (dynamic width) creates sparse, ultra-efficient architectures.