Inference latency in auto-regressive Large Language Models (LLMs) is strictly memory-bandwidth bound. During generation, loading tens of billions of model weights from GPU High-Bandwidth Memory (HBM) into SRAM yields a tiny arithmetic intensity of 1 FLOP per byte transferred when generating a single token per forward step.
Speculative Decoding and Medusa Multi-Head Decoding break this sequential bottleneck by decoupling token proposal from verification. By generating multiple candidate tokens in parallel and verifying them in a single GPU pass, inference speedups of 2x to 3x are achieved without modifying model accuracy or output distributions.
1. Memory-Bandwidth Bottleneck of Autoregressive Generation
For a 70B parameter FP16 model (140 GB weights), generating 1 token requires fetching 140 GB across the GPU memory bus. On an NVIDIA A100 GPU (2.0 TB/s bandwidth), the theoretical minimum time per token is:
However, GPUs perform matrix multiplication on candidate batches of $K$ tokens at virtually the same memory read latency. Speculative decoding exploits this hardware property by verifying $K$ draft tokens in parallel during a single target forward pass.
2. Standard Speculative Decoding Algorithm
Given a small draft model $\pi_D$ (e.g., Llama-3-8B-Draft) and a large target model $\pi_T$ (e.g., Llama-3-70B-Target):
- Draft Generation Phase: $\pi_D$ runs $K$ sequential autoregressive steps to propose a sequence of candidate tokens $\hat{x}_1, \hat{x}_2, \dots, \hat{x}_K$.
- Target Verification Phase: $\pi_T$ evaluates all $K$ tokens in a single parallel forward pass, computing output probability distributions $p(x_{t} | x_{
- Modified Rejection Sampling: For each candidate token $k \in \{1, \dots, K\}$, accept the token with probability:
If candidate $k$ is rejected, the sampling process halts, discards tokens $k+1 \dots K$, and resamples token $k$ from the adjusted distribution $p'(x) = \text{ReLU}(\pi_T(x) - \pi_D(x)) / \sum \text{ReLU}(\pi_T(x) - \pi_D(x))$.
3. The Medusa Architecture: Draft Heads Without a Separate Model
Maintaining a separate draft model requires complex memory management and multi-model synchronization. Medusa solves this by attaching multiple lightweight Feed-Forward heads directly onto the last hidden state of the target backbone model.
Instead of running a separate model, Medusa Head $k$ predicts token $t + k + 1$ simultaneously from the target model's final hidden state $h_t$:
Tree-Structured Attention Masks
Medusa constructs a candidate tree of candidate paths (e.g., top-3 tokens for head 1, top-2 for head 2). To evaluate 64 candidate paths concurrently in one forward pass, Medusa uses a custom 2D Tree Attention Mask matrix:
Join the Technical Discussion
Have questions about this architecture? Drop a comment below.