-
Notifications
You must be signed in to change notification settings - Fork 536
fix(megatron): restore untied output_layer quantization under Megatron-Bridge #2112
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
e3a455d
c56408d
f46cc40
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -248,6 +248,52 @@ def _incompatible_method(self, *args, **kwargs): | |
| return _incompatible_method | ||
|
|
||
|
|
||
| def _resolve_output_layer_untied(model: torch.nn.Module) -> bool | None: | ||
| """Whether ``output_layer`` weights are untied from the input embeddings, or None if unknown. | ||
|
|
||
| Megatron-Core models carry ``share_embeddings_and_output_weights`` (Megatron-Bridge sets it | ||
| from the HF config, Megatron-LM from ``--untie-embeddings-and-output-weights``), so reading | ||
| it off the model works under both frameworks. ``megatron.training.get_args()`` does not: | ||
| Bridge has no global args store, and defaulting to "tied" there silently drops the | ||
| ``output_layer`` weight-quantizer state from the sharded checkpoint. | ||
| """ | ||
| shared = getattr(model, "share_embeddings_and_output_weights", None) | ||
| if shared is not None: | ||
| return not bool(shared) | ||
| for name, module in model.named_modules(): | ||
| # Skip subtrees that do not own the language model's output_layer: the vision tower (never | ||
| # quantized here) and a distillation teacher, which may be tied differently from the | ||
| # student it is wrapped with. | ||
| if "vision_model" in name or "_teacher_model" in name: | ||
| continue | ||
| shared = getattr(module, "share_embeddings_and_output_weights", None) | ||
| if shared is not None: | ||
| return not bool(shared) | ||
| return None | ||
|
Comment on lines
+251
to
+272
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. [IMPORTANT Compatibility] "First module in Two concrete cases:
Since def _resolve_output_layer_untied(model):
shared = getattr(model, "share_embeddings_and_output_weights", None)
if shared is not None:
return not bool(shared)
for name, module in model.named_modules():
if "vision_model" in name:
continue
shared = getattr(module, "share_embeddings_and_output_weights", None)
if shared is not None:
return not bool(shared)
return NoneNote the existing consumers of this flag elsewhere in the codebase all read it off the specific model (
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Accepted — fixed in f46cc40.
Added unit coverage for all four cases (no signal, root wins over subtree, subtree fallback, vision/teacher skipped). |
||
|
|
||
|
|
||
| def _output_layer_untied(config) -> bool: | ||
| """Whether ``output_layer`` is untied, for use from ``sharded_state_dict``. | ||
|
|
||
| Precedence: the model-derived flag recorded by ``megatron_replace_quant_module_hook`` (the only | ||
| source available under Megatron-Bridge, which has no global args store), then Megatron-LM's | ||
| ``--untie-embeddings-and-output-weights``. The answer is cached back onto ``config`` so a model | ||
| carrying neither signal warns once instead of on every save and every load. | ||
| """ | ||
| untied = getattr(config, "modelopt_output_layer_untied", None) | ||
| if untied is not None: | ||
| return untied | ||
| try: | ||
| from megatron.training import get_args as _mlm_get_args | ||
|
|
||
| untied = bool(getattr(_mlm_get_args(), "untie_embeddings_and_output_weights", False)) | ||
| except Exception as e: | ||
| warn_rank_0(f"Failed to get Megatron arg untie_embeddings_and_output_weights: {e}") | ||
| untied = False | ||
| config.modelopt_output_layer_untied = untied | ||
| return untied | ||
|
|
||
|
|
||
| def megatron_replace_quant_module_hook(model: torch.nn.Module): | ||
| """Configure Megatron-Core model quantization support. | ||
|
|
||
|
|
@@ -260,6 +306,7 @@ def megatron_replace_quant_module_hook(model: torch.nn.Module): | |
| typing-matching the QuantModuleRegistry. | ||
| 3. For Attention modules, we configure them to use core_attention path for KV cache quantization. | ||
| """ | ||
| untied = _resolve_output_layer_untied(model) | ||
|
|
||
| def _configure_attention_for_kv_cache_quant(module: Attention): | ||
| """Configure Attention module for KV cache quantization compatibility.""" | ||
|
|
@@ -287,11 +334,17 @@ def _configure_attention_for_kv_cache_quant(module: Attention): | |
| def _register_extra_state_callbacks(model: torch.nn.Module): | ||
| for name, module in model.named_modules(): | ||
| if type(module) in QuantModuleRegistry: | ||
| # Skip output_layer w/o enabled weight_quantizer | ||
| if name.endswith("output_layer") and not getattr( | ||
| getattr(module, "weight_quantizer", None), "is_enabled", False | ||
| ): | ||
| continue | ||
| # Skip output_layer w/o enabled weight_quantizer. This hook also runs BEFORE | ||
| # QuantModule replacement (e.g. on restore), when ``weight_quantizer`` does not | ||
| # exist yet -- the old check then always skipped, so output_layer never received | ||
| # ModelOpt extra-state callbacks and its quantizer state (promotion to | ||
| # StaticBlockScaleQuantizer, ``_amax``, ``_global_amax``) was never restored. | ||
| # Fall back to the tying flag: an untied output_layer is quantizable. | ||
| if name.endswith("output_layer"): | ||
| _wq = getattr(module, "weight_quantizer", None) | ||
| _skip = not getattr(_wq, "is_enabled", False) if _wq is not None else not untied | ||
| if _skip: | ||
| continue | ||
| register_modelopt_extra_state_callbacks( | ||
|
Comment on lines
+337
to
348
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. [IMPORTANT Compatibility] The fallback registers extra-state callbacks for every untied Since this hook runs before
Suggestion: keep the fallback narrow so it only fires when
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Half accepted: I will document this in the PR description, but I do not think the code can be narrowed the way you suggest. You are right about the mechanism — at this point in the hook The problem with both proposed narrowings is that they need information this hook cannot have:
Scoping note on the blast radius: this hook only runs inside ModelOpt quantize/restore flows, so a plain Megatron checkpoint is unaffected; the change is confined to ModelOpt checkpoints where If a maintainer prefers, I am happy to thread the resolved quant config into the hook in a follow-up so registration can key off the layer's actual config rather than tiedness — that is the only version of this that is correct in both directions. Flagging it in the PR description meanwhile. |
||
| module, | ||
| quant_module_get_extra_state, | ||
|
|
@@ -307,6 +360,10 @@ def _register_extra_state_callbacks(model: torch.nn.Module): | |
| if "vision_model" not in name: | ||
| # We only enable hetereogenous_dist_checkpoint for language model, vision model is not quantized | ||
| module.config.hetereogenous_dist_checkpoint = True | ||
| if untied is not None: | ||
| # Carried on the config so _MegatronParallelLinear.sharded_state_dict can read | ||
| # it without Megatron-LM global args (absent under Megatron-Bridge). | ||
| module.config.modelopt_output_layer_untied = untied | ||
| _register_extra_state_callbacks(module) | ||
|
|
||
|
|
||
|
|
@@ -374,18 +431,59 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): | |
| # output_layer.input_quantizer._amax but TP-only does not. This lead to | ||
| # state_dict mismatch. | ||
| if prefix.endswith("output_layer."): | ||
| try: | ||
| from megatron.training import get_args as _mlm_get_args | ||
|
|
||
| _untied = bool( | ||
| getattr(_mlm_get_args(), "untie_embeddings_and_output_weights", False) | ||
| ) | ||
| except Exception as e: | ||
| warn_rank_0(f"Failed to get Megatron arg untie_embeddings_and_output_weights: {e}") | ||
| _untied = False | ||
| if not _untied: | ||
| if not _output_layer_untied(self.config): | ||
| return super().sharded_state_dict(prefix, sharded_offsets, metadata) | ||
|
Comment on lines
433
to
435
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. [SUGGESTION] The tiedness resolution now exists in two places with a precedence rule between them ( Minor related point: when
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Accepted — fixed in f46cc40. Both points are addressed by a single if not _output_layer_untied(self.config):
return super().sharded_state_dict(prefix, sharded_offsets, metadata)Unit-tested for both precedence and the caching behaviour. |
||
|
|
||
| # Materialize missing weight-quantizer scale buffers so their keys appear in the load | ||
| # plan -- the dist-checkpoint loader SILENTLY SKIPS any checkpoint key the model does | ||
| # not advertise, which leaves output_layer uncalibrated and exports it as BF16. | ||
| # ``_amax`` must be allocated FLAT ``[numel // block, 1]``: that is the in-memory | ||
| # layout every other block-quantized layer uses, and ``_process_quantizer_amax`` below | ||
| # exposes it to the checkpoint as a ``[out_features, blocks]`` VIEW sharing the same | ||
| # storage, so the loader writes straight through. Allocating the viewed shape instead | ||
| # loads fine but leaves the wrong in-memory shape, which breaks the export scale math. | ||
| # Only STATIC block quant owns these buffers: a dynamic quantizer derives its scales | ||
| # per forward and deliberately never holds an ``_amax`` (its ``amax`` property asserts | ||
| # ``not self._dynamic``), so materializing one there would be actively wrong. | ||
| _wq = getattr(self, "weight_quantizer", None) | ||
| if ( | ||
| _wq is not None | ||
| and getattr(_wq, "is_enabled", False) | ||
| and getattr(_wq, "is_static_block_quant", False) | ||
| ): | ||
| _block_sizes = getattr(_wq, "_block_sizes", None) or {} | ||
| _block = _block_sizes.get(-1) or _block_sizes.get(1) | ||
| # `_process_quantizer_amax` later does `v.view(weight.shape[0], -1)`, which | ||
| # requires in_features (not just numel) to divide evenly by the block size. | ||
| if _block and self.weight.shape[-1] % int(_block) == 0: | ||
| if getattr(_wq, "_amax", None) is None: | ||
| # Seed from the weights rather than zeros. For weight-only quantization | ||
| # this is exactly what max calibration produces, so a checkpoint that | ||
| # turns out not to carry these keys degrades to "recalibrated from | ||
| # weights" instead of to scale=0 (or NaN) at export and forward. | ||
| _wq.amax = ( | ||
| self.weight.detach() | ||
| .reshape(-1, int(_block)) | ||
| .abs() | ||
| .amax(dim=1, keepdim=True) | ||
| .float() | ||
| ) | ||
| # register_buffer directly: the ``global_amax`` property lives on | ||
| # StaticBlockScaleQuantizer, and on restore this is still a plain | ||
| # TensorQuantizer (promotion happens later), so the setter is unavailable. | ||
| if getattr(_wq, "_global_amax", None) is None: | ||
| _wq.register_buffer( | ||
| "_global_amax", _wq._amax.detach().max().float().clone() | ||
| ) | ||
|
Comment on lines
+448
to
+477
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. [CRITICAL Algorithm] This materialization is not gated on the quantizer being static block quant, and it runs on the save path as well as the load path.
Suggested fix: gate on static block quant, and only materialize when actually building a load plan (or at minimum verify post-load that the values are non-zero): _wq = getattr(self, "weight_quantizer", None)
if (
_wq is not None
and getattr(_wq, "is_enabled", False)
and getattr(_wq, "is_static_block_quant", False)
):
...and consider initializing
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Accepted, both halves — fixed in f46cc40.
|
||
| else: | ||
| # Leaving the buffers unallocated is the silent-drop failure this block exists | ||
| # to prevent, so say so rather than proceeding quietly. | ||
| warn_rank_0( | ||
| f"{prefix}weight_quantizer: cannot materialize scale buffers " | ||
| f"(block_size={_block}, in_features={self.weight.shape[-1]}); its " | ||
| "calibrated scales will not be restored from the checkpoint." | ||
| ) | ||
|
|
||
| quantizer_state_dict = {} | ||
| for k, v in self.state_dict(prefix="", keep_vars=True).items(): | ||
| if "_quantizer" in k and "_amax" in k: | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -1302,3 +1302,54 @@ def test_homogeneous_sharded_state_dict_te_spec(dist_workers, tmp_path): | |
| {"transformer_impl": "transformer_engine"}, | ||
| ), | ||
| ) | ||
|
|
||
|
|
||
| def test_resolve_output_layer_untied(): | ||
| """The tiedness signal is read off the model, not from Megatron-LM global args.""" | ||
| from modelopt.torch.quantization.plugins.megatron import _resolve_output_layer_untied | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 📐 Maintainability & Code Quality | 🟠 Major | ⚡ Quick win Move the plugin imports to module scope. These imports have no circular-dependency, optional-dependency, or deferred-heavy-import justification. Import errors should fail during test collection.
As per path instructions, imports inside test methods require an explicit justification and otherwise belong at the top of the file. 📍 Affects 1 file
🤖 Prompt for AI AgentsSources: Coding guidelines, Path instructions |
||
|
|
||
| class _Flagged(torch.nn.Module): | ||
| def __init__(self, shared): | ||
| super().__init__() | ||
| self.share_embeddings_and_output_weights = shared | ||
|
|
||
| # No signal anywhere -> unknown. | ||
| assert _resolve_output_layer_untied(torch.nn.Module()) is None | ||
|
|
||
| # The root's own flag wins over any subtree. | ||
| root = _Flagged(False) | ||
| root.inner = _Flagged(True) | ||
| assert _resolve_output_layer_untied(root) is True | ||
|
|
||
| # Otherwise fall back to a subtree scan. | ||
| root = torch.nn.Module() | ||
| root.language_model = _Flagged(True) | ||
| assert _resolve_output_layer_untied(root) is False | ||
|
|
||
| # Subtrees that do not own the language model's output_layer are skipped: the vision tower | ||
| # and a distillation teacher, either of which may be tied differently from the student. | ||
| root = torch.nn.Module() | ||
| root.vision_model = _Flagged(True) | ||
| root._teacher_model = _Flagged(True) | ||
| root.language_model = _Flagged(False) | ||
| assert _resolve_output_layer_untied(root) is True | ||
|
|
||
|
|
||
| def test_output_layer_untied_precedence_and_caching(): | ||
| """The model-derived flag wins over Megatron-LM args, and the answer is cached.""" | ||
| from modelopt.torch.quantization.plugins.megatron import _output_layer_untied | ||
|
|
||
| class _Config: | ||
| pass | ||
|
|
||
| config = _Config() | ||
| config.modelopt_output_layer_untied = True | ||
| assert _output_layer_untied(config) is True | ||
|
|
||
| # With no model-derived flag the args fallback runs, and its answer is cached back onto the | ||
| # config so a model carrying neither signal does not warn on every save and every load. | ||
| config = _Config() | ||
| resolved = _output_layer_untied(config) | ||
| assert isinstance(resolved, bool) | ||
| assert config.modelopt_output_layer_untied is resolved | ||
| assert _output_layer_untied(config) is resolved | ||
|
Comment on lines
+1349
to
+1355
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | ⚡ Quick win Test the Megatron-LM fallback result. The test passes when Patch 🤖 Prompt for AI AgentsSource: Coding guidelines |
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
[SUGGESTION] Two smaller points on placing this in
_forward:The comment documents the workaround as a workaround. It says outright that "the better fix is to promote it on the normal path" — and the other half of this PR is touching that normal path. Since promotion after
set_extra_stateis whatquant_module_set_extra_state→maybe_promote_nvfp4_static_quantizeralready does for every other module, it would be worth a follow-up issue reference here so this doesn't become permanent.Promotion mutates module classes and can register a new
_global_amaxbuffer at first forward, i.e. after the model has been wrapped in DDP / the distributed optimizer and afterparam_and_grad_bufferbucketing. Newly registered buffers aren't broadcast by DDP, so ranks rely on each computing the same value locally — which is the concern raised in theglobal_amax=comment above. If promotion instead happens at the end of restore (before wrapping), both problems disappear.Also: the per-run numbers in the comment (
already promoted 460, converted 1, skipped 0) will read as stale the first time someone runs a different model/recipe. Consider dropping the counts and keeping the "exactlyoutput_layerin practice" statement.There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Accepted the concrete parts — f46cc40.
output_layerin practice" statement.I have deliberately not moved the promotion in this PR — doing it at restore time touches the shared
quant_module_set_extra_statepath for every module, which is a larger change than theoutput_layerfix this PR is scoped to. Happy to open it as a follow-up issue.