# How MTPLX Handles Residual Correction in MTP Speculative Decoding

> Discover how MTPLX handles residual correction in MTP speculative decoding. Learn how it computes residual distributions for mathematically sound fallback sampling when draft tokens are rejected.

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

---

**MTPLX handles residual correction by computing a residual distribution that captures the probability mass of the target model not covered by the draft model, enabling mathematically sound fallback sampling when speculative tokens are rejected during MTP verification.**

The open-source MTPLX library implements Multi-Token Prediction (MTP) speculative decoding to accelerate large language model inference while maintaining statistical fidelity. When the draft model generates tokens that fail verification against the target model, MTPLX employs a residual correction mechanism to ensure probability conservation. This article examines the implementation details found in the `youssofal/MTPLX` repository, focusing on how residual distributions are calculated and integrated into the speculative decoding pipeline.

## Understanding MTP Speculative Decoding in MTPLX

MTPLX’s speculative decoding architecture uses a fast draft model to predict multiple tokens ahead, which are then verified in parallel by the slower target model. The process relies on two distinct probability distributions: the **draft distribution** generated by the lightweight draft model, and the **target distribution** produced by the full-scale target model.

When verification succeeds, the draft tokens are accepted and the inference accelerates significantly. However, when verification rejects a draft token, the system cannot simply discard the computation. Instead, it must sample a replacement token from the portion of the target distribution not explained by the draft. This is handled by the **residual correction** logic implemented in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) and integrated through [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py).

## Computing the Residual Distribution

At the heart of MTPLX’s residual correction lies the `residual_distribution` function in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py). This function computes the residual probability vector that represents the difference between the target model’s distribution and the draft model’s distribution for rejected tokens.

### The Mathematical Foundation

The residual calculation follows a precise algebraic formulation. Given the target probability vector `target` and the draft probability vector `draft`, MTPLX computes the residual as follows:

```python
residual = target - draft * mask
residual = max(residual, 0)          # clip negative values

residual = residual / residual.sum() # renormalise

```

The `mask` parameter zeroes out entries already selected by the draft, ensuring the residual contains only probability mass for tokens not present in the draft proposal. This guarantees that the residual distribution and the draft distribution together sum to the full target distribution, preserving statistical validity.

### Handling Numerical Edge Cases

Numerical stability is protected through explicit clipping and renormalization steps found in [`mtplx/mtp_batch_numerics.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_batch_numerics.py). If the draft distribution exhausts all probability mass (i.e., `target` ≤ `draft` for every token), the residual becomes a zero vector. In this scenario, MTPLX defaults to the target distribution directly to prevent division-by-zero errors.

The test suite in [`tests/test_sampling.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py) validates this behavior specifically, confirming that `residual_distribution` returns a proper probability vector even when handling edge cases like complete draft coverage or floating-point precision issues.

## Verification and Fallback Logic

The integration of residual correction with the speculative decoding loop occurs in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py). After each speculative step, MTPLX invokes `verify_one_token` to check draft validity against the target model.

If verification accepts the draft token, the token is emitted and the draft distribution proceeds to the next step. If verification rejects the token, MTPLX falls back to the residual distribution computed by `residual_distribution`. The rejected position is then sampled from this residual vector, ensuring the final token distribution respects the target model’s true probabilities without requiring a full forward pass recomputation.

This mechanism enables **smooth rollback**: when speculative tokens are invalid, MTPLX switches instantly to a mathematically sound distribution without re-running the expensive target model for the entire sequence.

## Practical Implementation Example

Below is a minimal illustration demonstrating how the residual correction logic operates within a custom decoding loop. While the actual MTPLX runtime performs these steps automatically via [`mtplx/runtime.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/runtime.py), this example shows the underlying mechanics:

```python
from mtplx.sampling import residual_distribution
from mtplx.runtime import MTPLXRuntime
from mtplx.mtp_patch import MTPContract

# Initialize a runtime with MTP enabled

runtime = MTPLXRuntime(contract=MTPContract())
session = runtime.new_session()

# Obtain draft and target logits for the next token

draft_logits, target_logits = session.get_logits(draft=True, verify=True)

# Convert logits to probabilities

draft_probs = draft_logits.softmax(dim=-1)
target_probs = target_logits.softmax(dim=-1)

# Compute the residual distribution

residual = residual_distribution(target_probs, draft_probs)

# Sample a token: if verification fails, sample from residual

if session.verify_one_token():
    token = draft_probs.sample()
else:
    token = residual.sample()

```

This pattern ensures **deterministic behavior** for testing while maintaining **probability conservation** across the speculative decoding pipeline.

## Summary

- **Residual correction** in MTPLX computes the difference between target and draft distributions to handle rejected speculative tokens.
- The `residual_distribution` function in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) implements the core calculation, clipping negative values and renormalizing to ensure valid probability distributions.
- Verification logic in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py) uses `verify_one_token` to decide when to fallback to residual sampling.
- Edge cases where the draft covers all probability mass are handled by defaulting to the target distribution, validated in [`tests/test_sampling.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sampling.py).
- The mechanism ensures statistical fidelity without requiring full recomputation of the target model when draft tokens are rejected.

## Frequently Asked Questions

### What is the purpose of residual correction in MTPLX speculative decoding?

Residual correction ensures that when the draft model proposes a token that the target model rejects, the system can still sample a valid replacement from the remaining probability mass. This preserves the exact statistical properties of the target model while allowing the speed benefits of speculative decoding.

### How does MTPLX prevent numerical instability in residual calculations?

MTPLX prevents numerical instability by clipping negative values to zero and renormalizing the residual vector. If the draft distribution covers all probability mass (resulting in a zero residual), the system defaults to the target distribution directly. These safeguards are implemented in [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) with numeric utilities from [`mtplx/mtp_batch_numerics.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_batch_numerics.py).

### Where is the residual correction logic integrated in the MTPLX pipeline?

The residual correction logic is integrated in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py), specifically within the verification loop that calls `verify_one_token`. When a token fails verification, the system invokes `residual_distribution` from [`mtplx/sampling.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/sampling.py) to compute the fallback distribution, orchestrated by the main runtime in [`mtplx/runtime.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/runtime.py).

### How does MTPLX ensure probability conservation when using residual correction?

MTPLX ensures probability conservation by constructing the residual distribution such that the sum of draft probabilities and residual probabilities equals the target distribution. The mathematical formulation `residual = target - draft * mask` guarantees that no probability mass is lost or created during the speculative verification process.