How MTPLX Verifies Drafted Tokens Using Batched Forward Passes

MTPLX verifies drafted tokens by generating draft predictions per stream, executing a single batched forward pass that evaluates both primary and draft tokens together, then accepting or repairing based on a target comparison—all within one decode cycle.

MTPLX implements Multi-Token-Parallel (MTP) decoding, a technique that accelerates autoregressive generation by predicting multiple tokens simultaneously. Central to this approach is draft verification, a three-stage pipeline that validates speculative predictions using efficient batched forward passes. This article examines the complete verification mechanism as implemented in the youssofal/MTPLX repository.

The Three-Stage Verification Pipeline

MTPLX structures draft verification as a repeating cycle executed every decode step. Understanding each stage reveals how the system maintains correctness while maximizing throughput.

Stage 1: Draft Generation

For every active stream, MTPLX first generates two outputs: a primary token (sampled from the main model head) and a draft token (sampled from the MTP head).

The function _sample_mtp_k1_row_proposal → _sample_mtp_k1_draft in [mtplx/batched_decode.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L66-L87) handles this sampling. The draft token represents the model's speculative prediction for what follows the primary token.

Stage 2: Batched Verification Forward

The critical efficiency gain comes from MTPLX's single forward pass verification. Rather than processing streams sequentially, the system constructs a [B, 2] input tensor—where B is the batch size—and evaluates both the primary token and draft token for all streams in one operation.

This batched forward occurs inside the driver loop (_run_foldin_loop or the regular speculative loop). The verification logits are returned as v_logits, and _eval_bundle reads the decision tensor with shape [x0, draft, x1, accept] using a single host-CPU synchronization point source.

The attention context management in mtplx/attention_context.py wraps this forward pass, ensuring proper KV-cache handling across the batch.

Stage 3: Accept/Repair Decision

After the verification forward completes, MTPLX compares the draft against the target token (x1) produced by that forward pass. This logic resides in _finish_mtp_k1_row_cycle, which calls verify_one_token from the sampling module in [mtplx/batched_decode.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L35-L52):

if sampler.temperature <= 0:
    target = _greedy_token_from_logits(verify_logits, sampler, token_counts=counts)
    accepted = draft == target
    second = draft if accepted else target
else:
    target_p = distribution_from_logits(np.asarray(verify_logits, dtype=np.float64),
                                       sampler, token_counts=counts)
    decision = verify_one_token(target_p, draft_q, draft, rng)
    accepted = bool(decision.accepted)
    second = int(decision.token_id)

The decision branches on two outcomes:

  • Accepted draft: When draft == target, the stream commits both tokens in one cycle, achieving 2× effective throughput.
  • Rejected draft: When the draft mismatches, the repair path activates. The entire batch rolls back to the previous valid state and re-forwards the correct [x0, x1] pair.

Notably, MTPLX uses uniform full-batch repair: when any stream rejects its draft, all streams participate in the repair forward source. This simplifies scheduling at the cost of occasional redundant computation.

Greedy vs. Stochastic Verification

MTPLX supports both deterministic and probabilistic verification modes, with different comparison logic in each.

Mode Verification Method Implementation Location
Greedy Direct token equality: draft == target _greedy_token_from_logits in sampling pipeline
Stochastic Acceptance sampling via verify_one_token mtplx/sampling.py

Stochastic mode uses the draft's proposal distribution (draft_q) and the verification distribution (target_p) to compute an acceptance probability, enabling temperature-dependent speculative decoding.

Cache State Management for Repair

The repair mechanism depends on mtplx/cache_state.py for snapshot and restore operations. Before each verification attempt, the system captures the current KV-cache state. If repair triggers, this snapshot restores the cache to its pre-speculation condition, ensuring coherent subsequent generation.

Practical Usage Examples

High-Level Batched Generation

from mtplx.batched_decode import generate_greedy_batched

# Assume `rt` is an MTP-enabled runtime (e.g., a Gemma-4 model)

prompts = [[1, 5, 23], [7, 12, 4]]  # token-id lists

result = generate_greedy_batched(
    rt,
    prompts,
    max_new_tokens=20,
    use_mtp_draft=True,          # enable draft generation

    reject_mode="repair",        # full-batch repair on mismatch

    serial=False,                # pipelined batched driver

)

print("Generated tokens per stream:", result.streams[0].tokens)
print("Total forward passes:", result.forwards)
print("Repair cycles (draft mismatches):", result.repair_cycles)

Low-Level Verification Step

from mtplx.batched_decode import _sample_mtp_k1_row_cycle

# `primary_logits`, `draft_logits`, `verify_logits`, `bonus_logits` are numpy arrays

row_cycle = _sample_mtp_k1_row_cycle(
    primary_logits,
    draft_logits,
    verify_logits,
    bonus_logits,
    sampler=visible_sampler,
    draft_sampler=draft_sampler,
    rng=np.random.default_rng(),
    history_tokens=history,
    omit_speculative_bonus=False,
)

print("Draft accepted?", row_cycle.accepted)
print("Second token (target or draft):", row_cycle.second_token)

Key Source Files

File Purpose
mtplx/batched_decode.py Core implementation of batched MTP decoding, draft sampling, and verification logic
mtplx/sampling.py Provides verify_one_token, distribution utilities, and penalty handling
mtplx/attention_context.py Supplies the attention_phase context manager for batched verification forwards
mtplx/cache_state.py Handles cache snapshot/restore for the repair path

Summary

  • MTPLX generates draft tokens per-stream using _sample_mtp_k1_draft in the batched decode pipeline
  • One batched forward pass processes primary and draft tokens together for all streams, producing verification logits
  • Greedy verification checks direct equality; stochastic verification uses verify_one_token with distribution-aware acceptance sampling
  • Accepted drafts commit two tokens per cycle; rejected drafts trigger uniform full-batch repair via cache rollback
  • The repair mechanism relies on mtplx/cache_state.py snapshots and _run_foldin_loop coordination

Frequently Asked Questions

How does MTPLX handle batch heterogeneity during verification?

MTPLX processes all streams through the same verification forward regardless of individual draft validity. When any stream rejects its draft, the entire batch enters repair mode. This uniform approach simplifies implementation and GPU scheduling, though it may perform extra computation for streams that accepted their drafts.

What determines whether a draft token gets accepted?

For greedy decoding, acceptance requires exact match between draft and target tokens. For stochastic decoding, verify_one_token computes an acceptance probability based on the ratio of target distribution probability to draft distribution probability at the proposed token, enabling temperature-aware speculative sampling.

Why use a single forward pass for verification instead of two sequential passes?

A single [B, 2] forward pass maximizes GPU utilization and minimizes memory bandwidth overhead. Processing primary and draft tokens together allows MTPLX to evaluate all speculative predictions in one kernel launch, reducing per-step latency compared to separate forwards.

Where does the actual token comparison happen in the codebase?

The final comparison occurs in _finish_mtp_k1_row_cycle within [mtplx/batched_decode.py](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L35-L52). This function coordinates with verify_one_token (stochastic mode) or direct equality checks (greedy mode) to produce the accepted boolean that drives subsequent repair decisions.

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 →