How MTPLX Ensures Output Distribution Matches Standard AR Decoding at Different Temperatures
MTPLX guarantees exact distributional equivalence with standard autoregressive decoding across any temperature through temperature-scaled softmax, exact acceptance/rejection probability, and residual distribution correction.
MTPLX is an open-source speculative decoding library that accelerates large language model inference without altering the output distribution. The primary challenge in speculative decoding is ensuring that draft tokens from a faster auxiliary model, when verified against the target model, produce exactly the same probability distribution as standard autoregressive (AR) sampling would—regardless of temperature scaling. This article examines how MTPLX solves this problem through three tightly-coupled mechanisms implemented in mtplx/sampling.py.
Temperature-Scaled Softmax Foundation
All probability computations in MTPLX begin with temperature-adjusted logits. The distribution_from_logits function, located at lines 80–92 of mtplx/sampling.py, implements numerically stable temperature scaling:
def distribution_from_logits(logits: np.ndarray, cfg: SamplerConfig) -> np.ndarray:
# Divide by temperature before softmax
scaled = logits / cfg.temperature
# Numerical stability: subtract max before exp
scaled = scaled - np.max(scaled)
exp_scaled = np.exp(scaled)
# Normalize to probability distribution
return exp_scaled / np.sum(exp_scaled)
When temperature approaches zero, the function returns a one-hot vector representing greedy decoding. For temperature > 0, the output is a proper probability distribution with entropy controlled by the temperature parameter. This scaling occurs before any speculative decoding steps, ensuring that all subsequent acceptance logic operates on correctly temperature-adjusted probabilities.
The default temperature configuration resides in SamplerConfig at lines 19–23 of the same file, allowing per-sampler customization while maintaining consistent behavior.
Exact Acceptance Probability Computation
For each draft token, MTPLX computes the probability of acceptance using the ratio of target to draft probabilities. The acceptance_probability function (lines 44–49) implements this core logic:
- Acceptance probability:
min(p_target / q_draft, 1.0)for tokens whereq_draft > 0 - Rejection guarantee: If
q_draft == 0andp_target > 0, the token is always rejected - Zero-zero case: If both probabilities are zero, acceptance is zero (undefined state)
This formulation prevents the marginal distribution from over-weighting tokens that the draft model never proposes. The acceptance logic preserves the target distribution exactly because accepted tokens retain their relative probabilities, while the probability mass for rejected tokens is redistributed through the residual mechanism below.
Residual Distribution Correction
When a draft token is rejected, MTPLX samples from a residual distribution that precisely captures the portion of the target distribution not covered by the draft. The residual_distribution function at lines 52–78 computes:
def residual_distribution(p_target: np.ndarray, p_draft: np.ndarray) -> np.ndarray:
# Element-wise difference, truncated at zero
residual = np.maximum(p_target - p_draft, 0.0)
# Normalize if non-empty residual exists
residual_sum = np.sum(residual)
if residual_sum > 0:
return residual / residual_sum
# Fallback: return full target distribution
return p_target
The residual represents "what the target model would sample that the draft missed." By construction, this correction ensures that the combined process—accept with probability p/q, otherwise sample residual—recovers the exact target distribution.
Both dense and sparse implementations are supported for computational efficiency with large vocabularies.
Mathematical Guarantee via Marginal Aggregation
The speculative_output_marginal function (lines 21–40) assembles the complete proof by aggregating all possible outcomes:
- For each draft token: Add
q_i × acceptance_probability(p_i, q_i)to the marginal - For rejection case: Add
rejection_probability × residual_distribution
Mathematically, this yields:
marginal_j = Σ_i [q_i × min(p_j/q_i, 1) × indicator(draft=i, target=j)]
+ (1 - Σ_i p_i/q_i × q_i) × residual_j
When simplified, the marginal equals p_j exactly. This function is re-exported from mtplx/speculative.py at line 10 as the primary interface for verification.
Practical Verification
The following example demonstrates distributional equivalence across temperatures:
import numpy as np
from mtplx.sampling import SamplerConfig, distribution_from_logits, speculative_output_marginal
# Test across multiple temperatures
for temp in [0.0, 0.5, 0.8, 1.0, 1.5]:
cfg = SamplerConfig(temperature=temp, top_p=1.0, top_k=0)
# Target model logits (high-quality, slow)
target_logits = np.array([2.0, 0.5, -1.0, 0.3])
# Draft model logits (lower-quality, fast)
draft_logits = np.array([1.8, 0.7, -0.8, 0.1])
# Compute distributions
p_target = distribution_from_logits(target_logits, cfg)
p_draft = distribution_from_logits(draft_logits, cfg)
# Verify equivalence
marginal = speculative_output_marginal(p_target, p_draft)
assert np.allclose(marginal, p_target, atol=1e-6)
print(f"✅ Temperature {temp}: distributional equivalence verified")
The assertion holds for all temperature values including 0.0 because the entire pipeline—softmax scaling, acceptance probability, and residual correction—operates on the temperature-adjusted distributions consistently.
Performance Implications
Temperature scaling in speculative decoding introduces no additional overhead beyond the initial softmax computation. The acceptance and residual operations are vectorized over the vocabulary dimension, with complexity O(V) per token regardless of temperature. Since temperature is applied uniformly before speculative steps, draft model speedups (typically 2–3×) are preserved across all sampling configurations.
Summary
- Temperature-scaled softmax (
mtplx/sampling.py:80-92) ensures all probability vectors respect the temperature parameter before any speculative logic - Exact acceptance probability (
mtplx/sampling.py:44-49) computesp/qratios that prevent distribution distortion - Residual distribution correction (
mtplx/sampling.py:52-78) captures target probability mass missed by the draft model - Marginal proof (
mtplx/sampling.py:21-40) mathematically guarantees the output distribution equals standard AR decoding
Frequently Asked Questions
Does MTPLX support temperature=0 (greedy decoding)?
Yes. When temperature <= 0, distribution_from_logits returns a one-hot vector at the argmax position. The acceptance logic reduces to deterministic verification: draft tokens are accepted only if they match the greedy target, otherwise the residual selects the true argmax. The marginal remains identical to greedy AR decoding.
How does MTPLX handle numerical stability with low temperatures?
The softmax implementation subtracts the maximum logit before exponentiation (lines 86–87 of mtplx/sampling.py). For very low temperatures, this prevents overflow when dividing by near-zero values. The ratio p/q in acceptance probability remains stable because both numerator and denominator are computed from identically scaled logits.
Can the draft model use a different temperature than the target?
No. For distributional equivalence, both models must use the same SamplerConfig. If the draft operated at a different temperature, the ratio p/q would not correctly represent the acceptance probability for the target distribution. The speculative_output_marginal function assumes matched configurations.
Why is residual correction necessary instead of simply resampling from the target?
Resampling from the full target distribution after rejection would re-introduce probability mass for the rejected token, violating the acceptance decision. The residual max(p-q, 0) precisely removes already-covered probability mass, ensuring that the combined accept/reject process is mathematically equivalent to direct target sampling.
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 →