Speculative Decoding accelerates inference by 2-3x without quality loss.
Core Idea
Instead of generating one token at a time, guess multiple tokens with a fast draft model. Then verify all with the big model at once.
Standard (Autoregressive):
Input: "Hello"
Step 1: β "world"
Step 2: β "how"
Step 3: β "are"
... (slow, 1 token per step)
Speculative Decoding:
Input: "Hello"
Draft Phase: Quick model guesses
β [world, how, are, you, today] (5 tokens at once!)
Verify Phase: Big model checks
β All correct?
β If yes: All accepted, save 5 steps!
β If no: Accept up to first mistake, regenerate 1
Draft-then-Verify Process
Step 1: Draft Phase
Fast model generates multiple candidates.
def draft_phase(input_ids, draft_model, num_tokens=5):
draft_tokens = []
for i in range(num_tokens):
logits = draft_model(input_ids + draft_tokens)
token = logits.argmax(dim=-1)[-1]
draft_tokens.append(token)
return draft_tokens # [token1, token2, token3, token4, token5]
Step 2: Verify Phase
Big model checks all draft tokens in parallel.
def verify_phase(input_ids, draft_tokens, verifier_model):
extended_ids = input_ids + draft_tokens
# Parallel verification over all positions!
logits = verifier_model(extended_ids)
# Output: (1, extended_len, vocab_size)
verified = []
for i, draft_token in enumerate(draft_tokens):
pos = len(input_ids) + i
verifier_distribution = logits[0, pos - 1]
draft_prob = verifier_distribution[draft_token]
if draft_prob > acceptance_threshold:
verified.append(draft_token)
else:
# First mistake: stop here, sample from verifier
true_token = sample_from(verifier_distribution)
verified.append(true_token)
break
return verified
Acceptance Criteria
Not all draft tokens are accepted. Multiple strategies exist.
1. Greedy Acceptance
Accept only if draft token is best prediction.
def greedy_acceptance(draft_token, verifier_logits):
best_token = verifier_logits.argmax()
return draft_token == best_token
2. Probabilistic Acceptance
Accept if draft probability >= Ξ± * verifier probability.
def probabilistic_acceptance(draft_token, draft_logits, verifier_logits, alpha=0.9):
draft_prob = softmax(draft_logits)[draft_token]
verifier_prob = softmax(verifier_logits)[draft_token]
if draft_prob >= alpha * verifier_prob:
return True
else:
return False
Practical Implementation
vLLM Speculative Decoding
from vllm import LLM
# Load models
model = LLM(model="meta-llama/Llama-2-70b")
draft_model = LLM(model="meta-llama/Llama-2-7b")
# Enable Speculative Decoding
outputs = model.generate(
prompt="Tell me a story:",
speculative_model=draft_model,
num_speculative_tokens=5,
temperature=0.8
)
Speedup Analysis
Without Speculative:
Throughput: 25 Tokens/Sec
Latency (100 tokens): 4 seconds
With Speculative (5 tokens):
Throughput: 65 Tokens/Sec (2.6x!)
Latency (100 tokens): 1.5 seconds
Key: Draft is very fast (7B), Verifier is valuable (70B)
Medusa and Eagle
Newer variants of Speculative Decoding.
Medusa (Meta)
Add extra heads to main model instead of separate draft.
Standard LLM: Input β Hidden States β Output Token
With Medusa: Input β Hidden States β Main Head: Token
ββ Medusa Head 1: Token[t+1]
ββ Medusa Head 2: Token[t+2]
ββ Medusa Head 3: Token[t+3]
ββ Medusa Head 4: Token[t+4]
Eagle (Microsoft)
Generates multiple tokens in parallel heads.
Practical Performance:
- Llama-70B: 2.4x Speedup
- Quality matches original
Best Practices
1. Draft Model Selection
Best: 70B Verifier + 7B Draft
β 2.5x Speedup (Goldstandard)
Bad: 70B Verifier + 0.5B Draft
β Draft rejected too often
β Only ~1.5x Speedup
2. Number of Speculative Tokens
draft_tokens = 1: ~1.1x Speedup (too few)
draft_tokens = 5: 2-3x Speedup (optimal)
draft_tokens = 10: 2.5-3x Speedup (diminishing returns)
Rule: Start with 5, optimize based on acceptance rate
3. Acceptance Threshold
alpha = 0.99: Too strict, many rejections
alpha = 0.9: Perfect balance (recommended)
alpha = 0.7: Too lenient, quality may suffer
