From c69ff4fedfa87e9462326fe463ff23a55b79ee82 Mon Sep 17 00:00:00 2001 From: Fridah-nv <201670829+Fridah-nv@users.noreply.github.com> Date: Wed, 26 Aug 2026 22:26:41 +0000 Subject: [PATCH] feat(export): support MTP models in per-layer export [prototype] The refusal said MTP exclusions and orphaned weights are applied after calibration, when every shard is already written. The first half stopped being true once the prefixes were derived from the checkpoint index before calibration: the pre-quantize exclusion loop already leaves MTP modules unquantized, and finalize() already calls _add_mtp_exclusions. That leaves the orphans -- MTP tensors with no slot in state_dict(), which the separate-file conventions produce. load_mtp_weights only fills existing slots and returns the rest, so it is safe to run before quantize; per-layer export now does that and stashes the leftovers on the model under MTP_EXTRA_STATE_ATTR, which finalize() feeds into its existing extra_state_dict path. The stash exists because calibration owns the finalize() call, so hf_ptq cannot pass them as an argument -- the same reason _mtp_layer_prefixes is already carried that way. The blanket refusal is replaced by a narrow one: if the post-calibration load finds tensors that were not staged, the run still fails, because the shards are written by then and they cannot be added. Prototype: covered by an inlined-convention test that fails without the stash pickup. The separate-file conventions (GLM-4.7 standalone mtp.safetensors, Qwen3-Next indexed tail shard) have no local fixture and are unverified. Signed-off-by: Fridah-nv <201670829+Fridah-nv@users.noreply.github.com> --- examples/hf_ptq/hf_ptq.py | 33 +++++++++++-------- modelopt/torch/export/layerwise_export.py | 13 ++++++++ .../gpu/torch/export/test_layerwise_export.py | 32 +++++++++++++++++- 3 files changed, 63 insertions(+), 15 deletions(-) diff --git a/examples/hf_ptq/hf_ptq.py b/examples/hf_ptq/hf_ptq.py index 0e3c8035bf8..90eec1d9560 100755 --- a/examples/hf_ptq/hf_ptq.py +++ b/examples/hf_ptq/hf_ptq.py @@ -78,6 +78,7 @@ has_spec_opt, save_expert_token_count_table, ) +from modelopt.torch.export.layerwise_export import MTP_EXTRA_STATE_ATTR from modelopt.torch.export.model_utils import get_language_model_from_vl, is_multimodal_model from modelopt.torch.quantization.config import need_calibration from modelopt.torch.quantization.plugins.accelerate import init_quantized_weights @@ -771,7 +772,7 @@ def mono_quantize( warnings.warn("Skipping quantization: model is already quantized.") -def assert_layerwise_export_compatible(args, full_model, mtp_layer_prefixes) -> None: +def assert_layerwise_export_compatible(args, full_model) -> None: """Refuse layerwise export before calibration starts, not after it writes a checkpoint. Layerwise export writes the finished checkpoint during calibration, so anything that @@ -786,13 +787,6 @@ def assert_layerwise_export_compatible(args, full_model, mtp_layer_prefixes) -> "overwrite config.json with the unquantized source config." ) - if mtp_layer_prefixes: - raise NotImplementedError( - f"layerwise.export_dir does not support models with MTP layers {mtp_layer_prefixes}: " - "their exclusions and any orphaned MTP weights are applied after calibration, by " - "which point every shard and the quant config are already written." - ) - if has_spec_opt(full_model): raise NotImplementedError( "layerwise.export_dir does not support speculative-decoding models: " @@ -933,12 +927,13 @@ def export_quantized( full_model._mtp_layer_prefixes = mtp_layer_prefixes if args.layerwise_export: - if mtp_state_dict: + staged = getattr(full_model, MTP_EXTRA_STATE_ATTR, None) or {} + if mtp_state_dict.keys() - staged.keys(): raise NotImplementedError( - "layerwise.export_dir does not support models with MTP weights: " - "they are loaded after calibration has already written every " - "shard, so they would be missing from the checkpoint. Export " - "without layerwise.export_dir." + "layerwise.export_dir found MTP weights it did not stage before " + f"calibration: {sorted(mtp_state_dict.keys() - staged.keys())[:4]}. " + "Every shard is already written, so they cannot be added now. " + "Export without layerwise.export_dir." ) # Calibration already wrote every shard, the index and the configs. print(f"Layerwise export already wrote the checkpoint to {export_path}") @@ -1353,10 +1348,20 @@ def quantize_main( quant_cfg["quant_cfg"].append({"quantizer_name": pattern, "enable": False}) print(f"Excluding MTP layer from quantization: {pattern}") + if args.layerwise_export and mtp_layer_prefixes: + # Per-layer export writes the checkpoint *during* calibration, so the MTP + # weights have to be in place first. load_mtp_weights only fills existing slots + # and hands back the rest, so running it early is safe; the orphans are stashed + # for finalize(), which owns the tail shard. + _, mtp_state_dict = load_mtp_weights(full_model, args.pyt_ckpt_path) + if mtp_state_dict: + setattr(full_model, MTP_EXTRA_STATE_ATTR, mtp_state_dict) + print(f"Layerwise export: staged {len(mtp_state_dict)} orphaned MTP tensors") + # Before resolve_checkpoint_dir, which hashes the config: with the placeholder # still in it, two --export_path values would share one checkpoint dir. if args.layerwise_export: - assert_layerwise_export_compatible(args, full_model, mtp_layer_prefixes) + assert_layerwise_export_compatible(args, full_model) quant_cfg = set_layerwise_export_dir(quant_cfg, args.export_path) print(f"Layerwise export enabled: writing quantized shards to {args.export_path}") # The shards are only a resume artifact if the manifest that names the resume diff --git a/modelopt/torch/export/layerwise_export.py b/modelopt/torch/export/layerwise_export.py index a196c9a9e52..92deb35076d 100644 --- a/modelopt/torch/export/layerwise_export.py +++ b/modelopt/torch/export/layerwise_export.py @@ -60,6 +60,12 @@ _TAIL_SHARD = "model-tail.safetensors" _INDEX_FILE = "model.safetensors.index.json" +#: Set by the caller on the model, holding MTP tensors that have no slot in +#: ``state_dict()``. finalize() runs inside calibration, so the caller cannot pass them as +#: an argument; this follows the ``_mtp_layer_prefixes`` convention already used to hand +#: MTP information across the same boundary. +MTP_EXTRA_STATE_ATTR = "_mtp_extra_state_dict" + def layer_shard_name(layer_idx: int) -> str: """Shard filename for one decoder layer, keyed by index so a re-export overwrites.""" @@ -354,6 +360,13 @@ def finalize(self) -> dict: continue self._collect(tail, name, tensor) + # Tensors the model never held (e.g. orphaned MTP weights), already in export + # form, so only the hub-name reversal applies. The stash is the layerwise route: + # calibration owns the finalize() call, so the caller cannot pass them directly. + for name, tensor in (getattr(model, MTP_EXTRA_STATE_ATTR, None) or {}).items(): + mapped = self._name_mapper(name) if self._name_mapper is not None else name + tail.setdefault(mapped, tensor.detach().contiguous().cpu()) + save_file(tail, str(self._export_dir / _TAIL_SHARD)) self._write_index() save_non_weight_artifacts(model, self._export_dir) diff --git a/tests/gpu/torch/export/test_layerwise_export.py b/tests/gpu/torch/export/test_layerwise_export.py index ca40f72e875..11d3b55f7e2 100644 --- a/tests/gpu/torch/export/test_layerwise_export.py +++ b/tests/gpu/torch/export/test_layerwise_export.py @@ -26,7 +26,11 @@ from safetensors.torch import load_file import modelopt.torch.quantization as mtq -from modelopt.torch.export.layerwise_export import LayerwiseExporter, layer_shard_name +from modelopt.torch.export.layerwise_export import ( + MTP_EXTRA_STATE_ATTR, + LayerwiseExporter, + layer_shard_name, +) from modelopt.torch.export.unified_export_hf import export_hf_checkpoint NUM_LAYERS = 4 @@ -429,6 +433,32 @@ def test_moe_export_matches(tmp_path): _assert_same_quant_config(baseline_dir, export_dir) +def test_orphaned_mtp_tensors_reach_the_tail_shard(tmp_path): + """MTP weights with no slot in state_dict() must still land in the checkpoint. + + finalize() runs inside calibration, so hf_ptq cannot pass them as an argument; it + stashes them on the model and the exporter picks them up. Without that they are + silently absent from a checkpoint that otherwise looks complete. + """ + export_dir = tmp_path / "fused" + model = _build_model() + orphans = { + "mtp.layers.0.weight": torch.ones(4, 4, dtype=torch.bfloat16), + "mtp.norm.weight": torch.ones(4, dtype=torch.bfloat16), + } + setattr(model, MTP_EXTRA_STATE_ATTR, orphans) + + mtq.quantize(model, _layerwise_cfg(export_dir, tmp_path / "ckpt"), _calib) + + exported = _load_checkpoint(export_dir) + for key, value in orphans.items(): + assert key in exported, f"{key} missing from the exported checkpoint" + assert torch.equal(exported[key].cpu(), value) + + weight_map = json.loads((export_dir / "model.safetensors.index.json").read_text())["weight_map"] + assert set(orphans) <= set(weight_map), "orphans written but left out of the index" + + def test_export_consumes_the_model_without_affecting_the_checkpoint(tmp_path): """Per-layer export converts each layer in place; the shard is written first.