# How MTPLX's Dedicated Injector Enables Step-3.5/3.7-Flash MTP Support

> Discover how MTPLX's dedicated injector enhances Step-3.5/3.7-Flash MTP support. Learn how it exposes pre-norm hidden states and corrects RMS norm weights for improved performance.

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

---

**MTPLX's dedicated injector for Step-3.5/3.7-Flash models replaces the generic DeepSeek-style MTP path with a specialized implementation in [`mtplx/step3p5_mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/step3p5_mtp_patch.py) that exposes pre-norm hidden states and corrects zero-centered RMS norm weights by applying a +1.0 shift at load time.**

The youssofal/MTPLX repository provides native Multi-Token Prediction (MTP) support for Step-3.5 and Step-3.7-Flash architectures through a dedicated injection system. Unlike the generic MTP implementation used for DeepSeek-style models, this injector handles two strict correctness requirements specific to the Step-Flash family: capturing hidden states before final RMS normalization and compensating for zero-centered norm weights.

## Why Step-3.5/3.7-Flash Requires a Dedicated Injector

Step-Flash models deviate from standard architectures in ways that break generic MTP implementations. The injector addresses these incompatibilities through targeted patches.

### The Pre-Norm Hidden State Requirement

Standard implementations return post-norm hidden states from the forward pass, but Step-Flash MTP layers require access to the **pre-norm hidden state** to compute predictions correctly. The injector overrides the original `Step3p5Model.__call__` method to expose both pre-norm and post-norm representations simultaneously.

### Zero-Centered RMS Norm Correction

Step-Flash uses `ZeroCenteredRMSNorm` layers where weights are centered around zero. During weight loading, the injector's `_apply_zero_centered_norm_shift` function detects these weights by checking if their mean is below 0.5, then applies a **+1.0 shift** to align them with standard RMS norm expectations. This mirrors the sanitization behavior found in `mlx-lm` but applies it specifically to the MTP submodule.

## The Injection Pipeline: From Config Detection to Model Patching

The injector follows a strict 11-stage pipeline defined in [`mtplx/step3p5_mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/step3p5_mtp_patch.py) to safely graft MTP capabilities onto existing Step-3.5 models.

### Configuration Detection and Validation

The process begins with `is_step3p5_mtp_config`, which inspects the model configuration for a valid `model_type` (`step3p5` or `step3p7`) and confirms that `num_nextn_predict_layers` is a positive integer. This prevents accidental injection into incompatible model architectures.

### Weight File Discovery

The `_candidate_weight_files` function implements a hierarchical search strategy:
- First, it searches for explicit `*.mtp.safetensors` files
- Falls back to parsing [`model.safetensors.index.json`](https://github.com/youssofal/MTPLX/blob/main/model.safetensors.index.json) for sharded checkpoints
- Finally globs any `model*.safetensors` files as a last resort

Once identified, `_load_raw_weights` loads these files using `mlx.core.mx.load` and aggregates them into a unified dictionary.

### Key Filtering and Rewriting

Not all weights in a Step-3.5 checkpoint belong to the MTP layers. The `_is_step_mtp_key` filter retains only keys prefixed with `mtp.` or those falling within the target layer range. Subsequently, `_rewrite_step_mtp_weights` transforms the original checkpoint keys into the internal `_StepMTP` module layout—for example, mapping `model.layers.{i}.enorm.weight` to `layers.{j}.enorm.weight`—while stripping shared embedding keys that would otherwise create conflicts.

### Module Construction and Weight Injection

The `_make_step_mtp_module` function constructs a lightweight `_StepMTP` hierarchy containing:
- **Zero-centered norms**: `enorm`, `hnorm`, and `shared_head_norm` instances
- **Projection layers**: The `eh_proj` linear transformation
- **Decoder blocks**: Full `Step3p5DecoderLayer` instances (`mtp_block`)
- **Output heads**: The `shared_head_head` projection

After construction, the rewritten weights are loaded into the module using `mtp.load_weights(..., strict=True)`, and `_validate_mtp_load_coverage` verifies that every parameter in the module tree received a corresponding weight.

### Model Class Patching and API Exposure

The injector dynamically subclasses the original model class as `_MTPLXStepModel`, overriding `__call__` to return both pre-norm and post-norm hidden states. It attaches four new methods to the model instance:
- `mtp_forward`: Executes the MTP block forward pass
- `mtp_update_cache`: Updates the MTP key-value cache
- `make_mtp_cache`: Initializes cache structures for MTP layers
- `make_cache`: Standard cache creation (delegated)

Finally, the injector stores the MTP module as `model.mtp` and sets auxiliary attributes (`_mtplx_hidden_variant`, `_mtplx_concat_order`) to enable transparent integration with the MTPLX generation pipeline.

## Implementation Example

The following example demonstrates loading a Step-3.5 checkpoint, injecting MTP support, and running inference:

```python
from pathlib import Path
import json
import mlx_lm
import mtplx

# Load base model

model_path = Path("/path/to/step3p5-flash-checkpoint")
config = json.loads((model_path / "config.json").read_text())
model = mlx_lm.load_model(model_path)

# Inject MTP support

injected = mtplx.inject_step3p5_mtp_support(
    model=model,
    model_path=model_path,
    config=config,
    contract=None,  # Optional: specify quantization bits here

)

# Create MTP cache and run forward pass

mtp_cache = model.make_mtp_cache()
logits, hidden = model.mtp_forward(
    hidden_states=model.embed_tokens([101]),
    next_token_ids=[102],
    mtp_cache=mtp_cache,
    return_hidden=True,
)

# Update cache for next prediction

new_hidden = model.mtp_update_cache(
    hidden_states=hidden,
    next_token_ids=[103],
    mtp_cache=mtp_cache,
)

```

## Optional Quantization Support

If the runtime contract specifies `mtp_quant_bits`, the injector invokes `_quantize_mtp_module` from [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py) to apply post-loading quantization to the MTP weights. This allows memory-efficient inference while maintaining the specialized Step-Flash corrections already applied during weight loading.

## Summary

- **Dedicated path**: Step-3.5/3.7-Flash models require [`mtplx/step3p5_mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/step3p5_mtp_patch.py) instead of the generic DeepSeek MTP implementation due to architectural constraints.
- **Pre-norm exposure**: The injector overrides `__call__` to provide hidden states before RMS normalization, satisfying Step-Flash's strict correctness requirements.
- **Weight correction**: Zero-centered RMS norm weights receive a +1.0 shift during loading via `_apply_zero_centered_norm_shift` to ensure numerical stability.
- **Hierarchical loading**: The system searches for `*.mtp.safetensors`, falls back to index files, and validates complete parameter coverage before finalizing injection.
- **Transparent integration**: Once injected, the model exposes `mtp_forward`, `mtp_update_cache`, and related methods that work seamlessly with MTPLX's generation pipeline.

## Frequently Asked Questions

### What makes Step-3.5/3.7-Flash MTP different from standard DeepSeek MTP?

Step-Flash architectures require access to hidden states **before** the final RMS normalization layer, whereas standard implementations only expose post-norm states. Additionally, Step-Flash uses zero-centered RMS norm weights that must be shifted by +1.0 during loading to match the expected distribution of standard RMS norms. The generic DeepSeek MTP path cannot satisfy these requirements, necessitating the dedicated injector in [`mtplx/step3p5_mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/step3p5_mtp_patch.py).

### How does the injector handle weight loading for sharded checkpoints?

The `_candidate_weight_files` function implements a three-tier discovery mechanism: it first looks for dedicated `*.mtp.safetensors` files, then parses [`model.safetensors.index.json`](https://github.com/youssofal/MTPLX/blob/main/model.safetensors.index.json) to locate sharded weights, and finally falls back to globbing all `model*.safetensors` files. The `_load_raw_weights` function then aggregates these using `mlx.core.mx.load` into a single dictionary for processing.

### Can I use quantization with the Step-3.5 MTP injector?

Yes. Pass a `contract` parameter with `mtp_quant_bits` specified to `inject_step3p5_mtp_support`. When provided, the injector calls `_quantize_mtp_module` from [`mtplx/mtp_patch.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_patch.py) to quantize the MTP weights after they have been loaded and corrected for zero-centered norms. This occurs after the +1.0 shift has been applied, preserving numerical correctness.

### What happens if the weight loading validation fails?

The injector uses `strict=True` when calling `mtp.load_weights()`, which raises an error if any parameter in the `_StepMTP` module tree lacks a corresponding weight. Additionally, `_validate_mtp_load_coverage` performs explicit validation against the module's parameter tree to ensure complete coverage before the injection is marked as successful.