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.safetensors files
  • Falls back to parsing 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:

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.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.

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:

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 →