# How MTPLX Verifies Drafted Tokens Using Batched Forward Passes

> Learn how MTPLX verifies drafted tokens efficiently. Discover its single batched forward pass for draft and primary token evaluation within one decode cycle.

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

---

**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](https://github.com/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)](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](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L1024-L1040).

The attention context management in [`mtplx/attention_context.py`](https://github.com/youssofal/MTPLX/blob/main/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)](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L35-L52):

```python
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](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py#L51-L55). 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`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/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

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

```python
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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/batched_decode.py) | Core implementation of batched MTP decoding, draft sampling, and verification logic |
| [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) | Provides `verify_one_token`, distribution utilities, and penalty handling |
| [`mtplx/attention_context.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/attention_context.py) | Supplies the `attention_phase` context manager for batched verification forwards |
| [`mtplx/cache_state.py`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/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)](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.