# How Vision Models Work with MTP in MTPLX: MROPE and Vision Splice Explained

> Discover how vision models in MTPLX leverage MTP, MROPE, and Vision Splice for seamless image and text token alignment in draft generation. Learn about embedding injection and 3-D positional encoding.

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

---

**Vision models in MTPLX integrate with Multi‑Token Prediction (MTP) through VisionSplice for embedding injection and MROPE for 3‑D positional encoding, ensuring image tokens align with text tokens during draft generation.**

MTPLX extends multimodal capabilities to speculative decoding by combining a vision tower with Multi‑Token Prediction (MTP). The system treats image tokens as first‑class citizens during generation through two core mechanisms: **VisionSplice** for streaming embeddings and **MROPE** (Multi‑axis RoPE) for spatial position encoding. These components work together in [`mtplx/vision/splice.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/splice.py) and [`mtplx/vision/mrope.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/mrope.py) to maintain alignment between the main trunk and the MTP draft model.

## Vision Splice: Streaming Vision Embeddings into MTP

The `VisionSplice` class in [`mtplx/vision/splice.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/splice.py) manages per‑request vision state, ensuring the MTP draft model consumes identical embeddings to the main trunk during both prefill and generation phases.

### Construction and State Management

When a request contains images, the backend instantiates `VisionSplice` with:
- `image_pad_token_id`: The placeholder token ID used in the prompt to reserve space for image features
- `embeddings`: An `mx.array` of shape `[total_pad_tokens, hidden]` containing pre‑computed vision features from the vision tower
- `image_grids`: Raw `(t, h, w)` tuples defining the temporal and spatial dimensions of each image
- `mrope_table` and `mrope_delta`: Position tables built by `mtplx.vision.mrope.build_mrope_positions`

### Embedding Splice During Prefill

The `spliced_chunk_embeddings` function handles the actual injection during the prefill phase:
- Scans input IDs for pad tokens (`ids == image_pad_token_id`)
- Replaces corresponding rows in the token‑embedding matrix with rows from `splice.embeddings`
- Advances `splice.cursor` to maintain strict alignment between pad tokens and vision features
- Raises `ValueError` if the request supplies fewer embedding rows than pad tokens, preventing silent misalignment

### Window Splice for MTP History Alignment

For subsequent windows (post‑prefill), `spliced_embeddings_for_window` reads specific rows via the `rows_before` argument without advancing the global cursor. This ensures the **MTP history alignment** sees the exact same visual embeddings that the trunk consumed, which is critical for maintaining coherent attention patterns across draft steps.

## MROPE: 3‑Dimensional Position Encoding for Vision Tokens

Standard Rotary Position Embeddings (RoPE) use 1‑D scalar offsets, but vision models like Qwen‑VL require a **3‑D grid** to represent image patches across temporal and spatial axes. [`mtplx/vision/mrope.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/mrope.py) implements this through the `build_mrope_positions` function.

### Building the Multi‑Axis Position Table

The function executes the following steps:
1. Scans `input_ids` for occurrences of the image pad token
2. Validates that the pad block size matches the grid dimensions `(t × h × w)` for each image
3. Generates a `(3, len)` integer table where each column represents `(t_idx, h_idx, w_idx)` coordinates
4. Computes a `delta` value that aligns the vision token positions with the scalar rope used for text tokens

### Delta Injection in MTP Attention

The `delta` is passed as `position_offset` to the MTP attention layer via `attn.rope(..., offset=position_offset)` in `_mtp_full_attention_layer`. This guarantees that vision tokens use the same rope function as text tokens after the image sequence, maintaining positional coherence during speculative decoding. Because the table is recomputed per request rather than cached, warm‑starts remain deterministic and pure functions of the prompt.

## MTP Injection Architecture

`inject_mtp_support` in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py) attaches draft modules to a loaded `mlx‑lm` model while preserving multimodal capabilities through the splice mechanism.

### Draft Module Construction

The injection process follows these steps:
1. **Load MTP weights** from side‑car `.safetensors` files specified by `expected_mtp_file` or from embedded checkpoints
2. **Quantize layers** according to the `MTPContract` (specifying `mtp_quant_bits` and `mtp_quant_policy`)
3. **Create `_MTPModule`** containing:
   - Two RMSNorm layers (`pre_fc_norm_hidden`, `pre_fc_norm_embedding`)
   - A linear projection concatenating vision embeddings (`e`) with hidden states (`h`)
   - A stack of `DecoderLayer` copies (`n_layers`) for draft processing

### Forward Pass Integration

The patched `_MTPLXTextModel` class:
- Calls the original trunk for prefill, which consumes the `VisionSplice`
- Stores the splice instance for the MTP core (`_mtp_core`) to ensure draft access to vision features
- Invokes `_mtp_full_attention_layer` with the MROPE delta as the `position_offset` parameter

During generation, `mtp_forward` and `mtp_update_cache` consume the same vision rows via the splice and apply the MROPE delta, producing draft hidden states that optionally mix with trunk states via `mix:<left>:<right>:<alpha>` parameters before emitting logits.

## End‑to‑End Implementation Examples

### Building a Vision Splice for Requests

```python
from mtplx.vision.splice import VisionSplice
from mtplx.vision.mrope import build_mrope_positions
import mlx.core as mx

PAD_ID = 99
embeds = mx.zeros((4, 2048))  # 4 image pad tokens, hidden dim 2048

splice = VisionSplice(
    image_pad_token_id=PAD_ID,
    embeddings=embeds,
    image_grids=[(1, 2, 2)],  # (temporal, height, width)

    image_digests=(0xdeadbeef,),
    pad_counts=(4,),
)

# Build MROPE table and delta

ids = [101, 102, PAD_ID, PAD_ID, PAD_ID, PAD_ID, 103]
table, delta = build_mrope_positions(
    ids,
    image_token_id=PAD_ID,
    image_grids=[(1, 2, 2)],
    spatial_merge_size=1,
)
splice.mrope_table = mx.array(table)
splice.mrope_delta = delta

```

### Splicing Embeddings During Prefill

```python
from mtplx.vision.splice import spliced_chunk_embeddings

def embed_tokens(ids):
    return model.embed_tokens(ids)

chunk = mx.array([101, 102, PAD_ID, PAD_ID, PAD_ID, PAD_ID])
chunk_embeds = spliced_chunk_embeddings(embed_tokens, chunk, splice)

# chunk_embeds now contains vision features in place of pad tokens

```

### Injecting MTP Support into the Model

```python
from mtplx.mtp_patch import inject_mtp_support, MTPContract

contract = MTPContract(
    hidden_variant="post_norm",
    mtp_position_mode="cache",
    mtp_quant_bits=4,
    mtp_quant_policy="all",
)
inject_mtp_support(model, model_path, config={}, contract=contract)

```

### Running MTP Generation with Vision Alignment

```python

# Initialize MTP cache

mtp_cache = model.language_model.make_mtp_cache()

logits, hidden = model.language_model.mtp_forward(
    hidden_states=initial_hidden,
    next_token_ids=mx.array([104]),
    mtp_cache=mtp_cache,
    concat_order="embedding_hidden",
    mtp_hidden_variant="post_norm",
    position_offset=splice.mrope_delta,  # Critical: aligns rope with vision

    vision_splice=splice,                # Provides same embeddings as trunk

)

```

## Summary

- **VisionSplice** manages per‑request image embeddings in [`mtplx/vision/splice.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/splice.py), injecting vision features into prefill chunks while maintaining cursor alignment between trunk and draft paths.
- **MROPE** in [`mtplx/vision/mrope.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/mrope.py) generates 3‑D position tables `[t, h, w]` and a positional delta that aligns vision tokens with the scalar rope used for text generation.
- **MTP injection** via `inject_mtp_support` in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py) creates draft modules that consume the same splice and delta, ensuring the MTP draft attends to images identically to the main model.
- The key invariant requires that `spliced_embeddings` and `position_offset` (MROPE delta) remain consistent between trunk prefill and MTP forward passes to prevent positional drift during multimodal speculative decoding.

## Frequently Asked Questions

### How does MTPLX prevent position misalignment between the trunk and MTP draft?

MTPLX computes a **MROPE delta** via `build_mrope_positions` in [`mtplx/vision/mrope.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/mrope.py), which is passed as `position_offset` to the MTP attention layers during `mtp_forward`. This delta aligns the 3‑D vision positions with the 1‑D scalar rope, ensuring both trunk and draft apply identical rotary embeddings to image tokens.

### What happens if the number of image embeddings does not match the pad token count?

The `spliced_chunk_embeddings` function in [`mtplx/vision/splice.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision/splice.py) validates alignment and raises a `ValueError` if the splice cursor exhausts the available embeddings before all pad tokens are filled. This prevents silent token misalignment that would corrupt the attention mechanism.

### Can VisionSplice handle multiple images in a single request?

Yes. The `VisionSplice` class accepts lists of `image_grids` and `pad_counts`, allowing it to manage multiple images by concatenating their embeddings and tracking independent `(t, h, w)` grids for MROPE table generation.

### Where are the MTP weights loaded from in a multimodal setup?

`inject_mtp_support` searches for side‑car `.safetensors` files specified by `expected_mtp_file` or embedded checkpoints within the model directory. These weights are then quantized and loaded into the `_MTPModule`, which operates alongside the vision tower restored via [`mtplx/vision_graft.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/vision_graft.py).