Files
Model-Optimizer/examples/megatron_bridge/prune_minitron.py
T
Keval Morabia 01c708e792 Add HybridModel MBridge support for nemo:26.08 (#2005)
### Description

MBridge main (nemo:26.08) will initialize Nemotron-H as HybridModel
instead of MambaModel (subclassed of HybridModel). Also make minimum
nemo container 26.04

### Testing 

Tested Nemotron-3-Nano PTQ with MBridge main (fails otherwise)

Tested locally `tests/gpu_megatron` and `tests/examples/megatron_bridge`
with `nemo:26.06.01` + Mount latest MBridge/Mcore

GH CICD tests will be added with nemo:26.08 release

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

* **New Features**
* Added TE HybridModel stack-spec support and enabled Megatron-Bridge
hybrid providers/models for export/import and runtime handling.
* **Updates**
* Dataset packing now oversamples raw text at **16x** and improves the
packed-mode underflow warning.
* Quantization: `--quant_cfg` now defaults to `None` unless explicitly
set (or via `--recipe`).
* Distillation example: validation settings are provided via a dedicated
top-level validation configuration.
* Improved plugin import warnings to report the originating call
location; model stats now support HybridModel.
* **Deprecations**
* Megatron-Bridge / Megatron-LM optimization features now require NeMo
container `nemo:26.04` or newer (`nemo:26.06` recommended).
  * The Mamba stack specification helper is deprecated.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
2026-07-23 18:26:15 +05:30

682 lines
30 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2024 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.
"""Example script for pruning a GPT / Mamba model using Minitron algorithm on a Megatron-Bridge model (load from HF).
Supports three NAS-based pruning targets (can be combined):
--prune_target_params Total parameter count (e.g. 6e9 for 6B total params)
--prune_target_active_params Active parameter count for MoE models (e.g. 3e9 for 3B active params)
--prune_target_memory_mb Memory footprint in MB (uses --seq_length for KV-cache estimate, assumes BF16)
Example usage to prune Qwen3-8B to 6B on 2-GPUs (Pipeline Parallelism = 2)
while skipping pruning of num_attention_heads using following defaults:
1024 samples from nemotron-post-training-dataset-v2 for calibration,
at-most 20% depth (num_layers) and 40% width is pruned per prunable hparam (hidden_size, ffn_hidden_size, ...),
top-10 candidates are evaluated for MMLU score (10% sampled data) to select the best model.
torchrun --nproc_per_node 2 prune_minitron.py \
--hf_model_name_or_path Qwen/Qwen3-8B \
--prune_target_params 6e9 \
--hparams_to_skip num_attention_heads \
--output_hf_path /tmp/Qwen3-8B-Pruned-6B
To see the full usage for advanced configurations, run:
torchrun --nproc_per_node 1 prune_minitron.py --help
See `README.md` in this directory for more details.
"""
# TODO: Test multi-node pruning
import argparse
import json
import os
import re
import torch
from megatron.bridge import AutoBridge
from megatron.bridge.models.mamba.mamba_provider import MambaModelProvider
try: # nemo:26.08+
from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider
# MambaModelProvider subclasses HybridModelProvider on nemo:26.08+, so the tuple covers both.
_HYBRID_PROVIDER_TYPES: tuple[type, ...] = (MambaModelProvider, HybridModelProvider)
except ImportError: # nemo:26.06 and earlier
_HYBRID_PROVIDER_TYPES = (MambaModelProvider,)
from transformers import (
AutoConfig,
AutoModelForCausalLM,
AutoModelForImageTextToText,
AutoProcessor,
)
import modelopt.torch.opt as mto
import modelopt.torch.prune as mtp
import modelopt.torch.utils.distributed as dist
from modelopt.torch.export import copy_hf_ckpt_remote_code
from modelopt.torch.utils import (
get_supported_datasets,
num2hrb,
print_args,
print_rank_0,
warn_rank_0,
)
from modelopt.torch.utils.plugins.mbridge import load_mbridge_model_from_hf
from modelopt.torch.utils.plugins.megatron_calibration import (
get_megatron_calibration_forward_loop,
get_megatron_vlm_calibration_forward_loop,
)
from modelopt.torch.utils.plugins.megatron_mmlu import megatron_mmlu
from modelopt.torch.utils.vlm_dataset_utils import get_supported_vlm_datasets
# isort: off
# Register Megatron-Bridge model-specific NAS/pruning plugins here to avoid a circular import
import modelopt.torch.nas.plugins.mbridge # noqa: F401
# isort: on
# Default calibration datasets when --calib_dataset_name is not set
DEFAULT_TEXT_CALIB_DATASET = "nemotron-post-training-dataset-v2"
DEFAULT_VLM_CALIB_DATASET = "nemotron_vlm_dataset_v2"
# HF config field names that enable MTP
_MTP_HF_CONFIG_FIELDS = ("num_nextn_predict_layers", "mtp_num_hidden_layers", "mtp_num_layers")
def _hf_config_has_mtp(hf_cfg) -> bool:
"""Whether an HF config declares MTP heads (checked top-level and under ``text_config``)."""
return any(
cfg is not None and getattr(cfg, field, 0)
for cfg in (getattr(hf_cfg, "text_config", None), hf_cfg)
for field in _MTP_HF_CONFIG_FIELDS
)
def get_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument("--hf_model_name_or_path", type=str, required=True)
parser.add_argument("--trust_remote_code", action="store_true")
target_group = parser.add_mutually_exclusive_group(required=True)
target_group.add_argument(
"--output_megatron_path",
type=str,
help="Path to save the pruned model in Megatron checkpoint format",
)
target_group.add_argument(
"--output_hf_path", type=str, help="Path to save the pruned model in HF checkpoint format"
)
# Parallelism arguments
parser.add_argument("--pp_size", type=int, default=1, help="Pipeline parallel size")
parser.add_argument(
"--num_layers_in_first_pipeline_stage",
type=int,
default=None,
help="Number of layers in the first pipeline stage (Uneven Pipeline Parallelism)",
)
parser.add_argument(
"--num_layers_in_last_pipeline_stage",
type=int,
default=None,
help="Number of layers in the last pipeline stage (Uneven Pipeline Parallelism)",
)
# Calibration dataset parameters
parser.add_argument(
"--calib_dataset_name",
type=str,
default=None,
help=(
"Calibration dataset. If unset, it is auto-selected by model type: a text dataset "
f"({DEFAULT_TEXT_CALIB_DATASET}) for language models, and an image-text dataset "
f"({DEFAULT_VLM_CALIB_DATASET}) for VLMs. Passing a text dataset for a VLM estimates importance from text "
f"only. Text dataset options: {get_supported_datasets()}; VLM (image) dataset options: "
f"{get_supported_vlm_datasets()}."
),
)
parser.add_argument(
"--calib_num_samples", type=int, default=1024, help="Number of samples for calibration"
)
# TODO: Add support for pre-training dataset (pre-tokenized)
parser.add_argument("--calib_batch_size", type=int, default=1, help="Calibration batch size")
parser.add_argument(
"--seq_length",
type=int,
default=4096,
help="Calibration sequence length (text only; ignored for image-text VLM calibration).",
)
# Pruning parameters
parser.add_argument(
"--prune_intermediate_ckpt",
type=str,
default=None,
help=(
"Directory to save/restore per-rank intermediate pruning scores for resuming / faster re-run. "
"If not provided, it will default to `<output_path>/modelopt_pruning_scores`"
),
)
parser.add_argument(
"--prune_export_config",
type=str,
help=(
'Target pruned config as JSON e.g., \'{"hidden_size": 512, "ffn_hidden_size": 2048}\'. '
f"Supported hyperparameters: {mtp.mcore_minitron.SUPPORTED_HPARAMS}. "
"Cannot be combined with NAS-based targets."
),
)
parser.add_argument(
"--prune_target_params",
type=float,
help=(
"Target total parameter count e.g., 6e9 for 6B params. "
"Uses NAS to find the best pruned model that maximizes --prune_score_func. "
"Can be combined with --prune_target_active_params and/or --prune_target_memory_mb. "
"For VLMs this targets the language-model tower only."
),
)
parser.add_argument(
"--prune_target_active_params",
type=float,
help=(
"Target active parameter count e.g., 3e9 for 3B active params (useful for MoE models). "
"Uses NAS to find the best pruned model that maximizes --prune_score_func. "
"Can be combined with --prune_target_params and/or --prune_target_memory_mb. "
"For VLMs this targets the language-model tower only."
),
)
parser.add_argument(
"--prune_target_memory_mb",
type=float,
help=(
"Target memory footprint in MB (weights + KV-cache estimated via seq_length and "
"--inference_batch_size; assumes BF16). "
"Uses NAS to find the best pruned model that maximizes --prune_score_func. "
"Can be combined with --prune_target_params and/or --prune_target_active_params. "
"For VLMs this targets the language-model tower only."
),
)
parser.add_argument(
"--inference_batch_size",
type=int,
default=None,
help=(
"Batch size used only for KV-cache sizing in --prune_target_memory_mb. "
"Defaults to --calib_batch_size when not set. "
"Use this to target an inference batch size that differs from the calibration batch size."
),
)
parser.add_argument(
"--prune_score_func",
type=str,
default="mmlu_10pct_bs1",
help=(
"Score function to use for NAS-based pruning. Only supports MMLU at the moment. "
"Format: mmlu_<N>pct_<bs> where <N> is the percentage of MMLU data to sample per subject and <bs> is "
"batch size for fast evaluation (default is mmlu_10pct_bs1)."
),
)
parser.add_argument(
"--ss_channel_divisor",
type=int,
default=None,
help=(
"hidden_size / ffn_hidden_size divisor for NAS-based pruning. "
"Leave as None to use default divisors."
),
)
parser.add_argument(
"--max_width_pruning",
type=float,
default=0.4,
help=(
f"Maximum width pruning percentage ({mtp.mcore_minitron.SUPPORTED_HPARAMS - {'num_layers'}}) "
"for NAS-based pruning"
),
)
parser.add_argument(
"--max_depth_pruning",
type=float,
default=0.2,
help="Maximum depth pruning percentage ('num_layers') for NAS-based pruning",
)
parser.add_argument(
"--hparams_to_skip",
nargs="*",
type=str,
default=[],
choices=mtp.mcore_minitron.SUPPORTED_HPARAMS,
help=(
"Space-separated list of hparams to skip for NAS-based pruning "
"e.g. dont prune 'num_attention_heads'"
),
)
parser.add_argument(
"--top_k",
type=int,
default=10,
help=(
"Number of top candidates to consider for NAS-based pruning. "
"Higher values will take longer to prune but may find a better model."
),
)
args = parser.parse_args()
# Validate pruning target arguments
_nas_targets = [
args.prune_target_params,
args.prune_target_active_params,
args.prune_target_memory_mb,
]
if args.prune_export_config and any(t is not None for t in _nas_targets):
parser.error("--prune_export_config cannot be combined with NAS-based targets.")
if not args.prune_export_config and not any(t is not None for t in _nas_targets):
parser.error(
"At least one of --prune_export_config, --prune_target_params,"
" --prune_target_active_params, or --prune_target_memory_mb is required."
)
# Post-process arguments
if args.prune_intermediate_ckpt is None:
if args.output_megatron_path:
args.prune_intermediate_ckpt = f"{args.output_megatron_path}/modelopt_pruning_scores"
elif args.output_hf_path:
args.prune_intermediate_ckpt = f"{args.output_hf_path}/modelopt_pruning_scores"
print_rank_0(
"No directory provided to cache per-rank intermediate pruning scores. "
f"Setting to: {args.prune_intermediate_ckpt}"
)
if args.prune_export_config:
try:
prune_export_config = json.loads(args.prune_export_config)
except json.JSONDecodeError as exc:
raise ValueError(
f"Invalid JSON for --prune_export_config: {args.prune_export_config}"
) from exc
if not isinstance(prune_export_config, dict):
raise ValueError("--prune_export_config must parse to a dictionary.")
args.prune_export_config = prune_export_config
if args.inference_batch_size is None:
args.inference_batch_size = args.calib_batch_size
print_args(args)
return args
def _log_vlm_param_breakdown(unwrapped_model, language_model, stage: str) -> None:
"""Log language-model / frozen-non-LM / total param counts for a VLM (rank 0)."""
def _local(module) -> int:
# De-dup weights shared within a rank (e.g. tied embedding/output on a single stage).
seen: set[int] = set()
n = 0
for p in module.parameters():
if id(p) not in seen:
seen.add(id(p))
n += p.numel()
return n
total = dist.allreduce(_local(unwrapped_model)) # sum across pipeline ranks
lm = dist.allreduce(_local(language_model))
# Under PP a tied embedding lives on both the first and last stage, so the sum double-counts it;
# subtract one copy (the allreduce over the first-stage-only ``word_embeddings`` gives exactly one).
if dist.size() > 1 and getattr(language_model, "share_embeddings_and_output_weights", False):
emb = dist.allreduce(
next(
(
p.numel()
for n, p in unwrapped_model.named_parameters()
if "word_embeddings" in n
),
0,
)
)
total -= emb
lm -= emb
print_rank_0(
f"[{stage}] language_model={num2hrb(lm)} (--prune_target_* applies here) | "
f"frozen non-language-model={num2hrb(total - lm)} | full model={num2hrb(total)}"
)
def main(args: argparse.Namespace):
assert dist.size() == args.pp_size, "Only Pipeline parallelism is supported for pruning."
if args.output_megatron_path and os.path.exists(
f"{args.output_megatron_path}/latest_checkpointed_iteration.txt"
):
warn_rank_0(f"\nPruned model already exists at {args.output_megatron_path}. Exiting...")
return
elif args.output_hf_path and os.path.exists(f"{args.output_hf_path}/config.json"):
warn_rank_0(f"\nPruned model already exists at {args.output_hf_path}. Exiting...")
return
bridge, provider, model, unwrapped_model, tokenizer = load_mbridge_model_from_hf(
hf_model_name_or_path=args.hf_model_name_or_path,
trust_remote_code=args.trust_remote_code,
provider_overrides={
"tensor_model_parallel_size": 1, # Tensor parallelism is not supported
"expert_tensor_parallel_size": 1, # Expert tensor parallelism is not supported
"pipeline_model_parallel_size": args.pp_size,
"num_layers_in_first_pipeline_stage": args.num_layers_in_first_pipeline_stage,
"num_layers_in_last_pipeline_stage": args.num_layers_in_last_pipeline_stage,
"pipeline_dtype": torch.bfloat16,
"seq_length": args.seq_length,
"mtp_num_layers": 0, # MTP is not supported during calibration
},
init_model_parallel=True,
moe_grouped_gemm=False,
)
# TODO: Support pruning with MTP heads enabled (e.g. Qwen3.5 mtp_num_hidden_layers=1).
# Requires ModelOpt fixes for gated-attention QKV under DynamicModule during MTP calibration,
# _DynamicMCoreLanguageModel conversion/export of MTP submodules, importance hooks on MTP
# layers, mcore_param_count including MTP in --prune_target_params, and a CI test with MTP.
if _hf_config_has_mtp(bridge.hf_pretrained.config):
warn_rank_0(
"Dropping Multi-Token Prediction (MTP): calibration does not yet support MTP. Exported "
"checkpoints will not contain MTP weights. Standard autoregressive inference is unaffected. To use "
"MTP speculative decoding later, run a separate SFT phase with mtp_num_layers=1 on the pruned model."
)
# For VLMs (e.g. Qwen3-VL), only the language model is pruned; the vision tower is left intact.
# hidden_size is shared with the vision->LM projector, so it is skipped
language_model = getattr(unwrapped_model, "language_model", unwrapped_model)
is_vlm = language_model is not unwrapped_model
if is_vlm:
warn_rank_0(
"VLM detected: pruning model.language_model only; all non-language-model components "
"(vision/audio encoders, projectors, etc.) are frozen and excluded. --prune_target_* "
"applies to the language-model tower, not the full model (hidden_size pruning is also "
"skipped -- it is shared with the projector)."
)
if args.prune_export_config and "hidden_size" in args.prune_export_config:
raise ValueError(
"Pruning 'hidden_size' is not supported for VLMs (shared with the vision projector)."
)
args.hparams_to_skip = sorted({*args.hparams_to_skip, "hidden_size"})
_log_vlm_param_breakdown(unwrapped_model, language_model, "before pruning")
# Auto-select the calibration dataset by model type when not explicitly provided.
if args.calib_dataset_name is None:
args.calib_dataset_name = (
DEFAULT_VLM_CALIB_DATASET if is_vlm else DEFAULT_TEXT_CALIB_DATASET
)
# Infer the calibration modality from the dataset: the known image-text datasets require a VLM, everything
# else is text. Passing a text dataset for a VLM estimates importance from text only (vision tower idle).
use_image_calib = args.calib_dataset_name in get_supported_vlm_datasets()
if use_image_calib and not is_vlm:
raise ValueError(
f"Calibration dataset '{args.calib_dataset_name}' is image-text and requires a VLM; "
"pass a text dataset for a language model."
)
if is_vlm and not use_image_calib:
warn_rank_0(
f"Text-only calibration on a VLM (dataset '{args.calib_dataset_name}'): the language "
"model's pruning importance will not see vision tokens."
)
print_rank_0(f"Using calibration dataset: {args.calib_dataset_name}")
# Estimate pruning importance for the language model: text-only on the LM for text datasets, or
# the full VLM forward over image-text pairs.
if use_image_calib:
processor = AutoProcessor.from_pretrained(
args.hf_model_name_or_path, trust_remote_code=args.trust_remote_code
)
forward_loop = get_megatron_vlm_calibration_forward_loop(
unwrapped_model, # full VLM (vision encoder + projector + language model)
processor,
dataset_name=args.calib_dataset_name,
num_samples=args.calib_num_samples,
batch_size=args.calib_batch_size,
)
else:
forward_loop = get_megatron_calibration_forward_loop(
tokenizer,
dataset_name=args.calib_dataset_name,
num_samples=args.calib_num_samples,
seq_length=args.seq_length,
batch_size=args.calib_batch_size,
pack=True, # Megatron pretraining-style global-stream document packing
)
pruning_config = {
"forward_loop": forward_loop,
"checkpoint": args.prune_intermediate_ckpt,
}
if args.prune_export_config is not None:
# Less restrictive search space for manual pruning
ss_config = mtp.mcore_minitron.get_mcore_minitron_config(
hidden_size_divisor=64,
ffn_hidden_size_divisor=64,
mamba_head_dim_divisor=8,
num_moe_experts_divisor=8,
num_layers_divisor=1,
)
pruning_constraints = {"export_config": args.prune_export_config}
else:
# NAS-based pruning: restrict search space to a smaller set of candidates.
# Allow more choices for MoE FFN as they are generally smaller.
# NOTE: Reduce divisors and increase config['top_k'] to potentially find a better model.
hidden_size_divisor = args.ss_channel_divisor or 256
ffn_hidden_size_divisor = args.ss_channel_divisor or (
256 if (provider.num_moe_experts or 0) > 0 else 512
)
ss_config = mtp.mcore_minitron.get_mcore_minitron_config(
hidden_size_divisor=hidden_size_divisor,
ffn_hidden_size_divisor=ffn_hidden_size_divisor,
mamba_head_dim_divisor=8,
num_moe_experts_divisor=8,
num_layers_divisor=2,
)
print_rank_0(f"Using search space config: {ss_config}")
pruning_constraints = {}
if args.prune_target_params is not None:
pruning_constraints["params"] = args.prune_target_params
if args.prune_target_active_params is not None:
pruning_constraints["active_params"] = args.prune_target_active_params
if args.prune_target_memory_mb is not None:
pruning_constraints["memory_mb"] = args.prune_target_memory_mb
print_rank_0(
f"Using NAS-based automatic pruning with score function: {args.prune_score_func}. "
"You can change this to be any other metric you want to maximize (e.g. negative validation loss)."
)
match = re.fullmatch(r"mmlu_(\d+)pct_bs(\d+)", args.prune_score_func)
legacy_match = re.fullmatch(r"mmlu_(\d+)pct", args.prune_score_func)
if match:
mmlu_frac = float(match.group(1)) / 100.0
batch_size = int(match.group(2))
elif legacy_match:
warn_rank_0(
f"Score function '{args.prune_score_func}' uses the deprecated format "
"'mmlu_<N>pct'. Use 'mmlu_<N>pct_bs<bs>' to specify the evaluation batch size. "
"Falling back to batch_size=1."
)
mmlu_frac = float(legacy_match.group(1)) / 100.0
batch_size = 1
else:
raise ValueError(
f"Invalid score function: {args.prune_score_func}. "
"Expected format: mmlu_<N>pct_bs<bs> (e.g. mmlu_10pct_bs1)"
)
def score_func(m):
return megatron_mmlu(
m, tokenizer, few_shots=0, fraction=mmlu_frac, batch_size=batch_size
)
pruning_config["score_func"] = score_func
pruning_config["max_width_pruning"] = args.max_width_pruning
pruning_config["max_depth_pruning"] = args.max_depth_pruning
pruning_config["hparams_to_skip"] = args.hparams_to_skip
pruning_config["top_k"] = args.top_k
# memory_mb constraint requires batch_size and seq_length
pruning_config["batch_size"] = args.inference_batch_size
pruning_config["seq_length"] = args.seq_length
print_rank_0(f"Pruning constraints: {pruning_constraints}")
# Prune the language model in place (for VLMs this mutates unwrapped_model.language_model, so the
# full wrapper is still saved below); for plain LMs language_model is unwrapped_model itself.
language_model, pruning_scores = mtp.prune( # in-place pruning
language_model,
mode=[("mcore_minitron", ss_config)], # type: ignore[arg-type]
constraints=pruning_constraints,
dummy_input=None,
config=pruning_config,
)
# Remove unnecessary modelopt_state since ckpt is homogeneous
if mto.ModeloptStateManager.has_state_for_mode_type("prune", model=language_model):
mto.ModeloptStateManager.remove_state(language_model)
if is_vlm:
_log_vlm_param_breakdown(unwrapped_model, language_model, "after pruning")
if isinstance(provider, _HYBRID_PROVIDER_TYPES):
hybrid_key = (
"hybrid_override_pattern"
if hasattr(unwrapped_model, "hybrid_override_pattern")
else "hybrid_layer_pattern"
)
setattr(provider, hybrid_key, getattr(unwrapped_model, hybrid_key))
# NOTE: Issue with NemotronH tokenizer's len() hence using use_fast=True as a WAR.
architectures = getattr(bridge.hf_pretrained.config, "architectures", None) or []
use_fast_tokenizer = "NemotronHForCausalLM" in architectures
tokenizer_kwargs = {"trust_remote_code": args.trust_remote_code, "use_fast": use_fast_tokenizer}
if args.output_megatron_path is not None:
bridge.save_megatron_model(
model,
args.output_megatron_path,
hf_tokenizer_path=args.hf_model_name_or_path,
hf_tokenizer_kwargs=tokenizer_kwargs,
)
print_rank_0(
f"Saved pruned model to {args.output_megatron_path} in Megatron checkpoint format"
)
else:
print_rank_0(f"Saving pruned model to {args.output_hf_path} in HF checkpoint format")
# [WAR] Save the pruned HF model by hand until Megatron-Bridge natively supports it.
# TODO: Replace this whole block with ``AutoBridge.from_auto_config(...).save_hf_weights(...)``
# once the Megatron-Bridge fix ships (nemo:26.08).
bridge.hf_pretrained.save_artifacts(args.output_hf_path)
hf_cfg = AutoConfig.from_pretrained(
args.output_hf_path, trust_remote_code=args.trust_remote_code
)
mcore_cfg = language_model.config
# For VLMs the language-model fields live under hf_cfg.text_config; write back there.
text_cfg = getattr(hf_cfg, "text_config", hf_cfg)
text_cfg.hidden_size = mcore_cfg.hidden_size
text_cfg.intermediate_size = mcore_cfg.ffn_hidden_size
text_cfg.num_attention_heads = mcore_cfg.num_attention_heads
text_cfg.head_dim = mcore_cfg.kv_channels
text_cfg.num_key_value_heads = mcore_cfg.num_query_groups
if hasattr(text_cfg, "mamba_num_heads"):
text_cfg.mamba_num_heads = mcore_cfg.mamba_num_heads
if hasattr(text_cfg, "mamba_head_dim"):
text_cfg.mamba_head_dim = mcore_cfg.mamba_head_dim
if hasattr(text_cfg, "moe_intermediate_size"):
text_cfg.moe_intermediate_size = mcore_cfg.moe_ffn_hidden_size
# HF names this field with or without the ``moe_`` prefix depending on the model
# (e.g. Qwen3.5-MoE uses ``shared_expert_intermediate_size``).
for shared_expert_field in (
"moe_shared_expert_intermediate_size",
"shared_expert_intermediate_size",
):
if hasattr(text_cfg, shared_expert_field):
setattr(
text_cfg, shared_expert_field, mcore_cfg.moe_shared_expert_intermediate_size
)
if hasattr(text_cfg, "num_experts"):
text_cfg.num_experts = mcore_cfg.num_moe_experts
if hasattr(text_cfg, "n_routed_experts"):
text_cfg.n_routed_experts = mcore_cfg.num_moe_experts
if hasattr(text_cfg, "n_shared_experts"):
text_cfg.n_shared_experts = (
mcore_cfg.moe_shared_expert_intermediate_size // mcore_cfg.moe_ffn_hidden_size
)
# Layers that survived depth pruning (1-indexed). sorted_layers is None when no layer scores
# were collected (no depth pruning) -> all layers kept.
sorted_layers = pruning_scores["sorted_layers"]
kept_layer_nums = (
set(sorted_layers[: mcore_cfg.num_layers])
if sorted_layers is not None
else set(range(1, mcore_cfg.num_layers + 1))
)
# layer_types is the HF per-layer attention-cadence field (mcore's linear_attention_freq /
# moe_layer_freq have no HF equivalent under those names, so only layer_types needs slicing).
if hasattr(text_cfg, "layer_types"):
text_cfg.layer_types = [
lt for i, lt in enumerate(text_cfg.layer_types) if i + 1 in kept_layer_nums
]
# Qwen3-VL injects deepstack vision features at specific LM layers; remap those indices to the
# surviving layers (a dropped one snaps to the nearest survivor below; count is preserved).
vision_cfg = getattr(hf_cfg, "vision_config", None)
ds_indices = getattr(vision_cfg, "deepstack_visual_indexes", None)
if vision_cfg is not None and ds_indices:
kept_sorted = sorted(kept_layer_nums)
vision_cfg.deepstack_visual_indexes = [
max(0, sum(k <= d + 1 for k in kept_sorted) - 1) for d in ds_indices
]
if any((d + 1) not in kept_layer_nums for d in ds_indices):
warn_rank_0(
"A deepstack vision-injection layer was dropped during depth pruning; its "
"feature was snapped to the nearest surviving layer. Text-only (LM) "
"distillation cannot recover this vision-path change -- consider full VLM "
"training/distillation instead of LM-only to recover vision quality."
)
if isinstance(provider, _HYBRID_PROVIDER_TYPES) and hasattr(
hf_cfg, "hybrid_override_pattern"
):
hf_cfg.hybrid_override_pattern = getattr(unwrapped_model, hybrid_key)
text_cfg.num_hidden_layers = mcore_cfg.num_layers
# Mark MTP as disabled on the HF text config written after pruning
for field in _MTP_HF_CONFIG_FIELDS:
if hasattr(text_cfg, field):
setattr(text_cfg, field, 0)
# Save dummy pruned HF model to get the correct bridge for saving pruned weights
dummy_model_cls = AutoModelForImageTextToText if is_vlm else AutoModelForCausalLM
dummy_model_cls.from_config(
hf_cfg, trust_remote_code=args.trust_remote_code
).save_pretrained(args.output_hf_path, trust_remote_code=args.trust_remote_code)
pruned_bridge = AutoBridge.from_hf_pretrained(
args.output_hf_path, trust_remote_code=args.trust_remote_code
)
pruned_bridge.save_hf_weights(model, args.output_hf_path)
copy_hf_ckpt_remote_code(args.hf_model_name_or_path, args.output_hf_path)
print_rank_0(f"Saved pruned model to {args.output_hf_path} in HF checkpoint format")
print_rank_0("Done!")
if __name__ == "__main__":
dist.setup()
args = get_args()
try:
main(args)
finally:
dist.cleanup()