Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ Changelog

*Megatron Framework (M-LM / M-Bridge)*

- Add optional Top-P (nucleus) truncation to the Megatron top-k logits KD loss (now ``TopLogitsKLLoss``) via ``logit_kl_top_p`` and ``logit_kl_top_p_min_k`` in ``DistillationConfig``: after the global Top-K selection, only the smallest prefix whose cumulative teacher probability reaches ``top_p`` (with a floor of ``min_k`` entries) contributes to the KL, mirroring ``--logits-save-top-p`` / ``--logits-save-top-p-min-k`` in Megatron-LM's logits saver.
- Add an end-to-end W4A4 NVFP4 PTQ and QAD tutorial for Qwen3.6-35B-A3B also covering evaluation and vLLM throughput benchmarking. See `examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md <https://github.com/NVIDIA/Model-Optimizer/tree/main/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/>`_ for details.
- Add ``--mlflow <tracking-uri>`` to the ``examples/megatron_bridge`` scripts that write a checkpoint -- ``prune_minitron.py``, ``quantize.py``, ``distill.py``, ``export_quantized_megatron_to_hf.py`` and ``export_distilled_megatron_to_hf.py`` (MLflow's own ``MLFLOW_TRACKING_URI`` is honoured too). Each run records the invocation, its arguments as searchable params and its log -- a distillation records its training metrics instead of the rank-0 log -- and writes ``.experiment.json`` into the checkpoint it produced, so a pruning, a quantization, the distillation that refines its checkpoint and the export that deploys it can be traced to one another; uploading the checkpoints themselves stays off unless ``--mlflow_log_checkpoints`` is passed.
- Add ``--ep_size`` to ``examples/megatron_bridge/export_quantized_megatron_to_hf.py`` so large MoE models with grouped-GEMM experts can be exported with their experts sharded across GPUs. Checkpoints built with ``--no_moe_grouped_gemm`` must still be exported at ``--ep_size 1``.
Expand All @@ -52,6 +53,8 @@ Changelog

**Backward Breaking Changes**

- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) is renamed to ``TopLogitsKLLoss``, keeping the old name as a deprecated alias. It now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and always appends a "ghost" token holding the probability mass outside the Top-K to both student and teacher (matching Megatron-LM's offline cached-logits KD loss). Loss values change for existing ``logit_kl_topk`` runs.
- ``LogitsAndIntermediatesLossBalancer`` (Megatron distillation plugin) no longer rescales the distillation loss to the magnitude of the LM loss when the LM loss is included. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (in [0, 1]), matching Megatron-LM's offline cached-logits KD. The default ``kd_loss_alpha=1.0`` skips the LM loss, as before. ``DistillationConfig.skip_lm_loss`` and ``kd_loss_scale`` are removed and raise ``ValueError`` if passed; set ``kd_loss_alpha`` instead. In ``examples/megatron_bridge/distill.py``, ``--kd_loss_alpha`` replaces ``--no_skip_lm_loss`` and ``--kd_loss_scale``.
- The Megatron-Core DeepSeek-V4 indexer (``CSAIndexer``) is now a quantization module and persists its quantizer state in the checkpoint as ``indexer._extra_state``. A DeepSeek-V4 model quantized with an earlier release resumes from its ``torch_dist`` checkpoint only with a non-strict load (``--dist-ckpt-strictness log_unexpected`` in Megatron-LM) until it is saved again.
- ``examples/hf_ptq`` no longer detects MTP layers by name. Weights the loader could not place -- an MTP head, an auxiliary tower -- are identified from Transformers' own accounting: the model is loaded with ``from_pretrained(..., output_loading_info=True)`` and the reported ``unexpected_keys`` (present in the checkpoint, not in the model's architecture) are recorded on the model and carried into the export unchanged. Everything the loader *did* place goes through the normal export path. This removes ``load_mtp_weights``, ``mtp_layer_prefixes_from_checkpoint`` and their support matrix of MTP storage conventions, along with ``_add_mtp_exclusions`` and the pre-quantization ``enable: False`` entries ``hf_ptq`` appended to the recipe's ``quant_cfg``. Two consequences: MTP layers now follow the recipe like any other module instead of being force-excluded by the script -- matching ``examples/megatron_bridge``, which has no MTP-specific code at all -- and ``quantization_config.ignore`` can no longer claim a layer is unquantized that the export in fact quantized. Recipes importing ``configs/ptq/units/default_disabled_quantizers`` still disable ``mtp.*``, so their behaviour is unchanged; a recipe omitting that unit will now quantize an MTP the model actually built.

Expand Down
32 changes: 24 additions & 8 deletions examples/megatron_bridge/distill.py
Comment thread
AAnoosheh marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -208,21 +208,36 @@ def get_args():
"--train_iters", type=int, required=True, help="Number of training iterations"
)
parser.add_argument(
"--no_skip_lm_loss", action="store_true", help="Disable skipping language model loss"
"--kd_loss_alpha",
type=float,
default=1.0,
help="KD loss weight alpha in (1 - alpha) * lm_loss + alpha * kd_loss. 1.0 skips the LM loss entirely.",
)
parser.add_argument("--kd_loss_scale", type=float, default=1.0, help="KD loss weight")
parser.add_argument(
"--no_async_save",
action="store_true",
help="Save checkpoints synchronously. Async saving spawns a worker that needs its own "
"CUDA context, which fails when the training process already fills the GPU.",
)
parser.add_argument(
"--logit_kl_topk",
"--logit_kl_top_k",
type=int,
default=None,
help="Restrict the logit KL loss to the teacher's top-k vocabulary entries, "
"replacing the full-vocab temporaries with [seq, k] ones.",
help="Restrict the logit KL loss to the teacher's top-k vocabulary entries plus a residual "
"bucket for the remaining probability mass (distributions are still normalized over the full vocab).",
)
parser.add_argument(
"--logit_kl_top_p",
type=float,
default=None,
help="Nucleus threshold in (0, 1] applied on top of --logit_kl_top_k: only the smallest prefix "
"of the sorted top-k whose cumulative teacher probability reaches this value is distilled.",
)
parser.add_argument(
"--logit_kl_top_p_min_k",
type=int,
default=1,
help="Minimum number of top-k entries kept per token when --logit_kl_top_p is active.",
)
parser.add_argument("--lr", type=float, default=1e-4, help="Peak learning rate")
parser.add_argument("--min_lr", type=float, default=1e-5, help="Minimum learning rate")
Expand Down Expand Up @@ -474,9 +489,10 @@ def _build_model_provider(hf_path, load_weights=True, moe_grouped_gemm=True):
)

kd_config = ModelOptDistillConfig(
skip_lm_loss=not args.no_skip_lm_loss,
kd_loss_scale=args.kd_loss_scale,
logit_kl_topk=args.logit_kl_topk,
kd_loss_alpha=args.kd_loss_alpha,
logit_kl_topk=args.logit_kl_top_k,
logit_kl_top_p=args.logit_kl_top_p,
logit_kl_top_p_min_k=args.logit_kl_top_p_min_k,
Comment thread
AAnoosheh marked this conversation as resolved.
)
Comment on lines 491 to 496

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[IMPORTANT Compatibility] ModelOptDistillConfig is imported from the external megatron.bridge.training.post_training.distillation (line 49), not from ModelOpt. This PR renames/removes fields on ModelOpt's DistillationConfig and then passes the new names to the bridge's class, so the example is only correct for a megatron-bridge version that tracks those fields.

Two concrete ways this breaks across the repo boundary:

  • If bridge's ModelOptDistillConfig is a standalone dataclass mirroring the fields (rather than a subclass of modelopt...DistillationConfig), then kd_loss_alpha=, logit_kl_top_p= and logit_kl_top_p_min_k= raise TypeError: __init__() got an unexpected keyword argument against every currently released bridge — the example dies at startup.
  • kd_loss_scale kept its name but changed type from float = 1.0 to float | None = None. Any bridge-side or Megatron-LM-side code that reads cfg.kd_loss_scale and arithmetically uses it now gets None. (skip_lm_loss is safe, since __post_init__ still assigns a bool to it.)

Nothing inside ModelOpt reads kd_loss_scale, so this is purely a cross-repo question. Please confirm against the bridge version this repo pins: that ModelOptDistillConfig inherits (or has been updated to mirror) the new fields, and that no bridge code reads kd_loss_scale. If the bridge change has not landed yet, this example needs to wait on it or guard the new kwargs — and the CHANGELOG.rst breaking-change entry should name the minimum megatron-bridge version, since users upgrading ModelOpt alone will hit it.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Checked against the pinned Megatron-Bridge: ModelOptDistillConfig is declared as class ModelOptDistillConfig(mtd_mcore.DistillationConfig): pass, a field-less subclass. It inherits kd_loss_alpha, logit_kl_top_p, and logit_kl_top_p_min_k automatically, so the example works with any Bridge version that keeps that subclass, and no minimum-version note is needed. Nothing in Bridge's or Megatron-LM's source reads kd_loss_scale or skip_lm_loss. Passing either field now raises a ValueError with migration guidance, so nothing can silently read a None. The only remaining references are Bridge docs and an example YAML, plus two Megatron-LM functional-test configs. Each is fixed by deleting those keys, since the default kd_loss_alpha=1.0 reproduces skip_lm_loss: true. I'll follow up with companion PRs there.


# HF VLM configs expose ``vision_config``; Megatron-Bridge nests the text model under
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,7 @@ srun ... python -u /opt/Model-Optimizer/examples/megatron_bridge/distill.py \
--gbs 512 \
--train_iters 500 \
--lr 1e-5 --min_lr 1e-6 --lr_warmup_iters 50 \
--logit_kl_topk 4096 \
--logit_kl_top_k 4096 \
--recompute_granularity full --recompute_method uniform --recompute_num_layers 1 \
--no_async_save \
--eval_iters 0 \
Expand All @@ -190,7 +190,7 @@ Non-default arguments:
- `--tp_size 1 --pp_size 1` — **required, not chosen** (see below). `--ep_size 8` must match the PTQ checkpoint. You may increase `--cp_size` to enable context parallelism for longer sequence lengths (`nemo:26.10` container onwards).
- `--seq_length 32768 --gbs 512` — 16.8M tokens/iteration, 1.7B per 100 iterations.
- `--lr 1e-5 --min_lr 1e-6` — an order of magnitude below typical distillation LRs: the job is to adapt weights to quantization, not to learn the task.
- `--logit_kl_topk 4096` — restricts the KD loss to the teacher's top-4096 vocab entries. With a 248,320-token vocabulary the dense `[seq, vocab]` fp32 logits are **30.31 GiB per tensor** at 32K, which OOMs on its own.
- `--logit_kl_top_k 4096` — restricts the KD loss to the teacher's top-4096 vocab entries. With a 248,320-token vocabulary the dense `[seq, vocab]` fp32 logits are **30.31 GiB per tensor** at 32K, which OOMs on its own.
- `--recompute_*` / `--no_async_save` / `--eval_iters 0` — all needed to fit. Async save spawns a worker needing its own CUDA context; the validation path computes full-vocab LM and MTP cross-entropy (top-k applies to training only), so eval OOMs at 32K even though training fits.

</details>
Expand Down Expand Up @@ -322,7 +322,7 @@ It is not verbosity. It is a **failure to terminate on a small fraction of sub-s
- The **median** also roughly doubles (+91.6%), so the whole distribution shifted right — this is not *only* a tail effect.
- Capped rate peaks at **iteration 50** (4.3%) and settles at 3.1% / 3.6% by 300 / 500; it is not gradual drift.

The obvious suspect — that `--logit_kl_topk 4096` leaves the stop tokens outside the loss — **did not hold up**. Probing the BF16 teacher over one runaway trace: `</think>` does fall outside top-4096 at 35% of positions overall, but *in the looping region* the teacher gives `<|im_end|>` a median rank of **5** and `</think>` ~570, both well inside top-k. The teacher is signalling "stop here" at positions the loss did cover, and the student still does not stop. More likely: the blend has few "the answer is written, now stop" positions in this style, and a teacher-forced loss never exercises free-running generation 10K+ tokens deep.
The obvious suspect — that `--logit_kl_top_k 4096` leaves the stop tokens outside the loss — **did not hold up**. Probing the BF16 teacher over one runaway trace: `</think>` does fall outside top-4096 at 35% of positions overall, but *in the looping region* the teacher gives `<|im_end|>` a median rank of **5** and `</think>` ~570, both well inside top-k. The teacher is signalling "stop here" at positions the loss did cover, and the student still does not stop. More likely: the blend has few "the answer is written, now stop" positions in this style, and a teacher-forced loss never exercises free-running generation 10K+ tokens deep.

**It is fixable at decode time.** Adding `presence_penalty: 1.5` (Qwen's own thinking-mode recommendation for this model) removes nearly all of it, with no retraining. Every cell is SciCode **without → with** the penalty, 8 runs per side. *Capped* = hit the 131,072-token limit with no stop token; almost all such sub-steps return nothing and score zero. Counts are pooled over all 8 runs, so the denominator is 338 × 8 = 2,704 sub-steps:

Expand Down
Loading
Loading