All articles LLM Architecture

Mixture of Depths (MoD): Dynamic Layer Routing & Compute Allocation

Routing compute dynamically per token across transformer layers: top-k routing, capacity bounds, and integrating MoD with Mixture-of-Experts (MoE).

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

  1. 50% Reduction in FLOPs: MoD cuts total forward pass compute in half while preserving pre-training loss parity with dense baselines.
  2. Increased Inference Step Speed: Fast-path token skipping reduces KV cache memory reads and increases generation speeds.
  3. MoE Synergies: Combining MoD (dynamic depth) with MoE (dynamic width) creates sparse, ultra-efficient architectures.
← Back to all articles