How MTPLX Ensures Exact Output Distribution with Speculative Decoding
MTPLX guarantees mathematically exact output distributions by combining draft token acceptance probabilities with residual distribution correction, ensuring the marginal distribution of speculative decoding matches the full target model exactly.
Speculative decoding promises faster LLM inference by using a small, fast "draft" model to propose tokens that a larger target model verifies. The critical challenge is maintaining output quality: if the verification step skews the distribution, you trade accuracy for speed. MTPLX solves this through three precisely engineered components in mtplx/sampling.py that together preserve the exact target distribution while still accepting draft tokens whenever possible.
The Three Mathematical Foundations of Exact Distribution
MTPLX's guarantee rests on three tightly coupled functions that transform target and draft distributions into provably correct speculative decoding behavior.
Acceptance Probability: When to Trust the Draft
The acceptance_probability function determines how likely a drafted token should be accepted. It implements the classic speculative decoding acceptance criterion:
# From mtplx/sampling.py, lines 44-49
def acceptance_probability(target_p, draft_q, token_id):
"""Compute probability of accepting a drafted token."""
q = draft_q[token_id]
if q == 0:
return 0.0 # Deterministic rejection if draft gives zero mass
return min(1.0, target_p[token_id] / q)
The ratio p/q (target probability divided by draft probability) capped at 1.0 ensures the acceptance decision respects the target distribution. When the draft underestimates a token's true probability, partial acceptance preserves the correct marginal.
Residual Distribution: Handling Rejections Correctly
When a draft token is rejected, MTPLX doesn't simply resample from the target. The residual_distribution function constructs a carefully normalized distribution that accounts for "where the draft went wrong":
# From mtplx/sampling.py, lines 52-78
def residual_distribution(target_p, draft_q):
"""Build residual probability mass for corrective sampling."""
residual = target_p - draft_q
residual = np.maximum(residual, 0) # Clip negative values
residual_sum = residual.sum()
if residual_sum == 0:
# Fallback: uniform over positions where target > 0
return target_p / target_p.sum()
return residual / residual_sum
This function subtracts the draft distribution from the target, clips negative values to zero, and normalizes. The result represents probability mass the draft failed to allocate correctly—mass that must be sampled from to maintain distributional correctness.
Speculative Output Marginal: The Exactness Proof
The centerpiece is speculative_output_marginal, which computes the theoretical output distribution of the full speculative decoding process. This function serves as both implementation guide and mathematical proof:
# From mtplx/sampling.py, lines 21-40
def speculative_output_marginal(target_p, draft_q):
"""Compute exact marginal distribution induced by speculative decoding."""
n = len(target_p)
marginal = np.zeros(n)
for i in range(n):
q_i = draft_q[i]
if q_i == 0:
continue
# Acceptance contribution
accept_p = min(1.0, target_p[i] / q_i)
marginal[i] += q_i * accept_p
# Rejection + residual contribution
if accept_p < 1.0:
residual = residual_distribution(target_p, draft_q)
marginal[i] += q_i * (1 - accept_p) * residual[i]
return marginal / marginal.sum()
This iterates over every possible draft token, adding:
- The accepted portion:
q * accept_p - The corrective portion:
q * (1-accept_p) * residual
The sum exactly reconstructs the target distribution, as verified by MTPLX's test suite.
How the Components Work Together in Practice
During actual generation, mtplx/speculative.py orchestrates these mathematical pieces into an efficient loop:
import numpy as np
from mtplx.sampling import (
acceptance_probability,
residual_distribution,
speculative_output_marginal,
)
# Target distribution from the full model (logits soft-maxed)
target = np.array([0.6, 0.3, 0.1])
# Draft distribution from a faster model
draft = np.array([0.4, 0.4, 0.2])
# 1️⃣ Acceptance probabilities for each draft token
accept_probs = [
acceptance_probability(target, draft, idx) for idx in range(len(draft))
]
# → [1.0, 0.75, 0.5]
# 2️⃣ Residual distribution (used when draft is rejected)
residual = residual_distribution(target, draft)
# → array([0.5, 0.0, 0.5])
# 3️⃣ Verify exact marginal matches target
exact_marginal = speculative_output_marginal(target, draft)
# → array([0.6, 0.3, 0.1]) # Matches original target exactly
The acceptance probabilities show intuitive behavior: token 0 is always accepted (draft underestimates it), token 1 has 75% acceptance, token 2 has 50%. Yet the final marginal is identical to the target—neither biased toward overconfident draft predictions nor unnecessarily conservative.
Why Mathematical Derivation Matters for Guaranteed Exactness
The key insight is that acceptance_probability and residual_distribution derive from the same target and draft distributions. This shared foundation ensures no probability mass is lost or double-counted. The speculative_output_marginal function explicitly enumerates all paths (accept or reject for each possible draft token), making the exactness guarantee transparent and verifiable.
MTPLX validates this property in tests/test_sampling.py with a unit test that enumerates every possible draft token and asserts exact distributional equality. This isn't merely empirical—it follows from the algebraic structure of the acceptance and residual formulas.
Summary
- Three-component design:
acceptance_probability,residual_distribution, andspeculative_output_marginalinmtplx/sampling.pyform a mathematically complete solution - Ratio-based acceptance: The
min(1, p/q)criterion ensures draft tokens are accepted proportionally to their agreement with the target - Residual correction: Rejected tokens trigger sampling from normalized residual mass, not raw target resampling
- Provable exactness:
speculative_output_marginalcomputes the theoretical output distribution, which unit tests verify matches the target exactly - Latency without loss: The full speculative decoding loop in
mtplx/speculative.pydelivers speedups while preserving output quality identical to full-model generation
Frequently Asked Questions
Does speculative decoding in MTPLX change the output distribution compared to the full model?
No. MTPLX's speculative decoding produces exactly the same output distribution as running the full target model. The speculative_output_marginal function proves this algebraically, and the test suite verifies it numerically. Speed comes from accepting draft tokens when safe, not from approximating the target distribution.
How does MTPLX handle cases where the draft model assigns zero probability to a high-probability target token?
The acceptance_probability function returns 0.0 when draft_q[token_id] == 0, causing deterministic rejection. The residual distribution then assigns appropriate probability mass to this token through residual_distribution, ensuring it can still be sampled with correct frequency despite the draft's blind spot.
What happens if the draft and target distributions are identical?
When target_p == draft_q, acceptance_probability returns 1.0 for all tokens, and speculative_output_marginal shows that all draft tokens are accepted. The speculative decoder reduces to simply running the draft model with zero rejection overhead, achieving maximum speedup while maintaining exactness trivially.
Can I verify the exactness guarantee myself?
Yes. Run pytest tests/test_sampling.py in the MTPLX repository to execute the unit test that enumerates draft tokens and asserts np.allclose(speculative_output_marginal(target, draft), target). The test exercises diverse distribution pairs including edge cases with zeros and mismatched supports.
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 →