How MTPLX Implements Multi-Token Prediction with Speculative Sampling

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

The foundation of MTPLX’s speculative engine resides in 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—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:

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:

Summary

  • MTPLX implements multi-token prediction through a draft-then-verify architecture located primarily in mtplx/sampling.py and orchestrated by 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, the high-level speculative loop that manages draft generation, batch verification, and state advancement is implemented in 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.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →