mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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 -->
[](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:
@@ -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 == {}
|
||||
Reference in New Issue
Block a user