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

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 and 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 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 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 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

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

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

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


# 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, injecting vision features into prefill chunks while maintaining cursor alignment between trunk and draft paths.
  • MROPE in 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 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, 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 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.

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 →