# How MTPLX Implements Multi-Token Prediction with Speculative Sampling

> Discover how MTPLX implements multi-token prediction with speculative sampling. Generate draft tokens, verify acceptance probabilities, and match target distributions exactly.

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

---

**MTPLX implements multi-token prediction with speculative sampling by generating draft tokens ahead of the target model, then verifying each token using acceptance probabilities and residual distributions to ensure the final output matches the target distribution exactly.**

The open-source MTPLX repository (`youssofal/MTPLX`) provides a reference implementation of speculative decoding that enables low-latency, high-throughput text generation. At its core, multi-token prediction relies on a draft model proposing several future tokens, which the target model then validates or corrects in a single forward pass using statistically exact sampling methods.

## Core Sampling Utilities in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py)

The foundation of MTPLX’s speculative engine resides in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py), which exports pure functions for probability manipulation and token verification. These utilities ensure that speculative acceleration never compromises the statistical properties of the target model.

### Acceptance Probability Calculation (`acceptance_probability`)

Before a draft token can be accepted, MTPLX calculates the per-token acceptance probability as the ratio of target to draft probabilities. The `acceptance_probability` function (lines 44–49) computes this value as **p / q**, where **p** represents the target model’s probability mass and **q** represents the draft model’s probability mass for the same token. This ratio determines the likelihood that the draft token came from the same distribution as the target.

### Residual Distribution Construction (`residual_distribution`)

When a draft token is rejected, the algorithm must sample a replacement that accounts for the probability mass missed by the draft. The `residual_distribution` function (lines 52–78) constructs this "re-sample" distribution by normalizing the remaining probability mass of the target distribution after removing the accepted portion. This ensures that the final output distribution exactly equals the target distribution, maintaining sampling correctness even during acceleration.

### Token Verification Logic (`verify_one_token`)

The decision to accept or reject a token occurs in `verify_one_token` (lines 13–19). This function compares the target and draft probabilities using the acceptance probability, then returns a `SpeculativeDecision` dataclass (lines 1–5) containing:
- `accepted`: Boolean flag indicating whether the draft token was kept
- `token_id`: The final token ID (either the original draft token or a resampled replacement)
- `acceptance_prob`: The calculated probability ratio for telemetry

## The Multi-Token Speculation Loop

The high-level orchestration layer—typically located in [`mtplx/speculative.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py)—repeatedly invokes the verification utilities to achieve multi-token prediction. The algorithm follows three distinct phases:

1. **Draft Generation**: The draft model generates **N** tokens (where **N** equals `speculative_depth`) ahead of the current position.
2. **Verification**: For each draft token, the system calls `verify_one_token(target_p, draft_q, token)`, checking acceptance probabilities. Accepted tokens are emitted immediately; rejected tokens trigger a draw from the `residual_distribution`.
3. **State Update**: The accepted token—whether original or corrected—updates the running token stream, and the loop continues until all draft tokens are processed or an early exit condition is met.

The following Python sketch illustrates this speculative decoding loop:

```python
def speculative_decode(target_logits, draft_logits, config, depth, rng=None):
    """
    Multi-token speculative decoding implementation.
    """
    # Convert logits to probability distributions

    target_p = distribution_from_logits(target_logits, config)
    draft_q = distribution_from_logits(draft_logits, config)
    
    # Generate draft token stream of length `depth`

    draft_tokens = [sample_from_distribution(draft_q, rng) for _ in range(depth)]
    
    # Verify each draft token

    output = []
    for token in draft_tokens:
        decision = verify_one_token(target_p, draft_q, token, rng)
        output.append(decision.token_id)
        
        # Update distributions for next position (simplified)

        # In production, target_p and draft_q would advance based on the emitted token

        
    return output

```

The "multi-token" advantage emerges from running this loop for `speculative_depth > 1`, allowing the draft model to stay several steps ahead while the target model verifies the sequence in parallelized forward passes.

## Statistical Correctness and Marginal Distributions

MTPLX guarantees exact sampling through the `speculative_output_marginal` function (lines 21–40), which mathematically proves that the composition of acceptance and residual sampling yields a marginal distribution identical to direct sampling from the target model. This correctness property holds regardless of the draft model's quality, ensuring that acceleration never introduces sampling bias.

## Test Coverage and Integration

The implementation includes comprehensive test suites that verify both unit correctness and end-to-end behavior:
- [`tests/test_sampling.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py) validates the mathematical properties of `acceptance_probability`, `residual_distribution`, and `verify_one_token`
- [`tests/test_tail_gemma4_stream_holdback.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_tail_gemma4_stream_holdback.py) demonstrates multi-token speculative decoding with a depth of 2 in a realistic streaming context

## Summary

- MTPLX implements multi-token prediction through a draft-then-verify architecture located primarily in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) and orchestrated by [`mtplx/speculative.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py).
- The `acceptance_probability` function (lines 44–49) calculates the acceptance ratio **p / q** for each draft token.
- Rejected tokens are replaced using samples from the `residual_distribution` (lines 52–78), ensuring exact target distribution matching.
- The `verify_one_token` utility (lines 13–19) encapsulates the acceptance decision logic, returning structured results via the `SpeculativeDecision` dataclass.
- Multi-token speculation achieves speedup by processing `speculative_depth` tokens in parallel while maintaining statistical correctness through marginal distribution guarantees (lines 21–40).

## Frequently Asked Questions

### How does MTPLX ensure the final distribution matches the target model?

MTPLX ensures distribution matching through the residual sampling mechanism. When a draft token is rejected, the system samples from the `residual_distribution` (lines 52–78), which contains exactly the probability mass that the draft model missed. This guarantees that the marginal distribution of accepted and corrected tokens equals the target distribution, as proven by the `speculative_output_marginal` analysis (lines 21–40).

### What is the role of the residual distribution in speculative sampling?

The residual distribution serves as the "correction" mechanism when the draft model proposes low-probability tokens. By computing the normalized remaining mass of the target distribution after accounting for the draft token's probability, `residual_distribution` provides a statistically exact alternative sampling pool that preserves the target model's output distribution.

### Where is the speculative loop orchestrated in the MTPLX codebase?

While the mathematical primitives reside in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py), the high-level speculative loop that manages draft generation, batch verification, and state advancement is implemented in [`mtplx/speculative.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/speculative.py). This orchestration layer coordinates the repeated calls to `verify_one_token` for each position in the speculative depth window.

### How does multi-token prediction differ from single-token speculative decoding?

Single-token speculation verifies one draft token per target model forward pass, while multi-token prediction extends the `speculative_depth` to verify **N** tokens simultaneously. This increases the probability of accepting multiple consecutive tokens before requiring another target model evaluation, reducing per-token latency while using the same `acceptance_probability` and `residual_distribution` verification primitives.