From b4112bdf82978c4318e74990c50d71d489a948d3 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 17 Sep 2026 19:01:48 +0200 Subject: [PATCH 01/11] Bring offline KD upgrades such as Ghost Token and Top-P to Megatron KD plugin Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 4 + examples/megatron_bridge/distill.py | 42 +++- modelopt/torch/distill/plugins/megatron.py | 214 ++++++++++++------ .../distill/plugins/test_distill_megatron.py | 190 +++++++++++++++- 4 files changed, 376 insertions(+), 74 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 46168b7fe0f..8b78b03a004 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -15,6 +15,7 @@ Changelog *Megatron Framework (M-LM / M-Bridge)* +- Add optional Top-P (nucleus) truncation to the Megatron ``TopKLogitsKLLoss`` 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 `_ for details. *Misc* @@ -23,6 +24,8 @@ Changelog **Backward Breaking Changes** +- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and by default 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; set ``logit_kl_ghost_token: false`` to drop the ghost token. +- ``LogitsAndIntermediatesLossBalancer`` (Megatron distillation plugin) no longer rescales the distillation loss to the magnitude of the LM loss. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (default ``0.9``, in [0, 1]), matching Megatron-LM's offline cached-logits KD. ``skip_lm_loss`` is now derived from ``kd_loss_alpha`` (skipped iff ``1.0``), so the LM loss is computed by default where it was previously skipped. ``examples/megatron_bridge/distill.py`` gains ``--kd_loss_alpha``. - Layerwise calibration now uses prior-layer QDQ activations by default (``layerwise.get_qdq_activations_from_prev_layer=True``). Set it to ``False`` to preserve full-precision activations for subsequent layers (the default behavior for @@ -32,6 +35,7 @@ Changelog **Deprecations** +- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``DeprecationWarning``; a user-provided ``skip_lm_loss`` is overridden by the value derived from ``kd_loss_alpha`` with a ``DeprecationWarning``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. - Rename the architecture-specific recipe tier from ``modelopt_recipes/huggingface/`` to ``modelopt_recipes/model_type/`` to clarify that it holds recipes shared across every checkpoint of a Hugging Face ``model_type``. Saved ``--recipe huggingface//...`` paths still resolve via a backward-compatibility alias but now emit a ``FutureWarning``, so update them to ``model_type//...`` as the ``huggingface/`` prefix is deprecated. - The single-format quantization CLI flags are deprecated in favour of ``--recipe`` and will be removed in a future release; passing one now emits a ``FutureWarning``. ``examples/hf_ptq``: ``--qformat`` and ``--kv_cache_qformat``. ``examples/megatron_bridge/quantize.py``: ``--quant_cfg``, ``--kv_cache_quant`` and ``--weight_only``. ``examples/torch_onnx/torch_quant_to_onnx.py``: ``--qformat``. A recipe carries the quantization config, the calibration algorithm and the KV-cache setting in one file, so they cannot drift apart the way separate flags can -- and ``--recipe`` already took precedence over all six, silently on ``hf_ptq`` and with a warning on ``megatron_bridge`` -- with one gap the recipe closes rather than inherits: a weight AutoQuantize recipe that omits ``kv_cache`` still falls back to ``--kv_cache_qformat``, so set ``kv_cache`` in the recipe when migrating. Use a recipe from ``modelopt_recipes/general/ptq/``, an architecture-specific one under ``modelopt_recipes/model_type//``, or a checkpoint-specific one under ``modelopt_recipes/models/``. The warning fires only when a flag is passed explicitly: ``--qformat`` defaults to ``fp8`` and ``--kv_cache_qformat`` to ``fp8_cast``, so warning on the defaults would fire on every run, including runs that correctly use ``--recipe``. ``examples/speculative_decoding/scripts/quantize_drafter.py`` keeps ``--qformat`` undeprecated: it has no ``--recipe`` alternative yet. - The TensorRT-LLM checkpoint export format is deprecated and will be removed in 0.49.0: ``export_tensorrt_llm_checkpoint`` and ``torch_to_tensorrt_llm_checkpoint`` now emit a ``DeprecationWarning`` on use. Use ``export_hf_checkpoint``, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang. Its implementation moved to ``modelopt.torch.export.trtllm``, so import those two functions from there and the ``ModelConfig`` dataclasses from ``modelopt.torch.export.trtllm.model_config``; both functions remain importable from ``modelopt.torch.export`` for this release only. diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index e1810c19d5e..89b44e6bc5f 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -172,9 +172,23 @@ 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" + "--no_skip_lm_loss", + action="store_true", + help="DEPRECATED and ignored. Whether the LM loss is skipped is derived from --kd_loss_alpha " + "(skipped iff alpha == 1.0).", + ) + parser.add_argument( + "--kd_loss_alpha", + type=float, + default=0.9, + 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=None, + help="DEPRECATED and ignored. Use --kd_loss_alpha.", ) - parser.add_argument("--kd_loss_scale", type=float, default=1.0, help="KD loss weight") parser.add_argument( "--no_async_save", action="store_true", @@ -188,6 +202,24 @@ def get_args(): help="Restrict the logit KL loss to the teacher's top-k vocabulary entries, " "replacing the full-vocab temporaries with [seq, k] ones.", ) + parser.add_argument( + "--logit_kl_top_p", + type=float, + default=None, + help="Nucleus threshold in (0, 1] applied on top of --logit_kl_topk: 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( + "--no_logit_kl_ghost_token", + action="store_true", + help="Disable the residual 'ghost' token (out-of-top-k probability mass) in the top-k KL loss.", + ) 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") parser.add_argument("--lr_warmup_iters", type=int, default=50, help="Number of LR warmup steps") @@ -435,9 +467,11 @@ 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, + kd_loss_alpha=args.kd_loss_alpha, logit_kl_topk=args.logit_kl_topk, + logit_kl_top_p=args.logit_kl_top_p, + logit_kl_top_p_min_k=args.logit_kl_top_p_min_k, + logit_kl_ghost_token=not args.no_logit_kl_ghost_token, ) # HF VLM configs expose ``vision_config``; Megatron-Bridge nests the text model under diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index c93f0961d1f..c6c5af8ab07 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -19,6 +19,7 @@ import logging import re +import warnings from abc import ABCMeta from collections.abc import Callable from dataclasses import dataclass, field @@ -56,18 +57,32 @@ class DistillationConfig: Args: intermediate_layer_pairs: List of tuples of intermediate layer names. logit_layers: Tuple of logit layer names. - skip_lm_loss: Whether to skip computing the standard language model loss (default: ``True``). - kd_loss_scale: Relative scaling factor for the distillation loss if ``skip_lm_loss`` is ``False``. + kd_loss_alpha: Weight of the distillation loss in the convex combination + ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. When ``1.0``, the standard + language model loss is skipped entirely (``skip_lm_loss`` is derived from this value). + skip_lm_loss: DEPRECATED. Derived from ``kd_loss_alpha`` (``True`` iff ``kd_loss_alpha == 1.0``); + any user-provided value is overridden with a warning. + kd_loss_scale: DEPRECATED and ignored. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. + logit_kl_top_p: Optional nucleus (top-P) threshold applied on top of the teacher's Top-K. + Only the smallest prefix of the (sorted) Top-K whose cumulative teacher probability + reaches this value contributes to the loss. Requires ``logit_kl_topk``. Must be in (0, 1]. + logit_kl_top_p_min_k: Minimum number of Top-K entries kept per token when top-P is active. + logit_kl_ghost_token: Whether ``TopKLogitsKLLoss`` appends a "ghost" token holding the + probability mass outside the kept entries to both distributions (default: ``True``). """ intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) logit_layers: tuple[str, str] = ("output_layer", "output_layer") - skip_lm_loss: bool = True - kd_loss_scale: float = 1.0 + kd_loss_alpha: float = 0.9 + skip_lm_loss: bool | None = None # deprecated, derived from kd_loss_alpha + kd_loss_scale: float | None = None # deprecated, ignored logit_kl_temperature: float = 1.0 logit_kl_topk: int | None = None + logit_kl_top_p: float | None = None + logit_kl_top_p_min_k: int = 1 + logit_kl_ghost_token: bool = True criterion: Criterion | None = None loss_balancer: mtd.DistillationLossBalancer | None = None @@ -76,8 +91,30 @@ def __post_init__(self): assert all(len(pair) in (2, 3) for pair in self.intermediate_layer_pairs), ( f"{self.intermediate_layer_pairs=}" ) - assert self.kd_loss_scale > 0, f"{self.kd_loss_scale=}" + assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" + if self.kd_loss_scale is not None: + warnings.warn( + "DistillationConfig.kd_loss_scale is deprecated and ignored. The distillation loss " + "is no longer rescaled to the LM loss magnitude; the total loss is now " + "(1 - kd_loss_alpha) * lm_loss + kd_loss_alpha * kd_loss. Set `kd_loss_alpha` instead.", + DeprecationWarning, + stacklevel=2, + ) + derived_skip_lm_loss = self.kd_loss_alpha == 1.0 + if self.skip_lm_loss is not None: + warnings.warn( + "DistillationConfig.skip_lm_loss is deprecated and is now derived from `kd_loss_alpha` " + f"(skip iff kd_loss_alpha == 1.0). Overriding skip_lm_loss={self.skip_lm_loss} with " + f"{derived_skip_lm_loss} (kd_loss_alpha={self.kd_loss_alpha}).", + DeprecationWarning, + stacklevel=2, + ) + self.skip_lm_loss = derived_skip_lm_loss assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" + if self.logit_kl_top_p is not None: + assert self.logit_kl_topk is not None, "logit_kl_top_p requires logit_kl_topk" + assert 0 < self.logit_kl_top_p <= 1, f"{self.logit_kl_top_p=}" + assert self.logit_kl_top_p_min_k >= 1, f"{self.logit_kl_top_p_min_k=}" @staticmethod def parse_intermediate_entry(entry: tuple[str, ...]) -> tuple[str, str, Callable]: @@ -130,7 +167,12 @@ def setup_distillation_config( # Use TopKLogitsKLLoss if logit_kl_topk is specified, otherwise use LogitsKLLoss if cfg.logit_kl_topk is not None: criterion[tuple(cfg.logit_layers)] = TopKLogitsKLLoss( - student_cfg, temperature=cfg.logit_kl_temperature, top_k=cfg.logit_kl_topk + student_cfg, + temperature=cfg.logit_kl_temperature, + top_k=cfg.logit_kl_topk, + top_p=cfg.logit_kl_top_p, + top_p_min_k=cfg.logit_kl_top_p_min_k, + add_ghost_token=cfg.logit_kl_ghost_token, ) else: criterion[tuple(cfg.logit_layers)] = LogitsKLLoss( @@ -156,7 +198,8 @@ def setup_distillation_config( if cfg.loss_balancer is None: cfg.loss_balancer = LogitsAndIntermediatesLossBalancer( - kd_loss_scale=cfg.kd_loss_scale, skip_original_loss=cfg.skip_lm_loss + kd_loss_alpha=cfg.kd_loss_alpha, + skip_original_loss=bool(cfg.skip_lm_loss), # always set by __post_init__ ) return cfg @@ -319,56 +362,47 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: """ predictions, targets = self.pre_forward(predictions, targets) - # Division by temp should happen prior to finding max for both student and teacher. - output_teacher = targets.float() / self._temperature - output_student = predictions.float() / self._temperature + # Temperature-scaled log probabilities (log softmax), globally normalized across TP vocab shards. + p = predictions.float() / self._temperature - self._tp_logsumexp(predictions) + q = targets.float() / self._temperature - self._tp_logsumexp(targets) - # Compute local softmax, and the reweight to compute global softmax. - if self._config.tensor_model_parallel_size > 1: - tp_group = parallel_state.get_tensor_model_parallel_group() + # KL divergence + if self._reverse: + p, q = q, p + loss = torch.sum(F.kl_div(p, q, reduction="none", log_target=True), dim=-1) - # Subtract maximum value along vocab dimension across all GPUs (for stability) - teacher_logits_max, _ = torch.max(output_teacher, dim=-1, keepdim=True) - torch.distributed.all_reduce( - teacher_logits_max, - op=torch.distributed.ReduceOp.MAX, - group=tp_group, - ) - output_teacher -= teacher_logits_max + return self.post_forward(loss, tp_reduce=True) - student_logits_max, _ = torch.max(output_student, dim=-1, keepdim=True) + def _tp_logsumexp(self, logits: Tensor, num_chunks: int = 4) -> Tensor: + """Log-sum-exp of ``logits / self._temperature`` over the vocab dim across all TP shards. + + Accumulates in fp32 over ``num_chunks`` vocab chunks so the transient fp32 working set is a + fraction of the (possibly bf16) full-vocab input. NOTE: for inputs requiring grad, autograd + still retains each chunk's ``exp`` output for backward, so the total saved activation + equals a full-vocab fp32 tensor. Returns shape ``[..., 1]``. + """ + # Max is exact under the monotonic temperature scaling, so take it in the native dtype + # and defer the temperature division to the centered values: lse(x/T) = max/T + log(sum(exp((x-max)/T))). + logits_max = logits.amax(dim=-1, keepdim=True).float() + if self._config.tensor_model_parallel_size > 1: + tp_group = parallel_state.get_tensor_model_parallel_group() torch.distributed.all_reduce( - student_logits_max, - op=torch.distributed.ReduceOp.MAX, - group=tp_group, + logits_max, op=torch.distributed.ReduceOp.MAX, group=tp_group ) - output_student -= student_logits_max.detach() - - # Compute global softmax denominators + logits_max = logits_max.detach() + + denom = None + chunk_size = -(-logits.size(-1) // num_chunks) # ceil division + for chunk in logits.split(chunk_size, dim=-1): + centered = (chunk.float() - logits_max) / self._temperature + partial = torch.exp(centered).sum(dim=-1, keepdim=True) + denom = partial if denom is None else denom + partial + if self._config.tensor_model_parallel_size > 1: # We can't use standard all_reduce function here since the computation # that follows it isn't identical across TP ranks. - denom_teacher = torch.sum(torch.exp(output_teacher), dim=-1, keepdim=True) - denom_teacher = dist_nn.functional.all_reduce(denom_teacher, group=tp_group) + denom = dist_nn.functional.all_reduce(denom, group=tp_group) - denom_student = torch.sum(torch.exp(output_student), dim=-1, keepdim=True) - denom_student = dist_nn.functional.all_reduce(denom_student, group=tp_group) - - # Compute log probabilities (log softmax) - teacher_log_prob = output_teacher - torch.log(denom_teacher) - student_log_prob = output_student - torch.log(denom_student) - - # KL divergence - p, q = student_log_prob, teacher_log_prob - else: - # Compute log probabilities - p, q = F.log_softmax(output_student, dim=-1), F.log_softmax(output_teacher, dim=-1) - - # KL divergence - if self._reverse: - p, q = q, p - loss = torch.sum(F.kl_div(p, q, reduction="none", log_target=True), dim=-1) - - return self.post_forward(loss, tp_reduce=True) + return logits_max / self._temperature + torch.log(denom) class TopKLogitsKLLoss(LogitsKLLoss): @@ -376,6 +410,16 @@ class TopKLogitsKLLoss(LogitsKLLoss): Calculates using the global Top-K entries without gathering full logits. NOTE: Will gather Top-K logits per rank, so mind the value of K for memory and communication. + + Both distributions are normalized over the *full* vocabulary (not re-normalized over the + Top-K), matching the offline cached-logits KD loss in Megatron-LM. Optional refinements: + + * **Top-P (nucleus)**: after sorting the Top-K by teacher probability, only the smallest prefix + whose cumulative teacher mass reaches ``top_p`` (with a floor of ``top_p_min_k`` entries) + contributes to the loss. + * **Ghost token**: a synthetic extra entry holding the probability mass outside the kept + entries, ``log(1 - sum(kept probs))``, is appended to both student and teacher so the loss + also penalizes mass the student places outside the teacher's nucleus. """ def __init__( @@ -384,6 +428,10 @@ def __init__( temperature: float = 1.0, reverse: bool = False, top_k: int = 1024, + *, + top_p: float | None = None, + top_p_min_k: int = 1, + add_ghost_token: bool = True, ): """Constructor. @@ -392,9 +440,19 @@ def __init__( temperature: Divide tensors by this value prior to calculating loss. reverse: Whether to reverse the loss as KLD(teacher, student) instead of KLD(student, teacher) top_k: The number of top vocabulary entries to keep from the teacher's distribution. + top_p: Optional nucleus threshold in (0, 1] applied on top of the Top-K selection. + top_p_min_k: Minimum number of entries kept per token when ``top_p`` is active. + add_ghost_token: Whether to append a residual "ghost" token holding the out-of-Top-K + probability mass to both distributions. """ super().__init__(model_config, temperature, reverse) + assert top_k >= 1, f"{top_k=}" + assert top_p is None or 0 < top_p <= 1, f"{top_p=}" + assert top_p_min_k >= 1, f"{top_p_min_k=}" self.top_k = top_k + self.top_p = top_p + self.top_p_min_k = top_p_min_k + self.add_ghost_token = add_ghost_token def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: """Forward function. @@ -445,14 +503,39 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: final_teacher_logits = top_teacher_vals final_student_logits = top_student_vals - # Standard (dense) Softmax + KL - p = F.log_softmax(final_student_logits, dim=-1) - q = F.log_softmax(final_teacher_logits, dim=-1) - - # KL divergence + # Log-probs of the Top-K entries under the full-vocab distributions, using global + # (full-vocab) log-normalizers so the entries carry true probabilities. + # NOTE: ``torch.topk`` returns entries sorted descending by teacher value. + teacher_logp = final_teacher_logits - self._tp_logsumexp(targets) + student_logp = final_student_logits - self._tp_logsumexp(predictions) + + # Top-P (nucleus) mask over the sorted Top-K: keep entry i iff cumulative mass *before* it + # is < p. This always keeps the entry that crosses the threshold (and thus top-1). + if self.top_p is not None: + teacher_probs = teacher_logp.exp() + mask = (teacher_probs.cumsum(dim=-1) - teacher_probs) < self.top_p + min_keep = min(self.top_p_min_k, teacher_logp.size(-1)) + mask |= torch.arange(teacher_logp.size(-1), device=mask.device) < min_keep + else: + mask = torch.ones_like(teacher_logp, dtype=torch.bool) + + # Ghost token: residual probability mass outside the kept entries, for both distributions. + if self.add_ghost_token: + eps = 1e-8 + student_kept_mass = (student_logp.exp() * mask).sum(dim=-1, keepdim=True) + teacher_kept_mass = (teacher_logp.exp() * mask).sum(dim=-1, keepdim=True) + student_residual = torch.log((1.0 - student_kept_mass).clamp(min=eps)) + teacher_residual = torch.log((1.0 - teacher_kept_mass).clamp(min=eps)) + student_logp = torch.cat([student_logp, student_residual], dim=-1) + teacher_logp = torch.cat([teacher_logp, teacher_residual], dim=-1) + mask = torch.cat([mask, mask.new_ones((*mask.shape[:-1], 1))], dim=-1) + + # Sparse KL divergence: sum_i q_i * (log q_i - log p_i) over kept entries. + p, q = student_logp, teacher_logp if self._reverse: p, q = q, p - loss = torch.sum(F.kl_div(p, q, reduction="none", log_target=True), dim=-1) + kl = q.exp() * (q - p) + loss = torch.sum(mask * kl, dim=-1) # No need to reduce since all ranks compute same global Top-K return self.post_forward(loss, tp_reduce=False) @@ -461,20 +544,23 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: class LogitsAndIntermediatesLossBalancer(mtd.DistillationLossBalancer): """LossBalancer implementation for Logit and Intermediate losses. - Dynamically weighs distillation and original losses to balance during training. + Intermediate losses are dynamically rescaled to the magnitude of the logits loss, then the + total distillation loss is combined with the original LM loss as a fixed convex combination + ``(1 - alpha) * lm_loss + alpha * kd_loss`` (matching Megatron-LM's offline cached-logits KD). """ - def __init__(self, kd_loss_scale: float = 1.0, skip_original_loss: bool = False): + def __init__(self, kd_loss_alpha: float = 0.9, skip_original_loss: bool = False): """Constructor. Args: - kd_loss_scale: Multiply distillation losses by this before weighing. - (Not used when `skip_original_loss` is True.) + kd_loss_alpha: Weight of the distillation loss in ``(1 - alpha) * lm + alpha * kd``. + Must be in [0, 1]. (Not used when `skip_original_loss` is True.) skip_original_loss: Used to signal whether the original loss should be used, regardless of whether it was passed into ``mtd.DistillationModel.compute_kd_loss()`` or not. """ super().__init__() - self._kd_loss_scale = kd_loss_scale + assert 0 <= kd_loss_alpha <= 1, f"{kd_loss_alpha=}" + self._kd_loss_alpha = kd_loss_alpha self._skip_original_loss = skip_original_loss def forward(self, loss_dict: dict[str, Tensor]) -> Tensor: @@ -500,13 +586,11 @@ def forward(self, loss_dict: dict[str, Tensor]) -> Tensor: intermediate_loss = logits_loss.new_tensor(intermediate_loss) intermediate_loss_scaled = intermediate_loss + kd_loss = logits_loss + intermediate_loss_scaled if self._skip_original_loss: - total_loss = logits_loss + intermediate_loss_scaled + total_loss = kd_loss else: - kd_loss = logits_loss + intermediate_loss_scaled - if kd_loss > 0 and original_loss > 0: # zero when one CP rank has only context tokens - kd_loss *= original_loss.detach() / kd_loss.detach() - total_loss = original_loss + kd_loss * self._kd_loss_scale + total_loss = (1 - self._kd_loss_alpha) * original_loss + self._kd_loss_alpha * kd_loss out_dict = { "kd_loss": total_loss, diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 7419bb49f86..7f381ae4fec 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -14,9 +14,12 @@ # limitations under the License. from functools import partial +from types import SimpleNamespace +import pytest import torch import torch.nn as nn +import torch.nn.functional as F from _test_utils.torch.megatron.models import get_mcore_gpt_model from _test_utils.torch.megatron.utils import run_mcore_inference_with_dummy_input from _test_utils.torch.misc import set_seed @@ -24,6 +27,9 @@ import modelopt.torch.distill as mtd from modelopt.torch.distill.plugins.megatron import ( DistillationConfig, + LogitsAndIntermediatesLossBalancer, + LogitsKLLoss, + TopKLogitsKLLoss, _mtp_excluded_from_quantization, adjust_distillation_model_for_mcore, setup_distillation_config, @@ -124,7 +130,7 @@ def _test_logits_kl_loss(rank, size): loss["kd_loss"].backward() -def _test_topk_logits_kl_loss(top_k, rank, size): +def _test_topk_logits_kl_loss(kd_kwargs, rank, size): """Test TopKLogitsKLLoss with simple forward/backward pass.""" set_seed(SEED) @@ -169,7 +175,7 @@ def _test_topk_logits_kl_loss(top_k, rank, size): # Setup distillation config with TopKLogitsKLLoss via logit_kl_topk argument distill_cfg = setup_distillation_config( - config_or_path=DistillationConfig(logit_kl_topk=top_k), + config_or_path=DistillationConfig(**kd_kwargs), student_cfg=student_model.config, teacher_cfg=teacher_model.config, ) @@ -211,6 +217,12 @@ def _test_topk_logits_kl_loss(top_k, rank, size): assert isinstance(loss, dict), "Loss should be a dictionary" assert "kd_loss" in loss, "Should contain kd_loss key" + # All TP ranks operate on the same global Top-K, so the loss must be identical across ranks. + gathered = [torch.empty_like(loss["kd_loss"]) for _ in range(size)] + torch.distributed.all_gather(gathered, loss["kd_loss"].detach()) + for other in gathered[1:]: + assert torch.allclose(gathered[0], other), "Top-K KD loss differs across TP ranks" + # Backward pass loss["kd_loss"].backward() @@ -260,7 +272,7 @@ def _test_skip_lm_loss_with_mtp(rank, size): ).cuda() distill_cfg = setup_distillation_config( - config_or_path=DistillationConfig(skip_lm_loss=True), + config_or_path=DistillationConfig(kd_loss_alpha=1.0), # skips LM loss student_cfg=student_model.config, teacher_cfg=teacher_model.config, ) @@ -313,9 +325,136 @@ def test_logits_kl_loss(dist_workers): dist_workers.run(_test_logits_kl_loss) -def test_topk_logits_kl_loss(dist_workers, top_k: int = 5): +@pytest.mark.parametrize( + ("top_p", "top_p_min_k", "ghost_token"), + [(None, 1, True), (None, 1, False), (0.9, 1, True), (0.9, 3, False)], +) +def test_topk_logits_kl_loss(dist_workers, top_p, top_p_min_k, ghost_token, top_k: int = 5): """Test TopKLogitsKLLoss with TP parallelism.""" - dist_workers.run(partial(_test_topk_logits_kl_loss, top_k)) + kd_kwargs = { + "logit_kl_topk": top_k, + "logit_kl_top_p": top_p, + "logit_kl_top_p_min_k": top_p_min_k, + "logit_kl_ghost_token": ghost_token, + } + dist_workers.run(partial(_test_topk_logits_kl_loss, kd_kwargs)) + + +def _make_loss_inputs(seq=4, batch=3, vocab=16): + torch.manual_seed(SEED) + student = torch.randn(seq, batch, vocab, requires_grad=True) + teacher = torch.randn(seq, batch, vocab) * 3 # peaky teacher so top-P actually truncates + return student, teacher + + +def test_topk_logits_kl_loss_numerics_full_vocab_matches_dense(): + """With K = vocab and ghost token, Top-K KL equals the dense full-vocab KL (residual ~0).""" + cfg = SimpleNamespace(tensor_model_parallel_size=1) + student, teacher = _make_loss_inputs() + dense = LogitsKLLoss(cfg)(student, teacher)[0] + topk = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=True)(student, teacher)[0] + assert torch.allclose(dense, topk, atol=1e-5) + # Without ghost token, the unnormalized Top-K KL over the full vocab is also the dense KL. + topk_no_ghost = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=False)( + student, teacher + )[0] + assert torch.allclose(dense, topk_no_ghost, atol=1e-5) + + +def test_topk_logits_kl_loss_numerics_ghost_token_reference(): + """Top-K + ghost token matches a hand-written reference on the K+1 bucketed distributions.""" + cfg = SimpleNamespace(tensor_model_parallel_size=1) + student, teacher = _make_loss_inputs() + k = 4 + loss = TopKLogitsKLLoss(cfg, top_k=k, add_ghost_token=True)(student, teacher)[0] + + q_full = F.log_softmax(teacher, dim=-1) + p_full = F.log_softmax(student, dim=-1) + _, idx = torch.topk(teacher, k, dim=-1) + q_k, p_k = q_full.gather(-1, idx), p_full.gather(-1, idx) + q_rest = torch.log1p(-q_k.exp().sum(-1, keepdim=True)) + p_rest = torch.log1p(-p_k.exp().sum(-1, keepdim=True)) + q = torch.cat([q_k, q_rest], -1) + p = torch.cat([p_k, p_rest], -1) + ref = (q.exp() * (q - p)).sum(-1).transpose(0, 1) + assert torch.allclose(loss, ref, atol=1e-5) + # Sanity: total mass within the K+1 buckets is 1 for both distributions. + assert torch.allclose(q.exp().sum(-1), torch.ones_like(q[..., 0]), atol=1e-5) + assert torch.allclose(p.exp().sum(-1), torch.ones_like(p[..., 0]), atol=1e-5) + + +@pytest.mark.parametrize("temperature", [0.5, 2.0, 3.7]) +def test_logits_kl_losses_temperature_scaling(temperature): + """Dense and Top-K losses match a plain ``log_softmax(x / T)`` reference at T != 1.""" + cfg = SimpleNamespace(tensor_model_parallel_size=1) + student, teacher = _make_loss_inputs() + q = F.log_softmax(teacher / temperature, dim=-1) + p = F.log_softmax(student / temperature, dim=-1) + + dense = LogitsKLLoss(cfg, temperature=temperature)(student, teacher)[0] + ref_dense = (q.exp() * (q - p)).sum(-1).transpose(0, 1) + assert torch.allclose(dense, ref_dense, atol=1e-5) + + k = 4 + topk = TopKLogitsKLLoss(cfg, temperature=temperature, top_k=k, add_ghost_token=True)( + student, teacher + )[0] + _, idx = torch.topk(teacher, k, dim=-1) + q_k, p_k = q.gather(-1, idx), p.gather(-1, idx) + q_rest = torch.log1p(-q_k.exp().sum(-1, keepdim=True)) + p_rest = torch.log1p(-p_k.exp().sum(-1, keepdim=True)) + qq = torch.cat([q_k, q_rest], -1) + pp = torch.cat([p_k, p_rest], -1) + ref_topk = (qq.exp() * (qq - pp)).sum(-1).transpose(0, 1) + assert torch.allclose(topk, ref_topk, atol=1e-5) + + +def test_topk_logits_kl_loss_top_p_masks_tail(): + """Top-P zeroes out-of-nucleus entries and honors the min_k floor.""" + cfg = SimpleNamespace(tensor_model_parallel_size=1) + student, teacher = _make_loss_inputs() + k = 8 + q_full = F.log_softmax(teacher, dim=-1) + q_k, idx = torch.topk(q_full, k, dim=-1) + p_k = F.log_softmax(student, dim=-1).gather(-1, idx) + probs = q_k.exp() + keep = (probs.cumsum(-1) - probs) < 0.5 + + # No ghost token: loss is exactly the masked partial KL sum. + loss = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5, add_ghost_token=False)(student, teacher)[0] + ref = (keep * probs * (q_k - p_k)).sum(-1).transpose(0, 1) + assert torch.allclose(loss, ref, atol=1e-5) + assert not keep.all(), "test inputs should produce some truncation" + + # min_k floor forces at least min_k entries even when nucleus is tiny. + min_k = 3 + loss_min = TopKLogitsKLLoss(cfg, top_k=k, top_p=1e-6, top_p_min_k=min_k, add_ghost_token=False)( + student, teacher + )[0] + ref_min = ((torch.arange(k) < min_k) * probs * (q_k - p_k)).sum(-1).transpose(0, 1) + assert torch.allclose(loss_min, ref_min, atol=1e-5) + + # Ghost token with top-P: residual is mass outside the kept nucleus, distributions sum to 1. + loss_ghost = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5, add_ghost_token=True)(student, teacher)[ + 0 + ] + q_rest = torch.log1p(-(probs * keep).sum(-1, keepdim=True)) + p_rest = torch.log1p(-(p_k.exp() * keep).sum(-1, keepdim=True)) + ref_ghost = ref + (q_rest.exp() * (q_rest - p_rest)).sum(-1).transpose(0, 1) + assert torch.allclose(loss_ghost, ref_ghost, atol=1e-5) + assert loss_ghost.shape == (student.size(1), student.size(0)) + loss_ghost.sum().backward() + assert student.grad is not None and torch.isfinite(student.grad).all() + + +def test_distillation_config_top_p_validation(): + with pytest.raises(AssertionError): + DistillationConfig(logit_kl_top_p=0.9) # requires logit_kl_topk + with pytest.raises(AssertionError): + DistillationConfig(logit_kl_topk=8, logit_kl_top_p=1.5) + with pytest.raises(AssertionError): + DistillationConfig(logit_kl_topk=8, logit_kl_top_p=0.9, logit_kl_top_p_min_k=0) + DistillationConfig(logit_kl_topk=8, logit_kl_top_p=1.0, logit_kl_top_p_min_k=2) def test_skip_lm_loss_with_mtp(dist_workers): @@ -356,3 +495,44 @@ def _model(*, with_mtp: bool, body_quant: bool, mtp_quant: bool) -> nn.Module: assert not _mtp_excluded_from_quantization( _model(with_mtp=False, body_quant=True, mtp_quant=False) ) + + +def test_loss_balancer_convex_combination(): + """Total loss is (1 - alpha) * lm + alpha * (logits + rescaled intermediate).""" + lm = torch.tensor(2.0) + logits = torch.tensor(0.5) + inter = torch.tensor(4.0) # rescaled to logits magnitude -> contributes 0.5 + key = mtd.loss_balancers.STUDENT_LOSS_KEY + + out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=0.25)( + {key: lm, "LogitsKLLoss_0": logits, "HiddenStateCosineLoss_0": inter} + ) + assert torch.allclose(out["kd_loss"], torch.tensor(0.75 * 2.0 + 0.25 * (0.5 + 0.5))) + assert torch.allclose(out["logits_loss"], logits) + + # alpha=1 ignores the LM loss entirely; skip_original_loss does the same regardless of alpha. + out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)({key: lm, "LogitsKLLoss_0": logits}) + assert torch.allclose(out["kd_loss"], logits) + out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=0.0, skip_original_loss=True)( + {key: lm, "LogitsKLLoss_0": logits} + ) + assert torch.allclose(out["kd_loss"], logits) + + with pytest.raises(AssertionError): + LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.5) + with pytest.raises(AssertionError): + DistillationConfig(kd_loss_alpha=-0.1) + + +def test_distillation_config_deprecations(): + """skip_lm_loss is derived from kd_loss_alpha; legacy fields warn.""" + assert DistillationConfig(kd_loss_alpha=1.0).skip_lm_loss is True + assert DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss is False + + with pytest.warns(DeprecationWarning, match="skip_lm_loss is deprecated"): + cfg = DistillationConfig(kd_loss_alpha=0.9, skip_lm_loss=True) + assert cfg.skip_lm_loss is False # user value overridden + + with pytest.warns(DeprecationWarning, match="kd_loss_scale is deprecated"): + cfg = DistillationConfig(kd_loss_scale=2.0) + assert cfg.kd_loss_alpha == 0.9 From 95c9aa7dad1072927544f3962a249144b4539729 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 17 Sep 2026 19:10:18 +0200 Subject: [PATCH 02/11] Don't expose ghost token in MBridge script Signed-off-by: Asha Anoosheh --- examples/megatron_bridge/distill.py | 6 ------ 1 file changed, 6 deletions(-) diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index 89b44e6bc5f..268dc855b3a 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -215,11 +215,6 @@ def get_args(): default=1, help="Minimum number of top-k entries kept per token when --logit_kl_top_p is active.", ) - parser.add_argument( - "--no_logit_kl_ghost_token", - action="store_true", - help="Disable the residual 'ghost' token (out-of-top-k probability mass) in the top-k KL loss.", - ) 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") parser.add_argument("--lr_warmup_iters", type=int, default=50, help="Number of LR warmup steps") @@ -471,7 +466,6 @@ def _build_model_provider(hf_path, load_weights=True, moe_grouped_gemm=True): logit_kl_topk=args.logit_kl_topk, logit_kl_top_p=args.logit_kl_top_p, logit_kl_top_p_min_k=args.logit_kl_top_p_min_k, - logit_kl_ghost_token=not args.no_logit_kl_ghost_token, ) # HF VLM configs expose ``vision_config``; Megatron-Bridge nests the text model under From 3f751f3385e1753cbd2c435497c86d5892f66d3a Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 17 Sep 2026 19:52:54 +0200 Subject: [PATCH 03/11] Address review comments Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 2 +- examples/megatron_bridge/distill.py | 12 ++ modelopt/torch/distill/plugins/megatron.py | 128 ++++++++++-------- .../distill/plugins/test_distill_megatron.py | 34 +++-- 4 files changed, 112 insertions(+), 64 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 8b78b03a004..3008c8e6f0f 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -35,7 +35,7 @@ Changelog **Deprecations** -- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``DeprecationWarning``; a user-provided ``skip_lm_loss`` is overridden by the value derived from ``kd_loss_alpha`` with a ``DeprecationWarning``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. +- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``FutureWarning``; an explicit ``skip_lm_loss=True`` is translated to ``kd_loss_alpha=1.0`` with a ``FutureWarning``, and otherwise ``skip_lm_loss`` is derived from ``kd_loss_alpha``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. - Rename the architecture-specific recipe tier from ``modelopt_recipes/huggingface/`` to ``modelopt_recipes/model_type/`` to clarify that it holds recipes shared across every checkpoint of a Hugging Face ``model_type``. Saved ``--recipe huggingface//...`` paths still resolve via a backward-compatibility alias but now emit a ``FutureWarning``, so update them to ``model_type//...`` as the ``huggingface/`` prefix is deprecated. - The single-format quantization CLI flags are deprecated in favour of ``--recipe`` and will be removed in a future release; passing one now emits a ``FutureWarning``. ``examples/hf_ptq``: ``--qformat`` and ``--kv_cache_qformat``. ``examples/megatron_bridge/quantize.py``: ``--quant_cfg``, ``--kv_cache_quant`` and ``--weight_only``. ``examples/torch_onnx/torch_quant_to_onnx.py``: ``--qformat``. A recipe carries the quantization config, the calibration algorithm and the KV-cache setting in one file, so they cannot drift apart the way separate flags can -- and ``--recipe`` already took precedence over all six, silently on ``hf_ptq`` and with a warning on ``megatron_bridge`` -- with one gap the recipe closes rather than inherits: a weight AutoQuantize recipe that omits ``kv_cache`` still falls back to ``--kv_cache_qformat``, so set ``kv_cache`` in the recipe when migrating. Use a recipe from ``modelopt_recipes/general/ptq/``, an architecture-specific one under ``modelopt_recipes/model_type//``, or a checkpoint-specific one under ``modelopt_recipes/models/``. The warning fires only when a flag is passed explicitly: ``--qformat`` defaults to ``fp8`` and ``--kv_cache_qformat`` to ``fp8_cast``, so warning on the defaults would fire on every run, including runs that correctly use ``--recipe``. ``examples/speculative_decoding/scripts/quantize_drafter.py`` keeps ``--qformat`` undeprecated: it has no ``--recipe`` alternative yet. - The TensorRT-LLM checkpoint export format is deprecated and will be removed in 0.49.0: ``export_tensorrt_llm_checkpoint`` and ``torch_to_tensorrt_llm_checkpoint`` now emit a ``DeprecationWarning`` on use. Use ``export_hf_checkpoint``, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang. Its implementation moved to ``modelopt.torch.export.trtllm``, so import those two functions from there and the ``ModelConfig`` dataclasses from ``modelopt.torch.export.trtllm.model_config``; both functions remain importable from ``modelopt.torch.export`` for this release only. diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index 268dc855b3a..da4cabb70ed 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -23,6 +23,7 @@ import argparse import contextlib import os +import warnings import torch from export_distilled_megatron_to_hf import export_llm_to_hf, save_vlm_to_hf @@ -461,6 +462,17 @@ def _build_model_provider(hf_path, load_weights=True, moe_grouped_gemm=True): f"sizes differ ({padded['student']} vs {padded['teacher']})." ) + if args.kd_loss_scale is not None: + warnings.warn( + "--kd_loss_scale is deprecated and ignored; use --kd_loss_alpha instead.", + FutureWarning, + ) + if args.no_skip_lm_loss: + warnings.warn( + "--no_skip_lm_loss is deprecated and ignored; whether the LM loss is skipped is derived " + "from --kd_loss_alpha (skipped iff 1.0).", + FutureWarning, + ) kd_config = ModelOptDistillConfig( kd_loss_alpha=args.kd_loss_alpha, logit_kl_topk=args.logit_kl_topk, diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index c6c5af8ab07..1296f399132 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -60,8 +60,8 @@ class DistillationConfig: kd_loss_alpha: Weight of the distillation loss in the convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. When ``1.0``, the standard language model loss is skipped entirely (``skip_lm_loss`` is derived from this value). - skip_lm_loss: DEPRECATED. Derived from ``kd_loss_alpha`` (``True`` iff ``kd_loss_alpha == 1.0``); - any user-provided value is overridden with a warning. + skip_lm_loss: DEPRECATED. Derived from ``kd_loss_alpha`` (``True`` iff ``kd_loss_alpha == 1.0``). + An explicit ``True`` is translated to ``kd_loss_alpha = 1.0`` with a warning. kd_loss_scale: DEPRECATED and ignored. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. @@ -97,19 +97,33 @@ def __post_init__(self): "DistillationConfig.kd_loss_scale is deprecated and ignored. The distillation loss " "is no longer rescaled to the LM loss magnitude; the total loss is now " "(1 - kd_loss_alpha) * lm_loss + kd_loss_alpha * kd_loss. Set `kd_loss_alpha` instead.", - DeprecationWarning, + FutureWarning, stacklevel=2, ) - derived_skip_lm_loss = self.kd_loss_alpha == 1.0 if self.skip_lm_loss is not None: - warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated and is now derived from `kd_loss_alpha` " - f"(skip iff kd_loss_alpha == 1.0). Overriding skip_lm_loss={self.skip_lm_loss} with " - f"{derived_skip_lm_loss} (kd_loss_alpha={self.kd_loss_alpha}).", - DeprecationWarning, - stacklevel=2, - ) - self.skip_lm_loss = derived_skip_lm_loss + if self.skip_lm_loss and self.kd_loss_alpha != 1.0: + warnings.warn( + "DistillationConfig.skip_lm_loss is deprecated; translating skip_lm_loss=True to " + "kd_loss_alpha=1.0. Set `kd_loss_alpha` directly instead.", + FutureWarning, + stacklevel=2, + ) + self.kd_loss_alpha = 1.0 + elif not self.skip_lm_loss and self.kd_loss_alpha == 1.0: + warnings.warn( + "DistillationConfig.skip_lm_loss is deprecated, and skip_lm_loss=False conflicts " + "with kd_loss_alpha=1.0, which skips the LM loss. Set `kd_loss_alpha` < 1.0 instead.", + FutureWarning, + stacklevel=2, + ) + else: + warnings.warn( + "DistillationConfig.skip_lm_loss is deprecated and is now derived from " + "`kd_loss_alpha` (skipped iff kd_loss_alpha == 1.0). Stop passing it.", + FutureWarning, + stacklevel=2, + ) + self.skip_lm_loss = self.kd_loss_alpha == 1.0 assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" if self.logit_kl_top_p is not None: assert self.logit_kl_topk is not None, "logit_kl_top_p requires logit_kl_topk" @@ -362,9 +376,13 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: """ predictions, targets = self.pre_forward(predictions, targets) - # Temperature-scaled log probabilities (log softmax), globally normalized across TP vocab shards. - p = predictions.float() / self._temperature - self._tp_logsumexp(predictions) - q = targets.float() / self._temperature - self._tp_logsumexp(targets) + # Division by temp should happen prior to finding max for both student and teacher. + output_teacher = targets.float() / self._temperature + output_student = predictions.float() / self._temperature + + # Log probabilities (log softmax), globally normalized across TP vocab shards. + p = output_student - self._tp_logsumexp(output_student) + q = output_teacher - self._tp_logsumexp(output_teacher) # KL divergence if self._reverse: @@ -373,36 +391,28 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: return self.post_forward(loss, tp_reduce=True) - def _tp_logsumexp(self, logits: Tensor, num_chunks: int = 4) -> Tensor: - """Log-sum-exp of ``logits / self._temperature`` over the vocab dim across all TP shards. + def _tp_logsumexp(self, logits: Tensor) -> Tensor: + """Log-sum-exp over the vocab dim across all TP shards (shape ``[..., 1]``). - Accumulates in fp32 over ``num_chunks`` vocab chunks so the transient fp32 working set is a - fraction of the (possibly bf16) full-vocab input. NOTE: for inputs requiring grad, autograd - still retains each chunk's ``exp`` output for backward, so the total saved activation - equals a full-vocab fp32 tensor. Returns shape ``[..., 1]``. + ``logits`` are expected to be fp32 and already temperature-scaled. """ - # Max is exact under the monotonic temperature scaling, so take it in the native dtype - # and defer the temperature division to the centered values: lse(x/T) = max/T + log(sum(exp((x-max)/T))). - logits_max = logits.amax(dim=-1, keepdim=True).float() - if self._config.tensor_model_parallel_size > 1: - tp_group = parallel_state.get_tensor_model_parallel_group() - torch.distributed.all_reduce( - logits_max, op=torch.distributed.ReduceOp.MAX, group=tp_group - ) + if self._config.tensor_model_parallel_size == 1: + return torch.logsumexp(logits, dim=-1, keepdim=True) + + tp_group = parallel_state.get_tensor_model_parallel_group() + + # Subtract maximum value along vocab dimension across all GPUs (for stability) + logits_max = logits.amax(dim=-1, keepdim=True) + torch.distributed.all_reduce(logits_max, op=torch.distributed.ReduceOp.MAX, group=tp_group) logits_max = logits_max.detach() - denom = None - chunk_size = -(-logits.size(-1) // num_chunks) # ceil division - for chunk in logits.split(chunk_size, dim=-1): - centered = (chunk.float() - logits_max) / self._temperature - partial = torch.exp(centered).sum(dim=-1, keepdim=True) - denom = partial if denom is None else denom + partial - if self._config.tensor_model_parallel_size > 1: - # We can't use standard all_reduce function here since the computation - # that follows it isn't identical across TP ranks. - denom = dist_nn.functional.all_reduce(denom, group=tp_group) + # Compute global softmax denominator. + # We can't use standard all_reduce function here since the computation + # that follows it isn't identical across TP ranks. + denom = torch.exp(logits - logits_max).sum(dim=-1, keepdim=True) + denom = dist_nn.functional.all_reduce(denom, group=tp_group) - return logits_max / self._temperature + torch.log(denom) + return logits_max + torch.log(denom) class TopKLogitsKLLoss(LogitsKLLoss): @@ -471,14 +481,15 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: f"top_k ({self.top_k}) is larger than total vocab size ({targets.size(-1) * tp_size})" ) - # Take K from each rank, then the global Top-K of those. Reduce before the fp32 cast: - # casting the full vocab first defeats the point. Selection is unchanged (widening is - # exact, temperature scaling monotonic). + # Divide by temperature first + output_teacher = targets.float() / self._temperature + output_student = predictions.float() / self._temperature + + # Extract local Top-K + # We take K from each rank and then find the global Top-K of all those. local_top_k = min(self.top_k, targets.size(-1)) - top_teacher_vals, top_idx = torch.topk(targets, local_top_k, dim=-1) - top_student_vals = torch.gather(predictions, dim=-1, index=top_idx) - top_teacher_vals = top_teacher_vals.float() / self._temperature - top_student_vals = top_student_vals.float() / self._temperature + top_teacher_vals, top_idx = torch.topk(output_teacher, local_top_k, dim=-1) + top_student_vals = torch.gather(output_student, dim=-1, index=top_idx) if tp_size > 1: tp_group = parallel_state.get_tensor_model_parallel_group() @@ -506,8 +517,8 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: # Log-probs of the Top-K entries under the full-vocab distributions, using global # (full-vocab) log-normalizers so the entries carry true probabilities. # NOTE: ``torch.topk`` returns entries sorted descending by teacher value. - teacher_logp = final_teacher_logits - self._tp_logsumexp(targets) - student_logp = final_student_logits - self._tp_logsumexp(predictions) + teacher_logp = final_teacher_logits - self._tp_logsumexp(output_teacher) + student_logp = final_student_logits - self._tp_logsumexp(output_student) # Top-P (nucleus) mask over the sorted Top-K: keep entry i iff cumulative mass *before* it # is < p. This always keeps the entry that crosses the threshold (and thus top-1). @@ -520,12 +531,18 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: mask = torch.ones_like(teacher_logp, dtype=torch.bool) # Ghost token: residual probability mass outside the kept entries, for both distributions. + # Computed in log space as log(1 - exp(log_kept)) = log(-expm1(log_kept)), which stays + # accurate and differentiable when the kept mass is close to 1. if self.add_ghost_token: - eps = 1e-8 - student_kept_mass = (student_logp.exp() * mask).sum(dim=-1, keepdim=True) - teacher_kept_mass = (teacher_logp.exp() * mask).sum(dim=-1, keepdim=True) - student_residual = torch.log((1.0 - student_kept_mass).clamp(min=eps)) - teacher_residual = torch.log((1.0 - teacher_kept_mass).clamp(min=eps)) + neg_tiny = -1e-7 # keep log(kept_mass) strictly below 0 so expm1 stays negative + student_log_kept = torch.logsumexp( + student_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True + ).clamp(max=neg_tiny) + teacher_log_kept = torch.logsumexp( + teacher_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True + ).clamp(max=neg_tiny) + student_residual = torch.log(-torch.expm1(student_log_kept)) + teacher_residual = torch.log(-torch.expm1(teacher_log_kept)) student_logp = torch.cat([student_logp, student_residual], dim=-1) teacher_logp = torch.cat([teacher_logp, teacher_residual], dim=-1) mask = torch.cat([mask, mask.new_ones((*mask.shape[:-1], 1))], dim=-1) @@ -580,7 +597,8 @@ def forward(self, loss_dict: dict[str, Tensor]) -> Tensor: intermediate_loss = sum(loss_dict.values()) / max(len(loss_dict), 1) if intermediate_loss > 0: - dynamic_scale = logits_loss.detach() / intermediate_loss.detach() + # abs(): the Top-K partial KL without a ghost token is not a true KL and can be negative. + dynamic_scale = logits_loss.detach().abs() / intermediate_loss.detach() intermediate_loss_scaled = intermediate_loss * dynamic_scale else: intermediate_loss = logits_loss.new_tensor(intermediate_loss) diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 7f381ae4fec..2101e937a49 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -353,7 +353,7 @@ def test_topk_logits_kl_loss_numerics_full_vocab_matches_dense(): student, teacher = _make_loss_inputs() dense = LogitsKLLoss(cfg)(student, teacher)[0] topk = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=True)(student, teacher)[0] - assert torch.allclose(dense, topk, atol=1e-5) + assert torch.allclose(dense, topk, atol=1e-6) # Without ghost token, the unnormalized Top-K KL over the full vocab is also the dense KL. topk_no_ghost = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=False)( student, teacher @@ -377,7 +377,7 @@ def test_topk_logits_kl_loss_numerics_ghost_token_reference(): q = torch.cat([q_k, q_rest], -1) p = torch.cat([p_k, p_rest], -1) ref = (q.exp() * (q - p)).sum(-1).transpose(0, 1) - assert torch.allclose(loss, ref, atol=1e-5) + assert torch.allclose(loss, ref, atol=1e-6) # Sanity: total mass within the K+1 buckets is 1 for both distributions. assert torch.allclose(q.exp().sum(-1), torch.ones_like(q[..., 0]), atol=1e-5) assert torch.allclose(p.exp().sum(-1), torch.ones_like(p[..., 0]), atol=1e-5) @@ -406,7 +406,7 @@ def test_logits_kl_losses_temperature_scaling(temperature): qq = torch.cat([q_k, q_rest], -1) pp = torch.cat([p_k, p_rest], -1) ref_topk = (qq.exp() * (qq - pp)).sum(-1).transpose(0, 1) - assert torch.allclose(topk, ref_topk, atol=1e-5) + assert torch.allclose(topk, ref_topk, atol=1e-6) def test_topk_logits_kl_loss_top_p_masks_tail(): @@ -441,7 +441,7 @@ def test_topk_logits_kl_loss_top_p_masks_tail(): q_rest = torch.log1p(-(probs * keep).sum(-1, keepdim=True)) p_rest = torch.log1p(-(p_k.exp() * keep).sum(-1, keepdim=True)) ref_ghost = ref + (q_rest.exp() * (q_rest - p_rest)).sum(-1).transpose(0, 1) - assert torch.allclose(loss_ghost, ref_ghost, atol=1e-5) + assert torch.allclose(loss_ghost, ref_ghost, atol=1e-6) assert loss_ghost.shape == (student.size(1), student.size(0)) loss_ghost.sum().backward() assert student.grad is not None and torch.isfinite(student.grad).all() @@ -510,6 +510,13 @@ def test_loss_balancer_convex_combination(): assert torch.allclose(out["kd_loss"], torch.tensor(0.75 * 2.0 + 0.25 * (0.5 + 0.5))) assert torch.allclose(out["logits_loss"], logits) + # A negative logits loss (possible for Top-K KL without ghost token) must not flip the sign of + # the intermediate-loss contribution. + out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)( + {key: lm, "LogitsKLLoss_0": -logits, "HiddenStateCosineLoss_0": inter} + ) + assert torch.allclose(out["kd_loss"], -logits + logits) # -0.5 + abs(-0.5) * (4.0 / 4.0) = 0 + # alpha=1 ignores the LM loss entirely; skip_original_loss does the same regardless of alpha. out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)({key: lm, "LogitsKLLoss_0": logits}) assert torch.allclose(out["kd_loss"], logits) @@ -525,14 +532,25 @@ def test_loss_balancer_convex_combination(): def test_distillation_config_deprecations(): - """skip_lm_loss is derived from kd_loss_alpha; legacy fields warn.""" + """skip_lm_loss is derived from kd_loss_alpha; legacy fields warn with FutureWarning.""" assert DistillationConfig(kd_loss_alpha=1.0).skip_lm_loss is True assert DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss is False - with pytest.warns(DeprecationWarning, match="skip_lm_loss is deprecated"): + # Explicit skip_lm_loss=True is translated to kd_loss_alpha=1.0 rather than overridden. + with pytest.warns(FutureWarning, match="translating skip_lm_loss=True"): cfg = DistillationConfig(kd_loss_alpha=0.9, skip_lm_loss=True) - assert cfg.skip_lm_loss is False # user value overridden + assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss is True + + # skip_lm_loss=False with alpha=1.0 is a conflict; alpha wins. + with pytest.warns(FutureWarning, match="conflicts with kd_loss_alpha=1.0"): + cfg = DistillationConfig(kd_loss_alpha=1.0, skip_lm_loss=False) + assert cfg.skip_lm_loss is True + + # Consistent but deprecated usage still warns. + with pytest.warns(FutureWarning, match="skip_lm_loss is deprecated"): + cfg = DistillationConfig(kd_loss_alpha=0.9, skip_lm_loss=False) + assert cfg.skip_lm_loss is False - with pytest.warns(DeprecationWarning, match="kd_loss_scale is deprecated"): + with pytest.warns(FutureWarning, match="kd_loss_scale is deprecated"): cfg = DistillationConfig(kd_loss_scale=2.0) assert cfg.kd_loss_alpha == 0.9 From df93c400a56ebd7d515e487714618651e924f854 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Tue, 22 Sep 2026 20:04:40 +0200 Subject: [PATCH 04/11] Make ghost_token not-togglable Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 2 +- modelopt/torch/distill/plugins/megatron.py | 52 ++++++--------- .../distill/plugins/test_distill_megatron.py | 66 +++++++------------ 3 files changed, 44 insertions(+), 76 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 3008c8e6f0f..0b053182883 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -24,7 +24,7 @@ Changelog **Backward Breaking Changes** -- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and by default 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; set ``logit_kl_ghost_token: false`` to drop the ghost token. +- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and by default 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. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (default ``0.9``, in [0, 1]), matching Megatron-LM's offline cached-logits KD. ``skip_lm_loss`` is now derived from ``kd_loss_alpha`` (skipped iff ``1.0``), so the LM loss is computed by default where it was previously skipped. ``examples/megatron_bridge/distill.py`` gains ``--kd_loss_alpha``. - Layerwise calibration now uses prior-layer QDQ activations by default (``layerwise.get_qdq_activations_from_prev_layer=True``). Set it to ``False`` to diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 1296f399132..604c71e2ca8 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -69,8 +69,6 @@ class DistillationConfig: Only the smallest prefix of the (sorted) Top-K whose cumulative teacher probability reaches this value contributes to the loss. Requires ``logit_kl_topk``. Must be in (0, 1]. logit_kl_top_p_min_k: Minimum number of Top-K entries kept per token when top-P is active. - logit_kl_ghost_token: Whether ``TopKLogitsKLLoss`` appends a "ghost" token holding the - probability mass outside the kept entries to both distributions (default: ``True``). """ intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) @@ -82,7 +80,6 @@ class DistillationConfig: logit_kl_topk: int | None = None logit_kl_top_p: float | None = None logit_kl_top_p_min_k: int = 1 - logit_kl_ghost_token: bool = True criterion: Criterion | None = None loss_balancer: mtd.DistillationLossBalancer | None = None @@ -186,7 +183,6 @@ def setup_distillation_config( top_k=cfg.logit_kl_topk, top_p=cfg.logit_kl_top_p, top_p_min_k=cfg.logit_kl_top_p_min_k, - add_ghost_token=cfg.logit_kl_ghost_token, ) else: criterion[tuple(cfg.logit_layers)] = LogitsKLLoss( @@ -422,14 +418,14 @@ class TopKLogitsKLLoss(LogitsKLLoss): NOTE: Will gather Top-K logits per rank, so mind the value of K for memory and communication. Both distributions are normalized over the *full* vocabulary (not re-normalized over the - Top-K), matching the offline cached-logits KD loss in Megatron-LM. Optional refinements: - - * **Top-P (nucleus)**: after sorting the Top-K by teacher probability, only the smallest prefix - whose cumulative teacher mass reaches ``top_p`` (with a floor of ``top_p_min_k`` entries) - contributes to the loss. - * **Ghost token**: a synthetic extra entry holding the probability mass outside the kept - entries, ``log(1 - sum(kept probs))``, is appended to both student and teacher so the loss - also penalizes mass the student places outside the teacher's nucleus. + Top-K), matching the offline cached-logits KD loss in Megatron-LM. A "ghost" token holding the + probability mass outside the kept entries, ``log(1 - sum(kept probs))``, is appended to both + student and teacher, so the KL is taken between two proper distributions over K + 1 buckets and + also penalizes mass the student places outside the teacher's kept entries. + + Optionally, **Top-P (nucleus)** truncation keeps only the smallest prefix of the teacher-sorted + Top-K whose cumulative teacher mass reaches ``top_p`` (with a floor of ``top_p_min_k`` entries); + the truncated entries' mass moves into the ghost token. """ def __init__( @@ -441,7 +437,6 @@ def __init__( *, top_p: float | None = None, top_p_min_k: int = 1, - add_ghost_token: bool = True, ): """Constructor. @@ -452,8 +447,6 @@ def __init__( top_k: The number of top vocabulary entries to keep from the teacher's distribution. top_p: Optional nucleus threshold in (0, 1] applied on top of the Top-K selection. top_p_min_k: Minimum number of entries kept per token when ``top_p`` is active. - add_ghost_token: Whether to append a residual "ghost" token holding the out-of-Top-K - probability mass to both distributions. """ super().__init__(model_config, temperature, reverse) assert top_k >= 1, f"{top_k=}" @@ -462,7 +455,6 @@ def __init__( self.top_k = top_k self.top_p = top_p self.top_p_min_k = top_p_min_k - self.add_ghost_token = add_ghost_token def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: """Forward function. @@ -533,19 +525,18 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: # Ghost token: residual probability mass outside the kept entries, for both distributions. # Computed in log space as log(1 - exp(log_kept)) = log(-expm1(log_kept)), which stays # accurate and differentiable when the kept mass is close to 1. - if self.add_ghost_token: - neg_tiny = -1e-7 # keep log(kept_mass) strictly below 0 so expm1 stays negative - student_log_kept = torch.logsumexp( - student_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True - ).clamp(max=neg_tiny) - teacher_log_kept = torch.logsumexp( - teacher_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True - ).clamp(max=neg_tiny) - student_residual = torch.log(-torch.expm1(student_log_kept)) - teacher_residual = torch.log(-torch.expm1(teacher_log_kept)) - student_logp = torch.cat([student_logp, student_residual], dim=-1) - teacher_logp = torch.cat([teacher_logp, teacher_residual], dim=-1) - mask = torch.cat([mask, mask.new_ones((*mask.shape[:-1], 1))], dim=-1) + neg_tiny = -1e-7 # keep log(kept_mass) strictly below 0 so expm1 stays negative + student_log_kept = torch.logsumexp( + student_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True + ).clamp(max=neg_tiny) + teacher_log_kept = torch.logsumexp( + teacher_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True + ).clamp(max=neg_tiny) + student_residual = torch.log(-torch.expm1(student_log_kept)) + teacher_residual = torch.log(-torch.expm1(teacher_log_kept)) + student_logp = torch.cat([student_logp, student_residual], dim=-1) + teacher_logp = torch.cat([teacher_logp, teacher_residual], dim=-1) + mask = torch.cat([mask, mask.new_ones((*mask.shape[:-1], 1))], dim=-1) # Sparse KL divergence: sum_i q_i * (log q_i - log p_i) over kept entries. p, q = student_logp, teacher_logp @@ -597,8 +588,7 @@ def forward(self, loss_dict: dict[str, Tensor]) -> Tensor: intermediate_loss = sum(loss_dict.values()) / max(len(loss_dict), 1) if intermediate_loss > 0: - # abs(): the Top-K partial KL without a ghost token is not a true KL and can be negative. - dynamic_scale = logits_loss.detach().abs() / intermediate_loss.detach() + dynamic_scale = logits_loss.detach() / intermediate_loss.detach() intermediate_loss_scaled = intermediate_loss * dynamic_scale else: intermediate_loss = logits_loss.new_tensor(intermediate_loss) diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 2101e937a49..f3a14e2a22c 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -326,16 +326,15 @@ def test_logits_kl_loss(dist_workers): @pytest.mark.parametrize( - ("top_p", "top_p_min_k", "ghost_token"), - [(None, 1, True), (None, 1, False), (0.9, 1, True), (0.9, 3, False)], + ("top_p", "top_p_min_k"), + [(None, 1), (0.9, 1), (0.9, 3)], ) -def test_topk_logits_kl_loss(dist_workers, top_p, top_p_min_k, ghost_token, top_k: int = 5): +def test_topk_logits_kl_loss(dist_workers, top_p, top_p_min_k, top_k: int = 5): """Test TopKLogitsKLLoss with TP parallelism.""" kd_kwargs = { "logit_kl_topk": top_k, "logit_kl_top_p": top_p, "logit_kl_top_p_min_k": top_p_min_k, - "logit_kl_ghost_token": ghost_token, } dist_workers.run(partial(_test_topk_logits_kl_loss, kd_kwargs)) @@ -352,13 +351,8 @@ def test_topk_logits_kl_loss_numerics_full_vocab_matches_dense(): cfg = SimpleNamespace(tensor_model_parallel_size=1) student, teacher = _make_loss_inputs() dense = LogitsKLLoss(cfg)(student, teacher)[0] - topk = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=True)(student, teacher)[0] + topk = TopKLogitsKLLoss(cfg, top_k=student.size(-1))(student, teacher)[0] assert torch.allclose(dense, topk, atol=1e-6) - # Without ghost token, the unnormalized Top-K KL over the full vocab is also the dense KL. - topk_no_ghost = TopKLogitsKLLoss(cfg, top_k=student.size(-1), add_ghost_token=False)( - student, teacher - )[0] - assert torch.allclose(dense, topk_no_ghost, atol=1e-5) def test_topk_logits_kl_loss_numerics_ghost_token_reference(): @@ -366,7 +360,7 @@ def test_topk_logits_kl_loss_numerics_ghost_token_reference(): cfg = SimpleNamespace(tensor_model_parallel_size=1) student, teacher = _make_loss_inputs() k = 4 - loss = TopKLogitsKLLoss(cfg, top_k=k, add_ghost_token=True)(student, teacher)[0] + loss = TopKLogitsKLLoss(cfg, top_k=k)(student, teacher)[0] q_full = F.log_softmax(teacher, dim=-1) p_full = F.log_softmax(student, dim=-1) @@ -396,9 +390,7 @@ def test_logits_kl_losses_temperature_scaling(temperature): assert torch.allclose(dense, ref_dense, atol=1e-5) k = 4 - topk = TopKLogitsKLLoss(cfg, temperature=temperature, top_k=k, add_ghost_token=True)( - student, teacher - )[0] + topk = TopKLogitsKLLoss(cfg, temperature=temperature, top_k=k)(student, teacher)[0] _, idx = torch.topk(teacher, k, dim=-1) q_k, p_k = q.gather(-1, idx), p.gather(-1, idx) q_rest = torch.log1p(-q_k.exp().sum(-1, keepdim=True)) @@ -410,40 +402,33 @@ def test_logits_kl_losses_temperature_scaling(temperature): def test_topk_logits_kl_loss_top_p_masks_tail(): - """Top-P zeroes out-of-nucleus entries and honors the min_k floor.""" + """Top-P moves out-of-nucleus mass into the ghost token and honors the min_k floor.""" cfg = SimpleNamespace(tensor_model_parallel_size=1) student, teacher = _make_loss_inputs() k = 8 q_full = F.log_softmax(teacher, dim=-1) q_k, idx = torch.topk(q_full, k, dim=-1) p_k = F.log_softmax(student, dim=-1).gather(-1, idx) + + def reference(keep): + partial = (keep * q_k.exp() * (q_k - p_k)).sum(-1) + q_rest = torch.log1p(-(q_k.exp() * keep).sum(-1)) + p_rest = torch.log1p(-(p_k.exp() * keep).sum(-1)) + return (partial + q_rest.exp() * (q_rest - p_rest)).transpose(0, 1) + probs = q_k.exp() keep = (probs.cumsum(-1) - probs) < 0.5 - - # No ghost token: loss is exactly the masked partial KL sum. - loss = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5, add_ghost_token=False)(student, teacher)[0] - ref = (keep * probs * (q_k - p_k)).sum(-1).transpose(0, 1) - assert torch.allclose(loss, ref, atol=1e-5) assert not keep.all(), "test inputs should produce some truncation" + loss = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5)(student, teacher)[0] + assert torch.allclose(loss, reference(keep), atol=1e-6) + assert loss.shape == (student.size(1), student.size(0)) - # min_k floor forces at least min_k entries even when nucleus is tiny. + # min_k floor forces at least min_k entries even when the nucleus is tiny. min_k = 3 - loss_min = TopKLogitsKLLoss(cfg, top_k=k, top_p=1e-6, top_p_min_k=min_k, add_ghost_token=False)( - student, teacher - )[0] - ref_min = ((torch.arange(k) < min_k) * probs * (q_k - p_k)).sum(-1).transpose(0, 1) - assert torch.allclose(loss_min, ref_min, atol=1e-5) - - # Ghost token with top-P: residual is mass outside the kept nucleus, distributions sum to 1. - loss_ghost = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5, add_ghost_token=True)(student, teacher)[ - 0 - ] - q_rest = torch.log1p(-(probs * keep).sum(-1, keepdim=True)) - p_rest = torch.log1p(-(p_k.exp() * keep).sum(-1, keepdim=True)) - ref_ghost = ref + (q_rest.exp() * (q_rest - p_rest)).sum(-1).transpose(0, 1) - assert torch.allclose(loss_ghost, ref_ghost, atol=1e-6) - assert loss_ghost.shape == (student.size(1), student.size(0)) - loss_ghost.sum().backward() + loss_min = TopKLogitsKLLoss(cfg, top_k=k, top_p=1e-6, top_p_min_k=min_k)(student, teacher)[0] + assert torch.allclose(loss_min, reference(torch.arange(k) < min_k), atol=1e-6) + + loss.sum().backward() assert student.grad is not None and torch.isfinite(student.grad).all() @@ -510,13 +495,6 @@ def test_loss_balancer_convex_combination(): assert torch.allclose(out["kd_loss"], torch.tensor(0.75 * 2.0 + 0.25 * (0.5 + 0.5))) assert torch.allclose(out["logits_loss"], logits) - # A negative logits loss (possible for Top-K KL without ghost token) must not flip the sign of - # the intermediate-loss contribution. - out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)( - {key: lm, "LogitsKLLoss_0": -logits, "HiddenStateCosineLoss_0": inter} - ) - assert torch.allclose(out["kd_loss"], -logits + logits) # -0.5 + abs(-0.5) * (4.0 / 4.0) = 0 - # alpha=1 ignores the LM loss entirely; skip_original_loss does the same regardless of alpha. out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)({key: lm, "LogitsKLLoss_0": logits}) assert torch.allclose(out["kd_loss"], logits) From 32be4ee3da0e9f8806fc09df2df2c991b5279aa0 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Wed, 23 Sep 2026 16:55:39 +0200 Subject: [PATCH 05/11] Address more reviews Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 2 +- examples/megatron_bridge/distill.py | 16 ++--- modelopt/torch/distill/plugins/megatron.py | 58 ++++++++++--------- .../distill/plugins/test_distill_megatron.py | 32 +++++----- 4 files changed, 60 insertions(+), 48 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 0b053182883..d47318357e8 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -35,7 +35,7 @@ Changelog **Deprecations** -- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``FutureWarning``; an explicit ``skip_lm_loss=True`` is translated to ``kd_loss_alpha=1.0`` with a ``FutureWarning``, and otherwise ``skip_lm_loss`` is derived from ``kd_loss_alpha``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. +- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``FutureWarning``; ``skip_lm_loss`` passed without ``kd_loss_alpha`` is translated to an equivalent ``kd_loss_alpha`` (``True`` to ``1.0``) with a ``FutureWarning``, and is otherwise ignored in favor of ``kd_loss_alpha``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. - Rename the architecture-specific recipe tier from ``modelopt_recipes/huggingface/`` to ``modelopt_recipes/model_type/`` to clarify that it holds recipes shared across every checkpoint of a Hugging Face ``model_type``. Saved ``--recipe huggingface//...`` paths still resolve via a backward-compatibility alias but now emit a ``FutureWarning``, so update them to ``model_type//...`` as the ``huggingface/`` prefix is deprecated. - The single-format quantization CLI flags are deprecated in favour of ``--recipe`` and will be removed in a future release; passing one now emits a ``FutureWarning``. ``examples/hf_ptq``: ``--qformat`` and ``--kv_cache_qformat``. ``examples/megatron_bridge/quantize.py``: ``--quant_cfg``, ``--kv_cache_quant`` and ``--weight_only``. ``examples/torch_onnx/torch_quant_to_onnx.py``: ``--qformat``. A recipe carries the quantization config, the calibration algorithm and the KV-cache setting in one file, so they cannot drift apart the way separate flags can -- and ``--recipe`` already took precedence over all six, silently on ``hf_ptq`` and with a warning on ``megatron_bridge`` -- with one gap the recipe closes rather than inherits: a weight AutoQuantize recipe that omits ``kv_cache`` still falls back to ``--kv_cache_qformat``, so set ``kv_cache`` in the recipe when migrating. Use a recipe from ``modelopt_recipes/general/ptq/``, an architecture-specific one under ``modelopt_recipes/model_type//``, or a checkpoint-specific one under ``modelopt_recipes/models/``. The warning fires only when a flag is passed explicitly: ``--qformat`` defaults to ``fp8`` and ``--kv_cache_qformat`` to ``fp8_cast``, so warning on the defaults would fire on every run, including runs that correctly use ``--recipe``. ``examples/speculative_decoding/scripts/quantize_drafter.py`` keeps ``--qformat`` undeprecated: it has no ``--recipe`` alternative yet. - The TensorRT-LLM checkpoint export format is deprecated and will be removed in 0.49.0: ``export_tensorrt_llm_checkpoint`` and ``torch_to_tensorrt_llm_checkpoint`` now emit a ``DeprecationWarning`` on use. Use ``export_hf_checkpoint``, which exports a unified Hugging Face checkpoint deployable on TensorRT-LLM, vLLM and SGLang. Its implementation moved to ``modelopt.torch.export.trtllm``, so import those two functions from there and the ``ModelConfig`` dataclasses from ``modelopt.torch.export.trtllm.model_config``; both functions remain importable from ``modelopt.torch.export`` for this release only. diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index da4cabb70ed..c9a7d8bc73f 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -178,18 +178,18 @@ def get_args(): help="DEPRECATED and ignored. Whether the LM loss is skipped is derived from --kd_loss_alpha " "(skipped iff alpha == 1.0).", ) - parser.add_argument( - "--kd_loss_alpha", - type=float, - default=0.9, - 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=None, help="DEPRECATED and ignored. Use --kd_loss_alpha.", ) + parser.add_argument( + "--kd_loss_alpha", + type=float, + default=0.9, + help="KD loss weight alpha in (1 - alpha) * lm_loss + alpha * kd_loss. 1.0 skips the LM loss entirely.", + ) parser.add_argument( "--no_async_save", action="store_true", @@ -200,8 +200,8 @@ def get_args(): "--logit_kl_topk", 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", diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 604c71e2ca8..7bdf1564b26 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -50,6 +50,9 @@ logger = logging.getLogger(__name__) +_DEFAULT_KD_LOSS_ALPHA = 0.9 + + @dataclass class DistillationConfig: """Knowledge-Distillation config. @@ -58,10 +61,11 @@ class DistillationConfig: intermediate_layer_pairs: List of tuples of intermediate layer names. logit_layers: Tuple of logit layer names. kd_loss_alpha: Weight of the distillation loss in the convex combination - ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. When ``1.0``, the standard - language model loss is skipped entirely (``skip_lm_loss`` is derived from this value). + ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. Default: ``0.9``. When ``1.0``, + the standard language model loss is skipped entirely (``skip_lm_loss`` is derived from it). skip_lm_loss: DEPRECATED. Derived from ``kd_loss_alpha`` (``True`` iff ``kd_loss_alpha == 1.0``). - An explicit ``True`` is translated to ``kd_loss_alpha = 1.0`` with a warning. + If passed without ``kd_loss_alpha``, it is translated (``True`` -> ``1.0``, ``False`` -> ``0.9``) + with a warning; if ``kd_loss_alpha`` is also set, it is ignored with a warning. kd_loss_scale: DEPRECATED and ignored. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. @@ -73,7 +77,7 @@ class DistillationConfig: intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) logit_layers: tuple[str, str] = ("output_layer", "output_layer") - kd_loss_alpha: float = 0.9 + kd_loss_alpha: float | None = None # resolved in __post_init__ (default 0.9) skip_lm_loss: bool | None = None # deprecated, derived from kd_loss_alpha kd_loss_scale: float | None = None # deprecated, ignored logit_kl_temperature: float = 1.0 @@ -88,7 +92,6 @@ def __post_init__(self): assert all(len(pair) in (2, 3) for pair in self.intermediate_layer_pairs), ( f"{self.intermediate_layer_pairs=}" ) - assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" if self.kd_loss_scale is not None: warnings.warn( "DistillationConfig.kd_loss_scale is deprecated and ignored. The distillation loss " @@ -98,28 +101,26 @@ def __post_init__(self): stacklevel=2, ) if self.skip_lm_loss is not None: - if self.skip_lm_loss and self.kd_loss_alpha != 1.0: - warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated; translating skip_lm_loss=True to " - "kd_loss_alpha=1.0. Set `kd_loss_alpha` directly instead.", - FutureWarning, - stacklevel=2, - ) - self.kd_loss_alpha = 1.0 - elif not self.skip_lm_loss and self.kd_loss_alpha == 1.0: + if self.kd_loss_alpha is None: + translated_alpha = 1.0 if self.skip_lm_loss else _DEFAULT_KD_LOSS_ALPHA warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated, and skip_lm_loss=False conflicts " - "with kd_loss_alpha=1.0, which skips the LM loss. Set `kd_loss_alpha` < 1.0 instead.", + "DistillationConfig.skip_lm_loss is deprecated; translating " + f"skip_lm_loss={self.skip_lm_loss} to kd_loss_alpha={translated_alpha}. " + "Set `kd_loss_alpha` directly instead.", FutureWarning, stacklevel=2, ) + self.kd_loss_alpha = translated_alpha else: warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated and is now derived from " - "`kd_loss_alpha` (skipped iff kd_loss_alpha == 1.0). Stop passing it.", + "DistillationConfig.skip_lm_loss is deprecated and ignored when `kd_loss_alpha` " + "is set (the LM loss is skipped iff kd_loss_alpha == 1.0). Stop passing it.", FutureWarning, stacklevel=2, ) + elif self.kd_loss_alpha is None: + self.kd_loss_alpha = _DEFAULT_KD_LOSS_ALPHA + assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" self.skip_lm_loss = self.kd_loss_alpha == 1.0 assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" if self.logit_kl_top_p is not None: @@ -208,7 +209,7 @@ def setup_distillation_config( if cfg.loss_balancer is None: cfg.loss_balancer = LogitsAndIntermediatesLossBalancer( - kd_loss_alpha=cfg.kd_loss_alpha, + kd_loss_alpha=cfg.kd_loss_alpha, # type: ignore[arg-type] # resolved by __post_init__ skip_original_loss=bool(cfg.skip_lm_loss), # always set by __post_init__ ) @@ -398,9 +399,9 @@ def _tp_logsumexp(self, logits: Tensor) -> Tensor: tp_group = parallel_state.get_tensor_model_parallel_group() # Subtract maximum value along vocab dimension across all GPUs (for stability) - logits_max = logits.amax(dim=-1, keepdim=True) + # Detached before the (non-autograd, in-place) collective: the max is only a stability shift. + logits_max = logits.amax(dim=-1, keepdim=True).detach() torch.distributed.all_reduce(logits_max, op=torch.distributed.ReduceOp.MAX, group=tp_group) - logits_max = logits_max.detach() # Compute global softmax denominator. # We can't use standard all_reduce function here since the computation @@ -585,13 +586,18 @@ def forward(self, loss_dict: dict[str, Tensor]) -> Tensor: if "Logits" in _key: # class name logits_key = _key # should only be one logits_loss = loss_dict.pop(logits_key) - intermediate_loss = sum(loss_dict.values()) / max(len(loss_dict), 1) - - if intermediate_loss > 0: - dynamic_scale = logits_loss.detach() / intermediate_loss.detach() + # Rescale intermediate losses to the logits-loss magnitude, without a host sync. + if loss_dict: + intermediate_loss = sum(loss_dict.values()) / len(loss_dict) + denom = intermediate_loss.detach() + dynamic_scale = torch.where( + denom > 0, + logits_loss.detach() / denom.clamp(min=torch.finfo(denom.dtype).tiny), + torch.zeros_like(denom), + ) intermediate_loss_scaled = intermediate_loss * dynamic_scale else: - intermediate_loss = logits_loss.new_tensor(intermediate_loss) + intermediate_loss = logits_loss.new_zeros(()) intermediate_loss_scaled = intermediate_loss kd_loss = logits_loss + intermediate_loss_scaled diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index f3a14e2a22c..ba228adaeda 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -495,6 +495,13 @@ def test_loss_balancer_convex_combination(): assert torch.allclose(out["kd_loss"], torch.tensor(0.75 * 2.0 + 0.25 * (0.5 + 0.5))) assert torch.allclose(out["logits_loss"], logits) + # Zero intermediate loss contributes nothing (and no division by zero). + out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)( + {key: lm, "LogitsKLLoss_0": logits, "HiddenStateCosineLoss_0": torch.tensor(0.0)} + ) + assert torch.allclose(out["kd_loss"], logits) + assert torch.isfinite(out["intermediate_loss"]) + # alpha=1 ignores the LM loss entirely; skip_original_loss does the same regardless of alpha. out = LogitsAndIntermediatesLossBalancer(kd_loss_alpha=1.0)({key: lm, "LogitsKLLoss_0": logits}) assert torch.allclose(out["kd_loss"], logits) @@ -511,23 +518,22 @@ def test_loss_balancer_convex_combination(): def test_distillation_config_deprecations(): """skip_lm_loss is derived from kd_loss_alpha; legacy fields warn with FutureWarning.""" + assert DistillationConfig().kd_loss_alpha == 0.9 assert DistillationConfig(kd_loss_alpha=1.0).skip_lm_loss is True assert DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss is False - # Explicit skip_lm_loss=True is translated to kd_loss_alpha=1.0 rather than overridden. - with pytest.warns(FutureWarning, match="translating skip_lm_loss=True"): - cfg = DistillationConfig(kd_loss_alpha=0.9, skip_lm_loss=True) + # skip_lm_loss alone is translated to an equivalent kd_loss_alpha. + with pytest.warns(FutureWarning, match="translating skip_lm_loss=True to kd_loss_alpha=1.0"): + cfg = DistillationConfig(skip_lm_loss=True) assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss is True - - # skip_lm_loss=False with alpha=1.0 is a conflict; alpha wins. - with pytest.warns(FutureWarning, match="conflicts with kd_loss_alpha=1.0"): - cfg = DistillationConfig(kd_loss_alpha=1.0, skip_lm_loss=False) - assert cfg.skip_lm_loss is True - - # Consistent but deprecated usage still warns. - with pytest.warns(FutureWarning, match="skip_lm_loss is deprecated"): - cfg = DistillationConfig(kd_loss_alpha=0.9, skip_lm_loss=False) - assert cfg.skip_lm_loss is False + with pytest.warns(FutureWarning, match="translating skip_lm_loss=False to kd_loss_alpha=0.9"): + cfg = DistillationConfig(skip_lm_loss=False) + assert cfg.kd_loss_alpha == 0.9 and cfg.skip_lm_loss is False + + # An explicit kd_loss_alpha always wins over the deprecated field. + with pytest.warns(FutureWarning, match="ignored when `kd_loss_alpha` is set"): + cfg = DistillationConfig(kd_loss_alpha=0.5, skip_lm_loss=True) + assert cfg.kd_loss_alpha == 0.5 and cfg.skip_lm_loss is False with pytest.warns(FutureWarning, match="kd_loss_scale is deprecated"): cfg = DistillationConfig(kd_loss_scale=2.0) From c1fb187110234fda57c76f5c6a62ea0c9573ba95 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Wed, 23 Sep 2026 17:09:57 +0200 Subject: [PATCH 06/11] Reduce student LM loss to a scalar in distillation GPU tests With kd_loss_alpha < 1 the LM loss is now mixed into the total, so the per-token student loss must be reduced first, as Megatron's loss function does during training. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Asha Anoosheh --- .../torch/distill/plugins/test_distill_megatron.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index ba228adaeda..375276fa21d 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -121,7 +121,9 @@ def _test_logits_kl_loss(rank, size): # Compute distillation loss loss = distillation_model.compute_kd_loss( - student_loss=student_loss, loss_reduction_fn=lambda x: x[0].mean() + # Reduce the per-token LM loss to a scalar, as Megatron's loss function does in training. + student_loss=student_loss.mean(), + loss_reduction_fn=lambda x: x[0].mean(), ) assert isinstance(loss, dict), "Loss should be a dictionary" assert "kd_loss" in loss, "Should contain kd_loss key" @@ -212,7 +214,9 @@ def _test_topk_logits_kl_loss(kd_kwargs, rank, size): # Compute distillation loss loss = distillation_model.compute_kd_loss( - student_loss=student_loss, loss_reduction_fn=lambda x: x[0].mean() + # Reduce the per-token LM loss to a scalar, as Megatron's loss function does in training. + student_loss=student_loss.mean(), + loss_reduction_fn=lambda x: x[0].mean(), ) assert isinstance(loss, dict), "Loss should be a dictionary" assert "kd_loss" in loss, "Should contain kd_loss key" From 51a8b606f1a3214037ef1d14f6922ea182de6af0 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Wed, 30 Sep 2026 12:10:05 +0200 Subject: [PATCH 07/11] Clarify that the ghost token is always appended in CHANGELOG Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 4df44c796b0..5e9625602fc 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -32,7 +32,7 @@ Changelog **Backward Breaking Changes** -- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) now normalizes both distributions over the full vocabulary instead of re-normalizing over the Top-K entries, and by default 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. +- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) 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. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (default ``0.9``, in [0, 1]), matching Megatron-LM's offline cached-logits KD. ``skip_lm_loss`` is now derived from ``kd_loss_alpha`` (skipped iff ``1.0``), so the LM loss is computed by default where it was previously skipped. ``examples/megatron_bridge/distill.py`` gains ``--kd_loss_alpha``. - ``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. From c1857f05d51d64cfed699ece2646785c0908d241 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 1 Oct 2026 15:58:11 +0200 Subject: [PATCH 08/11] Hard error on deprecated KDConfig args for simplicity Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 3 +- examples/megatron_bridge/distill.py | 26 +-------- modelopt/torch/distill/plugins/megatron.py | 57 +++++-------------- .../distill/plugins/test_distill_megatron.py | 27 +++------ 4 files changed, 24 insertions(+), 89 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 4f7fd6b0281..199fa85f404 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -44,7 +44,7 @@ Changelog **Backward Breaking Changes** - ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) 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. The total is now the fixed convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss`` with ``DistillationConfig.kd_loss_alpha`` (default ``0.9``, in [0, 1]), matching Megatron-LM's offline cached-logits KD. ``skip_lm_loss`` is now derived from ``kd_loss_alpha`` (skipped iff ``1.0``), so the LM loss is computed by default where it was previously skipped. ``examples/megatron_bridge/distill.py`` gains ``--kd_loss_alpha``. +- ``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``. - ``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. - ``examples/hf_ptq --vllm_fakequant_export`` now raises ``NotImplementedError`` when the checkpoint holds weights the model has no parameter for and a shard actually provides them (an MTP head, an auxiliary tower). The fake-quant exporter writes only model-backed state, so it would otherwise drop those weights silently -- and a fake-quant checkpoint is evaluated, where a missing head changes the score rather than failing loudly. Use the unified HF export, which carries them through. Buffers Transformers recomputes are not weights to lose: ``*.inv_freq`` is skipped even when a shard provides it, since older Llama/Mistral-lineage conversions do list it in the index and refusing an export over it would reject checkpoints that export correctly today. The check runs immediately after the model loads, not at export time, so an incompatible run fails before calibration rather than after it. @@ -81,7 +81,6 @@ Changelog **Deprecations** -- ``DistillationConfig.kd_loss_scale`` and ``DistillationConfig.skip_lm_loss`` (Megatron distillation plugin) are deprecated. ``kd_loss_scale`` is ignored with a ``FutureWarning``; ``skip_lm_loss`` passed without ``kd_loss_alpha`` is translated to an equivalent ``kd_loss_alpha`` (``True`` to ``1.0``) with a ``FutureWarning``, and is otherwise ignored in favor of ``kd_loss_alpha``. The ``--no_skip_lm_loss`` and ``--kd_loss_scale`` flags in ``examples/megatron_bridge/distill.py`` are likewise deprecated and ignored. - KV-domain searches in ``examples/hf_ptq`` now use ``--kv_auto_quantize_checkpoint``; ``--auto_quantize_checkpoint`` remains a deprecated fallback for a KV-primary recipe for one release. diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index 5aedd134b53..91aa6f37815 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -23,7 +23,6 @@ import argparse import contextlib import os -import warnings from pathlib import Path import torch @@ -208,22 +207,10 @@ def get_args(): parser.add_argument( "--train_iters", type=int, required=True, help="Number of training iterations" ) - parser.add_argument( - "--no_skip_lm_loss", - action="store_true", - help="DEPRECATED and ignored. Whether the LM loss is skipped is derived from --kd_loss_alpha " - "(skipped iff alpha == 1.0).", - ) - parser.add_argument( - "--kd_loss_scale", - type=float, - default=None, - help="DEPRECATED and ignored. Use --kd_loss_alpha.", - ) parser.add_argument( "--kd_loss_alpha", type=float, - default=0.9, + 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( @@ -501,17 +488,6 @@ def _build_model_provider(hf_path, load_weights=True, moe_grouped_gemm=True): f"sizes differ ({padded['student']} vs {padded['teacher']})." ) - if args.kd_loss_scale is not None: - warnings.warn( - "--kd_loss_scale is deprecated and ignored; use --kd_loss_alpha instead.", - FutureWarning, - ) - if args.no_skip_lm_loss: - warnings.warn( - "--no_skip_lm_loss is deprecated and ignored; whether the LM loss is skipped is derived " - "from --kd_loss_alpha (skipped iff 1.0).", - FutureWarning, - ) kd_config = ModelOptDistillConfig( kd_loss_alpha=args.kd_loss_alpha, logit_kl_topk=args.logit_kl_topk, diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 7bdf1564b26..39b64b4dc36 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -19,7 +19,6 @@ import logging import re -import warnings from abc import ABCMeta from collections.abc import Callable from dataclasses import dataclass, field @@ -50,9 +49,6 @@ logger = logging.getLogger(__name__) -_DEFAULT_KD_LOSS_ALPHA = 0.9 - - @dataclass class DistillationConfig: """Knowledge-Distillation config. @@ -61,12 +57,10 @@ class DistillationConfig: intermediate_layer_pairs: List of tuples of intermediate layer names. logit_layers: Tuple of logit layer names. kd_loss_alpha: Weight of the distillation loss in the convex combination - ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. Default: ``0.9``. When ``1.0``, - the standard language model loss is skipped entirely (``skip_lm_loss`` is derived from it). - skip_lm_loss: DEPRECATED. Derived from ``kd_loss_alpha`` (``True`` iff ``kd_loss_alpha == 1.0``). - If passed without ``kd_loss_alpha``, it is translated (``True`` -> ``1.0``, ``False`` -> ``0.9``) - with a warning; if ``kd_loss_alpha`` is also set, it is ignored with a warning. - kd_loss_scale: DEPRECATED and ignored. Use ``kd_loss_alpha`` instead. + ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. Default: ``1.0``. When ``1.0``, + the standard language model loss is skipped entirely. + skip_lm_loss: REMOVED; passing it raises. Set internally to ``kd_loss_alpha == 1.0``. + kd_loss_scale: REMOVED; passing it raises. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. logit_kl_top_p: Optional nucleus (top-P) threshold applied on top of the teacher's Top-K. @@ -77,9 +71,9 @@ class DistillationConfig: intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) logit_layers: tuple[str, str] = ("output_layer", "output_layer") - kd_loss_alpha: float | None = None # resolved in __post_init__ (default 0.9) - skip_lm_loss: bool | None = None # deprecated, derived from kd_loss_alpha - kd_loss_scale: float | None = None # deprecated, ignored + kd_loss_alpha: float = 1.0 + skip_lm_loss: bool | None = None # removed as an input; derived from kd_loss_alpha + kd_loss_scale: float | None = None # removed; kept only to raise a helpful error logit_kl_temperature: float = 1.0 logit_kl_topk: int | None = None logit_kl_top_p: float | None = None @@ -92,34 +86,13 @@ def __post_init__(self): assert all(len(pair) in (2, 3) for pair in self.intermediate_layer_pairs), ( f"{self.intermediate_layer_pairs=}" ) - if self.kd_loss_scale is not None: - warnings.warn( - "DistillationConfig.kd_loss_scale is deprecated and ignored. The distillation loss " - "is no longer rescaled to the LM loss magnitude; the total loss is now " - "(1 - kd_loss_alpha) * lm_loss + kd_loss_alpha * kd_loss. Set `kd_loss_alpha` instead.", - FutureWarning, - stacklevel=2, + if self.skip_lm_loss is not None or self.kd_loss_scale is not None: + raise ValueError( + "DistillationConfig `skip_lm_loss` and `kd_loss_scale` have been removed. Use " + "`kd_loss_alpha` instead: the total loss is (1 - kd_loss_alpha) * lm_loss + " + "kd_loss_alpha * kd_loss, and the LM loss is skipped when kd_loss_alpha == 1.0 " + "(the default, equivalent to the old skip_lm_loss=True)." ) - if self.skip_lm_loss is not None: - if self.kd_loss_alpha is None: - translated_alpha = 1.0 if self.skip_lm_loss else _DEFAULT_KD_LOSS_ALPHA - warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated; translating " - f"skip_lm_loss={self.skip_lm_loss} to kd_loss_alpha={translated_alpha}. " - "Set `kd_loss_alpha` directly instead.", - FutureWarning, - stacklevel=2, - ) - self.kd_loss_alpha = translated_alpha - else: - warnings.warn( - "DistillationConfig.skip_lm_loss is deprecated and ignored when `kd_loss_alpha` " - "is set (the LM loss is skipped iff kd_loss_alpha == 1.0). Stop passing it.", - FutureWarning, - stacklevel=2, - ) - elif self.kd_loss_alpha is None: - self.kd_loss_alpha = _DEFAULT_KD_LOSS_ALPHA assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" self.skip_lm_loss = self.kd_loss_alpha == 1.0 assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" @@ -209,7 +182,7 @@ def setup_distillation_config( if cfg.loss_balancer is None: cfg.loss_balancer = LogitsAndIntermediatesLossBalancer( - kd_loss_alpha=cfg.kd_loss_alpha, # type: ignore[arg-type] # resolved by __post_init__ + kd_loss_alpha=cfg.kd_loss_alpha, skip_original_loss=bool(cfg.skip_lm_loss), # always set by __post_init__ ) @@ -558,7 +531,7 @@ class LogitsAndIntermediatesLossBalancer(mtd.DistillationLossBalancer): ``(1 - alpha) * lm_loss + alpha * kd_loss`` (matching Megatron-LM's offline cached-logits KD). """ - def __init__(self, kd_loss_alpha: float = 0.9, skip_original_loss: bool = False): + def __init__(self, kd_loss_alpha: float = 1.0, skip_original_loss: bool = False): """Constructor. Args: diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 375276fa21d..087052af6e6 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -520,25 +520,12 @@ def test_loss_balancer_convex_combination(): DistillationConfig(kd_loss_alpha=-0.1) -def test_distillation_config_deprecations(): - """skip_lm_loss is derived from kd_loss_alpha; legacy fields warn with FutureWarning.""" - assert DistillationConfig().kd_loss_alpha == 0.9 - assert DistillationConfig(kd_loss_alpha=1.0).skip_lm_loss is True +def test_distillation_config_removed_fields(): + """skip_lm_loss is derived from kd_loss_alpha; the removed fields raise.""" + cfg = DistillationConfig() + assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss is True # pure KD by default assert DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss is False - # skip_lm_loss alone is translated to an equivalent kd_loss_alpha. - with pytest.warns(FutureWarning, match="translating skip_lm_loss=True to kd_loss_alpha=1.0"): - cfg = DistillationConfig(skip_lm_loss=True) - assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss is True - with pytest.warns(FutureWarning, match="translating skip_lm_loss=False to kd_loss_alpha=0.9"): - cfg = DistillationConfig(skip_lm_loss=False) - assert cfg.kd_loss_alpha == 0.9 and cfg.skip_lm_loss is False - - # An explicit kd_loss_alpha always wins over the deprecated field. - with pytest.warns(FutureWarning, match="ignored when `kd_loss_alpha` is set"): - cfg = DistillationConfig(kd_loss_alpha=0.5, skip_lm_loss=True) - assert cfg.kd_loss_alpha == 0.5 and cfg.skip_lm_loss is False - - with pytest.warns(FutureWarning, match="kd_loss_scale is deprecated"): - cfg = DistillationConfig(kd_loss_scale=2.0) - assert cfg.kd_loss_alpha == 0.9 + for removed in ({"skip_lm_loss": True}, {"skip_lm_loss": False}, {"kd_loss_scale": 2.0}): + with pytest.raises(ValueError, match="have been removed"): + DistillationConfig(**removed) From 2ad0e70ae89a2a9071332f23296556f5d055de71 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Thu, 1 Oct 2026 19:12:44 +0200 Subject: [PATCH 09/11] Address review: free teacher fp32 copy early, input-only skip_lm_loss, dtype-tied ghost floor - Compute each full-vocab normalizer once and free the teacher's fp32 copy before materializing the student's; document the memory/communication cost. - Stop writing back into the removed skip_lm_loss field so configs survive dataclasses.replace()/asdict() round-trips; read kd_loss_alpha == 1.0 instead. - Tie the ghost-token clamp floor to the dtype's eps and document its effects. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Asha Anoosheh --- modelopt/torch/distill/plugins/megatron.py | 40 ++++++++++++------- .../distill/plugins/test_distill_megatron.py | 20 ++++++---- 2 files changed, 38 insertions(+), 22 deletions(-) diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 39b64b4dc36..87b07264bd7 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -59,7 +59,7 @@ class DistillationConfig: kd_loss_alpha: Weight of the distillation loss in the convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. Default: ``1.0``. When ``1.0``, the standard language model loss is skipped entirely. - skip_lm_loss: REMOVED; passing it raises. Set internally to ``kd_loss_alpha == 1.0``. + skip_lm_loss: REMOVED; passing it raises. The LM loss is skipped iff ``kd_loss_alpha == 1.0``. kd_loss_scale: REMOVED; passing it raises. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. @@ -72,7 +72,7 @@ class DistillationConfig: intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) logit_layers: tuple[str, str] = ("output_layer", "output_layer") kd_loss_alpha: float = 1.0 - skip_lm_loss: bool | None = None # removed as an input; derived from kd_loss_alpha + skip_lm_loss: bool | None = None # removed; kept only to raise a helpful error kd_loss_scale: float | None = None # removed; kept only to raise a helpful error logit_kl_temperature: float = 1.0 logit_kl_topk: int | None = None @@ -94,7 +94,6 @@ def __post_init__(self): "(the default, equivalent to the old skip_lm_loss=True)." ) assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" - self.skip_lm_loss = self.kd_loss_alpha == 1.0 assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" if self.logit_kl_top_p is not None: assert self.logit_kl_topk is not None, "logit_kl_top_p requires logit_kl_topk" @@ -183,7 +182,7 @@ def setup_distillation_config( if cfg.loss_balancer is None: cfg.loss_balancer = LogitsAndIntermediatesLossBalancer( kd_loss_alpha=cfg.kd_loss_alpha, - skip_original_loss=bool(cfg.skip_lm_loss), # always set by __post_init__ + skip_original_loss=cfg.kd_loss_alpha == 1.0, ) return cfg @@ -389,7 +388,9 @@ class TopKLogitsKLLoss(LogitsKLLoss): """Calculates KL-Divergence loss restricted to the Teacher's Top-K vocabulary entries. Calculates using the global Top-K entries without gathering full logits. - NOTE: Will gather Top-K logits per rank, so mind the value of K for memory and communication. + NOTE: Will gather Top-K logits per rank, so mind the value of K for communication. The full-vocab + normalizers still allocate fp32 copies of the local logit shards (the teacher's is freed right + away; the student's is retained for backward), and add two TP all-reduces per distribution. Both distributions are normalized over the *full* vocabulary (not re-normalized over the Top-K), matching the offline cached-logits KD loss in Megatron-LM. A "ghost" token holding the @@ -447,15 +448,21 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: f"top_k ({self.top_k}) is larger than total vocab size ({targets.size(-1) * tp_size})" ) - # Divide by temperature first - output_teacher = targets.float() / self._temperature - output_student = predictions.float() / self._temperature - # Extract local Top-K # We take K from each rank and then find the global Top-K of all those. local_top_k = min(self.top_k, targets.size(-1)) + + # Teacher: full-vocab normalizer and local Top-K, then free its fp32 copy before the student's. + output_teacher = targets.float() / self._temperature + teacher_lse = self._tp_logsumexp(output_teacher) top_teacher_vals, top_idx = torch.topk(output_teacher, local_top_k, dim=-1) + del output_teacher + + # Student: the full-vocab normalizer is inherent to the ghost-token formulation. + output_student = predictions.float() / self._temperature + student_lse = self._tp_logsumexp(output_student) top_student_vals = torch.gather(output_student, dim=-1, index=top_idx) + del output_student if tp_size > 1: tp_group = parallel_state.get_tensor_model_parallel_group() @@ -483,8 +490,8 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: # Log-probs of the Top-K entries under the full-vocab distributions, using global # (full-vocab) log-normalizers so the entries carry true probabilities. # NOTE: ``torch.topk`` returns entries sorted descending by teacher value. - teacher_logp = final_teacher_logits - self._tp_logsumexp(output_teacher) - student_logp = final_student_logits - self._tp_logsumexp(output_student) + teacher_logp = final_teacher_logits - teacher_lse + student_logp = final_student_logits - student_lse # Top-P (nucleus) mask over the sorted Top-K: keep entry i iff cumulative mass *before* it # is < p. This always keeps the entry that crosses the threshold (and thus top-1). @@ -499,7 +506,11 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: # Ghost token: residual probability mass outside the kept entries, for both distributions. # Computed in log space as log(1 - exp(log_kept)) = log(-expm1(log_kept)), which stays # accurate and differentiable when the kept mass is close to 1. - neg_tiny = -1e-7 # keep log(kept_mass) strictly below 0 so expm1 stays negative + # The floor keeps log(kept_mass) strictly below 0 so expm1 stays negative. Deliberate side + # effects once kept mass exceeds 1 - eps: the ghost bucket is pinned at ~eps (so the loss can + # dip below 0 by O(eps)), and the clamp stops the ghost term's gradient. Both are negligible, + # and the latter is the correct limit when the kept entries hold all of the mass. + neg_tiny = -torch.finfo(student_logp.dtype).eps student_log_kept = torch.logsumexp( student_logp.masked_fill(~mask, float("-inf")), dim=-1, keepdim=True ).clamp(max=neg_tiny) @@ -657,7 +668,7 @@ def _sharded_state_dict(self, *args, **kwargs) -> "ShardedStateDict": # Skip `lm_loss` bypassing it when training if not needed for backprop. # Uses a per-forward call counter so that MTP head calls (which always precede the # main LM head call in _postprocess) still receive real CE loss even when - # skip_lm_loss=True — only the final main-head call is zeroed. + # kd_loss_alpha == 1.0 — only the final main-head call is zeroed. # An MTP head left out of quantization is exempt from that: there is no quantization # error to recover there, and its CE materialises an fp32 [seq, vocab] tensor. skip_mtp_loss = _mtp_excluded_from_quantization(model) @@ -666,7 +677,8 @@ def _compute_student_lm_loss(self, labels, logits) -> Tensor: self._lm_loss_call_count += 1 mtp_num_layers = self.config.mtp_num_layers or 0 is_mtp_call = self._lm_loss_call_count <= mtp_num_layers - if distill_cfg.skip_lm_loss and self.training and (not is_mtp_call or skip_mtp_loss): + skip_lm_loss = distill_cfg.kd_loss_alpha == 1.0 + if skip_lm_loss and self.training and (not is_mtp_call or skip_mtp_loss): return torch.zeros_like(labels, dtype=logits.dtype) return type(self).compute_language_model_loss(self, labels, logits) diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 087052af6e6..196c50f1390 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import dataclasses from functools import partial from types import SimpleNamespace @@ -232,7 +233,7 @@ def _test_topk_logits_kl_loss(kd_kwargs, rank, size): def _test_skip_lm_loss_with_mtp(rank, size): - """Test that skip_lm_loss only zeroes the main LM head and not MTP heads.""" + """Test that skipping the LM loss (kd_loss_alpha=1.0) only zeroes the main LM head and not MTP heads.""" set_seed(SEED) num_layers = 2 @@ -320,8 +321,8 @@ def _recording_loss(labels, logits): f"Expected {mtp_num_layers + 1} loss calls, got {len(recorded_losses)}" ) for i, loss in enumerate(recorded_losses[:-1]): - assert loss.any(), f"MTP head {i} loss should be non-zero with skip_lm_loss=True" - assert not recorded_losses[-1].any(), "Main LM head loss should be zero with skip_lm_loss=True" + assert loss.any(), f"MTP head {i} loss should be non-zero with kd_loss_alpha=1.0" + assert not recorded_losses[-1].any(), "Main LM head loss should be zero with kd_loss_alpha=1.0" def test_logits_kl_loss(dist_workers): @@ -447,7 +448,7 @@ def test_distillation_config_top_p_validation(): def test_skip_lm_loss_with_mtp(dist_workers): - """Test that skip_lm_loss only zeroes the main LM head, not MTP heads.""" + """Test that skipping the LM loss (kd_loss_alpha=1.0) only zeroes the main LM head, not MTP heads.""" dist_workers.run(_test_skip_lm_loss_with_mtp) @@ -521,11 +522,14 @@ def test_loss_balancer_convex_combination(): def test_distillation_config_removed_fields(): - """skip_lm_loss is derived from kd_loss_alpha; the removed fields raise.""" - cfg = DistillationConfig() - assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss is True # pure KD by default - assert DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss is False + """The removed fields raise, and a config can be rebuilt from itself.""" + assert DistillationConfig().kd_loss_alpha == 1.0 # pure KD by default for removed in ({"skip_lm_loss": True}, {"skip_lm_loss": False}, {"kd_loss_scale": 2.0}): with pytest.raises(ValueError, match="have been removed"): DistillationConfig(**removed) + + # Nothing is written back into the removed fields, so round-trips do not trip the check. + cfg = dataclasses.replace(DistillationConfig(kd_loss_alpha=0.9), logit_kl_topk=4) + assert cfg.kd_loss_alpha == 0.9 and cfg.logit_kl_topk == 4 + assert DistillationConfig(**dataclasses.asdict(cfg)).kd_loss_alpha == 0.9 From 02151b33737e9f8ec80bb76c9fa7161b093c6588 Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Mon, 5 Oct 2026 14:58:07 +0200 Subject: [PATCH 10/11] Rename TopKLogitsKLLoss to TopLogitsKLLoss and --logit_kl_topk to --logit_kl_top_k The loss now covers Top-K and optional Top-P truncation, so the class name no longer implies Top-K only; the old name remains as a deprecated alias. The Megatron-Bridge example flag matches --logit_kl_top_p. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Asha Anoosheh --- CHANGELOG.rst | 4 +-- examples/megatron_bridge/distill.py | 6 ++--- .../tutorials/Qwen3.6-35B-A3B/README.md | 6 ++--- modelopt/torch/distill/plugins/megatron.py | 25 +++++++++++++++---- tests/examples/megatron_bridge/test_qad.py | 2 +- .../distill/plugins/test_distill_megatron.py | 25 +++++++++++++------ 6 files changed, 46 insertions(+), 22 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 199fa85f404..d8f314075d0 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -33,7 +33,7 @@ Changelog *Megatron Framework (M-LM / M-Bridge)* -- Add optional Top-P (nucleus) truncation to the Megatron ``TopKLogitsKLLoss`` 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 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 `_ for details. - Add ``--mlflow `` 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. @@ -43,7 +43,7 @@ Changelog **Backward Breaking Changes** -- ``modelopt.torch.distill.plugins.megatron.TopKLogitsKLLoss`` (``logit_kl_topk`` in ``DistillationConfig``) 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. +- ``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``. - ``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. diff --git a/examples/megatron_bridge/distill.py b/examples/megatron_bridge/distill.py index 91aa6f37815..f71114cf1f3 100644 --- a/examples/megatron_bridge/distill.py +++ b/examples/megatron_bridge/distill.py @@ -220,7 +220,7 @@ def get_args(): "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 plus a residual " @@ -230,7 +230,7 @@ def get_args(): "--logit_kl_top_p", type=float, default=None, - help="Nucleus threshold in (0, 1] applied on top of --logit_kl_topk: only the smallest prefix " + 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( @@ -490,7 +490,7 @@ def _build_model_provider(hf_path, load_weights=True, moe_grouped_gemm=True): kd_config = ModelOptDistillConfig( kd_loss_alpha=args.kd_loss_alpha, - logit_kl_topk=args.logit_kl_topk, + 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, ) diff --git a/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md b/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md index 7a3545a59d6..24896e8e320 100644 --- a/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md +++ b/examples/megatron_bridge/tutorials/Qwen3.6-35B-A3B/README.md @@ -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 \ @@ -190,7 +190,7 @@ Non-default arguments: - `--tp_size 1 --pp_size 1 --cp_size 1` — **required, not chosen** (see below). `--ep_size 8` must match the PTQ checkpoint. - `--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. @@ -320,7 +320,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: `` 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 `` ~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: `` 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 `` ~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. **For the next QAD run**, three things follow: track the **length-capped rate** as a first-class metric alongside accuracy (a benchmark score can stay flat while 3.6% of responses return nothing); consider **top-p instead of top-k** for the KD loss so coverage adapts to the teacher's entropy rather than a fixed rank; and if memory allows, **full-vocab KL** — at 32K on this 248,320-token vocabulary the dense fp32 logits are 30.31 GiB per tensor, which is why top-k was used here, but more GPU memory or a smaller model or shorter sequence may afford it. diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 87b07264bd7..6c1d01db3e8 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -19,6 +19,7 @@ import logging import re +import warnings from abc import ABCMeta from collections.abc import Callable from dataclasses import dataclass, field @@ -62,7 +63,7 @@ class DistillationConfig: skip_lm_loss: REMOVED; passing it raises. The LM loss is skipped iff ``kd_loss_alpha == 1.0``. kd_loss_scale: REMOVED; passing it raises. Use ``kd_loss_alpha`` instead. logit_kl_temperature: Temperature for the logit KL-divergence loss. - logit_kl_topk: If not None, use TopKLogitsKLLoss instead of LogitsKLLoss with this top-k value. + logit_kl_topk: If not None, use TopLogitsKLLoss instead of LogitsKLLoss with this top-k value. logit_kl_top_p: Optional nucleus (top-P) threshold applied on top of the teacher's Top-K. Only the smallest prefix of the (sorted) Top-K whose cumulative teacher probability reaches this value contributes to the loss. Requires ``logit_kl_topk``. Must be in (0, 1]. @@ -148,9 +149,9 @@ def setup_distillation_config( if cfg.criterion is None: criterion = {} if parallel_state.is_pipeline_last_stage(): - # Use TopKLogitsKLLoss if logit_kl_topk is specified, otherwise use LogitsKLLoss + # Use TopLogitsKLLoss if logit_kl_topk is specified, otherwise use LogitsKLLoss if cfg.logit_kl_topk is not None: - criterion[tuple(cfg.logit_layers)] = TopKLogitsKLLoss( + criterion[tuple(cfg.logit_layers)] = TopLogitsKLLoss( student_cfg, temperature=cfg.logit_kl_temperature, top_k=cfg.logit_kl_topk, @@ -384,8 +385,8 @@ def _tp_logsumexp(self, logits: Tensor) -> Tensor: return logits_max + torch.log(denom) -class TopKLogitsKLLoss(LogitsKLLoss): - """Calculates KL-Divergence loss restricted to the Teacher's Top-K vocabulary entries. +class TopLogitsKLLoss(LogitsKLLoss): + """Calculates KL-Divergence loss restricted to the Teacher's Top-K (and optionally Top-P) entries. Calculates using the global Top-K entries without gathering full logits. NOTE: Will gather Top-K logits per rank, so mind the value of K for communication. The full-vocab @@ -534,6 +535,20 @@ def forward(self, predictions: Tensor, targets: Tensor) -> Tensor: return self.post_forward(loss, tp_reduce=False) +class TopKLogitsKLLoss(TopLogitsKLLoss): + """Deprecated alias of :class:`TopLogitsKLLoss`.""" + + def __init__(self, *args, **kwargs): + """Constructor. Emits a ``FutureWarning`` and forwards to :class:`TopLogitsKLLoss`.""" + warnings.warn( + "TopKLogitsKLLoss is deprecated and will be removed in a future release; " + "use TopLogitsKLLoss instead.", + FutureWarning, + stacklevel=2, + ) + super().__init__(*args, **kwargs) + + class LogitsAndIntermediatesLossBalancer(mtd.DistillationLossBalancer): """LossBalancer implementation for Logit and Intermediate losses. diff --git a/tests/examples/megatron_bridge/test_qad.py b/tests/examples/megatron_bridge/test_qad.py index ddaa6538ef2..6ff6205948f 100644 --- a/tests/examples/megatron_bridge/test_qad.py +++ b/tests/examples/megatron_bridge/test_qad.py @@ -95,7 +95,7 @@ def test_qad(tmp_path: Path, num_gpus, create_student): seq_length=16, mbs=1, gbs=4, - logit_kl_topk=8, + logit_kl_top_k=8, train_iters=train_iters, lr_warmup_iters=2, eval_interval=early_exit_iter, diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 196c50f1390..98ecb0f2eb3 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -31,6 +31,7 @@ LogitsAndIntermediatesLossBalancer, LogitsKLLoss, TopKLogitsKLLoss, + TopLogitsKLLoss, _mtp_excluded_from_quantization, adjust_distillation_model_for_mcore, setup_distillation_config, @@ -134,7 +135,7 @@ def _test_logits_kl_loss(rank, size): def _test_topk_logits_kl_loss(kd_kwargs, rank, size): - """Test TopKLogitsKLLoss with simple forward/backward pass.""" + """Test TopLogitsKLLoss with simple forward/backward pass.""" set_seed(SEED) num_layers = 2 @@ -176,7 +177,7 @@ def _test_topk_logits_kl_loss(kd_kwargs, rank, size): activation_func="squared_relu", ).cuda() - # Setup distillation config with TopKLogitsKLLoss via logit_kl_topk argument + # Setup distillation config with TopLogitsKLLoss via logit_kl_topk argument distill_cfg = setup_distillation_config( config_or_path=DistillationConfig(**kd_kwargs), student_cfg=student_model.config, @@ -335,7 +336,7 @@ def test_logits_kl_loss(dist_workers): [(None, 1), (0.9, 1), (0.9, 3)], ) def test_topk_logits_kl_loss(dist_workers, top_p, top_p_min_k, top_k: int = 5): - """Test TopKLogitsKLLoss with TP parallelism.""" + """Test TopLogitsKLLoss with TP parallelism.""" kd_kwargs = { "logit_kl_topk": top_k, "logit_kl_top_p": top_p, @@ -356,7 +357,7 @@ def test_topk_logits_kl_loss_numerics_full_vocab_matches_dense(): cfg = SimpleNamespace(tensor_model_parallel_size=1) student, teacher = _make_loss_inputs() dense = LogitsKLLoss(cfg)(student, teacher)[0] - topk = TopKLogitsKLLoss(cfg, top_k=student.size(-1))(student, teacher)[0] + topk = TopLogitsKLLoss(cfg, top_k=student.size(-1))(student, teacher)[0] assert torch.allclose(dense, topk, atol=1e-6) @@ -365,7 +366,7 @@ def test_topk_logits_kl_loss_numerics_ghost_token_reference(): cfg = SimpleNamespace(tensor_model_parallel_size=1) student, teacher = _make_loss_inputs() k = 4 - loss = TopKLogitsKLLoss(cfg, top_k=k)(student, teacher)[0] + loss = TopLogitsKLLoss(cfg, top_k=k)(student, teacher)[0] q_full = F.log_softmax(teacher, dim=-1) p_full = F.log_softmax(student, dim=-1) @@ -395,7 +396,7 @@ def test_logits_kl_losses_temperature_scaling(temperature): assert torch.allclose(dense, ref_dense, atol=1e-5) k = 4 - topk = TopKLogitsKLLoss(cfg, temperature=temperature, top_k=k)(student, teacher)[0] + topk = TopLogitsKLLoss(cfg, temperature=temperature, top_k=k)(student, teacher)[0] _, idx = torch.topk(teacher, k, dim=-1) q_k, p_k = q.gather(-1, idx), p.gather(-1, idx) q_rest = torch.log1p(-q_k.exp().sum(-1, keepdim=True)) @@ -424,13 +425,13 @@ def reference(keep): probs = q_k.exp() keep = (probs.cumsum(-1) - probs) < 0.5 assert not keep.all(), "test inputs should produce some truncation" - loss = TopKLogitsKLLoss(cfg, top_k=k, top_p=0.5)(student, teacher)[0] + loss = TopLogitsKLLoss(cfg, top_k=k, top_p=0.5)(student, teacher)[0] assert torch.allclose(loss, reference(keep), atol=1e-6) assert loss.shape == (student.size(1), student.size(0)) # min_k floor forces at least min_k entries even when the nucleus is tiny. min_k = 3 - loss_min = TopKLogitsKLLoss(cfg, top_k=k, top_p=1e-6, top_p_min_k=min_k)(student, teacher)[0] + loss_min = TopLogitsKLLoss(cfg, top_k=k, top_p=1e-6, top_p_min_k=min_k)(student, teacher)[0] assert torch.allclose(loss_min, reference(torch.arange(k) < min_k), atol=1e-6) loss.sum().backward() @@ -533,3 +534,11 @@ def test_distillation_config_removed_fields(): cfg = dataclasses.replace(DistillationConfig(kd_loss_alpha=0.9), logit_kl_topk=4) assert cfg.kd_loss_alpha == 0.9 and cfg.logit_kl_topk == 4 assert DistillationConfig(**dataclasses.asdict(cfg)).kd_loss_alpha == 0.9 + + +def test_topk_logits_kl_loss_deprecated_alias(): + """The old class name still works and warns.""" + cfg = SimpleNamespace(tensor_model_parallel_size=1) + with pytest.warns(FutureWarning, match="use TopLogitsKLLoss instead"): + loss_fn = TopKLogitsKLLoss(cfg, top_k=4) + assert isinstance(loss_fn, TopLogitsKLLoss) From 83c1a1dab478507294dd85258b0fff11ef2acd3f Mon Sep 17 00:00:00 2001 From: Asha Anoosheh Date: Mon, 5 Oct 2026 15:13:45 +0200 Subject: [PATCH 11/11] Reject removed DistillationConfig fields in __new__; make skip_lm_loss a property skip_lm_loss and kd_loss_scale are no longer dataclass fields, so configs round-trip through dataclasses.replace()/asdict() while passing either name still raises a ValueError with a migration hint. skip_lm_loss is now a read-only property derived from kd_loss_alpha == 1.0. Co-Authored-By: Claude Opus 5.5 (1M context) Signed-off-by: Asha Anoosheh --- modelopt/torch/distill/plugins/megatron.py | 44 ++++++++++++------- .../distill/plugins/test_distill_megatron.py | 7 ++- 2 files changed, 34 insertions(+), 17 deletions(-) diff --git a/modelopt/torch/distill/plugins/megatron.py b/modelopt/torch/distill/plugins/megatron.py index 343d0b5f269..3aac2666650 100644 --- a/modelopt/torch/distill/plugins/megatron.py +++ b/modelopt/torch/distill/plugins/megatron.py @@ -51,6 +51,10 @@ logger = logging.getLogger(__name__) +# Former DistillationConfig fields that now raise if passed (see ``DistillationConfig.__new__``). +_REMOVED_DISTILLATION_CONFIG_FIELDS = frozenset({"skip_lm_loss", "kd_loss_scale"}) + + @dataclass class DistillationConfig: """Knowledge-Distillation config. @@ -60,9 +64,7 @@ class DistillationConfig: logit_layers: Tuple of logit layer names. kd_loss_alpha: Weight of the distillation loss in the convex combination ``(1 - alpha) * lm_loss + alpha * kd_loss``. Must be in [0, 1]. Default: ``1.0``. When ``1.0``, - the standard language model loss is skipped entirely. - skip_lm_loss: REMOVED; passing it raises. The LM loss is skipped iff ``kd_loss_alpha == 1.0``. - kd_loss_scale: REMOVED; passing it raises. Use ``kd_loss_alpha`` instead. + the standard language model loss is skipped entirely (see :attr:`skip_lm_loss`). logit_kl_temperature: Temperature for the logit KL-divergence loss. logit_kl_topk: If not None, use TopLogitsKLLoss instead of LogitsKLLoss with this top-k value. logit_kl_top_p: Optional nucleus (top-P) threshold applied on top of the teacher's Top-K. @@ -74,8 +76,6 @@ class DistillationConfig: intermediate_layer_pairs: list[tuple[str, ...]] = field(default_factory=list) logit_layers: tuple[str, str] = ("output_layer", "output_layer") kd_loss_alpha: float = 1.0 - skip_lm_loss: bool | None = None # removed; kept only to raise a helpful error - kd_loss_scale: float | None = None # removed; kept only to raise a helpful error logit_kl_temperature: float = 1.0 logit_kl_topk: int | None = None logit_kl_top_p: float | None = None @@ -83,18 +83,32 @@ class DistillationConfig: criterion: Criterion | None = None loss_balancer: mtd.DistillationLossBalancer | None = None + def __new__(cls, *args, **kwargs): + """Reject removed fields with a migration hint before the dataclass ``__init__`` runs. + + Done here rather than in ``__init__`` so the check survives subclasses that re-apply + ``@dataclass`` (which regenerates ``__init__``). + """ + removed = _REMOVED_DISTILLATION_CONFIG_FIELDS & kwargs.keys() + if removed: + raise ValueError( + f"DistillationConfig {sorted(removed)} have been removed. Use `kd_loss_alpha` " + "instead: the total loss is (1 - kd_loss_alpha) * lm_loss + kd_loss_alpha * " + "kd_loss, and the LM loss is skipped when kd_loss_alpha == 1.0 (the default, " + "equivalent to the old skip_lm_loss=True)." + ) + return super().__new__(cls) + + @property + def skip_lm_loss(self) -> bool: + """Whether the standard LM loss is skipped, i.e. ``kd_loss_alpha == 1.0``.""" + return self.kd_loss_alpha == 1.0 + def __post_init__(self): assert len(self.logit_layers) == 2, f"{self.logit_layers=}" assert all(len(pair) in (2, 3) for pair in self.intermediate_layer_pairs), ( f"{self.intermediate_layer_pairs=}" ) - if self.skip_lm_loss is not None or self.kd_loss_scale is not None: - raise ValueError( - "DistillationConfig `skip_lm_loss` and `kd_loss_scale` have been removed. Use " - "`kd_loss_alpha` instead: the total loss is (1 - kd_loss_alpha) * lm_loss + " - "kd_loss_alpha * kd_loss, and the LM loss is skipped when kd_loss_alpha == 1.0 " - "(the default, equivalent to the old skip_lm_loss=True)." - ) assert 0 <= self.kd_loss_alpha <= 1, f"{self.kd_loss_alpha=}" assert self.logit_kl_temperature > 0, f"{self.logit_kl_temperature=}" if self.logit_kl_top_p is not None: @@ -184,7 +198,7 @@ def setup_distillation_config( if cfg.loss_balancer is None: cfg.loss_balancer = LogitsAndIntermediatesLossBalancer( kd_loss_alpha=cfg.kd_loss_alpha, - skip_original_loss=cfg.kd_loss_alpha == 1.0, + skip_original_loss=cfg.skip_lm_loss, ) return cfg @@ -684,11 +698,11 @@ def _sharded_state_dict(self, *args, **kwargs) -> "ShardedStateDict": # Skip `lm_loss` bypassing it when training if not needed for backprop. # Uses a per-forward call counter so that MTP head calls (which always precede the # main LM head call in _postprocess) still receive real CE loss even when - # kd_loss_alpha == 1.0 — only the final main-head call is zeroed. + # skip_lm_loss=True — only the final main-head call is zeroed. # An MTP head left out of quantization is exempt from that: there is no quantization # error to recover there, and its CE materialises an fp32 [seq, vocab] tensor. skip_mtp_loss = _mtp_excluded_from_quantization(model) - skip_lm_loss = distill_cfg.kd_loss_alpha == 1.0 + skip_lm_loss = distill_cfg.skip_lm_loss if skip_lm_loss and skip_mtp_loss: # Freeze the untrained MTP head: DDP's overlapped grad reduce asserts on params with no grad. warn_rank_0("MTP head is outside quantization and its loss is skipped: freezing it.") diff --git a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py index 9c1eafddabd..9985f9778b0 100644 --- a/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py +++ b/tests/gpu_megatron/torch/distill/plugins/test_distill_megatron.py @@ -538,13 +538,16 @@ def test_loss_balancer_convex_combination(): def test_distillation_config_removed_fields(): """The removed fields raise, and a config can be rebuilt from itself.""" - assert DistillationConfig().kd_loss_alpha == 1.0 # pure KD by default + cfg = DistillationConfig() + assert cfg.kd_loss_alpha == 1.0 and cfg.skip_lm_loss # pure KD by default + assert not DistillationConfig(kd_loss_alpha=0.9).skip_lm_loss + assert {"skip_lm_loss", "kd_loss_scale"}.isdisjoint(f.name for f in dataclasses.fields(cfg)) for removed in ({"skip_lm_loss": True}, {"skip_lm_loss": False}, {"kd_loss_scale": 2.0}): with pytest.raises(ValueError, match="have been removed"): DistillationConfig(**removed) - # Nothing is written back into the removed fields, so round-trips do not trip the check. + # skip_lm_loss is a derived property, not a field, so round-trips do not trip the check. cfg = dataclasses.replace(DistillationConfig(kd_loss_alpha=0.9), logit_kl_topk=4) assert cfg.kd_loss_alpha == 0.9 and cfg.logit_kl_topk == 4 assert DistillationConfig(**dataclasses.asdict(cfg)).kd_loss_alpha == 0.9