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 featuresembeddings: Anmx.arrayof shape[total_pad_tokens, hidden]containing pre‑computed vision features from the vision towerimage_grids: Raw(t, h, w)tuples defining the temporal and spatial dimensions of each imagemrope_tableandmrope_delta: Position tables built bymtplx.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.cursorto maintain strict alignment between pad tokens and vision features - Raises
ValueErrorif 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:
- Scans
input_idsfor occurrences of the image pad token - Validates that the pad block size matches the grid dimensions
(t × h × w)for each image - Generates a
(3, len)integer table where each column represents(t_idx, h_idx, w_idx)coordinates - Computes a
deltavalue 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:
- Load MTP weights from side‑car
.safetensorsfiles specified byexpected_mtp_fileor from embedded checkpoints - Quantize layers according to the
MTPContract(specifyingmtp_quant_bitsandmtp_quant_policy) - Create
_MTPModulecontaining:- 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
DecoderLayercopies (n_layers) for draft processing
- Two RMSNorm layers (
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_layerwith the MROPE delta as theposition_offsetparameter
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.pygenerates 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_supportinmtplx/mtp_patch.pycreates 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_embeddingsandposition_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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →