# Exact Speculative Sampling Algorithm in MTPLX: A Complete Technical Breakdown

> Explore the exact speculative sampling algorithm in MTPLX. Learn its four-step process for guaranteed identical output distributions matching your target model.

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

---

**MTPLX implements exact speculative sampling through a four-step probabilistic procedure—including acceptance probability calculation and residual distribution correction—to guarantee output distributions that identically match the target model.**

The **exact speculative sampling algorithm** used in MTPLX ensures that draft tokens from a smaller, faster model can accelerate inference without distorting the final probability distribution. According to the MTPLX source code, this implementation in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) follows the mathematical guarantees established by the speculative decoding research lineage, ensuring the marginal output distribution equals the target model's distribution exactly.

## The Four Steps of MTPLX Exact Speculative Sampling

The algorithm proceeds through tightly-coupled stages: distribution preparation, acceptance probability computation, residual distribution construction, and final token verification.

### Step 1: Prepare Identical Target and Draft Distributions

Both target and draft logits undergo identical preprocessing to produce comparable probability vectors. The `distribution_from_logits` function applies:

- **Penalty filtering** via `apply_penalties` (OpenAI-style)
- **Temperature scaling** followed by softmax normalization
- **Top-p and top-k filtering** via `apply_top_p_top_k`

```python

# From mtplx/sampling.py lines 15-24

target_p = distribution_from_logits(target_logits, SamplerConfig())
draft_q  = distribution_from_logits(draft_logits,  SamplerConfig())

```

This identical processing pipeline ensures that `target_p` and `draft_q` represent probability distributions under the same sampling semantics, making subsequent comparisons valid.

### Step 2: Compute Acceptance Probability

For each candidate token *t*, the acceptance probability is defined as `min(1, pₜ / qₜ)`. The `acceptance_probability` function implements this with proper handling of edge cases:

```python

# From mtplx/sampling.py lines 44-50

def acceptance_probability(target_p, draft_q, token):
    p_t = target_p[token]
    q_t = draft_q[token]
    if q_t == 0:
        return 0.0 if p_t > 0 else 1.0  # Reject unless both are zero

    return min(1.0, p_t / q_t)

```

If the draft assigns zero probability to a token, it is rejected unless the target also assigns zero probability—preventing division by zero while preserving correctness.

### Step 3: Build the Residual Distribution

When rejection occurs, the algorithm samples from the **residual distribution** containing probability mass the draft missed. This is defined as `max(p - q, 0)` and re-normalized:

```python

# From mtplx/sampling.py lines 52-78

def residual_distribution(target_p, draft_q):
    residual = np.maximum(target_p - draft_q, 0)
    # Handles both dense np.ndarray and sparse SparseDistribution

    return normalize(residual)

```

The implementation supports both dense `numpy` arrays and memory-efficient `SparseDistribution` representations for large vocabularies.

### Step 4: Execute Token Verification

The `verify_one_token` function orchestrates the final decision by drawing a uniform random number and selecting the appropriate branch:

```python

# From mtplx/sampling.py lines 13-19 (conceptual structure)

def verify_one_token(target_p, draft_q, draft_token, rng=None):
    accept_p = acceptance_probability(target_p, draft_q, draft_token)
    
    if rng.random() <= accept_p:
        # Accepted: keep the draft token

        return SpeculativeDecision(True, draft_token, accept_p)
    else:
        # Rejected: sample from residual distribution

        residual = residual_distribution(target_p, draft_q)
        corrected = sample_from_distribution(residual, rng)
        return SpeculativeDecision(False, corrected, accept_p)

```

The returned `SpeculativeDecision` records acceptance status, final token ID, and acceptance probability for downstream analysis.

## Mathematical Proof of Exactness

The correctness guarantee is encoded in `speculative_output_marginal`. This function enumerates all possible draft tokens, weights each by draft probability `qₜ`, and combines acceptance and residual contributions:

$$\sum_{t} q_t \left[ \underbrace{\frac{p_t}{q_t}}_{\text{accept}} + \left(1 - \frac{p_t}{q_t}\right) \underbrace{r(\cdot)}_{\text{residual}} \right] = p$$

The `speculative_output_marginal` function serves as a **correctness oracle** in unit tests, verifying that the speculative decoder's marginal distribution matches the target distribution within numerical tolerance.

## Practical Implementation Examples

### Verifying a Single Token Decision

```python
import numpy as np
from mtplx.sampling import verify_one_token, SamplerConfig, distribution_from_logits

# Dummy logits for a 5-token vocabulary

target_logits = np.array([2.0, 0.5, -1.0, 0.0, 1.2])
draft_logits  = np.array([1.8, 0.3, -0.8, 0.1, 1.0])

# Convert logits to probability vectors using identical sampler config

cfg = SamplerConfig(temperature=0.7, top_p=0.9, top_k=0)
target_p = distribution_from_logits(target_logits, cfg)
draft_q  = distribution_from_logits(draft_logits, cfg)

# Sample draft token (normally generated by draft model)

draft_token = np.argmax(draft_q)  # Deterministic for illustration

# Execute exact speculative decision

decision = verify_one_token(target_p, draft_q, draft_token)

print(decision)

# → SpeculativeDecision(accepted=True, token_id=0, accept_probability=0.93)

```

### Testing with the Marginal Oracle

```python
import numpy as np
from mtplx.sampling import (
    speculative_output_marginal,
    distribution_from_logits,
    SamplerConfig,
)

# Target model logits (ground truth distribution)

target_logits = np.random.randn(1024)

# Draft model logits (e.g., from smaller, faster model)

draft_logits = target_logits + np.random.normal(scale=0.2, size=1024)

cfg = SamplerConfig(temperature=1.0, top_p=0.95, top_k=20)

target_p = distribution_from_logits(target_logits, cfg)
draft_q  = distribution_from_logits(draft_logits, cfg)

# Compute exact marginal distribution

marginal = speculative_output_marginal(target_p, draft_q)

# Verify correctness: marginal must equal target distribution

assert np.allclose(marginal, target_p, atol=1e-6)

```

## Key Source Files and Components

| File | Purpose |
|------|---------|
| [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) | Core implementation: `acceptance_probability`, `residual_distribution`, `verify_one_token`, `speculative_output_marginal` |
| [`tests/test_sampling.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py) | Unit tests validating marginal correctness and end-to-end speculative flow |
| `mtplx/backends/` | Backend-specific draft models feeding logits to the sampler |
| [`mtplx/kpi/reference_vllm.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/kpi/reference_vllm.py) | vLLM reference implementation for cross-validation |

## Summary

- **Exact speculative sampling in MTPLX** guarantees the output distribution matches the target model through four precise steps: distribution preparation, acceptance probability calculation, residual distribution construction, and token verification.

- The **acceptance probability** `min(1, p/q)` and **residual distribution** `max(p-q, 0)` work together to ensure mathematical exactness.

- **`verify_one_token`** encapsulates the runtime decision logic, while **`speculative_output_marginal`** provides a testable correctness oracle.

- Both dense and sparse distribution representations are supported for memory efficiency at scale.

## Frequently Asked Questions

### What makes MTPLX's speculative sampling "exact" rather than approximate?

**Exact speculative sampling** means the marginal distribution of output tokens equals the target model's distribution precisely, not approximately. MTPLX achieves this by correcting rejected draft tokens through the residual distribution rather than fallback heuristics. The `speculative_output_marginal` function mathematically verifies this property by summing over all possible draft outcomes.

### How does MTPLX handle the case where draft probability is zero?

The `acceptance_probability` function explicitly checks `if q_t == 0` and returns `0.0` if `p_t > 0` (always reject) or `1.0` if both are zero (trivial acceptance). This prevents division by zero while maintaining the exactness property, as implemented in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) lines 44-50.

### Can the exact speculative sampling algorithm work with any draft model?

**Yes.** The algorithm's correctness does not depend on the draft model's quality—only its computational efficiency. A poor draft model results in low acceptance rates and frequent residual sampling, while a well-matched draft model achieves high acceptance. The target distribution remains exact regardless, as proven by the marginal calculation.

### What is the performance cost of computing the residual distribution?

The `residual_distribution` function handles both **dense `np.ndarray`** and **sparse `SparseDistribution`** formats. For large vocabularies, sparse representations avoid materializing the full residual vector. The overhead is typically small compared to forward passes through the target model, and the cost is amortized across accepted draft tokens.