Exact Speculative Sampling Algorithm in MTPLX: A Complete Technical Breakdown

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

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


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


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


# 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

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

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 Core implementation: acceptance_probability, residual_distribution, verify_one_token, speculative_output_marginal
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 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 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.

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 →