How the MTPContract Interface Works in MTPLX Backends
The MTPContract is a frozen dataclass defined in mtplx/mtp_patch.py that serves as an immutable configuration object controlling how Multi-Token Prediction (MTP) heads are injected into MLX-LM models, governing hidden state variants, quantization policies, and tensor concatenation order.
The MTPContract interface is the central configuration mechanism powering Multi-Token Prediction support in the youssofal/MTPLX framework. This immutable contract defines how MTP heads integrate with base language models, controlling everything from hidden state transformations to quantization strategies. Understanding how this interface works is essential for developers customizing inference pipelines or implementing research variants of MTP architectures.
What Is the MTPContract Interface?
The MTPContract class is defined as a frozen dataclass at lines 48-60 of mtplx/mtp_patch.py. The @dataclass(frozen=True) decorator ensures complete immutability, preventing accidental mutation of configuration parameters during the injection pipeline.
The contract exposes eleven configuration fields that control MTP head behavior:
- base_hidden_variant: Controls the trunk hidden state variant (e.g.,
"post_norm","pre_norm") - hidden_variant: Controls the final MTP output variant, supporting simple values or mix expressions like
"mix:pre_norm:post_norm:0p75" - concat_order: Determines tensor concatenation order for the MTP projection (
"embedding_hidden"or"hidden_embedding") - mtp_position_mode: Positional encoding strategy (
"cache","local", or"absolute") - mtp_quant_bits, mtp_quant_group_size, mtp_quant_mode: Quantization parameters for the MTP module
- mtp_prequantized, mtp_prequantized_modules, mtp_prequantized_module_specs: Configuration for loading pre-quantized weights
Validation and Immutability
The validate() method (lines 62-70) enforces constraint checking on all fields, raising ValueError immediately if configuration parameters contain invalid values. This guard runs early in the injection pipeline:
contract = contract or MTPContract()
contract.validate()
Because the dataclass is frozen, modification methods return new instances rather than mutating existing objects. The standard dataclass replace() method creates variations while preserving the original contract's integrity.
Configuration Merging Methods
The interface provides three key methods for populating contracts from external sources:
with_config_defaults
The with_config_defaults(self, config) method extracts defaults from configuration dictionaries (such as mtplx_mtp_quantization entries) and returns a new contract instance with those values applied.
with_metadata
The with_metadata(self, metadata, preserve_explicit=True) method merges checkpoint metadata (e.g., mtplx_mtp_contract fields) into the contract. When preserve_explicit=True, the method respects already-explicit fields rather than overwriting them.
with_runtime_metadata
The with_runtime_metadata method (lines 91-102) serves as a runtime-specific wrapper that searches for a top-level "mtp_contract" key in runtime metadata dictionaries, falling back to the entire dict if the key is absent.
Backend Integration
The inject_mtp_support function (lines 31-119 of mtplx/mtp_patch.py) consumes the contract to orchestrate MTP injection. When the loader in mtplx/runtime.py (line 55) invokes this function, it passes the contract through the following pipeline:
- Layer Detection: Determines the number of MTP layers via
_num_mtp_layers - Weight Resolution: Locates the correct MTP weight file using
expected_mtp_file - Quantization Handling: Loads and optionally quantizes weights through
_load_mtp_weightsand_quantize_mtp_module - Module Construction: Builds the internal
_MTPModuleusing contract fields to set behavior while keeping architecture static - Method Injection: Attaches helper methods (
mtp_forward,mtp_update_cache,make_mtp_cache) that forward contract parameters to_mixed_hidden
Hidden Variant Logic
The hidden_variant field drives the _mixed_hidden method (lines 62-96), which supports both simple variants and complex mix expressions. Simple variants include "fc", "pre_norm", "post_norm", "embedding", and "prev".
For research use-cases, the interface accepts mix variants formatted as mix:<left>:<right>:<alpha>, enabling linear blending of two hidden sources. The helper _valid_mtp_hidden_variant (lines 5-20) handles parsing and validation of these expressions.
The concat_order field determines how the embedding and hidden tensors combine before the MTP linear projection, accepting either "embedding_hidden" or "hidden_embedding".
Pre-Quantized Support
When mtp_prequantized=True, the contract triggers an optimized loading path that avoids dequantize-requantize cycles. The system:
- Derives tensor geometry via
_contract_with_prequantized_tensor_geometry - Collects module specifications through
_contract_with_prequantized_module_specs - Applies selective quantization in
_quantize_mtp_moduleusing custom predicates based onmtp_prequantized_modulesandmtp_prequantized_module_specs
This pathway preserves accuracy and reduces load time for weights stored in quantized formats.
Practical Implementation Examples
Below are concrete patterns for interacting with the MTPContract interface in production code:
# Basic default injection
from mtplx.mtp_patch import MTPContract, inject_mtp_support
contract = MTPContract()
inject_mtp_support(model, "model_name", config={}, contract=contract)
# Custom hidden variant with quantization
custom_contract = (
MTPContract()
.with_metadata({"hidden_variant": "mix:pre_norm:post_norm:0p75"})
.replace(mtp_quant_bits=4, mtp_quant_policy="cyankiwi")
)
inject_mtp_support(model, "model_name", config={}, contract=custom_contract)
# Pre-quantized weight loading
prequantized_contract = MTPContract(
mtp_prequantized=True,
mtp_prequantized_modules=("layers.0.fc", "layers.1.fc"),
mtp_prequantized_module_specs={
"layers.0.fc": {"bits": 4, "group_size": 64, "mode": "affine"},
},
)
inject_mtp_support(model, "model_name", config={}, contract=prequantized_contract)
Summary
- The
MTPContractinterface is an immutable configuration object defined inmtplx/mtp_patch.pythat controls MTP head injection. - Eleven fields govern hidden state variants, concatenation order, positional modes, and quantization strategies.
- The
validate()method enforces configuration constraints at lines 62-70, ensuring runtime stability. - Immutable design allows safe modification via
replace(),with_config_defaults(), andwith_metadata()methods. - Backends consume the contract through
inject_mtp_supportto determine layer counts, load weights, and construct_MTPModuleinstances. - Pre-quantized support avoids accuracy degradation by loading quantized weights directly without reconversion.
Frequently Asked Questions
What is the difference between base_hidden_variant and hidden_variant?
The base_hidden_variant field controls which hidden state representation the model trunk produces before MTP processing, while hidden_variant determines the final output representation of the MTP head itself. The base variant typically uses "post_norm" or "pre_norm", whereas the hidden variant supports additional options like "fc", "embedding", or mix expressions for research flexibility.
How do I use pre-quantized weights with the MTPContract?
Set mtp_prequantized=True and specify mtp_prequantized_modules as a tuple of module prefixes (e.g., ("layers.0.fc",)). Provide detailed specifications via mtp_prequantized_module_specs, mapping each module to its quantization parameters including bits, group_size, and mode. This configuration bypasses dequantization-requantization cycles, preserving original weight precision and reducing initialization time.
Can I modify a contract after creating it?
No, MTPContract is a frozen dataclass, making all instances immutable by design. To alter configuration, use the replace() method to generate a new instance with updated fields, or utilize helper methods like with_metadata() and with_config_defaults() which return new contract objects rather than mutating existing ones.
How does the concat_order parameter affect model behavior?
The concat_order parameter determines whether the embedding tensor precedes or follows the hidden tensor during the concatenation operation before the MTP linear projection. Setting it to "embedding_hidden" places embeddings first, while "hidden_embedding" reverses the order. This affects how the model combines semantic and positional information during multi-token prediction.
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 →