How MTPLX Ensures Mathematically Correct Output Distribution at Any Temperature
MTPLX guarantees mathematically exact output probabilities by separating probability computation from sampling heuristics and maintaining parity between a reference NumPy implementation and on-device MLX kernels.
Token sampling in large language models often sacrifices mathematical rigor for speed, introducing subtle distribution drift when temperature scaling or filtering is applied. MTPLX solves this by treating the output distribution as a first-class constraint, ensuring that any temperature value—including zero, negative, or extreme values—produces a distribution identical to the mathematically defined one. The repository achieves this through a unified probability pipeline, provably correct speculative sampling, and exhaustive test coverage that validates both reference and hardware-accelerated paths.
The Unified Probability Pipeline: Temperature-First, Filter-Second
MTPLX's sampling logic lives in [mtplx/sampling.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py), where all logits flow through a strictly ordered three-stage pipeline:
-
Penalty application (
apply_penalties) — OpenAI-style presence and frequency penalties modify raw logits before any stochastic operation. -
Temperature-scaled softmax — The
softmaxfunction (lines 80-92) divides logits by temperature:logits / temperature. Whentemperature ≤ 0, it returns a one-hot vector at the argmax index, ensuring deterministic behavior without numerical instability. -
Top-p / top-k filtering (
apply_top_p_top_k, lines 15-48) — Applied after softmax, preserving the exact probability mass of the filtered vocabulary through renormalization.
This ordering is critical. Applying penalties post-softmax would violate the temperature scaling definition; filtering pre-softmax would distort the intended probability distribution.
from mtplx.sampling import SamplerConfig, distribution_from_logits
import numpy as np
# Configure sampling with arbitrary temperature
cfg = SamplerConfig(temperature=0.3, top_p=0.9, top_k=5)
# Raw logits → mathematically correct distribution
logits = np.array([2.5, 1.0, -0.5, 0.2])
probs = distribution_from_logits(logits, cfg) # Exact softmax(temp) then filter
print("Normalized probabilities:", probs)
# Output: array guaranteed to sum to 1.0, respecting temperature=0.3
Provably Correct Speculative Sampling
MTPLX's speculative decoder uses a marginal recovery theorem proven in the source code. When a draft model proposes tokens faster than the target model, two distributions interact: the target target_p and the draft draft_q.
Acceptance Probability
The acceptance_probability function computes min(1, p/q) with a deterministic fallback when q=0:
from mtplx.sampling import acceptance_probability
target = np.array([0.7, 0.2, 0.1])
draft = np.array([0.5, 0.4, 0.1])
token_id = 0
accept_p = acceptance_probability(target, draft, token_id)
# accept_p = min(1, 0.7/0.5) = 1.0 (always accept when p > q)
If q=0 and p>0, acceptance is 1; if both are 0, acceptance is 0. This eliminates undefined behavior from division by zero.
Residual Distribution
The residual_distribution function (lines 52-88) builds a proper probability distribution from max(p - q, 0), renormalizing to ensure valid probabilities:
from mtplx.sampling import residual_distribution, speculative_output_marginal
# Verify mathematical correctness: marginal over all drafts recovers target
marginal = speculative_output_marginal(target, draft)
assert np.allclose(marginal, target) # Guaranteed by implementation
The speculative_output_marginal function (lines 21-41) proves that summing over every possible draft token—weighted by acceptance probability and residual correction—exactly equals the original target distribution. This is tested by [tests/test_sampling.py](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py), lines 52-57.
Device-Side Parity: MLX Kernels Match Reference Math
Mathematical guarantees mean nothing if hardware acceleration diverges. MTPLX solves this with [mtplx/fast_sampling.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/fast_sampling.py), where MLX kernels mirror the NumPy reference:
| Reference (NumPy) | MLX Kernel | Purpose |
|---|---|---|
distribution_from_logits |
_host_sparse_distribution |
Temperature-softmax-filter pipeline |
apply_penalties |
apply_penalties_mlx (lines 50-82) |
Identical penalty logic on-device |
deterministic_top_k_order |
_deterministic_mlx_top_k_support |
Exact ordering for reproducibility |
import mlx.core as mx
from mtplx.fast_sampling import sample_token_ids_from_mlx_logits
# MLX path: mathematically identical to NumPy reference
logits_mlx = mx.array([[4.0, 1.0, 0.5], [0.2, 3.0, 2.5]], dtype=mx.float32)
sampled_ids = sample_token_ids_from_mlx_logits(
logits_mlx,
SamplerConfig(temperature=0.6, top_p=0.01, top_k=2)
)
mx.eval(sampled_ids) # Forces execution on Apple Silicon GPU/NE
The top_p=0.01 parameter forces deterministic top-1 selection after temperature scaling, demonstrating that extreme filtering values are handled consistently across both paths.
Test-Driven Correctness Guarantees
[tests/test_sampling.py](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py) validates every mathematical claim:
test_distribution_from_logits_normalizes_after_filtering— Verifies that temperature-aware softmax produces valid probabilities even after aggressive top-p/top-k filteringtest_acceptance_probability_caps_at_one— Ensuresmin(1, p/q)never exceeds 1.0test_residual_distribution_nan_mass_falls_back_to_target— Handles degenerate cases where residual mass would produce NaNspeculative_output_marginal_recovers_target_distribution— The core theorem: speculative sampling preserves the target distribution
These tests run against both NumPy and MLX implementations, catching any divergence between reference math and hardware-optimized paths.
Summary
MTPLX ensures mathematically correct output distribution at any temperature through:
- Strict pipeline ordering: penalties → temperature-scaled softmax → filter — never reordered
- Rigorous temperature handling: exact softmax definition with deterministic zero-temperature behavior
- Proven speculative sampling: acceptance probability and residual distribution with marginal recovery guarantee
- Implementation parity: MLX kernels that byte-for-byte match NumPy reference semantics
- Exhaustive testing: unit tests covering edge cases from normal operation to numerical degeneracy
Frequently Asked Questions
What happens when temperature is zero or negative in MTPLX?
MTPLX returns a one-hot vector concentrated at the argmax index. The softmax function in [mtplx/sampling.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) (lines 80-92) explicitly handles temperature ≤ 0 by bypassing the exponential computation and setting probability 1.0 to the maximum logit index. This guarantees deterministic output without numerical overflow from division by zero.
How does MTPLX prevent distribution drift from top-p and top-k filtering?
MTPLX applies filters after softmax, then renormalizes the remaining probability mass. The apply_top_p_top_k function (lines 15-48) sorts by probability, applies the cutoff, and divides by the sum of surviving probabilities. This preserves the relative ratios of kept tokens and ensures the output sums to exactly 1.0.
Why does MTPLX use a NumPy reference implementation alongside MLX?
The NumPy reference serves as the ground-truth specification for mathematically correct behavior. The MLX kernels are verified against this reference, ensuring that GPU/NE acceleration never compromises correctness. This dual-path approach lets developers debug distribution issues on CPU with exact floating-point semantics before deploying to hardware.
How does speculative sampling in MTPLX maintain the target distribution?
Through the marginal recovery theorem implemented in speculative_output_marginal (lines 21-41). For each possible draft token, the algorithm considers: (a) accepting it with probability min(1, p/q), and (b) rejecting it and sampling from the residual max(p-q, 0) distribution. Summing over all draft outcomes yields exactly p, proving that speculative acceleration does not bias the output distribution.
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 →