Exact Speculative Sampling Algorithm in MTPLX: A Complete Technical Breakdown
MTPLX implements exact speculative sampling through a four-step probabilistic procedure—including acceptance probability calculation and residual distribution correction—to guarantee output distributions that identically match the target model.
The exact speculative sampling algorithm used in MTPLX ensures that draft tokens from a smaller, faster model can accelerate inference without distorting the final probability distribution. According to the MTPLX source code, this implementation in mtplx/sampling.py follows the mathematical guarantees established by the speculative decoding research lineage, ensuring the marginal output distribution equals the target model's distribution exactly.
The Four Steps of MTPLX Exact Speculative Sampling
The algorithm proceeds through tightly-coupled stages: distribution preparation, acceptance probability computation, residual distribution construction, and final token verification.
Step 1: Prepare Identical Target and Draft Distributions
Both target and draft logits undergo identical preprocessing to produce comparable probability vectors. The distribution_from_logits function applies:
- Penalty filtering via
apply_penalties(OpenAI-style) - Temperature scaling followed by softmax normalization
- Top-p and top-k filtering via
apply_top_p_top_k
# From mtplx/sampling.py lines 15-24
target_p = distribution_from_logits(target_logits, SamplerConfig())
draft_q = distribution_from_logits(draft_logits, SamplerConfig())
This identical processing pipeline ensures that target_p and draft_q represent probability distributions under the same sampling semantics, making subsequent comparisons valid.
Step 2: Compute Acceptance Probability
For each candidate token t, the acceptance probability is defined as min(1, pₜ / qₜ). The acceptance_probability function implements this with proper handling of edge cases:
# From mtplx/sampling.py lines 44-50
def acceptance_probability(target_p, draft_q, token):
p_t = target_p[token]
q_t = draft_q[token]
if q_t == 0:
return 0.0 if p_t > 0 else 1.0 # Reject unless both are zero
return min(1.0, p_t / q_t)
If the draft assigns zero probability to a token, it is rejected unless the target also assigns zero probability—preventing division by zero while preserving correctness.
Step 3: Build the Residual Distribution
When rejection occurs, the algorithm samples from the residual distribution containing probability mass the draft missed. This is defined as max(p - q, 0) and re-normalized:
# From mtplx/sampling.py lines 52-78
def residual_distribution(target_p, draft_q):
residual = np.maximum(target_p - draft_q, 0)
# Handles both dense np.ndarray and sparse SparseDistribution
return normalize(residual)
The implementation supports both dense numpy arrays and memory-efficient SparseDistribution representations for large vocabularies.
Step 4: Execute Token Verification
The verify_one_token function orchestrates the final decision by drawing a uniform random number and selecting the appropriate branch:
# From mtplx/sampling.py lines 13-19 (conceptual structure)
def verify_one_token(target_p, draft_q, draft_token, rng=None):
accept_p = acceptance_probability(target_p, draft_q, draft_token)
if rng.random() <= accept_p:
# Accepted: keep the draft token
return SpeculativeDecision(True, draft_token, accept_p)
else:
# Rejected: sample from residual distribution
residual = residual_distribution(target_p, draft_q)
corrected = sample_from_distribution(residual, rng)
return SpeculativeDecision(False, corrected, accept_p)
The returned SpeculativeDecision records acceptance status, final token ID, and acceptance probability for downstream analysis.
Mathematical Proof of Exactness
The correctness guarantee is encoded in speculative_output_marginal. This function enumerates all possible draft tokens, weights each by draft probability qₜ, and combines acceptance and residual contributions:
$$\sum_{t} q_t \left[ \underbrace{\frac{p_t}{q_t}}{\text{accept}} + \left(1 - \frac{p_t}{q_t}\right) \underbrace{r(\cdot)}{\text{residual}} \right] = p$$
The speculative_output_marginal function serves as a correctness oracle in unit tests, verifying that the speculative decoder's marginal distribution matches the target distribution within numerical tolerance.
Practical Implementation Examples
Verifying a Single Token Decision
import numpy as np
from mtplx.sampling import verify_one_token, SamplerConfig, distribution_from_logits
# Dummy logits for a 5-token vocabulary
target_logits = np.array([2.0, 0.5, -1.0, 0.0, 1.2])
draft_logits = np.array([1.8, 0.3, -0.8, 0.1, 1.0])
# Convert logits to probability vectors using identical sampler config
cfg = SamplerConfig(temperature=0.7, top_p=0.9, top_k=0)
target_p = distribution_from_logits(target_logits, cfg)
draft_q = distribution_from_logits(draft_logits, cfg)
# Sample draft token (normally generated by draft model)
draft_token = np.argmax(draft_q) # Deterministic for illustration
# Execute exact speculative decision
decision = verify_one_token(target_p, draft_q, draft_token)
print(decision)
# → SpeculativeDecision(accepted=True, token_id=0, accept_probability=0.93)
Testing with the Marginal Oracle
import numpy as np
from mtplx.sampling import (
speculative_output_marginal,
distribution_from_logits,
SamplerConfig,
)
# Target model logits (ground truth distribution)
target_logits = np.random.randn(1024)
# Draft model logits (e.g., from smaller, faster model)
draft_logits = target_logits + np.random.normal(scale=0.2, size=1024)
cfg = SamplerConfig(temperature=1.0, top_p=0.95, top_k=20)
target_p = distribution_from_logits(target_logits, cfg)
draft_q = distribution_from_logits(draft_logits, cfg)
# Compute exact marginal distribution
marginal = speculative_output_marginal(target_p, draft_q)
# Verify correctness: marginal must equal target distribution
assert np.allclose(marginal, target_p, atol=1e-6)
Key Source Files and Components
| File | Purpose |
|---|---|
mtplx/sampling.py |
Core implementation: acceptance_probability, residual_distribution, verify_one_token, speculative_output_marginal |
tests/test_sampling.py |
Unit tests validating marginal correctness and end-to-end speculative flow |
mtplx/backends/ |
Backend-specific draft models feeding logits to the sampler |
mtplx/kpi/reference_vllm.py |
vLLM reference implementation for cross-validation |
Summary
-
Exact speculative sampling in MTPLX guarantees the output distribution matches the target model through four precise steps: distribution preparation, acceptance probability calculation, residual distribution construction, and token verification.
-
The acceptance probability
min(1, p/q)and residual distributionmax(p-q, 0)work together to ensure mathematical exactness. -
verify_one_tokenencapsulates the runtime decision logic, whilespeculative_output_marginalprovides a testable correctness oracle. -
Both dense and sparse distribution representations are supported for memory efficiency at scale.
Frequently Asked Questions
What makes MTPLX's speculative sampling "exact" rather than approximate?
Exact speculative sampling means the marginal distribution of output tokens equals the target model's distribution precisely, not approximately. MTPLX achieves this by correcting rejected draft tokens through the residual distribution rather than fallback heuristics. The speculative_output_marginal function mathematically verifies this property by summing over all possible draft outcomes.
How does MTPLX handle the case where draft probability is zero?
The acceptance_probability function explicitly checks if q_t == 0 and returns 0.0 if p_t > 0 (always reject) or 1.0 if both are zero (trivial acceptance). This prevents division by zero while maintaining the exactness property, as implemented in mtplx/sampling.py lines 44-50.
Can the exact speculative sampling algorithm work with any draft model?
Yes. The algorithm's correctness does not depend on the draft model's quality—only its computational efficiency. A poor draft model results in low acceptance rates and frequent residual sampling, while a well-matched draft model achieves high acceptance. The target distribution remains exact regardless, as proven by the marginal calculation.
What is the performance cost of computing the residual distribution?
The residual_distribution function handles both dense np.ndarray and sparse SparseDistribution formats. For large vocabularies, sparse representations avoid materializing the full residual vector. The overhead is typically small compared to forward passes through the target model, and the cost is amortized across accepted draft tokens.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →