# How MTPLX Ensures Mathematically Correct Output Distribution at Any Temperature

> Discover how MTPLX ensures mathematically correct output distribution at any temperature. Explore its unique approach to probability computation and sampling heuristics.

- Repository: [Youssof Altoukhi/MTPLX](https://github.com/youssofal/MTPLX)
- Tags: deep-dive
- Published: 2026-09-05

---

**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](https://github.com/youssofal/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)](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py), where all logits flow through a **strictly ordered three-stage pipeline**:

1. **Penalty application** (`apply_penalties`) — OpenAI-style presence and frequency penalties modify raw logits *before* any stochastic operation.

2. **Temperature-scaled softmax** — The `softmax` function (lines 80-92) divides logits by temperature: `logits / temperature`. When `temperature ≤ 0`, it returns a one-hot vector at the argmax index, ensuring deterministic behavior without numerical instability.

3. **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.

```python
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`:

```python
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:

```python
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)](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)](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 |

```python
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)](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 filtering
- **`test_acceptance_probability_caps_at_one`** — Ensures `min(1, p/q)` never exceeds 1.0
- **`test_residual_distribution_nan_mass_falls_back_to_target`** — Handles degenerate cases where residual mass would produce NaN
- **`speculative_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)](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.