How MTPLX's Dedicated Injector Enables Step-3.5/3.7-Flash MTP Support
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 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 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.safetensorsfiles - Falls back to parsing
model.safetensors.index.jsonfor sharded checkpoints - Finally globs any
model*.safetensorsfiles 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, andshared_head_norminstances - Projection layers: The
eh_projlinear transformation - Decoder blocks: Full
Step3p5DecoderLayerinstances (mtp_block) - Output heads: The
shared_head_headprojection
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 passmtp_update_cache: Updates the MTP key-value cachemake_mtp_cache: Initializes cache structures for MTP layersmake_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:
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 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.pyinstead 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_shiftto 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.
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 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 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.
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 →