# How the MTPContract Interface Works in MTPLX Backends

> Discover how the MTPContract interface in MTPLX backends controls Multi-Token Prediction heads with immutable configurations for hidden states, quantization, and tensor order.

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

---

**The `MTPContract` is a frozen dataclass defined in [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/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:

```python
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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py)) consumes the contract to orchestrate MTP injection. When the loader in [`mtplx/runtime.py`](https://github.com/youssofal/MTPLX/blob/main/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:

```python

# Basic default injection

from mtplx.mtp_patch import MTPContract, inject_mtp_support

contract = MTPContract()
inject_mtp_support(model, "model_name", config={}, contract=contract)

```

```python

# 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)

```

```python

# 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`](https://github.com/youssofal/MTPLX/blob/main/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.