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:

  1. Layer Detection: Determines the number of MTP layers via _num_mtp_layers
  2. Weight Resolution: Locates the correct MTP weight file using expected_mtp_file
  3. Quantization Handling: Loads and optionally quantizes weights through _load_mtp_weights and _quantize_mtp_module
  4. Module Construction: Builds the internal _MTPModule using contract fields to set behavior while keeping architecture static
  5. 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:

  1. Derives tensor geometry via _contract_with_prequantized_tensor_geometry
  2. Collects module specifications through _contract_with_prequantized_module_specs
  3. Applies selective quantization in _quantize_mtp_module using custom predicates based on mtp_prequantized_modules and mtp_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 MTPContract interface is an immutable configuration object defined in mtplx/mtp_patch.py that 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(), and with_metadata() methods.
  • Backends consume the contract through inject_mtp_support to determine layer counts, load weights, and construct _MTPModule instances.
  • 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:

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 →