Files
Keval MorabiaandClaude Opus 5 0058a15537 [2/2] Track every Megatron-Bridge script with MLflow (#2514)
### What does this PR do?

Type of change: new feature

**[2/2] of a split. Based on #2544 — merge that first; this PR's diff is
only the Megatron-Bridge half.**

#2477 added MLflow tracking to `examples/megatron_bridge/quantize.py`.
It was one of five scripts in that directory that write a checkpoint;
the other four recorded nothing, so the provenance chain stopped at the
PTQ checkpoint and a deployed model could not be traced back to the run
that produced it.

All five now take the same `--mlflow` / `--mlflow_experiment` /
`--mlflow_run_name` flags, and **each declares what it records as a
`Tool` beside its own flags** — the shared `mlflow_utils.py` knows none
of them:

| Script | Records |
| --- | --- |
| `prune_minitron.py` | command, arguments, log, `prune_score` metric,
pointer |
| `quantize.py` (#2477, moved onto the shared `Tool` in #2544) | +
resolved recipe, quantizer summary |
| `distill.py` | + Megatron-Bridge's per-iteration metrics and resolved
config |
| `export_quantized_megatron_to_hf.py` | command, arguments, log,
pointer |
| `export_distilled_megatron_to_hf.py` | same, one pointer per exported
checkpoint |

Each writes `.experiment.json` into the checkpoint it produced, and each
tags what it consumed, so `prune → quantize → distill → export` is
walkable both from disk and by tag query on the server.

**`distill.py` opens the run and Megatron-Bridge joins it.** Its
`LoggerConfig` records per-iteration metrics and the full resolved
config — which a wrapper around `main()` cannot see — but nothing of
`distill.py`'s own arguments and no invocation. Megatron-Bridge takes
`mlflow.active_run()` when one exists, applies the tags and logs into
it, so `distill_run()` opens the run on the rank Megatron-Bridge looks
at (the **last** one) and the two share it. Its early exit is handled
explicitly: `train()` leaves through `sys.exit(0)` on `--exit_interval`,
which a blanket handler would record as `FAILED`.

**The library pieces that exist for that shared run land here with their
first caller**, rather than in [1/2] where they would have none:
`split_tracking_credentials`, so a URI handed to something which
*records* it carries no credential; `log_active_run_experiment_json`,
for pointing a checkpoint at a run this process did not open; and
`MlflowRunLogger._reattach`, because a co-owner can end the run first —
Megatron-Bridge does, as `KILLED`, when SIGTERM arrives mid-training.

Two of Megatron-Bridge's defaults are deliberately not inherited:
**checkpoint artifact upload stays off** unless
`--mlflow_log_checkpoints` (it pushes the whole checkpoint over HTTP
after every save), and **an untracked run passes no `mlflow_*` fields at
all**, since they landed in Megatron-Bridge 0.6 and sending them
unconditionally would break an untracked run on an older one.

### Usage

```bash
# Any of the five, same flags:
torchrun --nproc_per_node 8 prune_minitron.py  ... --mlflow https://<server>/
torchrun --nproc_per_node 8 quantize.py        ... --mlflow https://<server>/
torchrun --nproc_per_node 8 distill.py         ... --mlflow https://<server>/
torchrun --nproc_per_node 8 export_quantized_megatron_to_hf.py ... --mlflow https://<server>/

# Each checkpoint names the run that wrote it:
cat /output/qad/checkpoints/.experiment.json
```

Experiments default to
`$USER/megatron_bridge_{prune,quantize,distill,export,distill_export}/<model
basename>-<variant>`.

### Testing

- Real runs on a toy Qwen3 in one MLflow experiment covering all five
Megatron-Bridge scripts and `hf_ptq` — prune, quantize, QAD
distillation, quantized export, BF16 distillation, distilled export, HF
PTQ — each closing `FINISHED` with the invocation, its arguments as
params, its log, and a matching `.experiment.json` on disk. The chain
tags line up: each stage's `source_checkpoint_path` is the previous
stage's `checkpoint_path`.
- `tests/examples/megatron_bridge` in `nvcr.io/nvidia/nemo:26.08`, the
only lane that runs it: **76 passed**. Plus the three suites from #2544:
**195 pass**.
- `pre-commit run --files <changed>`: all hooks pass.
- Each fix from the review rounds has a test that fails with the fix
reverted: the resumed run, the foreign active run, the percent-decoded
credential, the credential that cannot be moved, the rank-dependent
`LoggerConfig`, the exit-callback guard, and the `iter_*` join.

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

- Is this change backward compatible?: ✅
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A
- Did you write any new necessary tests?: ✅
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅
- Did you get Claude approval on this PR?: several rounds; re-requested
on this head.

### Additional Information

Split from a single ~1150-line PR at review's request; #2544 carries the
library consolidation this builds on, and this branch is based on it.
Earlier review threads here show as outdated after the rebases — they
are all resolved and their fixes are in this branch.

One known gap, stated in the README rather than implied: `distill.py
--hf_export_path` writes a second HuggingFace checkpoint from rank 0,
which is not the rank that owns the run, so it carries no pointer yet.
For the same reason the uploaded `logs/distill.log` holds the last
rank's output — `print_rank_0` keeps the script's own lines on rank 0 —
which the README now says outright; carrying rank 0's log into a run
owned by another rank needs cross-rank upload and is a follow-up.

Two defects found on shared-run paths during review, both verified
against the installed Megatron-Bridge 0.6 rather than its docs.
Megatron-Bridge ends the run it shares with `distill.py` as `KILLED`
from its SIGTERM handler (`train.py:1413`) and then leaves through
`sys.exit()` (`train.py:805`), i.e. before `distill_run`'s `finally` —
and MLflow's fluent calls resolve their target by *opening* a run when
none is active, so a preempted distillation's log and metrics went to a
second, empty run and its `KILLED` status was overwritten. Separately,
an unreachable server disabled our logger but `logger_kwargs` still
handed Megatron-Bridge the same URI, and `state.py` calls
`set_experiment` unguarded from inside the training loop — so a
best-effort `$MLFLOW_TRACKING_URI` aborted the training instead of
degrading to untracked.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

---------

Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-09-28 21:29:25 +00:00

506 lines
22 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 post-training quantization (PTQ) of a GPT / Mamba model using ModelOpt on a
Megatron-Bridge model (loaded from HF).
The process is as follows:
1. Load a pretrained HuggingFace model into a Megatron-Core model via Megatron-Bridge.
2. Apply ModelOpt quantization (fake-quant) with calibration on a few samples from a dataset.
The quantization format is specified either by a short --quant_cfg alias or a --recipe YAML.
3. (Optional) Compress weights to a real low-bit representation.
4. Save the quantized model as a Megatron checkpoint (with ModelOpt state). The checkpoint can be
reloaded for further training (QAT / distillation) or converted to a HuggingFace (unified)
checkpoint for deployment with `export_quantized_megatron_to_hf.py` (for TensorRT-LLM / vLLM / SGLang).
Tensor / pipeline / expert parallelism are all supported here — the Megatron checkpoint is saved
sharded and can be re-sharded on load (e.g. `export_quantized_megatron_to_hf.py` reloads it at TP=1 for the HF export).
Example usage to quantize Qwen3-8B to NVFP4 on 2 GPUs (Tensor Parallelism = 2):
1024 samples from default dataset are used for calibration (sequence length = 4096).
torchrun --nproc_per_node 2 quantize.py \
--hf_model_name_or_path Qwen/Qwen3-8B \
--quant_cfg nvfp4 \
--tp_size 2 \
--calib_batch_size 1 \
--seq_length 4096 \
--export_megatron_path /tmp/Qwen3-8B-NVFP4-megatron
Equivalent run using a YAML recipe (authoritative for quant_cfg + algorithm + KV-cache config):
torchrun --nproc_per_node 2 quantize.py \
--hf_model_name_or_path Qwen/Qwen3-8B \
--recipe general/ptq/nvfp4_default-kv_fp8 \
--tp_size 2 \
--calib_batch_size 1 \
--seq_length 4096 \
--export_megatron_path /tmp/Qwen3-8B-NVFP4-megatron
To convert the saved Megatron checkpoint to a deployable HuggingFace checkpoint, use
`export_quantized_megatron_to_hf.py`.
To see the full usage for advanced configurations, run:
torchrun --nproc_per_node 1 quantize.py --help
See `README.md` in this directory for more details.
"""
import argparse
import copy
import gc
from pathlib import Path
import torch
from megatron.bridge.models.hf_pretrained.utils import is_safe_repo
from mlflow_utils import NON_PARAMS, add_mlflow_args, mlflow_run, resolve_mlflow_args
from transformers import AutoProcessor
import modelopt.torch.quantization as mtq
import modelopt.torch.utils.distributed as dist
from modelopt.recipe import ModelOptPTQRecipe, load_recipe
from modelopt.recipe.presets import (
KV_CACHE_NONE,
KV_QUANT_CFG_CHOICES,
QUANT_CFG_CHOICES,
RecipeSupersededAction,
)
from modelopt.torch.utils import print_args, print_rank_0, warn_rank_0
from modelopt.torch.utils.dataset_utils import get_supported_datasets
from modelopt.torch.utils.mlflow import Tool, masked_args, resolved_recipe_texts
from modelopt.torch.utils.plugins.mbridge import (
get_language_model,
load_mbridge_model_from_hf,
use_moe_grouped_gemm,
)
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_generate import megatron_generate
from modelopt.torch.utils.vlm_dataset_utils import get_supported_vlm_datasets
# Default calibration datasets when --calib_dataset_name is not set
DEFAULT_TEXT_CALIB_DATASET = "cnn_nemotron_v2_mix" # cnn_dailymail + nemotron-post-training-v2
DEFAULT_VLM_CALIB_DATASET = "nemotron_vlm_dataset_v2"
# The --quant_cfg / --kv_cache_quant CLI vocabularies are discovered from the preset
# YAMLs (shared with the hf_ptq examples via modelopt.recipe.presets). --quant_cfg
# additionally accepts any full config name from ``mtq.config.choices`` (e.g.
# ``FP8_DEFAULT_CFG``); see get_quant_config below.
# TODO: Add AutoQuantize (mtq.auto_quantize) support to automatically search a per-layer mix of
# quantization formats that meets a target compression / accuracy constraint, instead of applying a
# single fixed --quant_cfg / --recipe to the whole model.
QUANTIZE = Tool(
name="megatron_bridge_quantize",
tracks=(
"Track this run on an MLflow server, uploading the command, the resolved recipe, the "
"run log and the quantizer summary, and writing .experiment.json into "
"--export_megatron_path."
),
variant_help="recipe name, or --quant_cfg if no --recipe",
# ``or "none"``: neither flag is required, and a run without one fails in get_quant_config
# rather than while being named.
variant=lambda args: Path(args.recipe).stem if args.recipe else (args.quant_cfg or "none"),
model=lambda args: args.hf_model_name_or_path,
checkpoint=lambda args: args.export_megatron_path,
texts=lambda args: resolved_recipe_texts(args.recipe),
outputs=lambda args: {
"summary/quant_summary.txt": Path(args.export_megatron_path) / ".quant_summary.txt"
},
non_params=NON_PARAMS,
)
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")
parser.add_argument(
"--no_moe_grouped_gemm",
action="store_true",
help=(
"Force SequentialMLP for MoE experts instead of the fused TEGroupedMLP (grouped GEMM). "
"By default grouped GEMM is used unless the architecture cannot export it to "
"HuggingFace, in which case SequentialMLP is selected automatically."
),
)
parser.add_argument(
"--export_megatron_path",
type=str,
required=True,
help="Path to save the quantized model in Megatron checkpoint format (with ModelOpt state).",
)
# Parallelism arguments. Data parallelism is implicit: DP = world_size / (tp * pp * cp).
# e.g. `torchrun --nproc_per_node 8 quantize.py --tp_size 2` runs with DP=4.
parser.add_argument("--tp_size", type=int, default=1, help="Tensor parallel size")
parser.add_argument("--pp_size", type=int, default=1, help="Pipeline parallel size")
parser.add_argument("--ep_size", type=int, default=1, help="Expert parallel size")
parser.add_argument("--cp_size", type=int, default=1, help="Context parallel size")
# Quantization arguments
parser.add_argument(
"--recipe",
type=str,
default=None,
help=(
"PTQ recipe YAML file or builtin name (e.g. 'general/ptq/fp8_default-kv_fp8'). "
"When set, --quant_cfg, --kv_cache_quant and --weight_only are ignored; the recipe "
"is authoritative for quant_cfg, algorithm, and KV-cache config."
),
)
parser.add_argument(
"--quant_cfg",
action=RecipeSupersededAction,
type=str,
default=None,
help=(
"(deprecated: use --recipe) "
f"Quantization config. Preset names: {', '.join(QUANT_CFG_CHOICES)}. "
"You can also pass any full config name exposed by modelopt (e.g. FP8_DEFAULT_CFG). "
"Ignored when --recipe is set."
),
)
parser.add_argument(
"--kv_cache_quant",
action=RecipeSupersededAction,
type=str,
default=KV_CACHE_NONE,
choices=[KV_CACHE_NONE, *KV_QUANT_CFG_CHOICES],
help=(
"(deprecated: use --recipe) KV-cache quantization config to apply on top of "
"--quant_cfg. Ignored when --recipe is set."
),
)
parser.add_argument(
"--weight_only",
action=RecipeSupersededAction,
nargs=0,
const=True,
default=False,
help=(
"(deprecated: use --recipe) Disable input (activation) quantization, i.e. "
"weight-only quantization."
),
)
parser.add_argument(
"--compress",
action="store_true",
help="Compress weights to a real low-bit representation (instead of fake quantization).",
)
# Calibration dataset arguments (matched to hf_ptq.py)
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"
)
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).",
)
# Post-quantization generation (sanity check) arguments
parser.add_argument(
"--prompts",
type=str,
default="Hello!|Born in California, Soyer trained as a",
help="Prompts to sanity-check the quantized model. Use | to separate batches.",
)
parser.add_argument(
"--osl",
type=int,
default=32,
help="Output sequence length for the generation sanity check.",
)
parser.add_argument(
"--skip_generate",
action="store_true",
help="Skip the post-quantization generation sanity check.",
)
add_mlflow_args(parser, QUANTIZE)
args = parser.parse_args()
resolve_mlflow_args(args, parser, QUANTIZE)
print_args(masked_args(args))
# Flipped by main() once the Megatron checkpoint is on disk. The MLflow provenance
# pointer is gated on it rather than on --export_megatron_path existing, which proves
# nothing: print_quant_summary creates that directory before the save.
args.checkpoint_exported = False
return args
def get_quant_config(args: argparse.Namespace) -> dict:
"""Build the ModelOpt quantization config dict from the parsed arguments."""
if args.recipe is not None:
# A YAML recipe is authoritative: it encodes quant_cfg + algorithm + KV-cache config
# directly, so the --quant_cfg / --kv_cache_quant / --weight_only customizations below
# are skipped.
print_rank_0(f"Using recipe {args.recipe} for quantization")
if args.kv_cache_quant != KV_CACHE_NONE or args.weight_only:
warn_rank_0(
"--kv_cache_quant / --weight_only are ignored when --recipe is set; the recipe "
"is authoritative."
)
recipe = load_recipe(args.recipe)
if not isinstance(recipe, ModelOptPTQRecipe):
raise TypeError(
f"Expected a PTQ recipe but got {type(recipe).__name__} from {args.recipe}"
)
return recipe.quantize.model_dump()
if args.quant_cfg in QUANT_CFG_CHOICES:
mtq_config = QUANT_CFG_CHOICES[args.quant_cfg]
elif args.quant_cfg in mtq.config.choices:
mtq_config = getattr(mtq, args.quant_cfg)
else:
raise ValueError(
f"Unsupported --quant_cfg '{args.quant_cfg}'. Choose a preset name "
f"({', '.join(QUANT_CFG_CHOICES)}) or a full config name from {mtq.config.choices}."
)
# Deepcopy so we don't mutate a shared module-level config (the ``mtq.config.choices``
# full-name branch returns one; QUANT_CFG_CHOICES already hands back a fresh copy), and
# normalize the inner quant_cfg to the list format so we can safely append customizations below.
mtq_config = copy.deepcopy(mtq_config)
mtq_config["quant_cfg"] = mtq.normalize_quant_cfg_list(mtq_config["quant_cfg"])
if args.weight_only:
mtq_config["quant_cfg"].append({"quantizer_name": "*input_quantizer", "enable": False})
if args.kv_cache_quant != KV_CACHE_NONE:
kv_cache_quant_cfg = KV_QUANT_CFG_CHOICES[args.kv_cache_quant]["quant_cfg"]
mtq_config = mtq.utils.update_quant_cfg_with_kv_cache_quant(mtq_config, kv_cache_quant_cfg)
return mtq_config
def main(args: argparse.Namespace):
trust_remote_code = is_safe_repo(
trust_remote_code=args.trust_remote_code, hf_path=args.hf_model_name_or_path
)
moe_grouped_gemm = use_moe_grouped_gemm(
args.hf_model_name_or_path,
trust_remote_code=trust_remote_code,
force_sequential=args.no_moe_grouped_gemm,
)
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=trust_remote_code,
moe_grouped_gemm=moe_grouped_gemm,
provider_overrides={
"tensor_model_parallel_size": args.tp_size,
"pipeline_model_parallel_size": args.pp_size,
"expert_model_parallel_size": args.ep_size,
"context_parallel_size": args.cp_size,
"expert_tensor_parallel_size": 1, # Expert tensor parallelism is not supported
"pipeline_dtype": torch.bfloat16,
"seq_length": args.seq_length,
"gradient_accumulation_fusion": False, # not supported
},
init_model_parallel=True,
)
# Only the language model is quantized (vision tower + projector stay full precision)
language_model, is_vlm = get_language_model(unwrapped_model)
if is_vlm:
warn_rank_0(
"VLM detected: quantizing `model.language_model` only (vision tower left in full precision)."
)
# 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 calibration statistics will not see vision tokens."
)
print_rank_0(f"Using calibration dataset: {args.calib_dataset_name}")
mtq_config = get_quant_config(args)
# Quantize only the language model: disable quantizers on every top-level submodule that is not
# the language model (vision tower + projector). Skip aliases of language-model submodules (e.g.
# Qwen's ``self.decoder = language_model.decoder``) so the LM's own layers stay enabled.
if is_vlm:
lm_module_ids = {id(m) for m in language_model.modules()}
non_lm_children = sorted(
name
for name, child in unwrapped_model.named_children()
if name != "language_model" and id(child) not in lm_module_ids
)
for name in non_lm_children:
# Anchor to the child subtree (top-level child of the quantized root) so a short non-LM
# name cannot accidentally match a language-model quantizer path by substring.
mtq_config["quant_cfg"].append({"quantizer_name": f"{name}.*", "enable": False})
print_rank_0(f"Disabling quantizers on non-language-model submodules: {non_lm_children}")
# KV-cache quantization is incompatible with weight compression. Validate on the *resolved*
# config (KV-cache quantizers are named ``*[kv]_bmm_quantizer``) so this also covers
# recipe-driven KV-cache configs, not just the --kv_cache_quant flag.
if args.compress and any(
isinstance(entry, dict) and "bmm_quantizer" in str(entry.get("quantizer_name", ""))
for entry in mtq.normalize_quant_cfg_list(mtq_config["quant_cfg"])
):
raise ValueError("--compress cannot be combined with KV-cache quantization.")
print_rank_0(f"Quantizing the model with: {args.recipe or args.quant_cfg}")
if "awq" in str(mtq_config.get("algorithm")):
print_rank_0(
"AWQ calibration can take longer than other methods; reduce --calib_num_samples to speed it up."
)
# Dynamic and weight-only configs need no activation statistics, so skip both the
# (potentially expensive) calibration dataset download and the calibration forward pass.
if not mtq.need_calibration(mtq_config):
warn_rank_0("Dynamic or weight-only quantization detected; skipping calibration.")
forward_loop = None
elif not use_image_calib:
text_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
)
# Run text prefill on the language model: we quantize the root (a VLM root forward expects
# vision inputs), but text calibration must drive the inner LM. For plain LMs these are the same.
def forward_loop(_model=None):
text_forward_loop(language_model)
else:
# VLMs: drive the full VLM forward on image-text pairs so the language model's quantizers
# see vision-conditioned activations (we still quantize the LM only).
processor = AutoProcessor.from_pretrained(
args.hf_model_name_or_path, trust_remote_code=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,
)
if hasattr(unwrapped_model, "calibration_mode"):
# Some model wrappers (e.g. distillation/speculative) gate calibration behind a flag.
unwrapped_model.calibration_mode = True
mtq.quantize(unwrapped_model, mtq_config, forward_loop)
unwrapped_model.calibration_mode = False
else:
mtq.quantize(unwrapped_model, mtq_config, forward_loop)
# Free calibration/quantization memory before generate
gc.collect()
torch.cuda.empty_cache()
if args.compress:
mtq.compress(unwrapped_model)
print_rank_0("Weights are now compressed to low-bit!")
# Save the quantizer summary alongside the checkpoint for later inspection. Only the master
# rank writes the file to avoid a multi-rank race on the same path.
if dist.is_master():
mtq.print_quant_summary(unwrapped_model, args.export_megatron_path)
bridge.save_megatron_model(
model,
args.export_megatron_path,
hf_tokenizer_path=args.hf_model_name_or_path,
hf_tokenizer_kwargs={"trust_remote_code": trust_remote_code},
)
args.checkpoint_exported = True
if is_vlm:
print_rank_0(
f"\nSaved quantized VLM to {args.export_megatron_path} in Megatron format. To deploy this "
"model, convert it to a Unified HF ckpt with export_quantized_megatron_to_hf.py."
)
else:
print_rank_0(
f"\nSaved quantized model to {args.export_megatron_path} in Megatron format. To deploy this model "
"(TensorRT-LLM / vLLM / SGLang), convert it to a Unified HF ckpt with export_quantized_megatron_to_hf.py"
)
# Sanity-check generation with the fake-quantized model. Skipped when --compress is set: the
# weights are now real low-bit and megatron_generate may not support compressed forward for
# every quant format.
if args.compress and not args.skip_generate:
warn_rank_0(
"Skipping the post-quantization generation sanity check because --compress is set."
)
if not args.skip_generate and not args.compress:
print_rank_0("\nTesting quantized model with custom prompts...")
# Sanity-check text generation on the quantized language model.
language_model.eval()
for idx, prompt in enumerate(args.prompts.split("|")):
tokens = tokenizer(prompt, return_tensors="pt")
# enable_kv_cache=False avoids pre-allocating the static KV cache: this is a short sanity-check
# generation and the KV-cache allocation can OOM tight quantization runs on large MoE models.
generated_ids = megatron_generate(
language_model, tokens.input_ids.cuda(), osl=args.osl, enable_kv_cache=False
)
generated_texts = tokenizer.batch_decode(generated_ids)
print_rank_0(f"\nPrompt {idx + 1}: {prompt}\nGenerated: {generated_texts}")
print_rank_0("\nDone!")
if __name__ == "__main__":
dist.setup()
args = get_args()
try:
# Entered inside the try: opening the run is fatal by design, and the peers of a rank
# that exits without dist.abort() stay blocked on the first collective.
with mlflow_run(args, QUANTIZE):
main(args)
except BaseException:
dist.abort() # peers may be stuck in a collective this rank will never reach
finally:
dist.cleanup()