Where to Find Speculative Sampling Primitives in MTPLX: Complete Implementation Guide

The speculative sampling primitives in MTPLX are fully implemented in mtplx/sampling.py, which provides verify_one_token, acceptance_probability, and the SpeculativeDecision dataclass for draft token verification and correction.

The youssofal/MTPLX repository delivers a high-performance inference engine that accelerates large language model decoding through speculative execution. Understanding the location and structure of speculative sampling primitives is essential for optimizing draft-target token alignment or debugging acceptance probability calculations. All core mathematical operations for speculative decoding reside in a single Python module that implements the foundational acceptance and correction algorithms.

Core Implementation File: mtplx/sampling.py

This module houses every building block required for speculative decoding. Located at the root of the mtplx/ package, sampling.py defines the mathematical primitives that higher-level runtime loops invoke to determine whether draft tokens align with target model distributions.

Key Speculative Sampling Primitives

SpeculativeDecision Dataclass

At line 300, the SpeculativeDecision dataclass serves as the immutable container for verification results. It stores the accepted boolean flag, the chosen token ID, and the acceptance probability calculated during the verification step.

Acceptance Probability Calculation

The acceptance_probability function (line 44) implements the core statistical check using the formula $\min(1, p/q)$, where $p$ represents the target token probability and $q$ represents the draft token probability. This ratio determines whether the draft token sufficiently matches the target distribution to be accepted without correction.

Residual Distribution Construction

When draft tokens fail acceptance, the residual_distribution function (line 52) builds the correction distribution. This function computes the probability mass remaining after rejecting the draft token, ensuring that resampling maintains the exact target distribution properties required for unbiased speculative decoding.

Single-Step Verification Logic

The verify_one_token function (line 70) orchestrates the complete verification pipeline for a single token position. It calculates the acceptance probability, determines whether to keep the draft token, and falls back to sampling from the residual distribution when rejection occurs.

Marginal Distribution Oracle

For testing and validation, speculative_output_marginal (line 82) acts as a correctness oracle. This function computes the exact marginal distribution resulting from one-token speculative sampling, allowing unit tests to verify that the acceptance and correction logic preserves the target distribution's statistical properties.

Distribution Preparation Utilities

Before speculative verification begins, raw logits require normalization and filtering. The module provides several helper functions to prepare probability vectors:

  • distribution_from_logits (line 16) — Converts raw model logits into normalized probability distributions.
  • apply_top_p_top_k (line 15) — Applies nucleus (top-p) and top-k filtering to restrict the sampling space.
  • apply_penalties (line 61) — Adjusts probabilities based on repetition penalties or other custom constraints.

These utilities ensure both target and draft probability vectors undergo identical preprocessing before entering the speculative sampling pipeline.

Practical Implementation Example

The following example demonstrates the complete workflow: preparing distributions, verifying a draft token, and computing the marginal distribution for validation.

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

# -------------------------------------------------

# 1️⃣ Build target & draft probability vectors

# -------------------------------------------------

logits_target = np.array([1.2, 0.5, -0.3])   # model logits for the true (target) head

logits_draft  = np.array([0.9, 0.6, -0.1])   # logits from a faster draft model

cfg = SamplerConfig(temperature=0.7, top_p=0.9, top_k=0)
target_p = distribution_from_logits(logits_target, cfg)
draft_q  = distribution_from_logits(logits_draft,  cfg)

# -------------------------------------------------

# 2️⃣ Perform a single speculative verification step

# -------------------------------------------------

draft_token = int(np.argmax(draft_q))   # token the draft model would emit

decision = verify_one_token(target_p, draft_q, draft_token)

print(decision.accepted)          # True → we keep the draft token

print(decision.token_id)          # token actually emitted (draft or corrected)

print(decision.accept_probability)  # the acceptance probability used

# -------------------------------------------------

# 3️⃣ Compute the exact marginal distribution for a whole step

# -------------------------------------------------

marginal = speculative_output_marginal(target_p, draft_q)
print("Resulting marginal:", marginal)

Testing and Validation

The correctness of these speculative sampling primitives is validated through dedicated test suites:

Summary

  • The speculative sampling primitives in MTPLX reside entirely within mtplx/sampling.py.
  • verify_one_token at line 70 provides the main entry point for draft token verification.
  • Acceptance probability follows the $\min(1, p/q)$ formula implemented at line 44.
  • residual_distribution (line 52) handles correction sampling when draft tokens are rejected.
  • speculative_output_marginal (line 82) serves as the reference oracle for unit testing distribution correctness.
  • Helper functions like distribution_from_logits prepare probability vectors before speculative checks.

Frequently Asked Questions

Where are the speculative sampling primitives located in MTPLX?

All speculative sampling primitives are implemented in mtplx/sampling.py within the youssofal/MTPLX repository. This single module contains SpeculativeDecision, verify_one_token, acceptance_probability, and related utilities that form the mathematical foundation of the speculative decoding pipeline.

How does MTPLX calculate token acceptance probability?

MTPLX calculates acceptance probability using the acceptance_probability function at line 44 of mtplx/sampling.py. The implementation computes $\min(1, p/q)$, where $p$ is the target model's probability for the token and $q$ is the draft model's probability. This ratio determines whether the draft token can be accepted or must be rejected and resampled.

What is the purpose of the residual distribution in speculative sampling?

The residual_distribution function (line 52) constructs the correction distribution used when a draft token is rejected. It represents the remaining probability mass after removing the rejected draft token, ensuring that subsequent sampling maintains the exact statistical properties of the target distribution rather than introducing bias into the generation process.

Which utility functions prepare logits for speculative verification?

MTPLX provides distribution_from_logits (line 16) to convert raw logits into normalized probabilities, apply_top_p_top_k (line 15) for nucleus and top-k filtering, and apply_penalties (line 61) for adjusting distributions based on repetition constraints. These functions standardize probability vector preparation for both target and draft models before speculative verification occurs.

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 →