mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### 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>
682 lines
30 KiB
Python
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()
|