All articles Inference Optimization

Speculative Decoding & Medusa Heads: Accelerating LLM Token Generation

Scaling LLM inference throughput: draft-target model speculative sampling, tree-based verification, and multi-head parallel token prediction with Medusa.

Autoregressive Large Language Model (LLM) generation is strictly memory-bandwidth bound. For an 70B parameter model in 16-bit precision, generating a single output token requires loading 140GB of model parameters from GPU VRAM into registers, regardless of whether the batch size is 1 or 16. Speculative Decoding and Medusa Multi-Head Decoding break this memory latency wall by predicting multiple candidate tokens per forward pass and verifying them in parallel.

Autoregressive generation runs 1 token per forward pass; Speculative Decoding generates K tokens per forward pass with mathematical equivalence.

The memory bandwidth wall in autoregressive decoding

On an NVIDIA H100 GPU (3.3 TB/s memory bandwidth), reading a 70B parameter model (140GB in FP16) takes approximately 42 milliseconds per token. Arithmetic intensity (FLOPS / byte) during single-user batch-1 decoding is less than 2, leaving over 95% of the GPU's compute FLOPS idle!

Speculative decoding algorithm & acceptance probability

Speculative Decoding pairs a small, fast Draft Model (e.g. Llama-3 8B) with a large Target Model (e.g. Llama-3 70B):

  1. The Draft Model generates K candidate tokens speculatively in K rapid forward passes.
  2. The Target Model runs a single batch forward pass over all K tokens simultaneously.
  3. A modified rejection sampling rule determines how many tokens are accepted:
Acceptance Condition:
For each token i in 1..K:
r = Uniform(0, 1)
if r <= min(1, P_target(x_i) / P_draft(x_i)):
    Accept token x_i
else:
    Reject token x_i and sample next token from max(0, P_target(x) - P_draft(x))
    Break speculative verification loop

Because the Target Model runs a single forward pass over all K tokens, any accepted tokens effectively cost 0 extra memory bandwidth loads!

Medusa: Single-model speculative decoding with parallel heads

Medusa eliminates the need for a separate draft model altogether by attaching multiple lightweight ResNet/MLP decoding heads to the main model's top hidden state:

                       +----------------------+
                       | Base LLM Backbone    |
                       +----------+-----------+
                                  | (Hidden State H_t)
         +------------------------+------------------------+
         |                        |                        |
+--------v-------+       +--------v-------+       +--------v-------+
| Medusa Head 1  |       | Medusa Head 2  |       | Medusa Head 3  |
| Predicts t+1   |       | Predicts t+2   |       | Predicts t+3   |
+----------------+       +----------------+       +----------------+

Tree-based verification & custom attention masks

Rather than evaluating a single linear candidate sequence, Medusa constructs a tree of candidate token paths. Using custom 2D Tree Attention masks, the Target Model evaluates up to 64 candidate paths simultaneously in one forward pass.

Python implementation of Speculative Sampling verification

import torch

def verify_speculative_tokens(
    target_logits: torch.Tensor, # Shape: (K+1, Vocab_Size)
    draft_logits: torch.Tensor,  # Shape: (K, Vocab_Size)
    draft_tokens: torch.Tensor   # Shape: (K,)
) -> torch.Tensor:
    """
    Executes lossless rejection sampling for speculative verification.
    """
    accepted_tokens = []
    K = draft_tokens.shape[0]
    
    target_probs = torch.softmax(target_logits, dim=-1)
    draft_probs = torch.softmax(draft_logits, dim=-1)

    for i in range(K):
        token_id = draft_tokens[i].item()
        p_target = target_probs[i, token_id].item()
        p_draft = draft_probs[i, token_id].item()

        # Rejection sampling condition
        r = torch.rand(1).item()
        if r <= min(1.0, p_target / (p_draft + 1e-8)):
            accepted_tokens.append(token_id)
        else:
            # Resample replacement token from adjusted distribution
            diff_dist = torch.clamp(target_probs[i] - draft_probs[i], min=0)
            resampled_token = torch.multinomial(diff_dist / diff_dist.sum(), 1).item()
            accepted_tokens.append(resampled_token)
            return torch.tensor(accepted_tokens) # Terminate early on rejection

    # If all K tokens accepted, sample 1 extra token from target distribution
    final_token = torch.multinomial(target_probs[K], 1).item()
    accepted_tokens.append(final_token)
    return torch.tensor(accepted_tokens)

Inference optimization benchmarks

  1. 2x-3x Speedup on Batch 1: Speculative Decoding doubles generation throughput without modifying output distribution quality.
  2. Zero Quality Loss: The rejection sampling math guarantees exact distributional equivalence to standard greedy/sampled decoding.
  3. Medusa for Single-Model Deployments: Medusa heads add less than 2% parameter footprint while delivering 2.2x speedups without draft model overhead.
← Back to all articles