# How MTPLX Ensures Output Distribution Matches Standard AR Decoding at Different Temperatures

> Discover how MTPLX guarantees distributional equivalence with standard AR decoding at any temperature using temperature-scaled softmax and exact acceptance/rejection probability.

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

---

**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`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py), implements numerically stable temperature scaling:

```python
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 where `q_draft > 0`
- **Rejection guarantee**: If `q_draft == 0` and `p_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:

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

1. **For each draft token**: Add `q_i × acceptance_probability(p_i, q_i)` to the marginal
2. **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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py) at line 10 as the primary interface for verification.

## Practical Verification

The following example demonstrates distributional equivalence across temperatures:

```python
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`) computes `p/q` ratios 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`](https://github.com/youssofal/MTPLX/blob/main/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.