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:

$$t_{\text{step}} = \frac{140 \text{ GB}}{2,000 \text{ GB/s}} = 70 \text{ milliseconds} \implies \sim 14 \text{ tokens/sec}$$

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):

  1. 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$.
  2. Target Verification Phase: $\pi_T$ evaluates all $K$ tokens in a single parallel forward pass, computing output probability distributions $p(x_{t} | x_{
  3. Modified Rejection Sampling: For each candidate token $k \in \{1, \dots, K\}$, accept the token with probability:
$$P_{\text{accept}}(x_k) = \min \left( 1, \frac{\pi_T(x_k | x_{

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$:

$$\text{MedusaHead}_k(h_t) = \text{Softmax}(W_{k, 2} \cdot \text{SiLU}(W_{k, 1} \cdot 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:

import torch def create_medusa_tree_mask(tree_indices: torch.Tensor) -> torch.Tensor: """ Constructs 2D boolean tree attention mask for Medusa candidate verification """ K = tree_indices.size(0) mask = torch.zeros((K, K), dtype=torch.bool) for i in range(K): for j in range(k + 1): # A node can attend to itself and its ancestor nodes in the tree path if tree_indices[j] in tree_indices[:i+1]: mask[i, j] = True return mask

4. PyTorch Implementation of Rejection Sampling Verification

import torch def speculative_rejection_sampler( target_probs: torch.Tensor, # [K, Vocab] draft_probs: torch.Tensor, # [K, Vocab] draft_tokens: torch.Tensor # [K] ): """ Executes exact distribution-preserving rejection sampling for K speculative tokens. """ accepted_tokens = [] for k in range(draft_tokens.size(0)): tok = draft_tokens[k].item() p_target = target_probs[k, tok].item() p_draft = draft_probs[k, tok].item() # Calculate acceptance probability accept_prob = min(1.0, p_target / (p_draft + 1e-8)) r = torch.rand(1).item() if r < accept_prob: accepted_tokens.append(tok) else: # Token rejected! Resample replacement token from normalized difference diff = torch.relu(target_probs[k] - draft_probs[k]) norm_diff = diff / diff.sum() resampled_tok = torch.multinomial(norm_diff, num_samples=1).item() accepted_tokens.append(resampled_tok) break # Halt further token verification return accepted_tokens # Test verification if __name__ == "__main__": t_probs = torch.softmax(torch.randn(3, 1000), dim=-1) d_probs = torch.softmax(torch.randn(3, 1000), dim=-1) d_toks = torch.tensor([14, 882, 304]) accepted = speculative_rejection_sampler(t_probs, d_probs, d_toks) print("Accepted Token Sequence:", accepted)
🏢

About the Publisher: TechMind Editorial

TechMind is an independent engineering publication dedicated to systems architecture, LLM runtimes, and distributed infrastructure. Our editorial team comprises veteran systems engineers.

Join the Technical Discussion

Have questions about this architecture? Drop a comment below.