# How MTPLX Ensures Exact Output Distribution with Speculative Decoding

> Discover how MTPLX ensures exact output distribution with speculative decoding. Learn about draft token acceptance and residual distribution correction for precise results.

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

---

**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`](https://github.com/youssofal/MTPLX/blob/main/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:

```python

# 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":

```python

# 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:

```python

# 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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py) orchestrates these mathematical pieces into an efficient loop:

```python
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`](https://github.com/youssofal/MTPLX/blob/main/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`, and `speculative_output_marginal` in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) form 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_marginal` computes the theoretical output distribution, which unit tests verify matches the target exactly
- **Latency without loss**: The full speculative decoding loop in [`mtplx/speculative.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py) delivers 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.