# How MTPLX Handles Trainable LoRA Residuals Around MTP Proposer Modules

> Discover how MTPLX integrates trainable LoRA residuals with MTP proposer modules. Learn how LoRALinear wrappers enhance models while preserving base weights.

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

---

**MTPLX injects trainable LoRA (Low-Rank Adaptation) residuals into Multi-Token Proposer (MTP) pathways by wrapping target `Linear` layers with a `LoRALinear` wrapper that adds low-rank trainable tensors while keeping the original base model weights immutable.**

The MTPLX repository implements an efficient parameter-efficient fine-tuning strategy for Multi-Token Proposer (MTP) based language models. By introducing trainable LoRA residuals around MTP proposer modules, the system allows researchers to adapt speculative decoding proposers without modifying original pretrained weights. This implementation centers on a lightweight wrapper system defined in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) that intercepts forward passes through specific MTP linear projection layers.

## The LoRALinear Wrapper Architecture

At the core of MTPLX's approach is the **`LoRALinear`** class, which acts as a drop-in replacement for standard `nn.Linear` or `nn.QuantizedLinear` modules within the MTP proposer. The wrapper maintains a reference to the frozen base module (`self.base`) while introducing two trainable tensors, **`lora_a`** and **`lora_b`**, that represent the low-rank delta matrix.

During the forward pass, the wrapper computes the original linear transformation through the base module and adds the low-rank residual: `base_output + (x @ lora_a @ lora_b) * scaling`. This additive decomposition ensures that the trainable parameters constitute only a tiny fraction of the total model size. The implementation details for this wrapper can be found in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 43–67.

### Target Resolution in MTP Modules

MTPLX discovers which linear modules belong to the MTP proposer via path string resolution. The system identifies target layers such as `"layers.0.self_attn.q_proj"` or `"fc"` within the MTP sub-module structure. This resolution mechanism allows the adapter installation to precisely target the feed-forward and attention projection layers that constitute the MTP pathway, as implemented in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 75–83.

## Installing LoRA Adapters on the MTP Proposer

The **`install_mtp_lora_adapters`** function automates the process of wrapping MTP linear layers. This utility discovers the target modules, instantiates `LoRALinear` wrappers with the specified rank and alpha scaling, and replaces the original modules in the model hierarchy.

```python
from mtplx.mtp_adapters import install_mtp_lora_adapters

# `model` is any MTPLX model containing an MTP sub-module

installed_targets = install_mtp_lora_adapters(
    model,
    rank=16,                # low-rank dimension

    alpha=32,               # scaling factor (defaults to rank)

    depth_scales=[1.0, 0.5, 0.25],  # optional depth-gating coefficients

    trainable=True,        # freeze everything except LoRA tensors

)

print("LoRA adapters installed for:", installed_targets)

```

This installation process resolves default targets and applies the `LoRALinear` wrapper to each matching module, as detailed in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 86–102.

## Training Only the Residual Parameters

Once adapters are installed, MTPLX enforces a strict parameter freezing strategy through the **`freeze_for_mtp_adapter_training`** helper. The system first freezes the entire model, then selectively unfreezes only the `lora_a` and `lora_b` tensors (and optionally the depth-gating scales). This ensures that gradient updates affect **only the residuals** while the base pretrained weights remain completely immutable.

```python

# After installation, only LoRA tensors are trainable

optimizer = mx.optimizers.Adam(learning_rate=1e-4, params=model.parameters())

for batch in data_loader:
    loss = compute_loss(model, batch)
    loss.backward()
    optimizer.step()       # Updates only lora_a, lora_b, and depth_scales

    optimizer.zero_grad()

```

The freezing logic is implemented in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 34–44.

## Adapter Persistence and Deployment

### Saving Adapter State

Trained LoRA adapters can be serialized independently of the base model weights using **`save_mtp_lora_adapter`**. This function extracts all `lora_a` and `lora_b` tensors along with descriptive metadata (rank, alpha, target paths) into a compact archive file.

```python
from mtplx.mtp_adapters import save_mtp_lora_adapter

save_path = save_mtp_lora_adapter("my_mtp_adapter.npz", model,
                                 metadata={"run_id": "exp-001"})

```

The serialization logic is located in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 96–104.

### Loading and Reinstalling

For inference or continued training, **`install_saved_mtp_lora_adapter`** reinstalls the `LoRALinear` wrappers into a fresh model and injects the saved tensors into their respective positions. This loader handles adapter state restoration without requiring access to the training configuration.

```python
from mtplx.mtp_adapters import install_saved_mtp_lora_adapter

metadata = install_saved_mtp_lora_adapter(model, "my_mtp_adapter.npz",
                                          trainable=False)  # inference mode

```

The loading mechanism is implemented in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 57–66.

### Merging into Base Weights

For deployment scenarios where adapter overhead must be eliminated, MTPLX provides **`merge_installed_mtp_lora_adapters`**. This utility bakes the LoRA delta directly into the underlying base linear weights by computing `W_base = W_base + (lora_b @ lora_a) * scaling`, then removes the wrapper layers.

```python
from mtplx.mtp_adapters import merge_installed_mtp_lora_adapters

merge_info = merge_installed_mtp_lora_adapters(model)
print("Merged", merge_info["merged"], "targets")

```

The merge implementation is found in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) at lines 8–16.

## Complete Fine-Tuning Workflow

The following example demonstrates the end-to-end workflow for adding trainable LoRA residuals to an MTP proposer, fine-tuning on downstream data, and preparing for deployment:

```python
from mtplx.mtp_adapters import (
    install_mtp_lora_adapters,
    save_mtp_lora_adapter,
    merge_installed_mtp_lora_adapters
)
import mlx.optimizers as optim

# 1. Install adapters on the MTP pathway

targets = install_mtp_lora_adapters(
    model, rank=8, alpha=16, trainable=True
)

# 2. Train (only LoRA parameters update)

optimizer = optim.Adam(learning_rate=1e-4)
for epoch in range(3):
    for batch in train_data:
        loss = model(batch)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()

# 3. Save the trained residuals

save_mtp_lora_adapter("mtp_rank8.npz", model, metadata={"task": "domain_adapt"})

# 4. Optional: merge for inference-only deployment

merge_installed_mtp_lora_adapters(model)

```

## Summary

- **LoRALinear wrapper**: Replaces standard linear layers in the MTP proposer with a module that adds trainable low-rank tensors (`lora_a`, `lora_b`) to the forward pass while preserving frozen base weights.
- **Target resolution**: Identifies specific MTP pathway modules (attention projections, feed-forward layers) via string paths like `"layers.0.self_attn.q_proj"` for precise adapter placement.
- **Controlled training**: Freezes all base model parameters and unfreezes only LoRA tensors, ensuring efficient fine-tuning with minimal memory overhead.
- **State management**: Supports saving adapter checkpoints, loading into fresh models, and merging residuals back into base weights for inference optimization.

## Frequently Asked Questions

### What is the MTP proposer in MTPLX and why use LoRA with it?

The MTP (Multi-Token Proposer) is a speculative decoding module that predicts multiple future tokens in parallel to accelerate inference. MTPLX applies LoRA to this component to enable efficient domain adaptation of the proposer's predictions without requiring full fine-tuning of the base language model, reducing trainable parameters by over 99% in typical configurations.

### Which specific layers inside the MTP module get wrapped with LoRALinear?

MTPLX targets the linear projection layers within the MTP block's transformer layers, specifically the query/key/value projections (e.g., `q_proj`, `k_proj`, `v_proj`) and feed-forward network layers (`fc` or `gate_proj`). These are identified via path strings and wrapped according to the resolution logic in [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py) lines 75–83.

### How does MTPLX ensure only LoRA parameters are updated during training?

The `freeze_for_mtp_adapter_training` function iterates through all model parameters, sets `requires_grad=False` for everything, then explicitly sets `requires_grad=True` only for tensors named `lora_a`, `lora_b`, and optionally `depth_scales`. This guarantees that optimizer steps modify only the low-rank residuals while the base model remains completely frozen, as implemented in lines 34–44 of [`mtplx/mtp_adapters.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/mtp_adapters.py).

### Can trained LoRA adapters be merged back into the base model weights?

Yes. The `merge_installed_mtp_lora_adapters` function computes the effective weight update by multiplying `lora_b @ lora_a` (scaled by alpha/rank), adds this to the base linear weights, and removes the LoRA wrapper. This produces a standard model with no inference overhead while preserving the fine-tuned behavior, useful for production deployment scenarios.