fix: preserve inlined MTP layers for GLM5 (#1532)

### What does this PR do?

Type of change: Bug fix <!-- Use one of the following: Bug fix, new
feature, new example, new tests, documentation. -->

Extends `load_mtp_weights` to detect *inlined* MTP layers — keys
`model.layers.{i}.*` for `i in [num_hidden, num_hidden +
num_nextn_predict_layers)` — in addition to the existing `mtp.*`
separate-file convention.

**Bug.** `load_mtp_weights()` only matched the substring `"mtp"` in
safetensors keys. GLM-5.1 (`GlmMoeDsaForCausalLM`) stores MTP at
`model.layers.78.*` with no `mtp` substring, so detection returned `([],
{})`, `_mtp_layer_prefixes` was never set, and MTP tensors were silently
dropped from the exported safetensors (had to be re-added manually).

**Detection.**

1. **Detect** via `config.num_nextn_predict_layers` (the model's own
declaration of how many MTP layers exist).
2. **Compute** the inlined layer indices: `model.layers.{i}` for `i in
range(num_hidden, num_hidden + num_nextn)`.
3. **Load** matching tensors from the on-disk shards via `safe_open`
(walks `model.safetensors.index.json` if present,else falls back to the
single shard).
4. **Split** the loaded tensors by whether `model.state_dict()` has a
slot for them:
- keys present in `model.state_dict()` → `model.load_state_dict(...,
strict=False)` (DeepSeek-V3 case: HF instantiates the extra layers).
- keys absent from `model.state_dict()` → returned as
`not_in_state_dict` so the exporter routes them through
`extra_state_dict` (GLM-5.1
case: `GlmMoeDsaModel` in transformers ≥5.7 only builds `num_hidden`
decoders, leaving MTP keys orphaned at `from_pretrained` time).
The returned prefixes flow into the existing plumbing —
`_mtp_layer_prefixes` → `quant_cfg` disable
+`quantization_config.exclude_modules` — unchanged.

### Usage

```python
# Add a code snippet demonstrating how to use this
```

### Testing
Verified end-to-end on a mini GLM-5.1 fixture (4 hidden layers + 1
inlined MTP at `model.layers.4`, 7 synthesized MTP tensors mirroring the
full GLM-5.1 layout)
To be verified with full model

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅ / ❌ / N/A <!--- If ❌, explain
why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ / ❌ / N/A
<!--- Mandatory -->
- Did you write any new necessary tests?: ✅ / ❌ / N/A <!--- Mandatory
for new features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ / ❌ / N/A <!--- Only for new features, API changes, critical bug fixes
or backward incompatible changes. -->
- Did you get Claude approval on this PR?: ✅ / ❌ / N/A <!--- Run
`/claude review`. NVIDIA org members can self-trigger for complex
changes; orthogonal to CodeRabbit. -->

### Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **Improvements**
* Quantization utilities now stream safetensors and unify loading of
multi-token-prediction (MTP) weights from both inline and
separate/sharded conventions, reporting detected MTP prefixes and counts
of loaded vs orphaned tensors.
* **Tests**
* Added unit tests and a test import helper covering MTP discovery,
loading behaviors (inlined vs standalone/indexed shards), orphan
reporting, and non‑MTP checkpoints.

<!-- review_stack_entry_start -->

[![Review Change
Stack](https://storage.googleapis.com/coderabbit_public_assets/review-stack-in-coderabbit-ui.svg)](https://app.coderabbit.ai/change-stack/NVIDIA/Model-Optimizer/pull/1532?utm_source=github_walkthrough&utm_medium=github&utm_campaign=change_stack)

<!-- review_stack_entry_end -->
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Frida Hou <201670829+Fridah-nv@users.noreply.github.com>
This commit is contained in:
Frida Hou
2026-05-22 22:47:40 +00:00
committed by GitHub
parent 04f58166ab
commit 16a0130d5f
3 changed files with 291 additions and 82 deletions
+110 -82
View File
@@ -23,6 +23,7 @@ import os
import shutil
import sys
import warnings
from collections.abc import Callable, Iterable
from pathlib import Path
from typing import Any
@@ -30,7 +31,7 @@ import torch
import transformers
from accelerate import infer_auto_device_map, init_empty_weights
from accelerate.utils import get_max_memory
from safetensors.torch import load_file
from safetensors import safe_open
from transformers import (
AutoConfig,
AutoModel,
@@ -298,100 +299,127 @@ def get_processor(
return None
def get_inlined_mtp_prefixes(config: Any) -> list[str]:
"""Turn an HF config into the list of state-dict prefixes for inlined-MTP layers."""
# ``or 0``: some configs set num_nextn_predict_layers=None rather than omit it.
num_nextn = int(getattr(config, "num_nextn_predict_layers", 0) or 0)
if not num_nextn:
return []
num_hidden = config.num_hidden_layers
return [f"model.layers.{i}" for i in range(num_hidden, num_hidden + num_nextn)]
def _keys_to_prefixes(keys: Iterable[str]) -> set[str]:
"""Invert separate-file MTP keys into the prefixes the exporter needs for exclude_modules.
``"mtp.fc.weight"`` → ``{"mtp"}``; ``"mtp.layers.0.q_proj.weight"`` →
``{"mtp", "mtp.layers.0"}``. Caller must filter out inlined keys; otherwise
``"model.layers.78.eh_proj.weight"`` would emit ``"model"`` as a prefix.
"""
prefixes: set[str] = set()
for key in keys:
parts = key.split(".")
if parts:
prefixes.add(parts[0])
for i, part in enumerate(parts):
if part == "layers" and i + 1 < len(parts) and parts[i + 1].isdigit():
prefixes.add(".".join(parts[: i + 2]))
break
return prefixes
def _load_tensors_matching(
model_dir: Path, predicate: Callable[[str], bool]
) -> dict[str, torch.Tensor]:
"""Stream tensors satisfying ``predicate(key)`` from every safetensors
source in ``model_dir`` (indexed shards + standalone files, each opened
at most once).
"""
tensors: dict[str, torch.Tensor] = {}
seen_shards: set[str] = set()
index_file = model_dir / "model.safetensors.index.json"
if index_file.exists():
with open(index_file) as f:
weight_map = json.load(f)["weight_map"]
per_shard: dict[str, list[str]] = {}
for key, shard_name in weight_map.items():
if predicate(key):
per_shard.setdefault(shard_name, []).append(key)
for shard_name, keys in per_shard.items():
seen_shards.add(shard_name)
with safe_open(str(model_dir / shard_name), framework="pt", device="cpu") as f:
for k in keys:
tensors[k] = f.get_tensor(k)
for shard in sorted(model_dir.glob("*.safetensors")):
if shard.name in seen_shards:
continue
with safe_open(str(shard), framework="pt", device="cpu") as f:
for k in f.keys(): # noqa: SIM118 - safe_open is not iterable
if predicate(k):
tensors[k] = f.get_tensor(k)
return tensors
def _apply_to_model_state_dict(
model: torch.nn.Module, tensors: dict[str, torch.Tensor]
) -> dict[str, torch.Tensor]:
"""Load tensors with a slot in ``model.state_dict()`` in-place; return the
rest as orphans for ``extra_state_dict``.
"""
model_state = model.state_dict()
in_state_dict = {k: v for k, v in tensors.items() if k in model_state}
out_state_dict = {k: v for k, v in tensors.items() if k not in model_state}
if in_state_dict:
model.load_state_dict(in_state_dict, strict=False)
return out_state_dict
def load_mtp_weights(
model: torch.nn.Module, model_path: str
) -> tuple[list[str], dict[str, torch.Tensor]]:
"""Load MTP weights from the model checkpoint.
"""Detect and load MTP weights. Support matrix:
Some models store additional layers in separate safetensors files with non-standard
names (e.g., mtp.safetensors). HuggingFace's from_pretrained() may not load these
files even though they're referenced in model.safetensors.index.json.
Convention Architectures On-disk shape
------------- ------------------------ -------------------------------
inlined GLM-5.1 (``GlmMoeDsa``), ``model.layers.{N}.*``
DeepSeek-V3
separate-file GLM-4.7 standalone ``mtp.safetensors``
separate-file Qwen3-Next indexed ``mtp.*`` tail shard
This function detects such cases and explicitly loads the missing weights.
Inlined ``N`` in ``[num_hidden, num_hidden + num_nextn_predict_layers)``;
may be orphaned at ``from_pretrained`` time if the HF class only builds
``num_hidden`` decoders.
Args:
model: The loaded model that may be missing weights
model_path: Path to the model directory
Returns:
List of layer prefixes that were loaded from non-standard safetensors files.
These layers should typically be excluded from quantization.
Empty list if no additional weights were loaded.
Dictionary of MTP weights that were not loaded into the model state dict.
Returns ``(prefixes, not_in_state_dict)``: ``prefixes`` populates
``quantization_config.exclude_modules``; ``not_in_state_dict`` is fed to
``export_hf_checkpoint(extra_state_dict=...)``.
"""
model_path = Path(model_path)
index_file = model_path / "model.safetensors.index.json"
model_dir = Path(model_path)
if not index_file.exists():
inlined_prefixes = set(get_inlined_mtp_prefixes(model.config))
inlined_tuple = tuple(p + "." for p in inlined_prefixes)
# Combined predicate covering both conventions in one pass.
def predicate(key: str) -> bool:
return key.startswith(inlined_tuple) or "mtp" in key
tensors = _load_tensors_matching(model_dir, predicate)
if not tensors:
return [], {}
# Load the index to find all referenced safetensors files
index = json.load(open(index_file))
weight_map = index["weight_map"]
# Find all files in weight_map whose key or value contains "mtp"
mtp_weight_map = {}
for k, v in weight_map.items():
if "mtp" in k or "mtp" in v:
mtp_weight_map.setdefault(v, []).append(k)
separate_keys = [k for k in tensors if not k.startswith(inlined_tuple)]
prefixes = inlined_prefixes | _keys_to_prefixes(separate_keys)
if not mtp_weight_map:
return [], {}
not_in_state_dict = _apply_to_model_state_dict(model, tensors)
def _extract_layer_prefixes(keys):
mtp_layer_prefixes = set()
for key in keys:
parts = key.split(".")
# Capture the top-level MTP module prefix (e.g., "mtp" from "mtp.fc.weight")
# so that non-layer MTP weights like mtp.fc, mtp.norm are also excluded
if parts:
mtp_layer_prefixes.add(parts[0])
# Also capture specific layer prefixes (e.g., "mtp.layers.0")
for i, part in enumerate(parts):
if part == "layers" and i + 1 < len(parts) and parts[i + 1].isdigit():
prefix = ".".join(parts[: i + 2])
mtp_layer_prefixes.add(prefix)
break
print(
f"✓ Detected {len(tensors)} MTP tensors under {sorted(prefixes)} "
f"(loaded into model: {len(tensors) - len(not_in_state_dict)}, "
f"orphaned: {len(not_in_state_dict)})"
)
return mtp_layer_prefixes
# Flatten mtp_weight_map.values() (list of list of str) to a single list of str
mtp_keys = [k for keys in mtp_weight_map.values() for k in keys]
mtp_layer_prefixes = _extract_layer_prefixes(mtp_keys)
# Check which non-standard files exist and have missing weights
model_state = model.state_dict()
total_loaded = 0
not_in_state_dict = {}
for filename, mtp_keys in mtp_weight_map.items():
filepath = model_path / filename
if not filepath.exists():
continue
print(f"Loading {len(mtp_keys)} mtp weights from {filename}...")
weights = load_file(str(filepath), device="cpu")
weights = {k: v for k, v in weights.items() if k in mtp_keys}
# Load the MTP weights to the model state dict
in_state_dict = {k: weights[k] for k in weights if k in model_state}
not_in_state_dict = not_in_state_dict | {
k: weights[k] for k in weights if k not in model_state
}
if in_state_dict:
model.load_state_dict(in_state_dict, strict=False)
total_loaded += len(in_state_dict)
if total_loaded > 0:
print(
f"✓ Successfully loaded {total_loaded} MTP weights, "
f"{len(not_in_state_dict)} MTP weights not in model.state_dict"
)
if mtp_layer_prefixes:
print(f"✓ Detected MTP layers to exclude from quantization: {mtp_layer_prefixes}")
return list(mtp_layer_prefixes), not_in_state_dict
return sorted(prefixes), not_in_state_dict
def get_dtype(dtype):
@@ -0,0 +1,30 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Re-export ``examples/llm_ptq/example_utils`` so tests can import it via
``from _test_utils.examples.llm_ptq_example_utils import example_utils``
without per-file ``sys.path`` shims.
"""
import sys
from _test_utils.examples.run_command import MODELOPT_ROOT
_LLM_PTQ_DIR = MODELOPT_ROOT / "examples" / "llm_ptq"
if str(_LLM_PTQ_DIR) not in sys.path:
sys.path.insert(0, str(_LLM_PTQ_DIR))
import example_utils
__all__ = ["example_utils"]
@@ -0,0 +1,151 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""End-to-end unit tests for ``examples/llm_ptq/example_utils.load_mtp_weights``.
One test per supported on-disk MTP convention (inlined-orphaned, inlined-in-state-dict,
separate-file-standalone, separate-file-indexed) plus a negative case.
"""
import json
from types import SimpleNamespace
import torch
from _test_utils.examples.llm_ptq_example_utils import example_utils
from safetensors.torch import save_file
class _FakeModel:
"""Stub exposing only the surface ``load_mtp_weights`` touches."""
def __init__(self, config, state_dict_keys):
self.config = config
self._sd = {k: torch.zeros(1) for k in state_dict_keys}
self.loaded = {}
def state_dict(self):
return dict(self._sd)
def load_state_dict(self, state_dict, strict=True):
self.loaded.update(state_dict)
self._sd.update(state_dict)
def _write_safetensors(path, tensors):
save_file(tensors, str(path), metadata={"format": "pt"})
def test_load_mtp_weights_inlined_orphaned(tmp_path):
# GLM-5.1: HF builds only num_hidden decoders → MTP keys orphaned.
main_keys = ["model.embed_tokens.weight", "model.layers.0.x.weight"]
mtp_keys = ["model.layers.4.eh_proj.weight", "model.layers.4.enorm.weight"]
_write_safetensors(
tmp_path / "model.safetensors",
{k: torch.zeros(2, 2) for k in main_keys + mtp_keys},
)
cfg = SimpleNamespace(num_hidden_layers=4, num_nextn_predict_layers=1)
model = _FakeModel(cfg, state_dict_keys=main_keys)
prefixes, orphans = example_utils.load_mtp_weights(model, str(tmp_path))
assert prefixes == ["model.layers.4"]
assert set(orphans) == set(mtp_keys)
assert model.loaded == {} # nothing matched the (MTP-less) state_dict
def test_load_mtp_weights_inlined_in_state_dict(tmp_path):
# DeepSeek-V3 via trust_remote_code: MTP slots exist → keys loaded, no orphans.
main_keys = ["model.embed_tokens.weight"]
mtp_keys = ["model.layers.4.eh_proj.weight", "model.layers.4.enorm.weight"]
_write_safetensors(
tmp_path / "model.safetensors",
{k: torch.ones(2, 2) for k in main_keys + mtp_keys},
)
cfg = SimpleNamespace(num_hidden_layers=4, num_nextn_predict_layers=1)
model = _FakeModel(cfg, state_dict_keys=main_keys + mtp_keys)
prefixes, orphans = example_utils.load_mtp_weights(model, str(tmp_path))
assert prefixes == ["model.layers.4"]
assert orphans == {}
assert set(model.loaded) == set(mtp_keys)
def test_load_mtp_weights_separate_standalone_file(tmp_path):
# GLM-4.7: standalone mtp.safetensors with no shard index.
_write_safetensors(
tmp_path / "model.safetensors", {"model.embed_tokens.weight": torch.zeros(2, 2)}
)
_write_safetensors(
tmp_path / "mtp.safetensors",
{
"mtp.fc.weight": torch.zeros(2, 2),
"mtp.layers.0.q_proj.weight": torch.zeros(2, 2),
},
)
cfg = SimpleNamespace(num_hidden_layers=4, num_nextn_predict_layers=0)
model = _FakeModel(cfg, state_dict_keys=["model.embed_tokens.weight"])
prefixes, orphans = example_utils.load_mtp_weights(model, str(tmp_path))
assert set(prefixes) == {"mtp", "mtp.layers.0"}
assert set(orphans) == {"mtp.fc.weight", "mtp.layers.0.q_proj.weight"}
def test_load_mtp_weights_separate_indexed_shard(tmp_path):
# Qwen3-Next: mtp.* keys in a dedicated indexed tail shard (filename has no "mtp").
main_shard = "model-00001-of-00002.safetensors"
mtp_shard = "model-00002-of-00002.safetensors"
_write_safetensors(tmp_path / main_shard, {"model.embed_tokens.weight": torch.zeros(2, 2)})
mtp_tensors = {
"mtp.fc.weight": torch.zeros(2, 2),
"mtp.norm.weight": torch.zeros(2),
"mtp.layers.0.input_layernorm.weight": torch.zeros(2),
"mtp.layers.0.self_attn.q_proj.weight": torch.zeros(2, 2),
}
_write_safetensors(tmp_path / mtp_shard, mtp_tensors)
(tmp_path / "model.safetensors.index.json").write_text(
json.dumps(
{
"weight_map": {
"model.embed_tokens.weight": main_shard,
**dict.fromkeys(mtp_tensors, mtp_shard),
}
}
)
)
cfg = SimpleNamespace(num_hidden_layers=4, num_nextn_predict_layers=0)
model = _FakeModel(cfg, state_dict_keys=["model.embed_tokens.weight"])
prefixes, orphans = example_utils.load_mtp_weights(model, str(tmp_path))
assert set(prefixes) == {"mtp", "mtp.layers.0"}
assert set(orphans) == set(mtp_tensors)
def test_load_mtp_weights_no_mtp_returns_empty(tmp_path):
# Also pins the ``num_nextn_predict_layers=None`` regression: some configs
# set the field explicitly to None, which must not crash ``int(None)``.
_write_safetensors(
tmp_path / "model.safetensors",
{
"model.embed_tokens.weight": torch.zeros(2, 2),
"model.layers.0.x.weight": torch.zeros(2, 2),
},
)
cfg = SimpleNamespace(num_hidden_layers=4, num_nextn_predict_layers=None)
model = _FakeModel(cfg, state_dict_keys=["model.embed_tokens.weight"])
prefixes, orphans = example_utils.load_mtp_weights(model, str(tmp_path))
assert prefixes == []
assert orphans == {}