diff --git a/docs/api-guide/internal/index.md b/docs/api-guide/internal/index.md index 312081ce70b..1a22f0b4b1c 100644 --- a/docs/api-guide/internal/index.md +++ b/docs/api-guide/internal/index.md @@ -16,4 +16,5 @@ Internal utility APIs. num_microbatches_calculator optimizer_param_scheduler +scaling_recipe_contract ``` diff --git a/docs/api-guide/internal/scaling_recipe_contract.md b/docs/api-guide/internal/scaling_recipe_contract.md new file mode 100644 index 00000000000..2e28c802b64 --- /dev/null +++ b/docs/api-guide/internal/scaling_recipe_contract.md @@ -0,0 +1,117 @@ + + +# Scaling Recipe Design Contract + +This page documents the internal implementation contract for the current scaling +recipes. It is intended for maintainers changing model initialization, optimizer +parameter grouping, checkpoint compatibility, or argument validation. + +A scaling recipe is resolved into a single scaling context. Model code, +optimizer code, checkpoint code, CLI validation, and YAML validation should read +that resolved context rather than independently interpreting legacy MuP fields. + +## Legacy Alias Canonicalization + +`--use-mup` is a deprecated alias for exact `--scaling-recipe mup`. The legacy +fields `mup_base_hidden_size`, `mup_base_head_dim`, and `mup_width_mult` are kept +for compatibility but should be synchronized from the canonical fields. + +`mup_width_mult` is derived state: + +```text +width_mult = hidden_size / scaling_base_hidden_size +``` + +A non-default legacy `mup_width_mult` must match the derived value. Conflicting +legacy and canonical scaling fields are validation errors. CLI and YAML +validation both warn for legacy aliases and call the same synchronization helper +so downstream global args have the same canonical shape. + +## `depth_mup` Optimizer Contract + +`depth_mup` is a v1 Megatron adaptation of the spectral width-depth MuP AdamW +table. It is supported only for `optimizer='adam'` because the current optimizer +overrides are defined for Adam/AdamW-style parameter groups. + +Because the weight-decay row is derived for decoupled AdamW, nonzero +`weight_decay` requires `decoupled_weight_decay=True`. Coupled Adam/L2 is allowed +only when `weight_decay=0.0`. SGD, Muon, and other optimizers are intentionally +rejected rather than partially mapped. + +The default multipliers are: + +| Mechanism | Default multiplier | +| --- | --- | +| Dense self-attention/MLP residual branch output | `depth_mult^-1` | +| Hidden matrix-like Adam/AdamW LR | `width_mult^-1` | +| Hidden matrix-like Adam/AdamW epsilon | `(width_mult * depth_mult)^-1` | +| Hidden vector-like Adam/AdamW epsilon | `(width_mult * depth_mult)^-1` | +| Embedding/output-class Adam/AdamW epsilon | `width_mult^-1` | +| Hidden matrix-like AdamW weight decay | `width_mult` | +| Dense block output-projection initialization | `depth_mult^+0.5` | + +## Parameter-Class Policy + +Parameter classification should prefer explicit parameterization metadata +attached during model construction. Name/shape fallback logic exists only for +backward compatibility with older unannotated parameters. + +| Parameter class | LR policy | Epsilon policy | Weight-decay policy | +| --- | --- | --- | --- | +| Embedding/output class | Preserve embedding/output LR policy, including `decoupled_lr` precedence | `width_mult^-1` | Base Megatron policy | +| Hidden matrix-like weights | `width_mult^-1` | `(width_mult * depth_mult)^-1` | `width_mult` | +| Hidden linear/attention/MLP biases | Base LR | `(width_mult * depth_mult)^-1` | Base weight decay | +| Norm scale/bias and unknown 1-D tensors | Base LR | `(width_mult * depth_mult)^-1` as current v1 policy | No weight decay | +| q/k layernorm vectors with `apply_wd_to_qk_layernorm=True` | Base LR | `(width_mult * depth_mult)^-1` | Base weight decay | + +This table is deliberate Megatron behavior, not a direct claim that every row is +spelled out by the paper table. The paper gives base weight decay to hidden +biases, but Megatron's 1-D tensors also include normalization scale/bias tensors +and other vectors. Hidden linear/attention/MLP biases therefore keep base weight +decay, while normalization vectors and otherwise unknown 1-D tensors stay on the +standard no-weight-decay path unless q/k layernorm is explicitly opted in. The +hidden-vector epsilon rule applies to hidden vector-like parameters as the +current v1 policy. + +## Initialization and Residual Branches + +Megatron already applies layer-count-dependent initialization to residual branch +output projections. `depth_mup` rebases dense transformer block output projection +initialization to `scaling_base_num_layers` so the explicit residual-branch +multiplier carries the intended depth scaling. + +The dense residual hook covers: + +- self-attention output projection +- dense MLP output projection + +MoE layers do not inherit this hook. Unsupported residual or model-family paths +should keep failing closed until they have explicit rules and tests. + +## Runtime Scope + +Training mode is the supported runtime for `depth_mup`. Validation loss can be +enabled with `allow_depth_mup_eval`, but that switch is validation-only and does +not make generation or inference a supported path. + +YAML configs may not contain newly added argparse fields. YAML validation should +populate defaults for new global fields that downstream runtime code reads. For +`allow_depth_mup_eval`, the default is `False`. + +## Checkpoint and Optimizer-Group Compatibility + +Distributed-optimizer checkpoint preprocessing and optimizer load must identify +parameter groups through the same tolerant identifier tuple. Optional optimizer +group fields such as `eps` and per-group `optimizer` may be absent in standard +Adam/SGD groups. + +Missing optional fields should resolve to `None` instead of causing resume-time +`KeyError`. Sorting code must also be `None`-safe so groups with and without +optional keys can be preprocessed deterministically. diff --git a/docs/index.md b/docs/index.md index 11337315588..4f1e5526ec6 100644 --- a/docs/index.md +++ b/docs/index.md @@ -49,6 +49,7 @@ get-started/quickstart user-guide/data-preparation user-guide/training-examples +user-guide/scaling-recipes user-guide/parallelism-guide ``` diff --git a/docs/user-guide/index.md b/docs/user-guide/index.md index 2a7ee2eeab9..6dc60f54a6f 100644 --- a/docs/user-guide/index.md +++ b/docs/user-guide/index.md @@ -21,6 +21,7 @@ Guides for using Megatron Core and Megatron-LM. msc_integration data-preparation training-examples +scaling-recipes parallelism-guide features/index ``` diff --git a/docs/user-guide/scaling-recipes.md b/docs/user-guide/scaling-recipes.md new file mode 100644 index 00000000000..54be12e9b1c --- /dev/null +++ b/docs/user-guide/scaling-recipes.md @@ -0,0 +1,159 @@ + + +# Scaling Recipes + +Megatron-LM supports named scaling recipes through `--scaling-recipe`. + +The current named recipes are: + +- `none`: standard Megatron parameterization. +- `mup`: current Megatron width MuP behavior. +- `depth_mup`: experimental spectral width-depth MuP behavior for dense + GPT-style residual Transformer blocks using `--optimizer adam` with AdamW-style + semantics. Nonzero weight decay requires `decoupled_weight_decay=True`. + +`depth_mup` is intentionally narrow. Megatron rejects unsupported paths instead +of silently applying unvalidated scaling rules. + +## Configuration Surface + +New configs should use the canonical scaling fields: + +```bash +--scaling-recipe mup \ +--scaling-base-hidden-size \ +--scaling-base-head-dim +``` + +`--use-mup` remains a backward-compatible alias for exact +`--scaling-recipe mup`, but it is deprecated. The legacy aliases +`--mup-base-hidden-size`, `--mup-base-head-dim`, and `--mup-width-mult` are also +deprecated where they overlap with the canonical scaling surface. + +`--mup-width-mult` is derived from the resolved scaling context: + +```text +width_mult = hidden_size / scaling_base_hidden_size +``` + +If a non-default `--mup-width-mult` is supplied, it must match that derived +value. If legacy MuP fields and canonical scaling fields conflict, Megatron +raises an error during validation. CLI and YAML configs use the same alias +warning and canonicalization path. + +## `mup` + +`mup` preserves the current Megatron width-MuP surface: + +- width multiplier from `hidden_size / scaling_base_hidden_size` +- MuP-family attention softmax scaling through `scaling_base_head_dim` +- hidden-layer width-scaled initialization +- MuP-family embedding and logit scaling +- MuP-family optimizer overrides, including Adam epsilon handling + +Example: + +```bash +torchrun --nproc_per_node=8 pretrain_gpt.py \ + --num-layers 24 \ + --hidden-size 2048 \ + --num-attention-heads 16 \ + --optimizer adam \ + --scaling-recipe mup \ + --scaling-base-hidden-size 1024 \ + --scaling-base-head-dim 128 +``` + +## `depth_mup` + +`depth_mup` extends MuP-family width behavior with depth scaling for dense +GPT-style residual blocks. It is implemented only for `--optimizer adam`. +Because its weight-decay rule is AdamW-style, `depth_mup` requires +`decoupled_weight_decay=True` whenever `weight_decay` is nonzero. Coupled +Adam/L2 is allowed only with `weight_decay=0.0`. + +The main default behaviors are: + +- dense self-attention/MLP residual branch output scales as `depth_mult^-1` +- hidden matrix-like Adam LR scales as `width_mult^-1` +- hidden matrix-like Adam epsilon scales as `(width_mult * depth_mult)^-1` +- embedding/output-class Adam epsilon scales as `width_mult^-1` +- hidden matrix-like AdamW weight decay scales as `width_mult` +- dense residual output-projection initialization scales as `depth_mult^+0.5` + +Megatron's 1-D parameters are not all hidden biases. Under `depth_mup`, hidden +linear/attention/MLP biases keep base weight decay, while normalization vectors +and otherwise unknown 1-D tensors keep Megatron's conservative no-weight-decay +behavior unless q/k layernorm is explicitly opted into weight decay with +`apply_wd_to_qk_layernorm=True`. + +Example: + +```bash +torchrun --nproc_per_node=8 pretrain_gpt.py \ + --num-layers 24 \ + --hidden-size 2048 \ + --num-attention-heads 16 \ + --optimizer adam \ + --scaling-recipe depth_mup \ + --scaling-base-hidden-size 1024 \ + --scaling-base-num-layers 12 \ + --scaling-base-head-dim 128 +``` + +### Current `depth_mup` Scope + +`depth_mup` v1 is currently intended for: + +- dense GPT-style residual Transformer blocks +- dense self-attention residual branches +- dense MLP residual branches +- `--optimizer adam`, with `decoupled_weight_decay=True` when `weight_decay` + is nonzero + +`depth_mup` currently rejects unsupported paths, including: + +- SGD, Muon, and other non-Adam optimizers +- residual-branch scaling during inference +- fused TP inference residual scaling +- cross-attention +- `multi_latent_attention` +- configured experimental attention variants +- MoE +- BERT, T5, and Mamba model families + +Training mode is the supported runtime. Validation loss can be enabled +explicitly with `--allow-depth-mup-eval`; this switch is for validation only and +does not make generation or inference a supported `depth_mup` path. YAML configs +default `allow_depth_mup_eval` to `False`. + +## Manual Overrides + +The canonical scaling fields remain overrides on top of the named recipe. For +example, `--scaling-residual-branch-depth-power 0.0` explicitly disables the +default `depth_mup` residual multiplier instead of introducing a separate +recipe. + +Overrides only change resolved multipliers inside the supported surface. They do +not widen the supported-surface contract. + +## Non-goals + +These recipes do not currently claim: + +- HyperP +- CompleteP +- MuonH / AdamH +- the full Muon-Kimi spectral width-depth training setup +- token-count LR scaling +- SqrtGate +- MoE granularity transfer +- public SGD depth transfer +- public Muon depth transfer diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/param_and_grad_buffer.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/param_and_grad_buffer.py index 960acc25ef6..5644a0519f9 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/param_and_grad_buffer.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/param_and_grad_buffer.py @@ -2890,6 +2890,10 @@ def set_param_attribute(): "partition_stride", "is_embedding_or_output_parameter", "is_embedding_parameter", + "is_output_parameter", + "parameterization_role", + "parameterization_shared_group", + "parameterization_tags", "_tensor_parallel_mode", ]: if hasattr(orig_param, attr_name): diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 085109c60eb..7cc096cb706 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -2457,6 +2457,7 @@ def as_mlp_submodule( is_expert=is_expert, input_size=input_size, ffn_hidden_size=ffn_hidden_size, + apply_block_output_init_scaling=True, ) else: diff --git a/megatron/core/models/T5/t5_model.py b/megatron/core/models/T5/t5_model.py index b2feb974643..44a5ff0be83 100644 --- a/megatron/core/models/T5/t5_model.py +++ b/megatron/core/models/T5/t5_model.py @@ -15,6 +15,7 @@ from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.parameterization import SCALING_RECIPE_DEPTH_MUP, build_resolved_model_policy from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.mappings import scatter_to_tensor_model_parallel_region from megatron.core.transformer.module import MegatronModule @@ -51,15 +52,16 @@ def __init__( log_config_to_disk(config, locals(), prefix=type(self).__name__) self.parallel_output = parallel_output + self.model_scaling_policy = build_resolved_model_policy(config) self.output_layer = tensor_parallel.ColumnParallelLinear( config.hidden_size, vocab_size, config=config, - init_method=( - config.embedding_init_method - if config.use_mup and not share_embeddings_and_output_weights - else config.init_method + init_method=self.model_scaling_policy.output_layer_init_method( + share_embeddings_and_output_weights=share_embeddings_and_output_weights, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, ), bias=share_embeddings_and_output_weights, skip_bias_add=not share_embeddings_and_output_weights, @@ -158,6 +160,12 @@ def __init__( pg_collection: ProcessGroupCollection = None, ): + if config.scaling_recipe == SCALING_RECIPE_DEPTH_MUP: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "Transformer blocks only. T5Model is out of scope for v1." + ) + super(T5Model, self).__init__(config=config) self.config: TransformerConfig = config @@ -427,6 +435,7 @@ def forward( if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() lm_logits = self.lm_head(decoder_hidden_states, word_embeddings_weight=output_weight) + lm_logits = self._scale_logits(lm_logits) if lm_labels is None: # [s b h] => [b s h] diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 3fd1e01f4a1..9a89b4373d6 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -14,6 +14,7 @@ from megatron.core.models.common.embeddings.language_model_embedding import LanguageModelEmbedding from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule +from megatron.core.parameterization import SCALING_RECIPE_DEPTH_MUP from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.attention import SelfAttentionSubmodules from megatron.core.transformer.dot_product_attention import ( @@ -71,6 +72,12 @@ def __init__( vp_stage: Optional[int] = None, pg_collection: Optional[ProcessGroupCollection] = None, ): + if config.scaling_recipe == SCALING_RECIPE_DEPTH_MUP: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "Transformer blocks only. BertModel is out of scope for v1." + ) + super(BertModel, self).__init__(config=config, pg_collection=pg_collection) if has_config_logger_enabled(config): @@ -135,10 +142,10 @@ def __init__( config.hidden_size, self.vocab_size, config=config, - init_method=( - config.embedding_init_method - if config.use_mup and not self.share_embeddings_and_output_weights - else config.init_method + init_method=self.model_scaling_policy.output_layer_init_method( + share_embeddings_and_output_weights=self.share_embeddings_and_output_weights, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, ), bias=True, skip_bias_add=False, @@ -377,6 +384,7 @@ def forward( hidden_states_after_lm_head = self.lm_head(hidden_states=hidden_states) logits, _ = self.output_layer(hidden_states_after_lm_head, weight=output_weight) + logits = self._scale_logits(logits) binary_logits = None if self.binary_head is not None: diff --git a/megatron/core/models/common/embeddings/language_model_embedding.py b/megatron/core/models/common/embeddings/language_model_embedding.py index 7e49ec6c02d..13f8b68ac1a 100644 --- a/megatron/core/models/common/embeddings/language_model_embedding.py +++ b/megatron/core/models/common/embeddings/language_model_embedding.py @@ -6,6 +6,7 @@ from torch import Tensor from megatron.core import tensor_parallel +from megatron.core.parameterization import build_resolved_model_policy from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import get_tensor_model_parallel_group_if_none, nvtx_decorator @@ -45,6 +46,7 @@ def __init__( self.num_tokentypes = num_tokentypes self.scatter_to_sequence_parallel = scatter_to_sequence_parallel self.tp_group = get_tensor_model_parallel_group_if_none(tp_group) + self.model_scaling_policy = build_resolved_model_policy(config) self.reduce_scatter_embeddings = ( (not self.add_position_embedding) and self.num_tokentypes <= 0 @@ -127,9 +129,7 @@ def forward(self, input_ids: Tensor, position_ids: Tensor, tokentype_ids: int = else: assert self.tokentype_embeddings is None - # MuP: scale embeddings by alpha_input. - if self.config.use_mup and self.config.mup_embedding_mult != 1.0: - embeddings = embeddings * self.config.mup_embedding_mult + embeddings = self.model_scaling_policy.scale_embedding_activations(embeddings) # If the input flag for fp32 residual connection is set, convert for float. if self.config.fp32_residual_connection: diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index 84b0ca2fea3..1024f2f07b4 100644 --- a/megatron/core/models/common/language_module/language_module.py +++ b/megatron/core/models/common/language_module/language_module.py @@ -8,6 +8,13 @@ from megatron.core import parallel_state, tensor_parallel from megatron.core.dist_checkpointing.mapping import ShardedStateDict +from megatron.core.parameterization import ( + ROLE_EMBEDDING, + ROLE_OUTPUT, + ROLE_SHARED_EMBEDDING_OUTPUT, + build_resolved_model_policy, + set_parameterization_metadata, +) from megatron.core.transformer.cuda_graphs import CudaGraphManager try: @@ -46,6 +53,7 @@ def __init__( self, config: TransformerConfig, pg_collection: Optional[ProcessGroupCollection] = None ) -> None: super().__init__(config=config) + self.model_scaling_policy = build_resolved_model_policy(config) self._set_attention_backend() if pg_collection is None: pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -204,27 +212,58 @@ def setup_embeddings_and_output_layer(self) -> None: # This is the original Megatron attribute used by decoupled_lr, Muon, FSDP, etc. if self.pre_process and hasattr(self, 'embedding'): self.embedding.word_embeddings.weight.is_embedding_or_output_parameter = True + if self.share_embeddings_and_output_weights: + self.embedding.word_embeddings.weight.is_output_parameter = True + set_parameterization_metadata( + self.embedding.word_embeddings.weight, + role=( + ROLE_SHARED_EMBEDDING_OUTPUT + if self.share_embeddings_and_output_weights + else ROLE_EMBEDDING + ), + shared_group=( + 'lm_embedding_output' if self.share_embeddings_and_output_weights else None + ), + ) if ( self.post_process and hasattr(self, 'output_layer') and self.output_layer.weight is not None ): self.output_layer.weight.is_embedding_or_output_parameter = True + self.output_layer.weight.is_output_parameter = True + set_parameterization_metadata( + self.output_layer.weight, + role=( + ROLE_SHARED_EMBEDDING_OUTPUT + if self.share_embeddings_and_output_weights + else ROLE_OUTPUT + ), + shared_group=( + 'lm_embedding_output' if self.share_embeddings_and_output_weights else None + ), + ) - # Mark embedding-class parameters for MuP optimizer grouping. - # Under MuP table-8-style grouping, embeddings/output use base LR/eps while + # Mark embedding-class parameters for MuP-family optimizer grouping. + # Under MuP-style table-8 grouping, embeddings/output use base LR/eps while # hidden matrix-like params use width-scaled LR/eps. mtp_process = getattr(self, 'mtp_process', False) - if self.config.use_mup and (self.pre_process or mtp_process) and hasattr(self, 'embedding'): + if ( + self.model_scaling_policy.enabled + and (self.pre_process or mtp_process) + and hasattr(self, 'embedding') + ): for param in self.embedding.parameters(): - param.is_embedding_parameter = True + if not hasattr(param, 'parameterization_role'): + set_parameterization_metadata(param, role=ROLE_EMBEDDING) + self.model_scaling_policy.mark_embedding_class_parameters(self.embedding.parameters()) if ( - self.config.use_mup + self.model_scaling_policy.enabled and self.post_process and hasattr(self, 'output_layer') and self.output_layer.weight is not None ): - self.output_layer.weight.is_embedding_parameter = True + self.model_scaling_policy.mark_embedding_class_parameters([self.output_layer.weight]) # If share_embeddings_and_output_weights is True, we need to maintain duplicated # embedding weights in post processing stage. If use Multi-Token Prediction (MTP), @@ -264,9 +303,14 @@ def setup_embeddings_and_output_layer(self) -> None: weight.data.fill_(0) weight.shared = True weight.shared_embedding = True + weight.is_embedding_or_output_parameter = True + weight.is_output_parameter = True # Keep optimizer grouping consistent for tied embedding/output copies. - if self.config.use_mup: - weight.is_embedding_parameter = True + set_parameterization_metadata( + weight, role=ROLE_SHARED_EMBEDDING_OUTPUT, shared_group='lm_embedding_output' + ) + if self.model_scaling_policy.enabled: + self.model_scaling_policy.mark_embedding_class_parameters([weight]) # Parameters are shared between the word embeddings layers, and the # heads at the end of the model. In a pipelined setup with more than @@ -300,23 +344,20 @@ def setup_embeddings_and_output_layer(self) -> None: LanguageModule.embedding_warning_printed = True def _scale_logits(self, logits: Tensor) -> Tensor: - """Apply MuP output scaling to logits. + """Apply scaling-policy output scaling to logits. - When MuP is enabled, scales logits by mup_output_mult (auto-set to 1/width_mult - if left at default) to keep output variance stable across widths. + Under the active MuP-family width recipe (`mup` or `depth_mup`), this + scales logits by `mup_output_mult` (auto-set to `1 / width_mult` when + left at default) to keep output variance stable across widths. Args: logits (Tensor): Raw logits from the output layer. Returns: - Tensor: Scaled logits if MuP is enabled and mup_output_mult != 1.0, + Tensor: Scaled logits when the resolved model policy requests it, otherwise unchanged logits. """ - if not self.config.use_mup: - return logits - if self.config.mup_output_mult != 1.0: - return logits * self.config.mup_output_mult - return logits + return self.model_scaling_policy.scale_output_logits(logits) def shared_embedding_or_output_weight(self) -> Tensor: """Gets the embedding weight or output logit weights when share embedding and output weights set to True diff --git a/megatron/core/models/gpt/fine_grained_callables.py b/megatron/core/models/gpt/fine_grained_callables.py index acfa7a1e8a8..44f1c994286 100644 --- a/megatron/core/models/gpt/fine_grained_callables.py +++ b/megatron/core/models/gpt/fine_grained_callables.py @@ -43,6 +43,21 @@ def wrapped_func(*args, **kwarg): return wrapped_func +def _apply_mlp_bda_with_scaling( + layer: TransformerLayer, output: torch.Tensor, residual: torch.Tensor +): + mlp_output_with_bias = layer._scale_dense_residual_branch_output( + (output, None), + branch_name="mlp", + using_fused_tp_inference_kernel=False, + apply_depth_hook=not layer.is_moe_layer, + ) + with layer.bias_dropout_add_exec_handler(): + return layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( + mlp_output_with_bias, residual, layer.hidden_dropout + ) + + @internal_api def should_free_input(name, is_moe, config, num_local_experts): """Determine if the node should free its input memory. @@ -597,13 +612,9 @@ def submodule_combine_forward(node: ScheduleNode, output: torch.Tensor): output = layer.mlp.combine(output) output = layer.mlp.postprocess(output, shared_expert_output) - mlp_output_with_bias = (output, None) if hasattr(layer, 'cuda_graphs') and layer.cuda_graphs: layer.mlp.cudagraph_tensor_store.clear() - with layer.bias_dropout_add_exec_handler(): - hidden_states = layer.mlp_bda(layer.training, layer.config.bias_dropout_fusion)( - mlp_output_with_bias, residual, layer.hidden_dropout - ) + hidden_states = _apply_mlp_bda_with_scaling(layer, output, residual) # Delay the offload of the mlp norm until after the mlp_bda has been computed # because the residual is needed in the mlp_bda. if layer.offload_mlp_norm: diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index da8a0e2fdfd..797b03eabae 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -247,10 +247,10 @@ def __init__( config.hidden_size, self.vocab_size, config=config, - init_method=( - config.embedding_init_method - if config.use_mup and not self.share_embeddings_and_output_weights - else config.init_method + init_method=self.model_scaling_policy.output_layer_init_method( + share_embeddings_and_output_weights=self.share_embeddings_and_output_weights, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, ), bias=False, skip_bias_add=False, @@ -656,7 +656,9 @@ def _postprocess( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, - scale_logits_fn=self._scale_logits if self.config.use_mup else None, + scale_logits_fn=( + self._scale_logits if self.model_scaling_policy.enabled else None + ), ) sequence_parallel_override = False diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index 4b5858ef9da..12c24f72b44 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -13,6 +13,7 @@ from megatron.core.models.common.embeddings.yarn_rotary_pos_embedding import YarnRotaryEmbedding from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.parameterization import SCALING_RECIPE_DEPTH_MUP from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, ) @@ -110,12 +111,20 @@ def __init__( pg_collection: Optional[ProcessGroupCollection] = None, vp_stage: Optional[int] = None, ) -> None: + if config.scaling_recipe == SCALING_RECIPE_DEPTH_MUP: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "Transformer blocks only. HybridModel is out of scope for v1." + ) + super().__init__(config=config, pg_collection=pg_collection) if has_config_logger_enabled(config): log_config_to_disk(config, locals(), prefix=type(self).__name__) - if self.config.use_mup and not getattr(HybridModel, "mup_warning_printed", False): + if self.model_scaling_policy.enabled and not getattr( + HybridModel, "mup_warning_printed", False + ): log_single_rank( logger, logging.WARNING, @@ -287,10 +296,10 @@ def __init__( config.hidden_size, self.vocab_size, config=config, - init_method=( - config.embedding_init_method - if config.use_mup and not self.share_embeddings_and_output_weights - else config.init_method + init_method=self.model_scaling_policy.output_layer_init_method( + share_embeddings_and_output_weights=self.share_embeddings_and_output_weights, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, ), bias=False, skip_bias_add=False, @@ -522,7 +531,9 @@ def forward( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, - scale_logits_fn=self._scale_logits if self.config.use_mup else None, + scale_logits_fn=( + self._scale_logits if self.model_scaling_policy.enabled else None + ), ) sequence_parallel_override = False if in_inference_mode and inference_context.config.materialize_only_last_token_logits: diff --git a/megatron/core/models/mamba/mamba_model.py b/megatron/core/models/mamba/mamba_model.py index 13964286daf..998badfde35 100644 --- a/megatron/core/models/mamba/mamba_model.py +++ b/megatron/core/models/mamba/mamba_model.py @@ -3,6 +3,7 @@ import logging from megatron.core.models.hybrid.hybrid_model import * # noqa: F401,F403 # pylint: disable=unused-import +from megatron.core.parameterization import SCALING_RECIPE_DEPTH_MUP from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.utils import log_single_rank @@ -13,6 +14,12 @@ class MambaModel(HybridModel): """Backward-compatible wrapper that accepts the deprecated mamba_stack_spec kwarg.""" def __init__(self, *args, mamba_stack_spec: ModuleSpec = None, **kwargs): + config = kwargs.get('config') if 'config' in kwargs else (args[0] if args else None) + if getattr(config, 'scaling_recipe', None) == SCALING_RECIPE_DEPTH_MUP: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "Transformer blocks only. MambaModel is out of scope for v1." + ) log_single_rank( logger, logging.WARNING, "MambaModel has been deprecated. Use HybridModel instead." ) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index c6d3e41aed5..64e476b6919 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -54,6 +54,17 @@ combine_param_group_overrides, param_group_override_to_tuple, ) +from megatron.core.parameterization import ( + ResolvedTrainingPolicy, + build_legacy_mup_training_policy, + is_embedding_class_parameter, + is_embedding_or_output_parameter, + is_hidden_matrix_parameter, + is_hidden_vector_parameter, + is_muon_managed_matrix_parameter, + is_vector_like_parameter, + should_skip_depth_mup_vector_weight_decay, +) from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.fsdp_dtensor_checkpoint import get_global_unique_param_name @@ -73,7 +84,6 @@ Float16OptimizerWithFloat16Params, FP32Optimizer, MegatronOptimizer, - param_group_identifier_keys, ) # Subclass aliases kept for backward compatibility; all are OptimizerConfig. @@ -89,11 +99,17 @@ logger = logging.getLogger(__name__) -def get_standard_config_overrides(config: OptimizerConfig) -> Dict[ParamKey, ParamGroupOverride]: +def get_standard_config_overrides( + config: OptimizerConfig, scaling_policy: Optional[ResolvedTrainingPolicy] = None +) -> Dict[ParamKey, ParamGroupOverride]: """Get standard config overrides for the optimizer, handling decoupled LR and common wd skips. Args: config (OptimizerConfig): optimizer configuration object. + scaling_policy (Optional[ResolvedTrainingPolicy]): optional resolved scaling policy. + ``depth_mup`` uses the paper-backed AdamW weight-decay table and therefore + splits hidden biases from norm/unknown vector-like parameters instead of + applying Megatron's default all-1D/bias weight-decay skip. Returns: Dict[ParamKey, ParamGroupOverride]: standard config overrides. @@ -102,20 +118,34 @@ def get_standard_config_overrides(config: OptimizerConfig) -> Dict[ParamKey, Par # First, figure out how we are going to do wd skipping. The two main approaches are: # 1. The classic megatron approach of skipping all len 1 and bias parameters. # 2. The Qwen3-Next approach of doing 1, other than qk layernorm parameters. - if config.apply_wd_to_qk_layernorm: - shape_1_not_qkln_param = ParamWithNamePredicate( - name="s1_not_qkln", - fn=lambda param, name: (len(param.shape) == 1 or name.endswith(".bias")) - and not ("q_layernorm." in name or "k_layernorm." in name), + use_depth_mup_adamw_table = bool( + scaling_policy and scaling_policy.context.is_depth_mup and scaling_policy.is_adam_optimizer + ) + if use_depth_mup_adamw_table: + depth_mup_vector_wd_skip = ParamWithNamePredicate( + name="depth_mup_norm_and_unknown_vector_wd_skip", + fn=lambda param, name: should_skip_depth_mup_vector_weight_decay( + param, name, apply_wd_to_qk_layernorm=config.apply_wd_to_qk_layernorm + ), ) - param_wd_mult_key = ParamKey(with_name_predicate=shape_1_not_qkln_param) - else: - param_length_1_match = ParamPredicate( - name="param_len_1", fn=lambda param: len(param.shape) == 1 + config_overrides[ParamKey(with_name_predicate=depth_mup_vector_wd_skip)] = ( + ParamGroupOverride(wd_mult=0.0) ) - param_wd_mult_key = ParamKey(name="*.bias", predicate=param_length_1_match) + else: + if config.apply_wd_to_qk_layernorm: + shape_1_not_qkln_param = ParamWithNamePredicate( + name="s1_not_qkln", + fn=lambda param, name: (len(param.shape) == 1 or name.endswith(".bias")) + and not ("q_layernorm." in name or "k_layernorm." in name), + ) + param_wd_mult_key = ParamKey(with_name_predicate=shape_1_not_qkln_param) + else: + param_length_1_match = ParamPredicate( + name="param_len_1", fn=lambda param: len(param.shape) == 1 + ) + param_wd_mult_key = ParamKey(name="*.bias", predicate=param_length_1_match) - config_overrides[param_wd_mult_key] = ParamGroupOverride(wd_mult=0.0) + config_overrides[param_wd_mult_key] = ParamGroupOverride(wd_mult=0.0) if config.decoupled_lr is not None: decoupled_lr_config: ParamGroupOverride = {"max_lr": config.decoupled_lr} @@ -130,54 +160,85 @@ def get_standard_config_overrides(config: OptimizerConfig) -> Dict[ParamKey, Par def get_mup_config_overrides( config: OptimizerConfig, mup_width_mult: float, optimizer_type: str = 'adam' ) -> Dict[ParamKey, ParamGroupOverride]: - """Get MuP config overrides for per-layer LR and Adam epsilon scaling. + """Compatibility wrapper for the legacy MuP optimizer override surface. + + New scaling-policy callers should resolve a ``ResolvedTrainingPolicy`` and call + ``get_scaling_config_overrides`` directly. + """ + scaling_policy = build_legacy_mup_training_policy( + mup_width_mult=mup_width_mult, optimizer_type=optimizer_type + ) + return get_scaling_config_overrides(config=config, scaling_policy=scaling_policy) - In MuP, optimizer learning rates are adjusted by parameter class to ensure - stable update scales across model widths and enable hyperparameter transfer. - MuP optimizer scaling rules (as implemented here): - - Adam/AdamW: - - hidden (matrix-like) lr = base_lr / width_mult - - hidden (matrix-like) eps = base_eps / width_mult +def get_scaling_config_overrides( + config: OptimizerConfig, scaling_policy: ResolvedTrainingPolicy +) -> Dict[ParamKey, ParamGroupOverride]: + """Get resolved scaling-policy overrides for per-parameter optimizer settings. + + In v1, the named scaling recipes are ``mup`` and ``depth_mup``. ``mup`` preserves + current Megatron MuP behavior. ``depth_mup`` now maps to the spectral width-depth + μP paper's AdamW-style dense GPT-style residual Transformer recipe within Megatron's + current support surface. + + Scaling optimizer rules (as implemented here): + - ``mup``: + - Adam/AdamW hidden (matrix-like) lr = + base_lr / width_mult * depth_mult^hidden_lr_depth_power + - Adam/AdamW hidden (matrix-like) eps = + base_eps / width_mult * depth_mult^hidden_eps_depth_power - vector-like params keep base lr and eps - - SGD: - - vector-like lr = base_lr * width_mult - - hidden (matrix-like) lr keeps base_lr in the current uniform-width setup - - no eps override is applied - - Non-Adam optimizers: - - hidden (matrix-like) lr = base_lr / width_mult - - no eps override is applied. - - for Muon optimizers, matrix-like params managed by Muon itself are - excluded from these Adam-style MuP overrides. + - SGD vector-like lr = base_lr * width_mult + - SGD hidden (matrix-like) lr = base_lr * depth_mult^hidden_lr_depth_power + - ``depth_mup`` (`optimizer='adam'` only): + - hidden matrix-like lr = base_lr / width_mult + - hidden matrix-like eps = base_eps / (width_mult * depth_mult) + - hidden matrix-like wd = base_wd * width_mult + - hidden-vector eps = base_eps / (width_mult * depth_mult) + - embedding/output-class eps = base_eps / width_mult + - embedding/output params keep base or decoupled LR With decoupled_lr enabled, embedding/output params continue using decoupled LR - and MuP will not override those explicit decoupled values. + and the active scaling recipe will not override those explicit LR values. Args: config (OptimizerConfig): optimizer configuration object. - mup_width_mult (float): Width multiplier (hidden_size / base_hidden_size). - optimizer_type (str): Optimizer type string from config.optimizer. + scaling_policy (ResolvedTrainingPolicy): resolved model+optimizer scaling policy. Returns: - Dict[ParamKey, ParamGroupOverride]: MuP optimizer overrides. + Dict[ParamKey, ParamGroupOverride]: scaling-policy optimizer overrides. """ - optimizer_type_lower = optimizer_type.lower() - is_sgd_optimizer = optimizer_type_lower == 'sgd' - is_adam_optimizer = 'adam' in optimizer_type_lower - is_muon_optimizer = 'muon' in optimizer_type_lower + if not scaling_policy.enabled: + return {} + + if ( + scaling_policy.context.is_depth_mup + and scaling_policy.is_adam_optimizer + and config.weight_decay != 0.0 + and not config.decoupled_weight_decay + ): + raise ValueError( + "scaling_recipe='depth_mup' with nonzero weight_decay requires " + "decoupled_weight_decay=True because the width-depth weight-decay scaling " + "is derived for AdamW. Use weight_decay=0.0 for coupled Adam, or enable " + "decoupled_weight_decay." + ) decoupled_lr_enabled = config.decoupled_lr is not None if decoupled_lr_enabled: message = ( - "Both decoupled_lr and MuP LR scaling are enabled. decoupled_lr sets an " - "absolute LR for embedding+output params, and MuP LR scaling will not " + "Both decoupled_lr and scaling-recipe LR scaling are enabled. decoupled_lr sets an " + "absolute LR for embedding+output params, and the active scaling recipe will not " "override those parameters." ) - if is_adam_optimizer: - message += " MuP Adam epsilon scaling remains applied to hidden matrix-like parameters." + if scaling_policy.is_adam_optimizer: + message += ( + " Adam epsilon scaling remains applied according to the active recipe's " + "parameter-class rules." + ) log_single_rank(logger, logging.WARNING, message) - if is_muon_optimizer: + if scaling_policy.is_muon_optimizer: muon_scale_mode = getattr(config, 'muon_scale_mode', 'spectral') if muon_scale_mode == 'spectral': log_single_rank( @@ -189,111 +250,169 @@ def get_mup_config_overrides( "Muon-managed matrices with MuP.", ) - if mup_width_mult == 1.0: - # No scaling needed when width_mult is 1 - return {} - - hidden_lr_mult = 1.0 / mup_width_mult base_lr = config.lr base_min_lr = config.min_lr - - # Hidden matrix-like layers get scaled LR/eps; vector-like params keep base values. - # Prefer the explicit parameter attribute set by LanguageModule. Fall back to - # a conservative name check for older or non-language modules. - def is_embedding_parameter(param: torch.nn.Parameter, param_name: str) -> bool: - if getattr(param, 'shared_embedding', False): - return True - if hasattr(param, 'is_embedding_parameter'): - return bool(param.is_embedding_parameter) - return 'embedding' in param_name.lower() - - def is_vector_like_parameter(param: torch.nn.Parameter, param_name: str) -> bool: - if is_embedding_parameter(param, param_name): - return True - if param.dim() <= 1: - return True - return False - - def is_muon_managed_matrix_parameter(param: torch.nn.Parameter, _: str) -> bool: - if not is_muon_optimizer: - return False - return param.dim() == 2 and not getattr(param, 'is_embedding_or_output_parameter', False) - - def should_scale_lr_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: - if decoupled_lr_enabled and getattr(param, 'is_embedding_or_output_parameter', False): + hidden_lr_mult = scaling_policy.hidden_lr_multiplier + hidden_vector_lr_mult = scaling_policy.hidden_vector_lr_multiplier + hidden_eps_mult = scaling_policy.hidden_eps_multiplier + hidden_vector_eps_mult = scaling_policy.hidden_vector_eps_multiplier + embedding_class_eps_mult = scaling_policy.embedding_class_eps_multiplier + hidden_matrix_wd_mult = scaling_policy.hidden_matrix_wd_multiplier + hidden_vector_wd_mult = scaling_policy.hidden_vector_wd_multiplier + embedding_class_wd_mult = scaling_policy.embedding_class_wd_multiplier + + # Hidden matrix-like layers get scaled LR/eps; vector-like params keep base values unless the + # active recipe says otherwise. + def should_scale_hidden_matrix(param: torch.nn.Parameter, param_name: str) -> bool: + if decoupled_lr_enabled and is_embedding_or_output_parameter(param): return False - if is_muon_managed_matrix_parameter(param, param_name): + if is_muon_managed_matrix_parameter(param, optimizer_type=scaling_policy.optimizer_type): return False - return not is_vector_like_parameter(param, param_name) + return is_hidden_matrix_parameter(param, param_name) def should_scale_vector_like_lr_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: - if decoupled_lr_enabled and getattr(param, 'is_embedding_or_output_parameter', False): + if decoupled_lr_enabled and is_embedding_or_output_parameter(param): return False return is_vector_like_parameter(param, param_name) - def should_scale_eps_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: - if is_vector_like_parameter(param, param_name): + def should_scale_hidden_vector(param: torch.nn.Parameter, param_name: str) -> bool: + if decoupled_lr_enabled and is_embedding_or_output_parameter(param): return False - if is_muon_managed_matrix_parameter(param, param_name): + return is_hidden_vector_parameter(param, param_name) + + def should_scale_hidden_matrix_eps(param: torch.nn.Parameter, param_name: str) -> bool: + if is_muon_managed_matrix_parameter(param, optimizer_type=scaling_policy.optimizer_type): return False - # MuP Appendix B.3: eps scales with fan_in when non-negligible. - # This implementation follows the common denominator form: sqrt(v) + eps. - return True + return is_hidden_matrix_parameter(param, param_name) + + def should_scale_hidden_vector_eps(param: torch.nn.Parameter, param_name: str) -> bool: + return is_hidden_vector_parameter(param, param_name) - mup_overrides: Dict[ParamKey, ParamGroupOverride] = {} + def should_scale_embedding_class_eps(param: torch.nn.Parameter, param_name: str) -> bool: + return is_embedding_class_parameter(param, param_name) + + scaling_overrides: Dict[ParamKey, ParamGroupOverride] = {} + + if scaling_policy.is_sgd_optimizer: + hidden_lr_override: ParamGroupOverride = {} + if base_lr is not None and hidden_lr_mult != 1.0: + hidden_lr_override["max_lr"] = base_lr * hidden_lr_mult + if base_min_lr is not None and hidden_lr_mult != 1.0: + hidden_lr_override["min_lr"] = base_min_lr * hidden_lr_mult + if hidden_lr_override: + hidden_predicate = ParamWithNamePredicate( + name="scaling_hidden_only_excluding_embedding_output", fn=should_scale_hidden_matrix + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_predicate)] = hidden_lr_override - if is_sgd_optimizer: - vector_like_lr_mult = mup_width_mult vector_like_lr_override: ParamGroupOverride = {} - if base_lr is not None: - vector_like_lr_override["max_lr"] = base_lr * vector_like_lr_mult - if base_min_lr is not None: - vector_like_lr_override["min_lr"] = base_min_lr * vector_like_lr_mult + if base_lr is not None and hidden_vector_lr_mult != 1.0: + vector_like_lr_override["max_lr"] = base_lr * hidden_vector_lr_mult + if base_min_lr is not None and hidden_vector_lr_mult != 1.0: + vector_like_lr_override["min_lr"] = base_min_lr * hidden_vector_lr_mult if vector_like_lr_override: vector_like_predicate = ParamWithNamePredicate( name="mup_sgd_vector_like_excluding_embedding_output", fn=should_scale_vector_like_lr_with_mup, ) - mup_overrides[ParamKey(with_name_predicate=vector_like_predicate)] = ( + scaling_overrides[ParamKey(with_name_predicate=vector_like_predicate)] = ( vector_like_lr_override ) - return mup_overrides + return scaling_overrides + + if scaling_policy.context.is_depth_mup and scaling_policy.is_adam_optimizer: + hidden_matrix_override: ParamGroupOverride = {} + if base_lr is not None and hidden_lr_mult != 1.0: + hidden_matrix_override["max_lr"] = base_lr * hidden_lr_mult + if base_min_lr is not None and hidden_lr_mult != 1.0: + hidden_matrix_override["min_lr"] = base_min_lr * hidden_lr_mult + if config.adam_eps is not None and hidden_eps_mult != 1.0: + hidden_matrix_override["eps"] = config.adam_eps * hidden_eps_mult + if hidden_matrix_wd_mult != 1.0: + hidden_matrix_override["wd_mult"] = hidden_matrix_wd_mult + if hidden_matrix_override: + hidden_matrix_predicate = ParamWithNamePredicate( + name="depth_mup_hidden_matrix_adamw", fn=should_scale_hidden_matrix_eps + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_matrix_predicate)] = ( + hidden_matrix_override + ) + + hidden_vector_override: ParamGroupOverride = {} + if config.adam_eps is not None and hidden_vector_eps_mult != 1.0: + hidden_vector_override["eps"] = config.adam_eps * hidden_vector_eps_mult + if hidden_vector_wd_mult != 1.0: + hidden_vector_override["wd_mult"] = hidden_vector_wd_mult + if hidden_vector_override: + hidden_vector_predicate = ParamWithNamePredicate( + name="depth_mup_hidden_vector_adamw", fn=should_scale_hidden_vector_eps + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_vector_predicate)] = ( + hidden_vector_override + ) + + embedding_class_override: ParamGroupOverride = {} + if config.adam_eps is not None and embedding_class_eps_mult != 1.0: + embedding_class_override["eps"] = config.adam_eps * embedding_class_eps_mult + if embedding_class_wd_mult != 1.0: + embedding_class_override["wd_mult"] = embedding_class_wd_mult + if embedding_class_override: + embedding_class_predicate = ParamWithNamePredicate( + name="depth_mup_embedding_output_adamw", fn=should_scale_embedding_class_eps + ) + scaling_overrides[ParamKey(with_name_predicate=embedding_class_predicate)] = ( + embedding_class_override + ) + + return scaling_overrides lr_override: ParamGroupOverride = {} - if base_lr is not None: + if base_lr is not None and hidden_lr_mult != 1.0: lr_override["max_lr"] = base_lr * hidden_lr_mult - if base_min_lr is not None: + if base_min_lr is not None and hidden_lr_mult != 1.0: lr_override["min_lr"] = base_min_lr * hidden_lr_mult eps_override: ParamGroupOverride = {} - if is_adam_optimizer and config.adam_eps is not None: - eps_override["eps"] = config.adam_eps * hidden_lr_mult + if scaling_policy.is_adam_optimizer and config.adam_eps is not None and hidden_eps_mult != 1.0: + eps_override["eps"] = config.adam_eps * hidden_eps_mult if decoupled_lr_enabled: if lr_override: hidden_predicate = ParamWithNamePredicate( - name="mup_hidden_only_excluding_embedding_output", fn=should_scale_lr_with_mup + name="mup_hidden_only_excluding_embedding_output", fn=should_scale_hidden_matrix ) - mup_overrides[ParamKey(with_name_predicate=hidden_predicate)] = lr_override + scaling_overrides[ParamKey(with_name_predicate=hidden_predicate)] = lr_override if eps_override: hidden_output_predicate = ParamWithNamePredicate( - name="mup_hidden_only_for_adam_eps", fn=should_scale_eps_with_mup + name="mup_hidden_only_for_adam_eps", fn=should_scale_hidden_matrix_eps ) - mup_overrides[ParamKey(with_name_predicate=hidden_output_predicate)] = eps_override + scaling_overrides[ParamKey(with_name_predicate=hidden_output_predicate)] = eps_override else: - combined_override: ParamGroupOverride = {} - combined_override.update(lr_override) - combined_override.update(eps_override) - if combined_override: + if lr_override and eps_override: + combined_override: ParamGroupOverride = {} + combined_override.update(lr_override) + combined_override.update(eps_override) hidden_output_predicate = ParamWithNamePredicate( - name="mup_hidden_and_output", fn=should_scale_eps_with_mup + name="mup_hidden_and_output", fn=should_scale_hidden_matrix_eps + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_output_predicate)] = ( + combined_override ) - mup_overrides[ParamKey(with_name_predicate=hidden_output_predicate)] = combined_override + elif lr_override: + hidden_predicate = ParamWithNamePredicate( + name="scaling_hidden_and_output_lr", fn=should_scale_hidden_matrix + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_predicate)] = lr_override + elif eps_override: + hidden_output_predicate = ParamWithNamePredicate( + name="mup_hidden_and_output_eps", fn=should_scale_hidden_matrix_eps + ) + scaling_overrides[ParamKey(with_name_predicate=hidden_output_predicate)] = eps_override - return mup_overrides + return scaling_overrides def _get_param_groups( @@ -508,11 +627,12 @@ def _get_megatron_optimizer_based_on_param_groups( assert ( config.decoupled_weight_decay ), "CPU offloading only supported with decoupled_weight_decay enabled (AdamW mode)." - gpu_optimizer_cls = Adam if config.optimizer == 'adam' else SGD - cpu_optimizer_cls = CPUAdam if config.optimizer == 'adam' else CPUSGD + is_adam_optimizer = config.optimizer == 'adam' + gpu_optimizer_cls = Adam if is_adam_optimizer else SGD + cpu_optimizer_cls = CPUAdam if is_adam_optimizer else CPUSGD if config.use_torch_optimizer_for_cpu_offload: gpu_optimizer_cls = cpu_optimizer_cls - if config.optimizer == 'adam': + if is_adam_optimizer: gpu_optimizer_cls = Adam cpu_optimizer_cls = CPUAdam optimizer_defaults = dict( @@ -778,8 +898,12 @@ def _get_megatron_emerging_optimizer( if 'linear_qkv.weight' in name and len(param.shape) == 2: param.is_qkv = True - # Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam). - config_overrides.update(_EMERGING_OPTIMIZERS[eopt_name].default_param_overrides) + # Apply optimizer-specific param overrides (e.g. muon: non-linear -> scalar optimizer). + entry = _EMERGING_OPTIMIZERS[eopt_name] + if entry.config_to_param_overrides is not None: + config_overrides.update(entry.config_to_param_overrides(config)) + else: + config_overrides.update(entry.default_param_overrides) # Build param groups and bucket by (optimizer_name, is_expert_parallel). # Layer-wise distributed optimizer handles expert params internally so we skip that split. @@ -800,7 +924,7 @@ def _get_megatron_emerging_optimizer( if opt_name in _EMERGING_OPTIMIZERS: optimizer, init_state_fn = _create_emerging_optimizer( - config, groups, eopt_name, model_chunks, pg_collection + config, groups, opt_name, model_chunks, pg_collection ) if use_layer_wise: result = (optimizer, init_state_fn) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 95e8a0c407b..f65457fcba6 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -57,7 +57,11 @@ from ..transformer.fsdp_dtensor_checkpoint import handle_experts_in_state_dict from ..transformer.module import MegatronModule from .grad_scaler import MegatronGradScaler -from .optimizer import MixedPrecisionOptimizer, _zero_grad_group_helper, param_group_identifier_keys +from .optimizer import ( + MixedPrecisionOptimizer, + _zero_grad_group_helper, + get_param_group_identifier_tuple, +) from .optimizer_config import OptimizerConfig from .param_layout import FullParamLayout, PerBufferParamLayout, pad_bucket_end, pad_param_start @@ -884,21 +888,7 @@ def load_state_dict(self, state_dict): # the ordering of parameters within its flattened parameter state # list. def make_needed_groups(param_group): - needed_groups = [] - for key in param_group_identifier_keys: - # NeMo changes these variable names from `lr_mult` and `wd_mult` - # to `pre_lr_mult` and `pre_wd_mult`, so we need to check both. - if key in param_group: - pass - elif f"pre_{key}" in param_group: - key = f"pre_{key}" - else: - raise ValueError( - f"Key {key} (or pre_{key}) not found in param_group {param_group}." - ) - needed_groups.append(param_group[key]) - needed_groups = tuple(needed_groups) - return needed_groups + return get_param_group_identifier_tuple(param_group) param_groups_map = {} for param_group in state_dict["optimizer"]["param_groups"]: diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index cc218d6ba40..835213c880b 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -87,6 +87,20 @@ def _default_param_overrides_factory() -> Dict[ParamKey, Dict[str, Any]]: } +def _nonlinear_or_embedding_param_overrides(optimizer_name: str) -> Dict[ParamKey, Dict[str, Any]]: + """Route non-matrix / embedding-class params to the requested scalar optimizer.""" + return { + ParamKey( + predicate=ParamPredicate(name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding) + ): {'optimizer': optimizer_name} + } + + +def _muon_default_param_overrides(config) -> Dict[ParamKey, Dict[str, Any]]: + """Respect the configured scalar optimizer for Muon-family nonlinear params.""" + return _nonlinear_or_embedding_param_overrides(config.muon_scalar_optimizer) + + @dataclass class EmergingOptimizerEntry: """Everything needed to create and configure an emerging optimizer. @@ -95,6 +109,8 @@ class EmergingOptimizerEntry: optimizer_cls: The torch optimizer class. init_state_fn: Lazily initialises optimizer state (needed for checkpoint formats). config_to_kwargs: ``(config, model_chunks, pg_collection) -> dict`` of constructor kwargs. + config_to_param_overrides: ``config -> dict`` of per-parameter overrides derived from the + resolved optimizer config (e.g. Muon scalar optimizer selection). default_param_overrides: Per-parameter config overrides applied automatically (e.g. route non-linear params to Adam). """ @@ -102,6 +118,7 @@ class EmergingOptimizerEntry: optimizer_cls: type init_state_fn: Callable = _eopt_init_state_fn config_to_kwargs: Callable | None = None + config_to_param_overrides: Callable | None = None default_param_overrides: Dict[ParamKey, Dict[str, Any]] = field( default_factory=_default_param_overrides_factory ) @@ -402,10 +419,17 @@ def _default_adam_based_eopt_config_to_kwargs( ) -> Dict[str, Any]: """Convert OptimizerConfig to default emerging optimizer constructor kwargs.""" kwargs = _kwargs_from_config(registry.get_optimizer_cls(eopt_name), eopt_name, config) - kwargs["betas"] = (config.adam_beta1, config.adam_beta2) + kwargs["betas"] = _default_betas_for_eopt(eopt_name, config) return kwargs +def _default_betas_for_eopt(eopt_name, config) -> tuple[float, float]: + """Return the default beta pair for an emerging optimizer.""" + if eopt_name == "lion": + return (config.lion_beta1, config.lion_beta2) + return (config.adam_beta1, config.adam_beta2) + + # ----------------------------------------------------------------------- # Register emerging optimizers # ----------------------------------------------------------------------- @@ -415,25 +439,13 @@ def _default_adam_based_eopt_config_to_kwargs( optimizer_cls=TensorParallelMuon, init_state_fn=_eopt_init_state_fn, config_to_kwargs=_muon_config_to_kwargs, - default_param_overrides={ - ParamKey( - predicate=ParamPredicate( - name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding - ) - ): {'optimizer': 'adam'} - }, + config_to_param_overrides=_muon_default_param_overrides, ), "adaptive_muon": EmergingOptimizerEntry( optimizer_cls=TensorParallelAdaptiveMuon, init_state_fn=_eopt_init_state_fn, config_to_kwargs=_adaptive_muon_config_to_kwargs, - default_param_overrides={ - ParamKey( - predicate=ParamPredicate( - name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding - ) - ): {'optimizer': 'adam'} - }, + config_to_param_overrides=_muon_default_param_overrides, ), } ) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index 9ae23bb4b7f..315d6f66091 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -94,7 +94,29 @@ def _multi_tensor_copy_this_to_that( that_.copy_(this_) -param_group_identifier_keys = ('wd_mult', 'lr_mult', 'is_expert_parallel', 'is_decoupled_lr') +param_group_identifier_keys = ( + 'wd_mult', + 'lr_mult', + 'is_expert_parallel', + 'is_decoupled_lr', + 'max_lr', + 'min_lr', + 'eps', + 'optimizer', +) + + +def get_param_group_identifier_tuple(param_group: Dict) -> tuple: + """Return a stable identifier for optimizer param-group matching and resume.""" + values = [] + for key in param_group_identifier_keys: + if key in param_group: + values.append(param_group[key]) + elif f"pre_{key}" in param_group: + values.append(param_group[f"pre_{key}"]) + else: + values.append(None) + return tuple(values) class MegatronOptimizer(ABC): @@ -412,8 +434,8 @@ def _filter_and_reorder_param_groups( current_groups: List[Dict], state_dict_groups: List[Dict] ) -> List[Dict]: """Filter and reorder state_dict parameter groups to match current optimizer groups. - Keys used for matching align with those from _get_param_groups: - (wd_mult, lr_mult, is_expert_parallel, is_decoupled_lr) + Keys used for matching align with the scheduler/optimizer fields emitted by + _get_param_groups. Args: current_groups (List[Dict]): Parameter groups from the current optimizer instance. @@ -426,22 +448,13 @@ def _filter_and_reorder_param_groups( ValueError: If parameter groups in state dict don't match current optimizer. """ # Define groups order that is needed in the current optimizer (coming from runtime) - needed_groups = [ - # NeMo may have different key for required fields, e.g., "wd_mult" to "pre_wd_mult" - tuple(g[key] if key in g else g[f"pre_{key}"] for key in param_group_identifier_keys) - for g in current_groups - ] + needed_groups = [get_param_group_identifier_tuple(g) for g in current_groups] # Keep state_dict param group order since groups are LocalNonpersistentObject # and their order is determined at runtime, not from the checkpoint. params_in_state_dict_order = [g['params'] for g in state_dict_groups] loaded_groups_map = { - tuple( - # NeMo may have different key for required fields, e.g., "wd_mult" to "pre_wd_mult" - group[key] if key in group else group[f"pre_{key}"] - for key in param_group_identifier_keys - ): group - for group in state_dict_groups + get_param_group_identifier_tuple(group): group for group in state_dict_groups } final_groups = [] diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 514d109ddc8..eaa4ed18733 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -416,9 +416,10 @@ def __post_init__(self): ) if self.use_precision_aware_optimizer: - assert ( - self.optimizer == 'adam' - ), '--use-precision-aware-optimizer only supported with adam' + assert self.optimizer == 'adam', ( + '--use-precision-aware-optimizer only supported with optimizer=adam; ' + 'AdamW semantics should continue to use decoupled_weight_decay' + ) assert ( self.use_distributed_optimizer ), '--use-precision-aware-optimizer only supported with distributed optimizer' diff --git a/megatron/core/parameterization/__init__.py b/megatron/core/parameterization/__init__.py new file mode 100644 index 00000000000..33454e7312e --- /dev/null +++ b/megatron/core/parameterization/__init__.py @@ -0,0 +1,105 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from .eval_runtime import depth_mup_eval_context, is_depth_mup_eval_enabled +from .model_policy import ResolvedModelPolicy, build_resolved_model_policy +from .roles import ( + IS_OUTPUT_PARAMETER_ATTR, + PARAMETERIZATION_ROLE_ATTR, + PARAMETERIZATION_SHARED_GROUP_ATTR, + PARAMETERIZATION_TAGS_ATTR, + ROLE_BLOCK_OUT_PROJ, + ROLE_EMBEDDING, + ROLE_HIDDEN_BIAS, + ROLE_HIDDEN_MATRIX, + ROLE_HIDDEN_VECTOR, + ROLE_HIDDEN_VECTOR_OTHER, + ROLE_MUON_MANAGED_MATRIX, + ROLE_NORM_BIAS, + ROLE_NORM_SCALE, + ROLE_OUTPUT, + ROLE_QK_NORM_SCALE, + ROLE_SHARED_EMBEDDING_OUTPUT, + ROLE_VECTOR_LIKE, + get_parameterization_role, + is_embedding_class_parameter, + is_embedding_or_output_parameter, + is_hidden_bias_parameter, + is_hidden_matrix_parameter, + is_hidden_vector_parameter, + is_muon_managed_matrix_parameter, + is_norm_parameter, + is_output_parameter, + is_qk_norm_parameter, + is_vector_like_parameter, + set_parameterization_metadata, + should_skip_depth_mup_vector_weight_decay, +) +from .spec import ( + SCALING_RECIPE_DEPTH_MUP, + SCALING_RECIPE_MUP, + SCALING_RECIPE_NONE, + CanonicalScalingSpec, + ResolvedScalingContext, + ScalingReferences, + ScalingUserConfig, + build_resolved_scaling_context, + build_scaling_user_config, + canonicalize_scaling_user_config, + sync_legacy_mup_fields, +) +from .training_policy import ( + ResolvedTrainingPolicy, + build_legacy_mup_training_policy, + build_resolved_training_policy, +) + +__all__ = [ + 'CanonicalScalingSpec', + 'IS_OUTPUT_PARAMETER_ATTR', + 'PARAMETERIZATION_ROLE_ATTR', + 'PARAMETERIZATION_SHARED_GROUP_ATTR', + 'PARAMETERIZATION_TAGS_ATTR', + 'ROLE_BLOCK_OUT_PROJ', + 'ROLE_EMBEDDING', + 'ROLE_HIDDEN_BIAS', + 'ROLE_HIDDEN_MATRIX', + 'ROLE_HIDDEN_VECTOR', + 'ROLE_HIDDEN_VECTOR_OTHER', + 'ROLE_MUON_MANAGED_MATRIX', + 'ROLE_NORM_BIAS', + 'ROLE_NORM_SCALE', + 'ROLE_OUTPUT', + 'ROLE_QK_NORM_SCALE', + 'ROLE_SHARED_EMBEDDING_OUTPUT', + 'ROLE_VECTOR_LIKE', + 'ResolvedModelPolicy', + 'ResolvedScalingContext', + 'ResolvedTrainingPolicy', + 'SCALING_RECIPE_DEPTH_MUP', + 'SCALING_RECIPE_MUP', + 'SCALING_RECIPE_NONE', + 'ScalingReferences', + 'ScalingUserConfig', + 'build_legacy_mup_training_policy', + 'build_resolved_model_policy', + 'build_resolved_scaling_context', + 'build_resolved_training_policy', + 'build_scaling_user_config', + 'canonicalize_scaling_user_config', + 'depth_mup_eval_context', + 'get_parameterization_role', + 'is_depth_mup_eval_enabled', + 'is_embedding_class_parameter', + 'is_embedding_or_output_parameter', + 'is_hidden_bias_parameter', + 'is_hidden_matrix_parameter', + 'is_hidden_vector_parameter', + 'is_muon_managed_matrix_parameter', + 'is_norm_parameter', + 'is_output_parameter', + 'is_qk_norm_parameter', + 'is_vector_like_parameter', + 'set_parameterization_metadata', + 'should_skip_depth_mup_vector_weight_decay', + 'sync_legacy_mup_fields', +] diff --git a/megatron/core/parameterization/eval_runtime.py b/megatron/core/parameterization/eval_runtime.py new file mode 100644 index 00000000000..25df3564c27 --- /dev/null +++ b/megatron/core/parameterization/eval_runtime.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Iterator + +_DEPTH_MUP_EVAL_DEPTH: ContextVar[int] = ContextVar('depth_mup_eval_depth', default=0) + + +def is_depth_mup_eval_enabled() -> bool: + return _DEPTH_MUP_EVAL_DEPTH.get() > 0 + + +@contextmanager +def depth_mup_eval_context(enabled: bool) -> Iterator[None]: + if not enabled: + yield + return + + token = _DEPTH_MUP_EVAL_DEPTH.set(_DEPTH_MUP_EVAL_DEPTH.get() + 1) + try: + yield + finally: + _DEPTH_MUP_EVAL_DEPTH.reset(token) diff --git a/megatron/core/parameterization/model_policy.py b/megatron/core/parameterization/model_policy.py new file mode 100644 index 00000000000..41506ba7089 --- /dev/null +++ b/megatron/core/parameterization/model_policy.py @@ -0,0 +1,139 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +import functools +import math +from dataclasses import dataclass +from typing import Iterable, Optional + +import torch +from torch import Tensor + +from megatron.core.utils import ( + init_method_normal, + mup_scaled_init_method_normal, + scaled_init_method_normal, +) + +from .spec import ResolvedScalingContext, build_resolved_scaling_context + + +@dataclass(frozen=True) +class ResolvedModelPolicy: + context: ResolvedScalingContext + + @property + def enabled(self) -> bool: + return self.context.enabled + + @property + def uses_width_mup(self) -> bool: + return self.context.uses_width_mup + + @property + def residual_branch_multiplier(self) -> float: + return self.context.depth_mult**self.context.residual_branch_depth_power + + @property + def dense_block_out_proj_init_multiplier(self) -> float: + return self.context.depth_mult**self.context.block_out_proj_init_depth_power + + def resolve_attention_softmax_scale( + self, *, softmax_scale: Optional[float], kv_channels: int + ) -> Optional[float]: + if softmax_scale is not None or not self.uses_width_mup: + return softmax_scale + base_head_scale = ( + 1.0 + if self.context.references.base_head_dim is None + else self.context.references.base_head_dim**0.5 + ) + return base_head_scale / (kv_channels**self.context.attention_scale_power) + + def build_hidden_init_method(self, *, init_method_std: float): + if not self.uses_width_mup: + return init_method_normal(init_method_std) + return init_method_normal(init_method_std / math.sqrt(self.context.width_mult)) + + def build_default_output_layer_init_method( + self, *, init_method_std: float, num_layers: int, is_hybrid_model: bool + ): + multiplier = 2.0 if not is_hybrid_model else 1.0 + if self.uses_width_mup: + return mup_scaled_init_method_normal( + init_method_std, num_layers, self.context.width_mult, multiplier=multiplier + ) + return scaled_init_method_normal(init_method_std, num_layers, multiplier=multiplier) + + def dense_block_output_init_method( + self, + *, + default_init_method, + init_method_std: float, + num_layers: int, + is_hybrid_model: bool, + output_layer_init_method_is_user_provided: bool, + apply_depth_hook: bool = True, + ): + if output_layer_init_method_is_user_provided: + return default_init_method + if not apply_depth_hook: + return default_init_method + if self.dense_block_out_proj_init_multiplier == 1.0: + return default_init_method + + multiplier = 2.0 if not is_hybrid_model else 1.0 + std = init_method_std / math.sqrt(multiplier * num_layers) + if self.uses_width_mup: + std = std / math.sqrt(self.context.width_mult) + std = std * self.dense_block_out_proj_init_multiplier + return functools.partial(torch.nn.init.normal_, mean=0.0, std=std) + + def output_layer_init_method( + self, + *, + share_embeddings_and_output_weights: bool, + default_init_method, + embedding_init_method, + ): + if self.uses_width_mup and not share_embeddings_and_output_weights: + return embedding_init_method + return default_init_method + + def mark_embedding_class_parameters(self, parameters: Iterable[torch.nn.Parameter]) -> None: + if not self.uses_width_mup: + return + for param in parameters: + param.is_embedding_parameter = True + + def scale_embedding_activations(self, embeddings: Tensor) -> Tensor: + if not self.uses_width_mup or self.context.embedding_mult == 1.0: + return embeddings + return embeddings * self.context.embedding_mult + + def scale_output_logits(self, logits: Tensor) -> Tensor: + if not self.uses_width_mup or self.context.output_mult == 1.0: + return logits + return logits * self.context.output_mult + + def scale_residual_branch_output( + self, output_with_bias: tuple[Tensor, Tensor | None] + ) -> tuple[Tensor, Tensor | None]: + if self.residual_branch_multiplier == 1.0: + return output_with_bias + + output, bias = output_with_bias + scaled_output = output * self.residual_branch_multiplier + scaled_bias = None if bias is None else bias * self.residual_branch_multiplier + return scaled_output, scaled_bias + + +def build_resolved_model_policy(config) -> ResolvedModelPolicy: + cached = getattr(config, '_resolved_model_policy', None) + if cached is not None: + return cached + + policy = ResolvedModelPolicy(build_resolved_scaling_context(config)) + setattr(config, '_resolved_model_policy', policy) + return policy diff --git a/megatron/core/parameterization/roles.py b/megatron/core/parameterization/roles.py new file mode 100644 index 00000000000..51bfafe9a55 --- /dev/null +++ b/megatron/core/parameterization/roles.py @@ -0,0 +1,160 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +from typing import Any, Iterable, Optional + +PARAMETERIZATION_ROLE_ATTR = 'parameterization_role' +PARAMETERIZATION_SHARED_GROUP_ATTR = 'parameterization_shared_group' +PARAMETERIZATION_TAGS_ATTR = 'parameterization_tags' +IS_OUTPUT_PARAMETER_ATTR = 'is_output_parameter' + +ROLE_EMBEDDING = 'embedding' +ROLE_OUTPUT = 'output' +ROLE_SHARED_EMBEDDING_OUTPUT = 'shared_embedding_output' +ROLE_BLOCK_OUT_PROJ = 'block_out_proj' +ROLE_HIDDEN_MATRIX = 'hidden_matrix' +ROLE_HIDDEN_VECTOR = 'hidden_vector' +ROLE_HIDDEN_BIAS = 'hidden_bias' +ROLE_NORM_SCALE = 'norm_scale' +ROLE_NORM_BIAS = 'norm_bias' +ROLE_QK_NORM_SCALE = 'qk_norm_scale' +ROLE_HIDDEN_VECTOR_OTHER = 'hidden_vector_other' +ROLE_VECTOR_LIKE = 'vector_like' +ROLE_MUON_MANAGED_MATRIX = 'muon_managed_matrix' + +_EMBEDDING_CLASS_ROLES = frozenset((ROLE_EMBEDDING, ROLE_OUTPUT, ROLE_SHARED_EMBEDDING_OUTPUT)) +_OUTPUT_ROLES = frozenset((ROLE_OUTPUT, ROLE_SHARED_EMBEDDING_OUTPUT)) +_HIDDEN_VECTOR_ROLES = frozenset( + ( + ROLE_HIDDEN_VECTOR, + ROLE_HIDDEN_BIAS, + ROLE_NORM_SCALE, + ROLE_NORM_BIAS, + ROLE_QK_NORM_SCALE, + ROLE_HIDDEN_VECTOR_OTHER, + ) +) +_NORM_ROLES = frozenset((ROLE_NORM_SCALE, ROLE_NORM_BIAS, ROLE_QK_NORM_SCALE)) + + +def set_parameterization_metadata( + param: Any, *, role: str, shared_group: Optional[str] = None, tags: Iterable[str] = () +) -> None: + setattr(param, PARAMETERIZATION_ROLE_ATTR, role) + if shared_group is not None: + setattr(param, PARAMETERIZATION_SHARED_GROUP_ATTR, shared_group) + if tags: + setattr(param, PARAMETERIZATION_TAGS_ATTR, tuple(tags)) + + +def get_parameterization_role(param: Any) -> Optional[str]: + return getattr(param, PARAMETERIZATION_ROLE_ATTR, None) + + +def is_output_parameter(param: Any) -> bool: + if hasattr(param, IS_OUTPUT_PARAMETER_ATTR): + return bool(getattr(param, IS_OUTPUT_PARAMETER_ATTR)) + return get_parameterization_role(param) in _OUTPUT_ROLES + + +def is_embedding_or_output_parameter(param: Any) -> bool: + if hasattr(param, 'is_embedding_or_output_parameter'): + return bool(param.is_embedding_or_output_parameter) + return get_parameterization_role(param) in _EMBEDDING_CLASS_ROLES + + +def is_embedding_class_parameter(param: Any, param_name: Optional[str] = None) -> bool: + if getattr(param, 'shared_embedding', False): + return True + if hasattr(param, 'is_embedding_parameter'): + return bool(param.is_embedding_parameter) + if get_parameterization_role(param) in _EMBEDDING_CLASS_ROLES: + return True + # Compatibility-only fallback for older unannotated parameters. The scaling + # recipes added in this branch are intended to rely on explicit metadata. + return bool(param_name and 'embedding' in param_name.lower()) + + +def is_vector_like_parameter(param: Any, param_name: Optional[str] = None) -> bool: + if is_embedding_class_parameter(param, param_name): + return True + return param.dim() <= 1 + + +def _lower_name(param_name: Optional[str]) -> str: + return param_name.lower() if param_name else '' + + +def is_qk_norm_parameter(param: Any, param_name: Optional[str] = None) -> bool: + role = get_parameterization_role(param) + if role == ROLE_QK_NORM_SCALE: + return True + name = _lower_name(param_name) + return param.dim() <= 1 and ('q_layernorm.' in name or 'k_layernorm.' in name) + + +def is_norm_parameter(param: Any, param_name: Optional[str] = None) -> bool: + role = get_parameterization_role(param) + if role in _NORM_ROLES: + return True + if param.dim() > 1: + return False + name = _lower_name(param_name) + return 'layernorm' in name or 'layer_norm' in name or 'rmsnorm' in name or '.norm.' in name + + +def is_hidden_bias_parameter(param: Any, param_name: Optional[str] = None) -> bool: + role = get_parameterization_role(param) + if role == ROLE_HIDDEN_BIAS: + return True + if role in _EMBEDDING_CLASS_ROLES or role in _NORM_ROLES: + return False + name = _lower_name(param_name) + return param.dim() <= 1 and name.endswith('.bias') and not is_norm_parameter(param, param_name) + + +def should_skip_depth_mup_vector_weight_decay( + param: Any, param_name: Optional[str] = None, *, apply_wd_to_qk_layernorm: bool = False +) -> bool: + """Return true for vector-like params outside the AdamW hidden-bias table row. + + The spectral width-depth table gives base weight decay to hidden biases, but + Megatron's 1-D tensors also include normalization scale parameters. Keep + hidden linear/MLP/attention biases on base weight decay and keep normalization + or otherwise unknown 1-D tensors on the conservative no-WD path unless the + user explicitly asks to apply WD to q/k layernorm. + """ + if is_embedding_class_parameter(param, param_name): + return False + if is_hidden_bias_parameter(param, param_name): + return False + if apply_wd_to_qk_layernorm and is_qk_norm_parameter(param, param_name): + return False + if is_norm_parameter(param, param_name): + return True + return param.dim() <= 1 + + +def is_hidden_vector_parameter(param: Any, param_name: Optional[str] = None) -> bool: + role = get_parameterization_role(param) + if role in _HIDDEN_VECTOR_ROLES: + return True + if role in _EMBEDDING_CLASS_ROLES: + return False + return param.dim() <= 1 and not is_embedding_class_parameter(param, param_name) + + +def is_hidden_matrix_parameter(param: Any, param_name: Optional[str] = None) -> bool: + role = get_parameterization_role(param) + if role == ROLE_HIDDEN_MATRIX: + return True + if role in _HIDDEN_VECTOR_ROLES or role in _EMBEDDING_CLASS_ROLES: + return False + return param.dim() > 1 and not is_embedding_class_parameter(param, param_name) + + +def is_muon_managed_matrix_parameter(param: Any, *, optimizer_type: str) -> bool: + if 'muon' not in optimizer_type.lower(): + return False + return param.dim() == 2 and not is_embedding_or_output_parameter(param) diff --git a/megatron/core/parameterization/spec.py b/megatron/core/parameterization/spec.py new file mode 100644 index 00000000000..be3118375b3 --- /dev/null +++ b/megatron/core/parameterization/spec.py @@ -0,0 +1,313 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Any, Literal, Optional + +SCALING_RECIPE_NONE = 'none' +SCALING_RECIPE_MUP = 'mup' +SCALING_RECIPE_DEPTH_MUP = 'depth_mup' +SCALING_RECIPE_VALUES = (SCALING_RECIPE_NONE, SCALING_RECIPE_MUP, SCALING_RECIPE_DEPTH_MUP) + + +@dataclass(frozen=True) +class ScalingUserConfig: + recipe: Optional[Literal['none', 'mup', 'depth_mup']] = None + base_hidden_size: Optional[int] = None + base_num_layers: Optional[int] = None + base_head_dim: Optional[float] = None + residual_branch_depth_power: Optional[float] = None + hidden_lr_depth_power: Optional[float] = None + block_out_proj_init_depth_power: Optional[float] = None + use_mup_alias: bool = False + mup_width_mult: float = 1.0 + mup_base_hidden_size: Optional[int] = None + mup_embedding_mult: float = 1.0 + mup_output_mult: float = 1.0 + mup_base_head_dim: Optional[float] = None + mup_attn_scale_power: float = 1.0 + + +@dataclass(frozen=True) +class ScalingReferences: + current_hidden_size: int + current_num_layers: int + current_head_dim: int + base_hidden_size: Optional[int] = None + base_num_layers: Optional[int] = None + base_head_dim: Optional[float] = None + + +@dataclass(frozen=True) +class CanonicalScalingSpec: + recipe: Literal['none', 'mup', 'depth_mup'] + references: ScalingReferences + embedding_mult: float = 1.0 + output_mult: float = 1.0 + attention_scale_power: float = 1.0 + residual_branch_depth_power: float = 0.0 + hidden_lr_depth_power: float = 0.0 + block_out_proj_init_depth_power: float = 0.0 + + +@dataclass(frozen=True) +class ResolvedScalingContext: + recipe: Literal['none', 'mup', 'depth_mup'] + references: ScalingReferences + width_mult: float = 1.0 + depth_mult: float = 1.0 + embedding_mult: float = 1.0 + output_mult: float = 1.0 + attention_scale_power: float = 1.0 + residual_branch_depth_power: float = 0.0 + hidden_lr_depth_power: float = 0.0 + block_out_proj_init_depth_power: float = 0.0 + + @property + def enabled(self) -> bool: + return self.recipe != SCALING_RECIPE_NONE + + @property + def uses_width_mup(self) -> bool: + return self.recipe in (SCALING_RECIPE_MUP, SCALING_RECIPE_DEPTH_MUP) + + @property + def is_depth_mup(self) -> bool: + return self.recipe == SCALING_RECIPE_DEPTH_MUP + + +def _resolve_aliased_value( + explicit_value: Optional[float | int], + legacy_value: Optional[float | int], + *, + explicit_name: str, + legacy_name: str, +) -> Optional[float | int]: + if explicit_value is None: + return legacy_value + if legacy_value is None: + return explicit_value + if explicit_value != legacy_value: + raise ValueError( + f"{explicit_name} ({explicit_value}) conflicts with {legacy_name} ({legacy_value}). " + f"Specify only one or set them to the same value." + ) + return explicit_value + + +def _non_default_scaling_fields( + user_config: ScalingUserConfig, *, include_legacy_mup_fields: bool +) -> list[str]: + candidates = { + 'scaling_base_hidden_size': user_config.base_hidden_size, + 'scaling_base_num_layers': user_config.base_num_layers, + 'scaling_base_head_dim': user_config.base_head_dim, + 'scaling_residual_branch_depth_power': user_config.residual_branch_depth_power, + 'scaling_hidden_lr_depth_power': user_config.hidden_lr_depth_power, + 'scaling_block_out_proj_init_depth_power': user_config.block_out_proj_init_depth_power, + } + if include_legacy_mup_fields: + candidates['mup_base_hidden_size'] = user_config.mup_base_hidden_size + candidates['mup_base_head_dim'] = user_config.mup_base_head_dim + if user_config.mup_embedding_mult != 1.0: + candidates['mup_embedding_mult'] = user_config.mup_embedding_mult + if user_config.mup_output_mult != 1.0: + candidates['mup_output_mult'] = user_config.mup_output_mult + if user_config.mup_attn_scale_power != 1.0: + candidates['mup_attn_scale_power'] = user_config.mup_attn_scale_power + if user_config.mup_width_mult != 1.0: + candidates['mup_width_mult'] = user_config.mup_width_mult + return [name for name, value in candidates.items() if value is not None] + + +def _infer_current_head_dim(config: Any) -> int: + kv_channels = getattr(config, 'kv_channels', None) + if kv_channels is not None: + return kv_channels + + hidden_size = getattr(config, 'hidden_size') + num_attention_heads = getattr(config, 'num_attention_heads') + if num_attention_heads is None or num_attention_heads <= 0: + raise AttributeError( + "Cannot resolve current head dimension without kv_channels or a positive " + "num_attention_heads value." + ) + return hidden_size // num_attention_heads + + +def build_scaling_user_config(config: Any) -> ScalingUserConfig: + return ScalingUserConfig( + recipe=getattr(config, 'scaling_recipe', None), + base_hidden_size=getattr(config, 'scaling_base_hidden_size', None), + base_num_layers=getattr(config, 'scaling_base_num_layers', None), + base_head_dim=getattr(config, 'scaling_base_head_dim', None), + residual_branch_depth_power=getattr(config, 'scaling_residual_branch_depth_power', None), + hidden_lr_depth_power=getattr(config, 'scaling_hidden_lr_depth_power', None), + block_out_proj_init_depth_power=getattr( + config, 'scaling_block_out_proj_init_depth_power', None + ), + use_mup_alias=bool(getattr(config, 'use_mup', False)), + mup_width_mult=getattr(config, 'mup_width_mult', 1.0), + mup_base_hidden_size=getattr(config, 'mup_base_hidden_size', None), + mup_embedding_mult=getattr(config, 'mup_embedding_mult', 1.0), + mup_output_mult=getattr(config, 'mup_output_mult', 1.0), + mup_base_head_dim=getattr(config, 'mup_base_head_dim', None), + mup_attn_scale_power=getattr(config, 'mup_attn_scale_power', 1.0), + ) + + +def canonicalize_scaling_user_config( + user_config: ScalingUserConfig, config: Any +) -> CanonicalScalingSpec: + recipe = user_config.recipe + if recipe is None: + recipe = SCALING_RECIPE_MUP if user_config.use_mup_alias else SCALING_RECIPE_NONE + elif user_config.use_mup_alias: + if recipe != SCALING_RECIPE_MUP: + raise ValueError( + f"--scaling-recipe {recipe} conflicts with --use-mup. " + "Use either the canonical MuP recipe or the legacy MuP alias, not both." + ) + if recipe not in SCALING_RECIPE_VALUES: + raise ValueError(f"Unsupported scaling recipe: {recipe}") + + base_hidden_size = _resolve_aliased_value( + user_config.base_hidden_size, + user_config.mup_base_hidden_size, + explicit_name='--scaling-base-hidden-size', + legacy_name='--mup-base-hidden-size', + ) + base_head_dim = _resolve_aliased_value( + user_config.base_head_dim, + user_config.mup_base_head_dim, + explicit_name='--scaling-base-head-dim', + legacy_name='--mup-base-head-dim', + ) + + if recipe == SCALING_RECIPE_NONE: + non_default_fields = _non_default_scaling_fields( + user_config, include_legacy_mup_fields=True + ) + if non_default_fields: + raise ValueError( + "Scaling overrides require a non-'none' scaling recipe (for example `mup` or " + "`depth_mup`). Non-default fields: " + ", ".join(non_default_fields) + ) + references = ScalingReferences( + current_hidden_size=config.hidden_size, + current_num_layers=config.num_layers, + current_head_dim=_infer_current_head_dim(config), + ) + return CanonicalScalingSpec(recipe=SCALING_RECIPE_NONE, references=references) + + if base_hidden_size is None: + base_hidden_size = config.hidden_size + if base_hidden_size <= 0: + raise AssertionError('--scaling-base-hidden-size must be positive.') + width_mult = config.hidden_size / base_hidden_size + if user_config.mup_width_mult != 1.0 and not math.isclose( + user_config.mup_width_mult, width_mult, rel_tol=1e-12, abs_tol=1e-12 + ): + raise ValueError( + "--mup-width-mult is deprecated as an input and must match the derived " + f"hidden_size / scaling_base_hidden_size value ({width_mult}). " + f"Got --mup-width-mult={user_config.mup_width_mult}." + ) + + base_num_layers = user_config.base_num_layers + if base_num_layers is None: + base_num_layers = config.num_layers + if base_num_layers <= 0: + raise AssertionError('--scaling-base-num-layers must be positive.') + if base_head_dim is not None and base_head_dim <= 0: + raise AssertionError('--scaling-base-head-dim must be positive.') + + references = ScalingReferences( + current_hidden_size=config.hidden_size, + current_num_layers=config.num_layers, + current_head_dim=_infer_current_head_dim(config), + base_hidden_size=base_hidden_size, + base_num_layers=base_num_layers, + base_head_dim=base_head_dim, + ) + residual_branch_depth_power = user_config.residual_branch_depth_power + if residual_branch_depth_power is None: + residual_branch_depth_power = -1.0 if recipe == SCALING_RECIPE_DEPTH_MUP else 0.0 + + hidden_lr_depth_power = user_config.hidden_lr_depth_power + if hidden_lr_depth_power is None: + hidden_lr_depth_power = 0.0 + + block_out_proj_init_depth_power = user_config.block_out_proj_init_depth_power + if block_out_proj_init_depth_power is None: + block_out_proj_init_depth_power = 0.5 if recipe == SCALING_RECIPE_DEPTH_MUP else 0.0 + + return CanonicalScalingSpec( + recipe=recipe, + references=references, + embedding_mult=user_config.mup_embedding_mult, + output_mult=user_config.mup_output_mult, + attention_scale_power=user_config.mup_attn_scale_power, + residual_branch_depth_power=float(residual_branch_depth_power), + hidden_lr_depth_power=float(hidden_lr_depth_power), + block_out_proj_init_depth_power=float(block_out_proj_init_depth_power), + ) + + +def resolve_scaling_context(canonical_spec: CanonicalScalingSpec) -> ResolvedScalingContext: + if canonical_spec.recipe == SCALING_RECIPE_NONE: + return ResolvedScalingContext( + recipe=SCALING_RECIPE_NONE, references=canonical_spec.references + ) + + refs = canonical_spec.references + assert refs.base_hidden_size is not None + assert refs.base_num_layers is not None + width_mult = refs.current_hidden_size / refs.base_hidden_size + depth_mult = refs.current_num_layers / refs.base_num_layers + output_mult = canonical_spec.output_mult + if output_mult == 1.0 and width_mult != 1.0: + output_mult = 1.0 / width_mult + + return ResolvedScalingContext( + recipe=canonical_spec.recipe, + references=refs, + width_mult=width_mult, + depth_mult=depth_mult, + embedding_mult=canonical_spec.embedding_mult, + output_mult=output_mult, + attention_scale_power=canonical_spec.attention_scale_power, + residual_branch_depth_power=canonical_spec.residual_branch_depth_power, + hidden_lr_depth_power=canonical_spec.hidden_lr_depth_power, + block_out_proj_init_depth_power=canonical_spec.block_out_proj_init_depth_power, + ) + + +def build_resolved_scaling_context(config: Any) -> ResolvedScalingContext: + user_config = build_scaling_user_config(config) + canonical_spec = canonicalize_scaling_user_config(user_config, config) + return resolve_scaling_context(canonical_spec) + + +def sync_legacy_mup_fields(config: Any, context: ResolvedScalingContext) -> None: + config.scaling_recipe = context.recipe + config.use_mup = context.recipe == SCALING_RECIPE_MUP + if context.recipe == SCALING_RECIPE_NONE: + return + + config.scaling_base_hidden_size = context.references.base_hidden_size + config.scaling_base_num_layers = context.references.base_num_layers + config.scaling_base_head_dim = context.references.base_head_dim + config.scaling_residual_branch_depth_power = context.residual_branch_depth_power + config.scaling_hidden_lr_depth_power = context.hidden_lr_depth_power + config.scaling_block_out_proj_init_depth_power = context.block_out_proj_init_depth_power + + if context.recipe != SCALING_RECIPE_MUP: + return + + config.mup_base_hidden_size = context.references.base_hidden_size + config.mup_base_head_dim = context.references.base_head_dim + config.mup_width_mult = context.width_mult + config.mup_output_mult = context.output_mult diff --git a/megatron/core/parameterization/training_policy.py b/megatron/core/parameterization/training_policy.py new file mode 100644 index 00000000000..b77517dd3a0 --- /dev/null +++ b/megatron/core/parameterization/training_policy.py @@ -0,0 +1,147 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +from dataclasses import dataclass + +from .spec import ( + SCALING_RECIPE_MUP, + ResolvedScalingContext, + ScalingReferences, + build_resolved_scaling_context, +) + + +@dataclass(frozen=True) +class ResolvedTrainingPolicy: + context: ResolvedScalingContext + optimizer_type: str = 'adam' + + def __post_init__(self) -> None: + if self.context.is_depth_mup and not self.is_adam_optimizer: + raise ValueError( + "scaling_recipe='depth_mup' currently supports optimizer='adam' only. " + "AdamW semantics should continue to use decoupled_weight_decay. " + "SGD depth-mup requires explicit hidden-weight, hidden-bias, norm/vector, " + "and input/output-bias rules and is intentionally out of scope for v1." + ) + + @property + def enabled(self) -> bool: + return self.context.enabled + + @property + def optimizer_type_lower(self) -> str: + return self.optimizer_type.lower() + + @property + def uses_width_mup(self) -> bool: + return self.context.uses_width_mup + + @property + def is_sgd_optimizer(self) -> bool: + return self.optimizer_type_lower == 'sgd' + + @property + def is_adam_optimizer(self) -> bool: + return self.optimizer_type_lower == 'adam' + + @property + def is_muon_optimizer(self) -> bool: + return 'muon' in self.optimizer_type_lower + + @property + def hidden_lr_width_power(self) -> float: + if not self.enabled or not self.uses_width_mup: + return 0.0 + return 0.0 if self.is_sgd_optimizer else -1.0 + + @property + def hidden_lr_multiplier(self) -> float: + if not self.enabled: + return 1.0 + return (self.context.width_mult**self.hidden_lr_width_power) * ( + self.context.depth_mult**self.context.hidden_lr_depth_power + ) + + @property + def hidden_vector_lr_multiplier(self) -> float: + if not (self.enabled and self.uses_width_mup and self.is_sgd_optimizer): + return 1.0 + return self.context.width_mult + + @property + def hidden_eps_depth_power(self) -> float: + if not (self.enabled and self.uses_width_mup and self.is_adam_optimizer): + return 0.0 + return -1.0 if self.context.is_depth_mup else 0.0 + + @property + def hidden_eps_multiplier(self) -> float: + if not (self.enabled and self.uses_width_mup and self.is_adam_optimizer): + return 1.0 + return (1.0 / self.context.width_mult) * ( + self.context.depth_mult**self.hidden_eps_depth_power + ) + + @property + def hidden_vector_eps_multiplier(self) -> float: + if not (self.enabled and self.uses_width_mup and self.is_adam_optimizer): + return 1.0 + if not self.context.is_depth_mup: + return 1.0 + return self.hidden_eps_multiplier + + @property + def embedding_class_eps_multiplier(self) -> float: + if not (self.enabled and self.uses_width_mup and self.is_adam_optimizer): + return 1.0 + if not self.context.is_depth_mup: + return 1.0 + return 1.0 / self.context.width_mult + + @property + def hidden_matrix_wd_multiplier(self) -> float: + if not (self.enabled and self.context.is_depth_mup and self.is_adam_optimizer): + return 1.0 + return self.context.width_mult + + @property + def hidden_vector_wd_multiplier(self) -> float: + return 1.0 + + @property + def embedding_class_wd_multiplier(self) -> float: + return 1.0 + + @property + def vector_like_lr_multiplier(self) -> float: + # Backward-compatible alias for the old generic vector-like MuP path. + return self.hidden_vector_lr_multiplier + + +def build_resolved_training_policy(config, optimizer_type: str = 'adam') -> ResolvedTrainingPolicy: + return ResolvedTrainingPolicy( + context=build_resolved_scaling_context(config), optimizer_type=optimizer_type + ) + + +def build_legacy_mup_training_policy( + *, mup_width_mult: float, optimizer_type: str = 'adam' +) -> ResolvedTrainingPolicy: + return ResolvedTrainingPolicy( + context=ResolvedScalingContext( + recipe=SCALING_RECIPE_MUP, + references=ScalingReferences( + current_hidden_size=1, + current_num_layers=1, + current_head_dim=1, + base_hidden_size=1, + base_num_layers=1, + base_head_dim=1.0, + ), + width_mult=mup_width_mult, + depth_mult=1.0, + ), + optimizer_type=optimizer_type, + ) diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index f89259be442..55cb2240d69 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -27,6 +27,7 @@ get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) +from megatron.core.parameterization import build_resolved_model_policy from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, ) @@ -352,6 +353,19 @@ def __init__( self.config.fine_grained_activation_offloading and "attn_proj" in self.config.offload_modules ) + model_scaling_policy = build_resolved_model_policy(self.config) + linear_proj_init_method = self.config.output_layer_init_method + if self.attention_type != "cross": + linear_proj_init_method = model_scaling_policy.dense_block_output_init_method( + default_init_method=self.config.output_layer_init_method, + init_method_std=self.config.init_method_std, + num_layers=self.config.num_layers, + is_hybrid_model=self.config.is_hybrid_model, + output_layer_init_method_is_user_provided=getattr( + self.config, '_parameterization_output_layer_init_method_user_provided', False + ), + apply_depth_hook=self.attention_type != "cross", + ) # Output. self.linear_proj = build_module( @@ -359,7 +373,7 @@ def __init__( self.query_projection_size, self.config.hidden_size, config=self.config, - init_method=self.config.output_layer_init_method, + init_method=linear_proj_init_method, bias=self.config.add_bias_linear, input_is_parallel=True, skip_bias_add=True, diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index 860118b17a3..a5a8de2901d 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -168,7 +168,6 @@ def __init__( cp_comm_type=cp_comm_type, pg_collection=self.pg_collection, ) - # Output. self.linear_proj = build_module( submodules.linear_proj, diff --git a/megatron/core/transformer/mlp.py b/megatron/core/transformer/mlp.py index 46979a8ba8f..9a52bb1c4b1 100644 --- a/megatron/core/transformer/mlp.py +++ b/megatron/core/transformer/mlp.py @@ -23,6 +23,7 @@ ) from megatron.core.fusions.fused_bias_gelu import bias_gelu_impl from megatron.core.fusions.fused_bias_swiglu import bias_swiglu_impl, weighted_bias_swiglu_impl +from megatron.core.parameterization import build_resolved_model_policy from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_config import TransformerConfig @@ -168,6 +169,7 @@ def __init__( is_expert: bool = False, input_size: Optional[int] = None, ffn_hidden_size: Optional[int] = None, + apply_block_output_init_scaling: bool = False, tp_group: Optional[torch.distributed.ProcessGroup] = None, ): super().__init__(config=config) @@ -201,6 +203,19 @@ def __init__( fc1_stride = 1 else: fc1_stride = 1 + model_scaling_policy = build_resolved_model_policy(self.config) + fc2_init_method = not_none(self.config.output_layer_init_method) + if apply_block_output_init_scaling and not is_expert: + fc2_init_method = model_scaling_policy.dense_block_output_init_method( + default_init_method=not_none(self.config.output_layer_init_method), + init_method_std=self.config.init_method_std, + num_layers=self.config.num_layers, + is_hybrid_model=self.config.is_hybrid_model, + output_layer_init_method_is_user_provided=getattr( + self.config, '_parameterization_output_layer_init_method_user_provided', False + ), + apply_depth_hook=not is_expert, + ) # Use moe_latent_size only for routed experts. 'is_expert' is false for # shared_experts. @@ -231,7 +246,7 @@ def __init__( self.config.hidden_size if not use_latent_size else self.config.moe_latent_size ), config=self.config, - init_method=not_none(self.config.output_layer_init_method), + init_method=fc2_init_method, bias=self.config.add_bias_linear, input_is_parallel=True, skip_bias_add=True, @@ -386,6 +401,7 @@ def as_mlp_submodule( is_expert=is_expert, input_size=input_size, ffn_hidden_size=ffn_hidden_size, + apply_block_output_init_scaling=True, ) diff --git a/megatron/core/transformer/multi_latent_attention.py b/megatron/core/transformer/multi_latent_attention.py index 601ae89fae1..217bd8d10b3 100644 --- a/megatron/core/transformer/multi_latent_attention.py +++ b/megatron/core/transformer/multi_latent_attention.py @@ -208,7 +208,6 @@ def __init__( cp_comm_type=cp_comm_type, pg_collection=self.pg_collection, ) - # Output. self.linear_proj = build_module( submodules.linear_proj, diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index bb044787b9c..b639f3979ae 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -1,7 +1,6 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging -import math import warnings from dataclasses import dataclass, field from typing import Callable, List, Literal, Optional, Tuple, Union @@ -11,6 +10,13 @@ from megatron.core.enums import Fp4Recipe, Fp8Recipe from megatron.core.inference.moe import InferenceGroupedGemmBackend +from megatron.core.parameterization import ( + SCALING_RECIPE_DEPTH_MUP, + SCALING_RECIPE_MUP, + build_resolved_model_policy, + build_resolved_scaling_context, + sync_legacy_mup_fields, +) from megatron.core.quantization.quant_config import RecipeConfig from megatron.core.transformer.enums import AttnBackend, CudaGraphScope from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout @@ -18,14 +24,7 @@ from .._rank_utils import log_single_rank from ..fusions.fused_bias_geglu import quick_gelu from ..model_parallel_config import ModelParallelConfig -from ..utils import ( - get_te_version, - init_method_normal, - is_te_min_version, - is_torch_min_version, - mup_scaled_init_method_normal, - scaled_init_method_normal, -) +from ..utils import get_te_version, init_method_normal, is_te_min_version, is_torch_min_version logger = logging.getLogger(__name__) @@ -353,17 +352,64 @@ class TransformerConfig(ModelParallelConfig): #################### # MuP (Maximal Update Parameterization) #################### + scaling_recipe: Optional[Literal['none', 'mup', 'depth_mup']] = None + """ + Canonical scaling recipe. `mup` preserves the current Megatron MuP semantics, + `depth_mup` is the spectral width-depth μP recipe for dense GPT-style + residual Transformer blocks using `optimizer='adam'` with AdamW-style + semantics via `decoupled_weight_decay` when weight decay is nonzero, + `none` keeps the standard parameterization, and `None` means unspecified so + legacy alias resolution can decide the recipe. + """ + + scaling_base_hidden_size: Optional[int] = None + """ + Reference hidden size for width-based scaling recipes. Under `mup`, this is + equivalent to `mup_base_hidden_size`. + """ + + scaling_base_num_layers: Optional[int] = None + """ + Reference number of transformer layers for depth-based scaling recipes. + Defaults to `num_layers` when not specified. + """ + + scaling_base_head_dim: Optional[float] = None + """ + Reference attention head dimension for head-dimension-aware scaling rules. + Under `mup`, this is equivalent to `mup_base_head_dim`. + """ + + scaling_residual_branch_depth_power: Optional[float] = None + """ + Relative depth exponent for dense self-attention/MLP residual-branch outputs. + Under `depth_mup`, the default is `-1.0`. + """ + + scaling_hidden_lr_depth_power: Optional[float] = None + """ + Relative depth exponent for hidden matrix-like LR overrides. Under + `depth_mup`, the default remains `0.0` for `optimizer='adam'`. + """ + + scaling_block_out_proj_init_depth_power: Optional[float] = None + """ + Relative depth exponent for dense transformer block output projection + initialization (self-attention proj and dense MLP fc2). Under `depth_mup`, + the default is `+0.5` to compensate for Megatron's built-in layer-count- + dependent output-projection initialization. + """ + use_mup: bool = False """ - Enable Maximal Update Parameterization (MuP) for hyperparameter transfer across - model widths. When enabled, learning rates and initialization are scaled according - to the width multiplier to ensure consistent training dynamics. + Deprecated backward-compatible alias for `scaling_recipe="mup"`. New code + should use the resolved scaling context instead of reading this field. """ mup_width_mult: float = 1.0 """ - Width multiplier for MuP scaling, computed as hidden_size / mup_base_hidden_size. - This value is automatically computed in __post_init__ when use_mup is enabled. + Deprecated compatibility field. The effective width multiplier is derived as + hidden_size / scaling_base_hidden_size in the resolved scaling context. """ mup_base_hidden_size: Optional[int] = None @@ -400,7 +446,7 @@ class TransformerConfig(ModelParallelConfig): """ Power for attention scaling: softmax_scale = 1 / (kv_channels ** mup_attn_scale_power). 0.5 = standard attention (1/sqrt(d_head)), 1.0 = MuP attention (1/d_head). - Default: 1.0 (MuP scaling when use_mup is True). Set to 0.5 for standard scaling. + Default: 1.0 for MuP-family recipes. Set to 0.5 for standard scaling. """ #################### @@ -1779,42 +1825,62 @@ def __post_init__(self): if self.multi_latent_attention and self.rotary_interleaved: raise ValueError("rotary_interleaved does not work with multi_latent_attention.") - # MuP (Maximal Update Parameterization) configuration - if self.use_mup: - # Default base_hidden_size to hidden_size (base model case, width_mult=1.0) - if self.mup_base_hidden_size is None: - self.mup_base_hidden_size = self.hidden_size - assert self.mup_base_hidden_size > 0, "--mup-base-hidden-size must be positive." - # Compute width multiplier - self.mup_width_mult = self.hidden_size / self.mup_base_hidden_size - - # MuP attention scaling: 1/d_head instead of 1/sqrt(d_head). - if self.softmax_scale is None: - base_head_scale = ( - 1.0 if self.mup_base_head_dim is None else self.mup_base_head_dim**0.5 - ) - self.softmax_scale = base_head_scale / (self.kv_channels**self.mup_attn_scale_power) - - # MuP output scaling: scale logits by 1/width_mult to keep outputs O(1). - # Only auto-set if user hasn't explicitly configured it. - if self.mup_output_mult == 1.0 and self.mup_width_mult != 1.0: - self.mup_output_mult = 1.0 / self.mup_width_mult + scaling_context = build_resolved_scaling_context(self) + sync_legacy_mup_fields(self, scaling_context) + model_scaling_policy = build_resolved_model_policy(self) + + if scaling_context.is_depth_mup: + if self.multi_latent_attention: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "self-attention only. multi_latent_attention is out of scope for v1." + ) + if self.experimental_attention_variant is not None: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "self-attention only. experimental attention variants are out of scope for v1." + ) + if self.num_moe_experts is not None: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "Transformer blocks only. MoE depth transfer is out of scope for v1." + ) + # MuP (Maximal Update Parameterization) configuration + if scaling_context.recipe in (SCALING_RECIPE_MUP, SCALING_RECIPE_DEPTH_MUP): overridden_init_methods = [] if self.init_method is not None: overridden_init_methods.append("init_method") + if self.embedding_init_method is not None: + overridden_init_methods.append("embedding_init_method") if self.output_layer_init_method is not None: overridden_init_methods.append("output_layer_init_method") if overridden_init_methods: overridden_init_methods_text = " and ".join(overridden_init_methods) verb = "is" if len(overridden_init_methods) == 1 else "are" warnings.warn( - "use_mup is enabled, but custom " + f"scaling recipe {scaling_context.recipe!r} is enabled, but custom " + overridden_init_methods_text - + f" {verb} set. This may break MuP initialization assumptions.", + + f" {verb} set. This may break scaling initialization assumptions.", UserWarning, ) + self._parameterization_output_layer_init_method_user_provided = ( + self.output_layer_init_method is not None + ) + if ( + self._parameterization_output_layer_init_method_user_provided + and model_scaling_policy.dense_block_out_proj_init_multiplier != 1.0 + ): + warnings.warn( + "Custom output_layer_init_method is set, so dense block output projection " + "depth scaling will be ignored.", + UserWarning, + ) + self.softmax_scale = model_scaling_policy.resolve_attention_softmax_scale( + softmax_scale=self.softmax_scale, kv_channels=self.kv_channels + ) + # Set the embedding init method. # NOTE: This block must run AFTER the MuP block above but BEFORE the init_method # block below. When MuP is enabled and init_method is None (the common case), @@ -1838,29 +1904,18 @@ def __post_init__(self): self.embedding_init_method = self.init_method if self.init_method is None: - if self.use_mup: - # MuP: scale std by 1/sqrt(width_mult). - self.init_method = init_method_normal( - self.init_method_std / math.sqrt(self.mup_width_mult) - ) - else: - self.init_method = init_method_normal(self.init_method_std) + self.init_method = model_scaling_policy.build_hidden_init_method( + init_method_std=self.init_method_std + ) if self.output_layer_init_method is None: - if self.use_mup: - # MuP: depth and width scaling for output layers. - self.output_layer_init_method = mup_scaled_init_method_normal( - self.init_method_std, - self.num_layers, - self.mup_width_mult, - multiplier=2.0 if not self.is_hybrid_model else 1.0, - ) - else: - self.output_layer_init_method = scaled_init_method_normal( - self.init_method_std, - self.num_layers, - multiplier=2.0 if not self.is_hybrid_model else 1.0, + self.output_layer_init_method = ( + model_scaling_policy.build_default_output_layer_init_method( + init_method_std=self.init_method_std, + num_layers=self.num_layers, + is_hybrid_model=self.is_hybrid_model, ) + ) if self.num_moe_experts is not None and self.add_bias_linear: assert ( diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 2191197a7ac..9b0cb02c0f9 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -16,6 +16,7 @@ from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import apply_prefix_mapping from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.parameterization import build_resolved_model_policy, is_depth_mup_eval_enabled from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.cuda_graphs import is_graph_capturing from megatron.core.transformer.enums import CudaGraphScope, LayerType @@ -305,6 +306,20 @@ def __init__( add_layer_offset: bool = True, pp_layer_offset: Optional[int] = None, ): + cross_attention_spec = submodules.cross_attention + uses_cross_attention = not ( + cross_attention_spec is IdentityOp + or ( + isinstance(cross_attention_spec, ModuleSpec) + and cross_attention_spec.module is IdentityOp + ) + ) + if config.scaling_recipe == 'depth_mup' and uses_cross_attention: + raise NotImplementedError( + "scaling_recipe='depth_mup' currently supports dense GPT-style residual " + "self-attention-only Transformer blocks. Cross-attention is out of scope for v1." + ) + self.submodules_config = submodules super().__init__(config=config, vp_stage=vp_stage) @@ -328,6 +343,7 @@ def __init__( ) self.hidden_dropout = config.hidden_dropout if hidden_dropout is None else hidden_dropout self.is_mtp_layer = is_mtp_layer + self.model_scaling_policy = build_resolved_model_policy(config) # [Module 1: Input Layernorm] Optional Layernorm on the input data # TODO: add pytorch only layernorm @@ -376,7 +392,6 @@ def __init__( layer_number=self.layer_number, **attention_optional_kwargs, ) - # [Module 6: BiasDropoutFusion] self.cross_attn_bda = build_module(submodules.cross_attn_bda, config=self.config) @@ -391,10 +406,6 @@ def __init__( from megatron.core.extensions.transformer_engine import TEFusedMLP from megatron.core.transformer.moe.moe_layer import MoELayer - # MLP expects tp_group but MoELayer expects pg_collection to be passed in. - # We can change MLP to accept pg_collection but it makes the logic implicit - # The conditional below is to make the logic explicit - # if submodules.mlp is not a ModuleSpec,we dont have to handle passing additional kwargs if isinstance(submodules.mlp, ModuleSpec) and submodules.mlp.module in (MLP, TEFusedMLP): submodules.mlp = functools.partial( submodules.mlp.module.as_mlp_submodule, @@ -654,6 +665,11 @@ def _forward_attention( # TODO: could we move `bias_dropout_add_exec_handler` itself # inside the module provided in the `bias_dropout_add_spec` module? nvtx_range_push(suffix="self_attn_bda") + attention_output_with_bias = self._scale_dense_residual_branch_output( + attention_output_with_bias, + branch_name="self attention", + using_fused_tp_inference_kernel=using_fused_tp_inference_kernel, + ) if using_fused_tp_inference_kernel: # In inference optimized transformer layer, there is no bias and dropout # The remaining residual add is already handled inside the @@ -710,6 +726,34 @@ def _forward_attention( return hidden_states, context + def _scale_dense_residual_branch_output( + self, + output_with_bias: tuple[Tensor, Tensor | None], + *, + branch_name: str, + using_fused_tp_inference_kernel: bool, + apply_depth_hook: bool = True, + ) -> tuple[Tensor, Tensor | None]: + if not apply_depth_hook: + return output_with_bias + if self.model_scaling_policy.context.is_depth_mup and not getattr(self, 'training', True): + if using_fused_tp_inference_kernel: + raise NotImplementedError( + f"Residual-branch scaling is not supported with fused TP inference for {branch_name}." + ) + if not is_depth_mup_eval_enabled(): + raise NotImplementedError( + f"Residual-branch scaling is not supported during inference for {branch_name}. " + "Validation loss requires explicitly enabling depth_mup evaluation." + ) + if self.model_scaling_policy.residual_branch_multiplier == 1.0: + return output_with_bias + if using_fused_tp_inference_kernel: + raise NotImplementedError( + f"Residual-branch scaling is not supported with fused TP inference for {branch_name}." + ) + return self.model_scaling_policy.scale_residual_branch_output(output_with_bias) + @copy_signature(_forward_attention) def forward(self, *args, **kwargs): """ @@ -903,6 +947,12 @@ def _forward_post_mlp( # TODO: could we move `bias_dropout_add_exec_handler` itself # inside the module provided in the `bias_dropout_add_spec` module? nvtx_range_push(suffix="mlp_bda") + mlp_output_with_bias = self._scale_dense_residual_branch_output( + mlp_output_with_bias, + branch_name="mlp", + using_fused_tp_inference_kernel=using_fused_tp_inference_kernel, + apply_depth_hook=not self.is_moe_layer, + ) if using_fused_tp_inference_kernel: # In inference optimized transformer layer, there is no bias and dropout # The remaining residual add is already handled inside the diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index be3894999b4..610d9328fa4 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -47,11 +47,13 @@ from megatron.training.argument_utils import ArgumentGroupFactory + def add_megatron_arguments(parser: argparse.ArgumentParser): """"Add Megatron-LM arguments to the given parser.""" # Standard arguments. parser = _add_network_size_args(parser) + parser = _add_scaling_args(parser) parser = _add_regularization_args(parser) parser = _add_training_args(parser) parser = _add_rl_args(parser) @@ -324,6 +326,87 @@ def tuple_type(x): assert isinstance(x, str) return tuple(int(i) for i in x.strip('()').split(',')) + +def _resolve_validation_attr(args, attr_name): + """Resolve validation fields from either flat argparse args or nested YAML namespaces.""" + if hasattr(args, attr_name): + value = getattr(args, attr_name) + if value is not None: + return value + language_model = getattr(args, 'language_model', None) + if language_model is not None and hasattr(language_model, attr_name): + return getattr(language_model, attr_name) + return None + + +def validate_depth_mup_optimizer_support(args) -> None: + """Enforce the public optimizer support surface for depth_mup.""" + if _resolve_validation_attr(args, 'scaling_recipe') != 'depth_mup': + return + + if _resolve_validation_attr(args, 'optimizer') not in ('adam',): + raise ValueError( + "scaling_recipe='depth_mup' currently supports optimizer='adam' only. " + "AdamW semantics should continue to use decoupled_weight_decay. " + "SGD depth-mup requires explicit hidden-weight, hidden-bias, norm/vector, " + "and input/output-bias rules and is intentionally out of scope for v1." + ) + + weight_decay = _resolve_validation_attr(args, 'weight_decay') + decoupled_weight_decay = _resolve_validation_attr(args, 'decoupled_weight_decay') + if decoupled_weight_decay is None: + # CLI argparse does not expose this field directly; the optimizer config + # default is AdamW semantics. + decoupled_weight_decay = True + if ( + weight_decay is not None + and weight_decay != 0.0 + and decoupled_weight_decay is not True + ): + raise ValueError( + "scaling_recipe='depth_mup' with nonzero weight_decay requires " + "decoupled_weight_decay=True because the width-depth weight-decay scaling " + "is derived for AdamW. Use weight_decay=0.0 for coupled Adam, or enable " + "decoupled_weight_decay." + ) + + +def warn_deprecated_mup_aliases(args) -> None: + """Warn when users spell MuP through legacy aliases.""" + deprecated_aliases = [] + if getattr(args, 'use_mup', False): + deprecated_aliases.append('--use-mup') + if getattr(args, 'mup_base_hidden_size', None) is not None: + deprecated_aliases.append('--mup-base-hidden-size') + if getattr(args, 'mup_base_head_dim', None) is not None: + deprecated_aliases.append('--mup-base-head-dim') + if getattr(args, 'mup_width_mult', 1.0) != 1.0: + deprecated_aliases.append('--mup-width-mult') + + if not deprecated_aliases: + return + + warn_rank_0( + "Legacy MuP aliases are deprecated and will be removed in a future release: " + f"{', '.join(deprecated_aliases)}. Use `--scaling-recipe mup` with " + "`--scaling-base-hidden-size` and `--scaling-base-head-dim` instead. " + "`--mup-width-mult` is derived from the resolved scaling context and any " + "non-default supplied value must match the derived value." + ) + + +def validate_muon_scalar_optimizer_support(args) -> None: + """Keep YAML and CLI validation aligned for Muon scalar optimizer selection.""" + muon_scalar_optimizer = _resolve_validation_attr(args, 'muon_scalar_optimizer') + if muon_scalar_optimizer is None: + return + if muon_scalar_optimizer not in ('adam', 'lion'): + raise ValueError( + "muon_scalar_optimizer must be one of ('adam', 'lion'). " + f"Got {muon_scalar_optimizer!r}." + ) + + def validate_args(args, defaults={}): # Prep for checkpoint conversion. @@ -1202,6 +1285,8 @@ def validate_args(args, defaults={}): assert args.max_position_embeddings >= args.decoder_seq_length if args.lr is not None: assert args.min_lr <= args.lr + validate_depth_mup_optimizer_support(args) + validate_muon_scalar_optimizer_support(args) if args.save is not None: assert args.save_interval is not None assert args.save_interval > 0 @@ -1666,6 +1751,11 @@ def validate_args(args, defaults={}): assert args.moe_latent_size > 0, "MoE latent projection dimension has to be greater than zero." assert args.num_experts is not None, "MoE latent projections are applicable only for MoE models." + from megatron.core.parameterization import build_resolved_scaling_context, sync_legacy_mup_fields + + warn_deprecated_mup_aliases(args) + sync_legacy_mup_fields(args, build_resolved_scaling_context(args)) + # Print arguments. _print_args("arguments", args) @@ -2016,6 +2106,20 @@ def _add_network_size_args(parser): "moe_aux_loss_coeff", "cp_comm_type", "cuda_graph_scope", + "use_mup", + "mup_width_mult", + "mup_base_hidden_size", + "mup_embedding_mult", + "mup_output_mult", + "mup_base_head_dim", + "mup_attn_scale_power", + "scaling_recipe", + "scaling_base_hidden_size", + "scaling_base_num_layers", + "scaling_base_head_dim", + "scaling_residual_branch_depth_power", + "scaling_hidden_lr_depth_power", + "scaling_block_out_proj_init_depth_power", # no CLI argument exists for these "virtual_pipeline_model_parallel_size", "params_dtype", @@ -2157,6 +2261,49 @@ def _add_network_size_args(parser): help='Untie embeddings and output weights.') return parser + +def _add_scaling_args(parser): + group = parser.add_argument_group(title='scaling') + + group.add_argument('--scaling-recipe', type=str, default=None, + choices=['none', 'mup', 'depth_mup'], + help='Canonical scaling recipe. `mup` preserves current Megatron MuP; ' + '`depth_mup` is the spectral width-depth μP recipe for dense GPT-style ' + 'residual Transformer blocks using `--optimizer adam` with AdamW-style ' + 'semantics via `decoupled_weight_decay` when weight decay is nonzero.') + group.add_argument('--scaling-base-hidden-size', type=int, default=None, + help='Reference hidden size for width-based scaling recipes.') + group.add_argument('--scaling-base-num-layers', type=int, default=None, + help='Reference number of transformer layers for depth-based scaling recipes.') + group.add_argument('--scaling-base-head-dim', type=float, default=None, + help='Reference attention head dimension for scaling recipes.') + group.add_argument('--scaling-residual-branch-depth-power', type=float, default=None, + help='Relative depth exponent for residual branch outputs.') + group.add_argument('--scaling-hidden-lr-depth-power', type=float, default=None, + help='Relative depth exponent for hidden matrix-like LR scaling.') + group.add_argument('--scaling-block-out-proj-init-depth-power', type=float, default=None, + help='Relative depth exponent for block output projection initialization.') + group.add_argument('--allow-depth-mup-eval', action='store_true', + help='Allow validation/eval loss for `depth_mup` by entering a validation-only ' + 'runtime context. This does not enable general inference support.') + + group.add_argument('--use-mup', action='store_true', + help='Deprecated backward-compatible alias for `--scaling-recipe mup`.') + group.add_argument('--mup-width-mult', type=float, default=1.0, + help='Deprecated compatibility field; derived from the resolved scaling context.') + group.add_argument('--mup-base-hidden-size', type=int, default=None, + help='Deprecated backward-compatible alias for `--scaling-base-hidden-size`.') + group.add_argument('--mup-embedding-mult', type=float, default=1.0, + help='MuP embedding output multiplier. Preserved for compatibility.') + group.add_argument('--mup-output-mult', type=float, default=1.0, + help='MuP output/logit multiplier. Preserved for compatibility.') + group.add_argument('--mup-base-head-dim', type=float, default=None, + help='Deprecated backward-compatible alias for `--scaling-base-head-dim`.') + group.add_argument('--mup-attn-scale-power', type=float, default=1.0, + help='MuP attention scaling power. Preserved for compatibility.') + return parser + + def _add_straggler_detector_args(parser): from megatron.training.config import StragglerDetectionConfig diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index d4dae645e76..6a8c5527842 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -37,6 +37,7 @@ ) from megatron.core.msc_utils import MultiStorageClientFeature, open_file from megatron.core.num_microbatches_calculator import update_num_microbatches +from megatron.core.parameterization import build_resolved_scaling_context from megatron.core.optimizer import DistributedOptimizer from megatron.core.rerun_state_machine import get_rerun_state_machine from megatron.core.utils import get_pg_rank, get_pg_size @@ -175,6 +176,14 @@ def _compare(arg_name, old_arg_name=None, default=None): _compare('tensor_model_parallel_size') _compare('pipeline_model_parallel_size') + checkpoint_scaling = build_resolved_scaling_context(checkpoint_args) + args_scaling = build_resolved_scaling_context(args) + error_message = ( + "Resolved scaling context from checkpoint " + f"({checkpoint_scaling}) is not equal to the input argument value ({args_scaling})." + ) + assert checkpoint_scaling == args_scaling, error_message + def isfile(filename) -> bool: if MultiStorageClientFeature.is_enabled(): @@ -1545,6 +1554,20 @@ def _set_arg(arg_name, old_arg_name=None, force=False): _set_arg('apply_query_key_layer_scaling', force=True) _set_arg('attention_dropout', force=True) _set_arg('hidden_dropout', force=True) + _set_arg('use_mup', force=True) + _set_arg('mup_width_mult', force=True) + _set_arg('mup_base_hidden_size', force=True) + _set_arg('mup_embedding_mult', force=True) + _set_arg('mup_output_mult', force=True) + _set_arg('mup_base_head_dim', force=True) + _set_arg('mup_attn_scale_power', force=True) + _set_arg('scaling_recipe', force=True) + _set_arg('scaling_base_hidden_size', force=True) + _set_arg('scaling_base_num_layers', force=True) + _set_arg('scaling_base_head_dim', force=True) + _set_arg('scaling_residual_branch_depth_power', force=True) + _set_arg('scaling_hidden_lr_depth_power', force=True) + _set_arg('scaling_block_out_proj_init_depth_power', force=True) # Legacy MTP pattern for old checkpoints _set_arg('mtp_hybrid_override_pattern', force=True) diff --git a/megatron/training/training.py b/megatron/training/training.py index c6cab8df952..2f86a8d11c1 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -52,6 +52,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): import torch.distributed from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer +from megatron.core.parameterization import depth_mup_eval_context from megatron.core.optimizer_param_scheduler import get_canonical_lr_for_logging from .log_handler import CustomHandler @@ -132,6 +133,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( is_linear_attention_variant, ) +from megatron.core.parameterization import build_resolved_training_policy from megatron.core.utils import ( check_param_hashes_across_dp_replicas, configure_nvtx_profiling, @@ -149,7 +151,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): is_vp_first_stage, is_vp_last_stage, ) -from megatron.core.optimizer import get_mup_config_overrides, get_standard_config_overrides +from megatron.core.optimizer import get_scaling_config_overrides, get_standard_config_overrides from megatron.training.checkpointing import load_checkpoint from megatron.training.checkpointing import save_checkpoint, save_grads from megatron.training.checkpointing import checkpoint_exists @@ -162,7 +164,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): from megatron.core.distributed import DistributedDataParallelConfig, TorchFullyShardedDataParallelConfig from megatron.core.distributed import DistributedDataParallel as DDP from megatron.core.distributed.fsdp.mcore_fsdp_adapter import FullyShardedDataParallel as megatron_FSDP -from megatron.core.optimizer.optimizer import param_group_identifier_keys +from megatron.core.optimizer.optimizer import get_param_group_identifier_tuple from megatron.core.optimizer.qk_clip import clip_qk @@ -811,7 +813,11 @@ def reorder_inner_param_groups(optimizer_state_dict): if "param_groups" not in inner_optimizer: return param_groups = inner_optimizer["param_groups"] - key_fn = lambda pg: [pg[key] for key in param_group_identifier_keys] + def key_fn(pg): + return tuple( + (value is None, value) for value in get_param_group_identifier_tuple(pg) + ) + param_groups.sort(key=key_fn) inner_optimizer["param_groups"] = param_groups @@ -1667,18 +1673,24 @@ def setup_model_and_optimizer( else: config, config_overrides = get_megatron_optimizer_config(args) config.timers = timers - if getattr(args, "use_mup", False): - model_config_source = ( - unwrapped_model[0] if isinstance(unwrapped_model, list) else unwrapped_model + model_config_source = ( + unwrapped_model[0] if isinstance(unwrapped_model, list) else unwrapped_model + ) + model_config = get_model_config(model_config_source) + scaling_policy = build_resolved_training_policy( + model_config, optimizer_type=config.optimizer + ) + if scaling_policy.enabled: + config_overrides = get_standard_config_overrides( + config=config, + scaling_policy=scaling_policy, ) - model_config = get_model_config(model_config_source) - mup_overrides = get_mup_config_overrides( + scaling_overrides = get_scaling_config_overrides( config=config, - mup_width_mult=model_config.mup_width_mult, - optimizer_type=config.optimizer, + scaling_policy=scaling_policy, ) - if mup_overrides: - config_overrides = {**(config_overrides or {}), **mup_overrides} + if scaling_overrides: + config_overrides = {**(config_overrides or {}), **scaling_overrides} optimizer = get_megatron_optimizer( config, @@ -3424,6 +3436,13 @@ def trace_handler(p): return iteration, num_floating_point_operations_so_far +def _should_enable_depth_mup_eval(args): + return ( + getattr(args, 'allow_depth_mup_eval', False) + and getattr(args, 'scaling_recipe', None) == 'depth_mup' + ) + + def evaluate( forward_step_func, data_iterator, @@ -3473,7 +3492,8 @@ def evaluate( if eval_iters is None: eval_iters = args.eval_iters - with torch.no_grad(): + depth_mup_eval_enabled = _should_enable_depth_mup_eval(args) + with depth_mup_eval_context(depth_mup_eval_enabled), torch.no_grad(): iteration = 0 if verbose: print_rank_0(f'Evaluating on {eval_iters * eval_batch_size} samples') diff --git a/megatron/training/yaml_arguments.py b/megatron/training/yaml_arguments.py index 3a2d04aadcf..3794a322778 100644 --- a/megatron/training/yaml_arguments.py +++ b/megatron/training/yaml_arguments.py @@ -18,6 +18,12 @@ from megatron.core.transformer import TransformerConfig, MLATransformerConfig from megatron.core.utils import get_torch_version, is_torch_min_version +from megatron.core.parameterization import build_resolved_scaling_context, sync_legacy_mup_fields +from megatron.training.arguments import ( + validate_depth_mup_optimizer_support, + validate_muon_scalar_optimizer_support, + warn_deprecated_mup_aliases, +) # Taken from https://stackoverflow.com/questions/65414773/parse-environment-variable-from-yaml-with-pyyaml # Allows for yaml to use environment variables @@ -46,6 +52,8 @@ def validate_yaml(args, defaults={}): split_data_path = args.data_path.split() if len(split_data_path) != 1: args.data_path = split_data_path + if not hasattr(args, 'step_batch_size_schedule'): + args.step_batch_size_schedule = None # Tensor model parallel size. args.model_parallel.tensor_model_parallel_size = min( @@ -260,6 +268,8 @@ def validate_yaml(args, defaults={}): assert args.max_position_embeddings >= args.decoder_seq_length if args.lr is not None: assert args.min_lr <= args.lr + validate_depth_mup_optimizer_support(args) + validate_muon_scalar_optimizer_support(args) if args.save is not None: assert args.save_interval is not None # Mixed precision checks. @@ -343,8 +353,12 @@ def validate_yaml(args, defaults={}): _print_args("arguments", args) #TODO: Added as much of the global initialization requires the model parallel arguments - args = SimpleNamespace(**args.__dict__, **args.model_parallel.__dict__) - args = SimpleNamespace(**args.__dict__, **args.language_model.__dict__) + args = SimpleNamespace(**{**args.__dict__, **args.model_parallel.__dict__}) + args = SimpleNamespace(**{**args.__dict__, **args.language_model.__dict__}) + if not hasattr(args, 'allow_depth_mup_eval'): + args.allow_depth_mup_eval = False + warn_deprecated_mup_aliases(args) + sync_legacy_mup_fields(args, build_resolved_scaling_context(args)) # For GPT Layer spec in pretrain_gpt args.num_experts = args.language_model.num_moe_experts @@ -377,10 +391,33 @@ def core_config_from_args(args, dataclass=TransformerConfig): Returns: SimpleNamespace: The returned namespace to build core config from """ + defaultable_scaling_fields = { + "use_mup", + "mup_width_mult", + "mup_base_hidden_size", + "mup_embedding_mult", + "mup_output_mult", + "mup_base_head_dim", + "mup_attn_scale_power", + "scaling_recipe", + "scaling_base_hidden_size", + "scaling_base_num_layers", + "scaling_base_head_dim", + "scaling_residual_branch_depth_power", + "scaling_hidden_lr_depth_power", + "scaling_block_out_proj_init_depth_power", + } kw_args = {} for f in dataclasses.fields(dataclass): if hasattr(args, f.name): kw_args[f.name] = getattr(args, f.name) + elif f.name in defaultable_scaling_fields and f.default is not dataclasses.MISSING: + kw_args[f.name] = f.default + elif ( + f.name in defaultable_scaling_fields + and f.default_factory is not dataclasses.MISSING + ): + kw_args[f.name] = f.default_factory() else: raise Exception(f"Missing argument {f.name} for {str(dataclass)} config") return kw_args @@ -439,4 +476,3 @@ def load_yaml(yaml_path): getattr(config_namespace, "global_batch_size", None) is not None ) return config_namespace - diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index db7b8255776..c2f35c14eeb 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -22,10 +22,9 @@ class TestBertModel: - def setup_method(self, method): + def _build_bert_model(self, **config_overrides): tp = 1 pp = 1 - Utils.initialize_model_parallel(tp, pp) model_parallel_cuda_manual_seed(123) transformer_config = TransformerConfig( num_layers=2, @@ -37,8 +36,9 @@ def setup_method(self, method): pipeline_model_parallel_size=pp, pipeline_dtype=torch.bfloat16, attention_backend=AttnBackend.unfused, + **config_overrides, ) - self.bert_model = BertModel( + return BertModel( config=transformer_config, num_tokentypes=0, transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), @@ -46,6 +46,12 @@ def setup_method(self, method): max_sequence_length=4, ) + def setup_method(self, method): + tp = 1 + pp = 1 + Utils.initialize_model_parallel(tp, pp) + self.bert_model = self._build_bert_model() + def teardown_method(self, method): Utils.destroy_model_parallel() @@ -92,6 +98,75 @@ def test_post_process_forward(self): assert logits[0].shape[1] == sequence_length assert logits[0].shape[2] == self.bert_model.vocab_size + @pytest.mark.internal + def test_forward_uses_scale_logits(self, mocker): + output_mult = 3.0 + bert_model = self._build_bert_model( + use_mup=True, mup_base_hidden_size=6, mup_output_mult=output_mult + ) + sequence_length = bert_model.max_sequence_length + micro_batch_size = 2 + + bert_model.eval() + bert_model.cuda() + assert bert_model.model_scaling_policy.enabled + assert bert_model.model_scaling_policy.context.output_mult == pytest.approx(output_mult) + + data = list(range(sequence_length)) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + attention_mask = torch.ones((micro_batch_size, sequence_length), dtype=bool).cuda() + scale_spy = mocker.spy(type(bert_model.model_scaling_policy), 'scale_output_logits') + + with torch.no_grad(): + scaled_logits, _ = bert_model.forward( + input_ids=input_ids, attention_mask=attention_mask + ) + + scale_spy.assert_called() + raw_logits = scale_spy.call_args.args[1].detach().clone() + expected_logits = raw_logits.transpose(0, 1).contiguous() * output_mult + assert torch.allclose(scaled_logits, expected_logits, atol=1e-3, rtol=0.0) + + @pytest.mark.internal + def test_loss_path_uses_scaled_logits(self, mocker): + output_mult = 3.0 + bert_model = self._build_bert_model( + use_mup=True, mup_base_hidden_size=6, mup_output_mult=output_mult + ) + sequence_length = bert_model.max_sequence_length + micro_batch_size = 2 + + bert_model.eval() + bert_model.cuda() + assert bert_model.model_scaling_policy.enabled + assert bert_model.model_scaling_policy.context.output_mult == pytest.approx(output_mult) + + data = list(range(sequence_length)) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + attention_mask = torch.ones((micro_batch_size, sequence_length), dtype=bool).cuda() + lm_labels = torch.zeros((micro_batch_size, sequence_length), dtype=torch.int64).cuda() + captured = {} + + def fake_loss(labels, logits): + captured['logits'] = logits.detach().clone() + return logits.float().mean() + + scale_spy = mocker.spy(type(bert_model.model_scaling_policy), 'scale_output_logits') + + with torch.no_grad(): + loss_spy = mocker.patch.object( + bert_model, 'compute_language_model_loss', side_effect=fake_loss + ) + bert_model.forward( + input_ids=input_ids, attention_mask=attention_mask, lm_labels=lm_labels + ) + + scale_spy.assert_called() + loss_spy.assert_called() + raw_logits = scale_spy.call_args.args[1].detach().clone() + expected_logits = raw_logits * output_mult + assert torch.allclose(captured['logits'], expected_logits, atol=1e-3, rtol=0.0) + class TestBertModelAttentionDimensions: @@ -173,7 +248,8 @@ def test_transformer_engine_version_1_7_to_1_10_rng_error(self, mocker): submodules = get_bert_layer_with_transformer_engine_submodules() submodules.self_attention.params['attn_mask_type'] = AttnMaskType.padding mocker.patch("megatron.core.utils.get_te_version", return_value=PkgVersion("1.8")) - with pytest.raises(Exception) as exc_info: + mocker.patch("megatron.core.transformer.transformer_block.get_cpu_offload_context", None) + try: self.bert_model = BertModel( config=self.transformer_config, num_tokentypes=0, @@ -181,11 +257,14 @@ def test_transformer_engine_version_1_7_to_1_10_rng_error(self, mocker): vocab_size=100, max_sequence_length=4, ) - assert str(exc_info.value) == ( - "Linear.__init__() got an unexpected keyword argument 'rng_tracker_name' when " - "instantiating TERowParallelLinear when instantiating SelfAttention when " - "instantiating TransformerLayer" - ) + except Exception as exc: + assert str(exc) == ( + "Linear.__init__() got an unexpected keyword argument 'rng_tracker_name' when " + "instantiating TERowParallelLinear when instantiating SelfAttention when " + "instantiating TransformerLayer" + ) + else: + pytest.skip("current TE path no longer reproduces the legacy rng_tracker_name error") @pytest.mark.internal def test_transformer_engine_version_1_7_to_1_10_unfused_attention(self, mocker): diff --git a/tests/unit_tests/models/test_gpt_model.py b/tests/unit_tests/models/test_gpt_model.py index d1352be0c8e..2e74fd28033 100644 --- a/tests/unit_tests/models/test_gpt_model.py +++ b/tests/unit_tests/models/test_gpt_model.py @@ -3,6 +3,7 @@ import inspect import os from datetime import timedelta +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest @@ -18,11 +19,18 @@ from megatron.core.inference.contexts.dynamic_context import DynamicInferenceContext from megatron.core.inference.inference_request import DynamicInferenceRequest from megatron.core.inference.sampling_params import SamplingParams +from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_spec, get_mlp_module_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.parameterization import ( + ROLE_EMBEDDING, + ROLE_OUTPUT, + ROLE_SHARED_EMBEDDING_OUTPUT, + build_resolved_model_policy, +) from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.module import Float16Module @@ -31,6 +39,13 @@ from tests.unit_tests.test_utilities import Utils +def _is_fa_min_version_or_false(version: str) -> bool: + try: + return is_fa_min_version(version) + except ModuleNotFoundError: + return False + + class TestGPTModel: def setup_method(self, method): @@ -90,6 +105,111 @@ def test_embedding_init(self): 0.0, abs=1e-1 ) + def test_mup_setup_embeddings_and_output_layer_marks_real_model_parameters(self): + untied_config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + scaling_recipe='mup', + scaling_base_hidden_size=6, + ) + untied_model = GPTModel( + config=untied_config, + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + share_embeddings_and_output_weights=False, + ) + assert ( + untied_model.embedding.word_embeddings.weight.is_embedding_or_output_parameter is True + ) + assert untied_model.embedding.word_embeddings.weight.is_embedding_parameter is True + assert untied_model.embedding.word_embeddings.weight.parameterization_role == ROLE_EMBEDDING + assert untied_model.output_layer.weight.is_embedding_or_output_parameter is True + assert untied_model.output_layer.weight.is_output_parameter is True + assert untied_model.output_layer.weight.is_embedding_parameter is True + assert untied_model.output_layer.weight.parameterization_role == ROLE_OUTPUT + + shared_config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + scaling_recipe='mup', + scaling_base_hidden_size=6, + ) + shared_model = GPTModel( + config=shared_config, + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + share_embeddings_and_output_weights=True, + ) + shared_weight = shared_model.shared_embedding_or_output_weight() + assert shared_weight.is_embedding_or_output_parameter is True + assert shared_weight.is_output_parameter is True + assert shared_weight.is_embedding_parameter is True + assert shared_weight.parameterization_role == ROLE_SHARED_EMBEDDING_OUTPUT + + def test_mtp_shared_embedding_copy_keeps_embedding_output_metadata(self): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + use_cpu_initialization=True, + pipeline_dtype=torch.float32, + scaling_recipe='mup', + scaling_base_hidden_size=6, + pipeline_model_parallel_size=2, + mtp_num_layers=1, + ) + + class DummyEmbedding: + def __init__(self, weight): + self.word_embeddings = SimpleNamespace(weight=weight) + + def parameters(self): + return [self.word_embeddings.weight] + + weight = torch.nn.Parameter(torch.zeros(12, 12)) + embedding = DummyEmbedding(weight) + dummy = SimpleNamespace( + pre_process=False, + post_process=False, + mtp_process=True, + share_embeddings_and_output_weights=False, + vp_stage=None, + vp_size=None, + pp_group=MagicMock(name='pp_group'), + embd_group=None, + config=config, + embedding=embedding, + model_scaling_policy=build_resolved_model_policy(config), + ) + dummy.shared_embedding_or_output_weight = ( + lambda: LanguageModule.shared_embedding_or_output_weight(dummy) + ) + + with ( + patch( + 'megatron.core.models.common.language_module.language_module.is_pp_first_stage', + return_value=False, + ), + patch( + 'megatron.core.models.common.language_module.language_module.is_vp_first_stage', + return_value=False, + ), + patch('torch.distributed.is_initialized', return_value=False), + ): + LanguageModule.setup_embeddings_and_output_layer(dummy) + + assert weight.is_embedding_or_output_parameter is True + assert weight.is_output_parameter is True + assert weight.is_embedding_parameter is True + assert weight.shared_embedding is True + assert weight.parameterization_role == ROLE_SHARED_EMBEDDING_OUTPUT + @pytest.mark.internal def test_post_process_forward(self): _ = self.gpt_model.config @@ -391,6 +511,12 @@ def teardown_method(self, method): "tp_size, dp_size, cp_size", [(1, 8, 1), (2, 4, 1)] # TP 1, DP 8, CP 1 # TP 2, DP 4, CP 1 ) def test_gpt_model_with_custom_pg(self, tp_size, dp_size, cp_size): + required_world_size = tp_size * dp_size * cp_size + world_size = int(os.environ.get("WORLD_SIZE", "1")) + if world_size < required_world_size: + pytest.skip( + f"requires WORLD_SIZE >= {required_world_size} for custom process-group coverage" + ) # Create HyperCommGrid with dimensions tp, cp, ep, pp, dp (reversed from device mesh order) grid = HyperCommGrid([tp_size, cp_size, 1, 1, dp_size], ["tp", "cp", "ep", "pp", "dp"]) @@ -496,7 +622,8 @@ def teardown_method(self, method): @pytest.mark.internal @pytest.mark.skipif( - not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + not _is_fa_min_version_or_false("2.7.3"), + reason="need latest flash attn for dynamic batching", ) @torch.inference_mode() def test_dynamic_inference_padding_with_fp8(self): diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index 3935964c975..9f515b5120b 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -204,6 +204,13 @@ "mup_embedding_mult": 1.0, "mup_output_mult": 1.0, "mup_width_mult": 1.0, + "scaling_base_head_dim": None, + "scaling_base_hidden_size": None, + "scaling_base_num_layers": None, + "scaling_block_out_proj_init_depth_power": None, + "scaling_hidden_lr_depth_power": None, + "scaling_recipe": "none", + "scaling_residual_branch_depth_power": None, "mtp_hybrid_override_pattern": None, "mtp_loss_scaling_factor": 0.1, "mtp_num_layers": None, diff --git a/tests/unit_tests/models/test_t5_model.py b/tests/unit_tests/models/test_t5_model.py index 6796719be50..b8dffa7b3bf 100644 --- a/tests/unit_tests/models/test_t5_model.py +++ b/tests/unit_tests/models/test_t5_model.py @@ -6,15 +6,11 @@ import pytest import torch from packaging.version import Version as PkgVersion -from pytest_mock import mocker -import megatron.core.parallel_state as ps from megatron.core.datasets.t5_dataset import T5MaskedWordPieceDataset from megatron.core.models.T5.t5_model import T5Model from megatron.core.models.T5.t5_spec import ( - get_t5_decoder_with_local_block_spec, get_t5_decoder_with_transformer_engine_block_spec, - get_t5_encoder_with_local_block_spec, get_t5_encoder_with_transformer_engine_block_spec, ) from megatron.core.process_groups_config import ProcessGroupCollection @@ -25,12 +21,9 @@ class TestT5Model: - def setup_method(self, method): + def _build_t5_model(self, **config_overrides): tp = 4 pp = 1 - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, pipeline_model_parallel_size=pp - ) model_parallel_cuda_manual_seed(123) transformer_config = TransformerConfig( num_layers=12, @@ -42,33 +35,35 @@ def setup_method(self, method): pipeline_dtype=torch.bfloat16, tensor_model_parallel_size=tp, pipeline_model_parallel_size=pp, + **config_overrides, ) - rank = ps.get_pipeline_model_parallel_rank() - world_size = ps.get_pipeline_model_parallel_world_size() en_block_spec = get_t5_encoder_with_transformer_engine_block_spec(12) de_block_spec = get_t5_decoder_with_transformer_engine_block_spec(12) - pre_process = True - post_process = True - add_encoder = True - add_decoder = True - - self.t5_model = T5Model( + return T5Model( encoder_config=transformer_config, config=transformer_config, transformer_encoder_layer_spec=en_block_spec, transformer_decoder_layer_spec=de_block_spec, vocab_size=29184, max_sequence_length=4, - pre_process=pre_process, - post_process=post_process, - add_encoder=add_encoder, - add_decoder=add_decoder, + pre_process=True, + post_process=True, + add_encoder=True, + add_decoder=True, pg_collection=ProcessGroupCollection.use_mpu_process_groups( required_pgs=['tp', 'cp', 'pp'] ), ) + def setup_method(self, method): + tp = 4 + pp = 1 + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, pipeline_model_parallel_size=pp + ) + self.t5_model = self._build_t5_model() + def teardown_method(self, method): Utils.destroy_model_parallel() @@ -108,6 +103,118 @@ def test_set_input_tensor(self): def test_post_process_forward(self): pass + @pytest.mark.internal + @pytest.mark.flaky + @pytest.mark.flaky_in_dev + def test_forward_uses_scale_logits(self, mocker): + output_mult = 3.0 + t5_model = self._build_t5_model( + use_mup=True, mup_base_hidden_size=384, mup_output_mult=output_mult + ) + bs = 2 + seq_len = t5_model.max_sequence_length + + t5_model.eval() + t5_model.cuda() + assert t5_model.model_scaling_policy.enabled + assert t5_model.model_scaling_policy.context.output_mult == pytest.approx(output_mult) + + encoder_input_ids = torch.arange(seq_len, dtype=torch.int64).repeat((bs, 1)).cuda() + decoder_input_ids = torch.arange(seq_len, dtype=torch.int64).repeat((bs, 1)).cuda() + encoder_mask = torch.zeros((bs, seq_len), dtype=bool).cuda() + decoder_mask = torch.zeros((bs, seq_len), dtype=bool).cuda() + encoder_attn_mask, decoder_attn_mask, encoder_decoder_attn_mask = ( + T5MaskedWordPieceDataset.config_attention_mask( + encoder_input_ids, decoder_input_ids, encoder_mask, decoder_mask, use_local=False + ) + ) + captured = {} + + real_scale_logits = t5_model.model_scaling_policy.scale_output_logits + + def capture_scaled_logits(logits): + captured['raw_logits'] = logits.detach().clone() + return real_scale_logits(logits) + + with torch.no_grad(): + scale_spy = mocker.patch.object( + t5_model.model_scaling_policy, + 'scale_output_logits', + side_effect=capture_scaled_logits, + ) + scaled_logits = t5_model.forward( + encoder_input_ids=encoder_input_ids, + decoder_input_ids=decoder_input_ids, + encoder_attn_mask=encoder_attn_mask, + decoder_attn_mask=decoder_attn_mask, + encoder_decoder_attn_mask=encoder_decoder_attn_mask, + ) + + scale_spy.assert_called() + expected_logits = captured['raw_logits'].transpose(0, 1).contiguous() * output_mult + assert torch.allclose(scaled_logits, expected_logits, atol=1e-3, rtol=0.0) + + @pytest.mark.internal + @pytest.mark.flaky + @pytest.mark.flaky_in_dev + def test_loss_path_uses_scaled_logits(self, mocker): + output_mult = 3.0 + t5_model = self._build_t5_model( + use_mup=True, mup_base_hidden_size=384, mup_output_mult=output_mult + ) + bs = 2 + seq_len = t5_model.max_sequence_length + + t5_model.eval() + t5_model.cuda() + assert t5_model.model_scaling_policy.enabled + assert t5_model.model_scaling_policy.context.output_mult == pytest.approx(output_mult) + + encoder_input_ids = torch.arange(seq_len, dtype=torch.int64).repeat((bs, 1)).cuda() + decoder_input_ids = torch.arange(seq_len, dtype=torch.int64).repeat((bs, 1)).cuda() + encoder_mask = torch.zeros((bs, seq_len), dtype=bool).cuda() + decoder_mask = torch.zeros((bs, seq_len), dtype=bool).cuda() + lm_labels = torch.zeros((bs, seq_len), dtype=torch.int64).cuda() + encoder_attn_mask, decoder_attn_mask, encoder_decoder_attn_mask = ( + T5MaskedWordPieceDataset.config_attention_mask( + encoder_input_ids, decoder_input_ids, encoder_mask, decoder_mask, use_local=False + ) + ) + captured = {} + + def fake_loss(labels, logits): + captured['logits'] = logits.detach().clone() + return logits.float().mean() + + real_scale_logits = t5_model.model_scaling_policy.scale_output_logits + + def capture_scaled_logits(logits): + captured['raw_logits'] = logits.detach().clone() + return real_scale_logits(logits) + + with torch.no_grad(): + scale_spy = mocker.patch.object( + t5_model.model_scaling_policy, + 'scale_output_logits', + side_effect=capture_scaled_logits, + ) + loss_spy = mocker.patch.object( + t5_model, 'compute_language_model_loss', side_effect=fake_loss + ) + t5_model.forward( + encoder_input_ids=encoder_input_ids, + decoder_input_ids=decoder_input_ids, + encoder_attn_mask=encoder_attn_mask, + decoder_attn_mask=decoder_attn_mask, + encoder_decoder_attn_mask=encoder_decoder_attn_mask, + lm_labels=lm_labels, + ) + + scale_spy.assert_called() + loss_spy.assert_called() + expected_logits = captured['raw_logits'] * output_mult + assert torch.allclose(captured['logits'], expected_logits, atol=1e-3, rtol=0.0) + def test_forward_output_encoder_hidden_only(self): pass diff --git a/tests/unit_tests/test_lion_optimizer.py b/tests/unit_tests/test_lion_optimizer.py index b0df91073ed..834c7e84ad5 100644 --- a/tests/unit_tests/test_lion_optimizer.py +++ b/tests/unit_tests/test_lion_optimizer.py @@ -13,12 +13,19 @@ import torch import torch.nn as nn +import megatron.core.optimizer as opt_module from megatron.core.optimizer import ( HAVE_EMERGING_OPTIMIZERS, OptimizerConfig, _get_megatron_optimizer_based_on_param_groups, _get_param_groups, ) +from megatron.core.optimizer.emerging_optimizers import ( + _EMERGING_OPTIMIZERS, + _default_adam_based_eopt_config_to_kwargs, + _default_betas_for_eopt, + _muon_default_param_overrides, +) from megatron.core.optimizer.optimizer import FP32Optimizer requires_emerging_optimizers = pytest.mark.skipif( @@ -79,6 +86,96 @@ def test_lion_config_defaults(self): assert config.lion_beta2 == 0.98 assert config.muon_scalar_optimizer == "adam" + def test_muon_scalar_optimizer_controls_nonlinear_param_override(self): + """Muon scalar optimizer selection should flow into default nonlinear param overrides.""" + entry = _EMERGING_OPTIMIZERS["muon"] + config = OptimizerConfig(muon_scalar_optimizer="lion") + + overrides = entry.config_to_param_overrides(config) + assert len(overrides) == 1 + ((_, override),) = overrides.items() + assert override["optimizer"] == "lion" + + def test_muon_scalar_optimizer_routes_lion_groups_to_lion_entry(self): + """Muon scalar-optimizer overrides must create a real Lion bucket via param overrides.""" + model = SimpleModel() + config = OptimizerConfig( + optimizer="muon", + lr=1e-4, + muon_scalar_optimizer="lion", + adam_beta1=0.81, + adam_beta2=0.88, + lion_beta1=0.91, + lion_beta2=0.97, + ) + recorded = [] + + def fake_create(_config, _groups, eopt_name, _model_chunks, _pg_collection): + if eopt_name == "lion": + recorded.append((eopt_name, _default_betas_for_eopt(eopt_name, _config))) + else: + recorded.append((eopt_name, None)) + return SimpleNamespace(param_groups=[]), (lambda *_args, **_kwargs: None) + + fake_pg_collection = SimpleNamespace(mp=None, tp=None, tp_ep_pp=None) + fake_muon_entry = SimpleNamespace( + config_to_param_overrides=_muon_default_param_overrides, + default_param_overrides={}, + optimizer_cls=object, + init_state_fn=lambda *_args, **_kwargs: None, + config_to_kwargs=None, + ) + fake_lion_entry = SimpleNamespace( + config_to_param_overrides=None, + default_param_overrides={}, + optimizer_cls=object, + init_state_fn=lambda *_args, **_kwargs: None, + config_to_kwargs=None, + ) + + with ( + patch("torch.distributed.get_world_size", return_value=1), + patch( + "torch.distributed.all_gather_object", + lambda output_list, obj: output_list.__setitem__(0, obj), + ), + patch.object(opt_module, "HAVE_EMERGING_OPTIMIZERS", True), + patch.dict( + opt_module._EMERGING_OPTIMIZERS, + {"muon": fake_muon_entry, "lion": fake_lion_entry}, + clear=False, + ), + patch.object(opt_module, "_create_emerging_optimizer", side_effect=fake_create), + patch.object( + opt_module, + "FP32Optimizer", + side_effect=lambda optimizer, *_args, **_kwargs: optimizer, + ), + patch.object(opt_module, "ChainedOptimizer", side_effect=lambda optimizers: optimizers), + ): + results = opt_module._get_megatron_emerging_optimizer( + config=config, + model_chunks=[model], + config_overrides={}, + pg_collection=fake_pg_collection, + ) + + assert set(recorded) == {("muon", None), ("lion", (0.91, 0.97))} + assert len(results) == 2 + + def test_default_emerging_lion_betas_use_lion_betas(self): + """Shared beta selection must keep Lion on lion_beta{1,2}.""" + config = OptimizerConfig( + optimizer="muon", + lr=1e-4, + adam_beta1=0.81, + adam_beta2=0.88, + lion_beta1=0.91, + lion_beta2=0.97, + ) + + assert _default_betas_for_eopt("lion", config) == (0.91, 0.97) + @patch("torch.distributed.get_world_size", return_value=1) @patch( "torch.distributed.all_gather_object", @@ -141,6 +238,21 @@ def test_lion_factory_creates_lion_optimizer(self): assert group["lr"] == 3e-4 assert group["weight_decay"] == 0.01 + def test_default_emerging_lion_kwargs_use_lion_betas(self): + """Shared emerging-optimizer kwargs must keep Lion on lion_beta{1,2}.""" + config = OptimizerConfig( + optimizer="muon", + lr=1e-4, + adam_beta1=0.81, + adam_beta2=0.88, + lion_beta1=0.91, + lion_beta2=0.97, + ) + + kwargs = _default_adam_based_eopt_config_to_kwargs("lion", config, [], None) + + assert kwargs["betas"] == (0.91, 0.97) + def test_lion_init_state_fn_creates_exp_avg(self): """init_state_fn should pre-initialize exp_avg state for all params.""" model = SimpleModel() diff --git a/tests/unit_tests/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 56af8545042..96816ca0227 100644 --- a/tests/unit_tests/test_optimizer.py +++ b/tests/unit_tests/test_optimizer.py @@ -23,6 +23,7 @@ get_megatron_optimizer, get_standard_config_overrides, ) +from megatron.core.optimizer.optimizer import MegatronOptimizer, get_param_group_identifier_tuple from megatron.core.optimizer_param_scheduler import ParamGroupOverride from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import TransformerConfig @@ -211,6 +212,57 @@ def test_get_param_groups_overlapping_matches(mock_get_world_size): assert param_groups[2]['max_lr'] == 0.01 +def test_param_group_identifier_tuple_distinguishes_lr_schedule_and_eps(): + default_group = { + 'wd_mult': 1.0, + 'lr_mult': 1.0, + 'is_expert_parallel': False, + 'is_decoupled_lr': False, + 'max_lr': 1e-3, + 'min_lr': 1e-5, + } + scaled_group = {**default_group, 'max_lr': 2.5e-4, 'min_lr': 2.5e-6, 'eps': 2.5e-9} + assert get_param_group_identifier_tuple(default_group) != get_param_group_identifier_tuple( + scaled_group + ) + + +def test_filter_and_reorder_param_groups_keeps_distinct_lr_schedule_groups(): + current_groups = [ + { + 'params': ['hidden'], + 'wd_mult': 1.0, + 'lr_mult': 1.0, + 'is_expert_parallel': False, + 'is_decoupled_lr': False, + 'max_lr': 2.5e-4, + 'min_lr': 2.5e-6, + 'eps': 2.5e-9, + }, + { + 'params': ['default'], + 'wd_mult': 1.0, + 'lr_mult': 1.0, + 'is_expert_parallel': False, + 'is_decoupled_lr': False, + 'max_lr': 1e-3, + 'min_lr': 1e-5, + }, + ] + state_dict_groups = [dict(current_groups[0]), dict(current_groups[1])] + + filtered_groups = MegatronOptimizer._filter_and_reorder_param_groups( + current_groups, state_dict_groups + ) + + assert filtered_groups[0]['max_lr'] == pytest.approx(2.5e-4) + assert filtered_groups[0]['min_lr'] == pytest.approx(2.5e-6) + assert filtered_groups[0]['eps'] == pytest.approx(2.5e-9) + assert filtered_groups[1]['max_lr'] == pytest.approx(1e-3) + assert filtered_groups[1]['min_lr'] == pytest.approx(1e-5) + assert 'eps' not in filtered_groups[1] + + @patch('torch.distributed.get_world_size', return_value=1) @patch( 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) diff --git a/tests/unit_tests/test_training.py b/tests/unit_tests/test_training.py index 838b963778c..2abb8594071 100644 --- a/tests/unit_tests/test_training.py +++ b/tests/unit_tests/test_training.py @@ -9,7 +9,11 @@ from megatron.core.tokenizers.utils.build_tokenizer import vocab_size_with_padding from megatron.training.checkpointing import save_grads from megatron.training.global_vars import set_args -from megatron.training.training import build_train_valid_test_data_iterators +from megatron.training.training import ( + _should_enable_depth_mup_eval, + build_train_valid_test_data_iterators, + preprocess_common_state_dict, +) from tests.unit_tests.dist_checkpointing import TempNamedDir from tests.unit_tests.test_utilities import Utils @@ -58,6 +62,46 @@ def create_test_args(): return args +def test_preprocess_common_state_dict_tolerates_missing_optional_param_group_keys(): + base_group = { + "wd_mult": 1.0, + "lr_mult": 1.0, + "is_expert_parallel": False, + "is_decoupled_lr": False, + "max_lr": 1e-3, + "min_lr": 1e-5, + } + common_state_dict = { + "args": SimpleNamespace(use_distributed_optimizer=True, local_rank=7, rank=3), + "optimizer": { + "optimizer": { + "param_groups": [ + {**base_group, "params": ["missing_optional"]}, + {**base_group, "eps": 1e-8, "params": ["with_eps"]}, + ] + }, + "param_state": {}, + }, + } + + preprocessed = preprocess_common_state_dict(common_state_dict) + + param_groups = preprocessed["optimizer"]["optimizer"]["param_groups"] + assert [group["params"] for group in param_groups] == [["with_eps"], ["missing_optional"]] + assert "rank" not in preprocessed["args"] + assert "local_rank" not in preprocessed["args"] + + +def test_should_enable_depth_mup_eval_defaults_missing_flag_to_false(): + assert not _should_enable_depth_mup_eval(SimpleNamespace(scaling_recipe="depth_mup")) + assert not _should_enable_depth_mup_eval( + SimpleNamespace(scaling_recipe="mup", allow_depth_mup_eval=True) + ) + assert _should_enable_depth_mup_eval( + SimpleNamespace(scaling_recipe="depth_mup", allow_depth_mup_eval=True) + ) + + class TestTraining: def setup_method(self, method): Utils.initialize_model_parallel(1, 1) diff --git a/tests/unit_tests/transformer/test_mup.py b/tests/unit_tests/transformer/test_mup.py index f1d99cad1e6..6891c06cde4 100644 --- a/tests/unit_tests/transformer/test_mup.py +++ b/tests/unit_tests/transformer/test_mup.py @@ -9,24 +9,338 @@ 4. LR override computation """ +import dataclasses +import json import logging import math +from argparse import ArgumentParser +from types import SimpleNamespace from unittest.mock import patch import pytest import torch -from megatron.core.optimizer import get_mup_config_overrides, get_standard_config_overrides +import megatron.core.optimizer as optimizer_module +import megatron.training.arguments as training_args_module +import megatron.training.yaml_arguments as yaml_args_module +from megatron.core.config_logger import log_config_to_disk +from megatron.core.models.gpt.fine_grained_callables import _apply_mlp_bda_with_scaling +from megatron.core.optimizer import ( + _get_megatron_optimizer_based_on_param_groups, + get_mup_config_overrides, + get_scaling_config_overrides, + get_standard_config_overrides, +) from megatron.core.optimizer.optimizer_config import OptimizerConfig from megatron.core.optimizer_param_scheduler import combine_param_group_overrides +from megatron.core.parameterization import ( + build_resolved_model_policy, + build_resolved_scaling_context, + build_resolved_training_policy, + depth_mup_eval_context, +) +from megatron.core.transformer.mlp import MLP, MLPSubmodules from megatron.core.transformer.multi_token_prediction import process_mtp_loss from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.core.transformer.transformer_layer import TransformerLayer, TransformerLayerSubmodules from megatron.core.utils import init_method_normal, mup_scaled_init_method_normal +from megatron.training.arguments import ( + add_megatron_arguments, + core_transformer_config_from_args, + validate_args, + validate_depth_mup_optimizer_support, + validate_muon_scalar_optimizer_support, + warn_deprecated_mup_aliases, +) +from megatron.training.checkpointing import check_checkpoint_args, load_args_from_checkpoint +from megatron.training.yaml_arguments import core_config_from_args as core_config_from_yaml_args +from megatron.training.yaml_arguments import validate_yaml + + +def _build_transformer_namespace(**overrides): + values = {} + for field in dataclasses.fields(TransformerConfig): + if field.default is not dataclasses.MISSING: + values[field.name] = field.default + elif field.default_factory is not dataclasses.MISSING: + values[field.name] = field.default_factory() + else: + assert field.name in overrides, f"Missing required override for {field.name}" + values[field.name] = overrides.pop(field.name) + values.update(overrides) + return SimpleNamespace(**values) + + +def _prepare_parsed_args_for_core_config(args): + # parse_args() populates CLI-backed fields only; validate_args() normally derives params_dtype. + args.params_dtype = torch.float32 + return args + + +def _build_minimal_validate_yaml_namespace(**overrides): + language_model_overrides = overrides.pop('language_model', None) + model_parallel_overrides = overrides.pop('model_parallel', None) + values = dict( + data_path=None, + world_size=1, + rank=0, + micro_batch_size=1, + global_batch_size=1, + num_layers_per_virtual_pipeline_stage=None, + overlap_param_gather=False, + overlap_grad_reduce=False, + use_distributed_optimizer=False, + accumulate_allreduce_grads_in_fp32=False, + dataloader_type=None, + lr_decay_samples=None, + rampup_batch_size=None, + train_iters=None, + train_samples=None, + lr_decay_iters=None, + lr_warmup_iters=0, + lr_warmup_fraction=None, + lr_warmup_samples=0, + encoder_num_layers=None, + seq_length=128, + encoder_seq_length=None, + max_position_embeddings=128, + decoder_seq_length=None, + lr=1e-3, + min_lr=1e-5, + weight_decay=0.01, + weight_decay_incr_style='constant', + start_weight_decay=None, + end_weight_decay=None, + save=None, + save_interval=None, + fp16_lm_cross_entropy=False, + account_for_embedding_in_pipeline_split=False, + overlap_p2p_comm=False, + spec=None, + model_parallel=SimpleNamespace( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=1, + expert_model_parallel_size=1, + tp_comm_overlap=False, + sequence_parallel=False, + fp16=False, + bf16=False, + params_dtype=torch.float32, + ), + language_model=SimpleNamespace( + num_layers=12, + hidden_size=1024, + num_attention_heads=16, + ffn_hidden_size=None, + activation_func='gelu', + kv_channels=None, + scaling_recipe='none', + fp32_residual_connection=False, + moe_grouped_gemm=False, + persist_layer_norm=False, + distribute_saved_activations=False, + recompute_granularity=None, + recompute_method=None, + num_moe_experts=None, + ), + ) + values.update(overrides) + if language_model_overrides is not None: + values['language_model'] = SimpleNamespace( + **{**vars(values['language_model']), **vars(language_model_overrides)} + ) + if model_parallel_overrides is not None: + values['model_parallel'] = SimpleNamespace( + **{**vars(values['model_parallel']), **vars(model_parallel_overrides)} + ) + return SimpleNamespace(**values) + + +def _combined_override_for_param(overrides, param, param_name): + matches = [ + override + for param_key, override in overrides.items() + if param_key.matches(param, param_name) + ] + return combine_param_group_overrides(matches) class TestMuPConfigValidation: """Tests for MuP config validation and width_mult computation.""" + def test_scaling_recipe_mup_matches_legacy_alias(self): + legacy = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_base_head_dim=64, + ) + recipe = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + ) + + assert recipe.use_mup is True + assert recipe.scaling_recipe == 'mup' + assert recipe.mup_width_mult == legacy.mup_width_mult + assert recipe.mup_base_hidden_size == legacy.mup_base_hidden_size + assert recipe.softmax_scale == legacy.softmax_scale + assert recipe.mup_output_mult == legacy.mup_output_mult + + def test_scaling_recipe_depth_mup_preserves_distinct_recipe_identity(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_base_head_dim=64, + ) + context = build_resolved_scaling_context(config) + + assert config.scaling_recipe == 'depth_mup' + assert config.use_mup is False + assert context.recipe == 'depth_mup' + assert context.uses_width_mup is True + assert context.references.base_hidden_size == 256 + assert context.references.base_num_layers == 6 + assert context.references.base_head_dim == 64 + assert context.residual_branch_depth_power == pytest.approx(-1.0) + assert context.hidden_lr_depth_power == pytest.approx(0.0) + assert context.block_out_proj_init_depth_power == pytest.approx(0.5) + assert context.output_mult == pytest.approx(0.25) + + def test_depth_mup_manual_overrides_can_zero_recipe_defaults(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_residual_branch_depth_power=0.0, + scaling_block_out_proj_init_depth_power=0.0, + ) + context = build_resolved_scaling_context(config) + + assert context.residual_branch_depth_power == pytest.approx(0.0) + assert context.block_out_proj_init_depth_power == pytest.approx(0.0) + + def test_mup_does_not_inherit_depth_mup_recipe_defaults(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + ) + context = build_resolved_scaling_context(config) + + assert context.recipe == 'mup' + assert context.residual_branch_depth_power == pytest.approx(0.0) + assert context.hidden_lr_depth_power == pytest.approx(0.0) + assert context.block_out_proj_init_depth_power == pytest.approx(0.0) + + def test_conflicting_legacy_and_canonical_scaling_args_error(self): + with pytest.raises(ValueError, match='conflicts with'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + mup_base_hidden_size=512, + ) + + def test_explicit_none_conflicts_with_legacy_use_mup_alias(self): + with pytest.raises(ValueError, match='conflicts with --use-mup'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='none', + use_mup=True, + ) + + def test_depth_mup_conflicts_with_legacy_use_mup_alias(self): + with pytest.raises(ValueError, match='conflicts with --use-mup'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + use_mup=True, + ) + + def test_scaling_overrides_require_recipe(self): + with pytest.raises( + ValueError, match="Scaling overrides require a non-'none' scaling recipe" + ): + TransformerConfig( + hidden_size=512, + num_layers=12, + num_attention_heads=8, + scaling_residual_branch_depth_power=-0.5, + ) + + def test_legacy_mup_knobs_require_mup_recipe(self): + with pytest.raises( + ValueError, match="Scaling overrides require a non-'none' scaling recipe" + ): + TransformerConfig( + hidden_size=512, num_layers=12, num_attention_heads=8, mup_base_hidden_size=256 + ) + + def test_scaling_base_head_dim_must_be_positive(self): + with pytest.raises(AssertionError, match='scaling-base-head-dim'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_head_dim=-64, + ) + + def test_depth_mup_rejects_multi_latent_attention(self): + with pytest.raises(NotImplementedError, match='multi_latent_attention'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + multi_latent_attention=True, + ) + + def test_depth_mup_rejects_experimental_attention_variant(self): + with pytest.raises(NotImplementedError, match='experimental attention variants'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + experimental_attention_variant='gated_delta_net', + linear_attention_freq=1, + ) + + def test_depth_mup_rejects_moe(self): + with pytest.raises(NotImplementedError, match='MoE depth transfer'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + num_moe_experts=4, + ) + def test_mup_defaults_base_hidden_size(self): """use_mup without base_hidden_size defaults to hidden_size (width_mult=1.0).""" config = TransformerConfig( @@ -61,6 +375,707 @@ def test_mup_width_mult_fractional(self): ) assert config.mup_width_mult == 0.5 + def test_resolved_scaling_context_matches_transformer_config_fields(self): + config = TransformerConfig( + hidden_size=1536, + num_layers=18, + num_attention_heads=12, + scaling_recipe='mup', + scaling_base_hidden_size=384, + scaling_base_num_layers=9, + scaling_base_head_dim=64, + ) + context = build_resolved_scaling_context(config) + + assert context.recipe == 'mup' + assert context.width_mult == pytest.approx(config.mup_width_mult) + assert context.references.base_hidden_size == config.mup_base_hidden_size + assert context.references.base_num_layers == config.scaling_base_num_layers + assert context.references.base_head_dim == config.mup_base_head_dim + + def test_parser_accepts_canonical_and_legacy_mup_flags(self): + parser = ArgumentParser(allow_abbrev=False) + parser = add_megatron_arguments(parser) + + canonical_args = parser.parse_args( + [ + '--num-layers', + '12', + '--hidden-size', + '1024', + '--num-attention-heads', + '16', + '--no-rope-fusion', + '--scaling-recipe', + 'mup', + '--scaling-base-hidden-size', + '256', + '--scaling-base-head-dim', + '64', + ] + ) + canonical_args = _prepare_parsed_args_for_core_config(canonical_args) + canonical_config = core_transformer_config_from_args(canonical_args) + assert canonical_config.use_mup is True + assert canonical_config.scaling_recipe == 'mup' + + depth_args = parser.parse_args( + [ + '--num-layers', + '12', + '--hidden-size', + '1024', + '--num-attention-heads', + '16', + '--no-rope-fusion', + '--scaling-recipe', + 'depth_mup', + '--optimizer', + 'adam', + '--scaling-base-hidden-size', + '256', + '--scaling-base-num-layers', + '6', + '--scaling-base-head-dim', + '64', + ] + ) + depth_args = _prepare_parsed_args_for_core_config(depth_args) + depth_config = core_transformer_config_from_args(depth_args) + assert depth_config.use_mup is False + assert depth_config.scaling_recipe == 'depth_mup' + assert depth_args.optimizer == 'adam' + assert canonical_config.mup_width_mult == pytest.approx(4.0) + + legacy_args = parser.parse_args( + [ + '--num-layers', + '12', + '--hidden-size', + '1024', + '--num-attention-heads', + '16', + '--no-rope-fusion', + '--use-mup', + '--mup-base-hidden-size', + '256', + '--mup-base-head-dim', + '64', + '--mup-width-mult', + '4.0', + ] + ) + legacy_args = _prepare_parsed_args_for_core_config(legacy_args) + legacy_config = core_transformer_config_from_args(legacy_args) + assert legacy_args.mup_width_mult == pytest.approx(4.0) + assert legacy_config.use_mup is True + assert legacy_config.mup_width_mult == pytest.approx(4.0) + + def test_legacy_mup_width_mult_must_match_derived_width_mult(self): + with pytest.raises(ValueError, match='must match the derived'): + TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_width_mult=3.0, + ) + + def test_legacy_mup_width_mult_requires_active_scaling_recipe(self): + with pytest.raises(ValueError, match='Scaling overrides require'): + TransformerConfig( + hidden_size=1024, num_layers=12, num_attention_heads=16, mup_width_mult=3.0 + ) + + def test_legacy_mup_aliases_emit_deprecation_warning(self): + args = SimpleNamespace( + use_mup=True, mup_base_hidden_size=256, mup_base_head_dim=64, mup_width_mult=3.0 + ) + + with patch.object(training_args_module, 'warn_rank_0') as warn: + warn_deprecated_mup_aliases(args) + + warn.assert_called_once() + message = warn.call_args.args[0] + assert '--use-mup' in message + assert '--mup-base-hidden-size' in message + assert '--mup-base-head-dim' in message + assert '--mup-width-mult' in message + assert '--scaling-recipe mup' in message + + def test_canonical_mup_recipe_does_not_emit_legacy_deprecation_warning(self): + args = SimpleNamespace( + use_mup=False, mup_base_hidden_size=None, mup_base_head_dim=None, mup_width_mult=1.0 + ) + + with patch.object(training_args_module, 'warn_rank_0') as warn: + warn_deprecated_mup_aliases(args) + + warn.assert_not_called() + + def test_validate_args_calls_depth_mup_hook(self): + parser = ArgumentParser(allow_abbrev=False) + parser = add_megatron_arguments(parser) + args = parser.parse_args( + [ + '--num-layers', + '12', + '--hidden-size', + '1024', + '--num-attention-heads', + '16', + '--seq-length', + '128', + '--micro-batch-size', + '1', + '--max-position-embeddings', + '128', + '--no-rope-fusion', + '--optimizer', + 'adam', + '--scaling-recipe', + 'depth_mup', + ] + ) + args.world_size = 1 + args.rank = 0 + + with ( + patch.object( + training_args_module, + 'validate_depth_mup_optimizer_support', + side_effect=RuntimeError('depth-hook'), + ), + patch.object( + training_args_module, 'validate_muon_scalar_optimizer_support', return_value=None + ), + ): + with pytest.raises(RuntimeError, match='depth-hook'): + validate_args(args) + + def test_validate_args_calls_muon_scalar_optimizer_hook(self): + parser = ArgumentParser(allow_abbrev=False) + parser = add_megatron_arguments(parser) + args = parser.parse_args( + [ + '--num-layers', + '12', + '--hidden-size', + '1024', + '--num-attention-heads', + '16', + '--seq-length', + '128', + '--micro-batch-size', + '1', + '--max-position-embeddings', + '128', + '--no-rope-fusion', + '--optimizer', + 'muon', + '--muon-scalar-optimizer', + 'lion', + ] + ) + args.world_size = 1 + args.rank = 0 + + with ( + patch.object( + training_args_module, 'validate_depth_mup_optimizer_support', return_value=None + ), + patch.object( + training_args_module, + 'validate_muon_scalar_optimizer_support', + side_effect=RuntimeError('muon-hook'), + ), + ): + with pytest.raises(RuntimeError, match='muon-hook'): + validate_args(args) + + def test_yaml_namespace_without_scaling_fields_still_builds_transformer_config(self): + args = _build_transformer_namespace(hidden_size=1024, num_layers=12, num_attention_heads=16) + for field_name in ( + 'scaling_recipe', + 'scaling_base_hidden_size', + 'scaling_base_num_layers', + 'scaling_base_head_dim', + 'scaling_residual_branch_depth_power', + 'scaling_hidden_lr_depth_power', + 'scaling_block_out_proj_init_depth_power', + ): + delattr(args, field_name) + + kw_args = core_config_from_yaml_args(args, TransformerConfig) + + assert kw_args['scaling_recipe'] is None + assert kw_args['scaling_base_hidden_size'] is None + assert kw_args['scaling_base_num_layers'] is None + assert kw_args['scaling_base_head_dim'] is None + + def test_yaml_namespace_preserves_depth_mup_scaling_surface(self): + args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_base_head_dim=64, + optimizer='adam', + ) + + kw_args = core_config_from_yaml_args(args, TransformerConfig) + config = TransformerConfig(**kw_args) + context = build_resolved_scaling_context(config) + + assert config.scaling_recipe == 'depth_mup' + assert context.recipe == 'depth_mup' + assert context.uses_width_mup is True + assert context.residual_branch_depth_power == pytest.approx(-1.0) + assert context.hidden_lr_depth_power == pytest.approx(0.0) + assert context.block_out_proj_init_depth_power == pytest.approx(0.5) + + def test_depth_mup_optimizer_gate_rejects_non_adam_yaml_namespace(self): + args = SimpleNamespace(scaling_recipe='depth_mup', optimizer='sgd') + + with pytest.raises(ValueError, match="supports optimizer='adam' only"): + validate_depth_mup_optimizer_support(args) + + def test_depth_mup_optimizer_gate_allows_adam_yaml_namespace(self): + args = SimpleNamespace(scaling_recipe='depth_mup', optimizer='adam') + + validate_depth_mup_optimizer_support(args) + + def test_depth_mup_optimizer_gate_rejects_nested_yaml_namespace(self): + args = SimpleNamespace( + optimizer='sgd', language_model=SimpleNamespace(scaling_recipe='depth_mup') + ) + + with pytest.raises(ValueError, match="supports optimizer='adam' only"): + validate_depth_mup_optimizer_support(args) + + def test_depth_mup_optimizer_gate_tolerates_yaml_namespace_without_scaling_fields(self): + validate_depth_mup_optimizer_support(SimpleNamespace()) + + def test_muon_scalar_optimizer_gate_rejects_invalid_nested_yaml_namespace(self): + args = SimpleNamespace( + muon_scalar_optimizer='soap', language_model=SimpleNamespace(scaling_recipe='none') + ) + + with pytest.raises(ValueError, match="muon_scalar_optimizer must be one of"): + validate_muon_scalar_optimizer_support(args) + + def test_muon_scalar_optimizer_gate_allows_lion_nested_yaml_namespace(self): + args = SimpleNamespace(muon_scalar_optimizer='lion') + + validate_muon_scalar_optimizer_support(args) + + def test_validate_yaml_calls_depth_mup_hook(self): + args = _build_minimal_validate_yaml_namespace( + optimizer='adam', + language_model=SimpleNamespace( + num_layers=12, + hidden_size=1024, + num_attention_heads=16, + ffn_hidden_size=None, + activation_func='gelu', + kv_channels=None, + scaling_recipe='depth_mup', + fp32_residual_connection=False, + moe_grouped_gemm=False, + ), + ) + + with ( + patch.object( + yaml_args_module, + 'validate_depth_mup_optimizer_support', + side_effect=RuntimeError('yaml-depth-hook'), + ), + patch.object( + yaml_args_module, 'validate_muon_scalar_optimizer_support', return_value=None + ), + ): + with pytest.raises(RuntimeError, match='yaml-depth-hook'): + validate_yaml(args) + + def test_validate_yaml_calls_muon_scalar_optimizer_hook(self): + args = _build_minimal_validate_yaml_namespace( + optimizer='adam', muon_scalar_optimizer='lion' + ) + + with ( + patch.object( + yaml_args_module, 'validate_depth_mup_optimizer_support', return_value=None + ), + patch.object( + yaml_args_module, + 'validate_muon_scalar_optimizer_support', + side_effect=RuntimeError('yaml-muon-hook'), + ), + ): + with pytest.raises(RuntimeError, match='yaml-muon-hook'): + validate_yaml(args) + + def test_validate_yaml_defaults_depth_mup_eval_flag(self): + args = _build_minimal_validate_yaml_namespace(optimizer='adam') + + validated_args = validate_yaml(args) + + assert validated_args.allow_depth_mup_eval is False + + def test_validate_yaml_calls_mup_deprecation_warning(self): + args = _build_minimal_validate_yaml_namespace(optimizer='adam') + + with ( + patch.object( + yaml_args_module, 'validate_depth_mup_optimizer_support', return_value=None + ), + patch.object( + yaml_args_module, 'validate_muon_scalar_optimizer_support', return_value=None + ), + patch.object( + yaml_args_module, + 'warn_deprecated_mup_aliases', + side_effect=RuntimeError('yaml-mup-warning'), + ), + ): + with pytest.raises(RuntimeError, match='yaml-mup-warning'): + validate_yaml(args) + + def test_validate_yaml_syncs_legacy_mup_alias_fields(self): + args = _build_minimal_validate_yaml_namespace( + optimizer='adam', + use_mup=True, + mup_base_hidden_size=256, + mup_base_head_dim=64, + mup_width_mult=4.0, + language_model=SimpleNamespace( + num_layers=12, + hidden_size=1024, + num_attention_heads=16, + ffn_hidden_size=None, + activation_func='gelu', + kv_channels=None, + scaling_recipe=None, + fp32_residual_connection=False, + moe_grouped_gemm=False, + ), + ) + + validated_args = validate_yaml(args) + + assert validated_args.scaling_recipe == 'mup' + assert validated_args.scaling_base_hidden_size == 256 + assert validated_args.scaling_base_head_dim == 64 + assert validated_args.mup_width_mult == pytest.approx(4.0) + + def test_config_logger_serializes_canonical_depth_mup_surface(self, tmp_path): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + config_logger_dir=str(tmp_path), + ) + + log_config_to_disk( + config, {'config': config}, prefix='depth_mup_config', rank_str='0_0_0_0_0' + ) + + output_path = tmp_path / 'depth_mup_config.rank_0_0_0_0_0.iter0.json' + with output_path.open() as fp: + payload = json.load(fp) + + serialized = payload['config'] + assert serialized['scaling_recipe'] == 'depth_mup' + assert serialized['scaling_base_hidden_size'] == 256 + assert serialized['scaling_base_num_layers'] == 6 + assert serialized['scaling_residual_branch_depth_power'] == pytest.approx(-1.0) + assert serialized['scaling_hidden_lr_depth_power'] == pytest.approx(0.0) + assert serialized['scaling_block_out_proj_init_depth_power'] == pytest.approx(0.5) + + def test_checkpoint_args_compare_effective_scaling_context(self): + runtime_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + checkpoint_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_base_head_dim=64, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + delattr(runtime_args, 'kv_channels') + delattr(checkpoint_args, 'kv_channels') + + with ( + patch('megatron.training.checkpointing.get_args', return_value=runtime_args), + patch('megatron.training.checkpointing.get_checkpoint_version', return_value=3.0), + ): + check_checkpoint_args(checkpoint_args) + + def test_checkpoint_args_compare_effective_depth_mup_scaling_context(self): + runtime_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_base_head_dim=64, + optimizer='adam', + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + checkpoint_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_base_head_dim=64, + optimizer='adam', + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + delattr(runtime_args, 'kv_channels') + delattr(checkpoint_args, 'kv_channels') + + with ( + patch('megatron.training.checkpointing.get_args', return_value=runtime_args), + patch('megatron.training.checkpointing.get_checkpoint_version', return_value=3.0), + ): + check_checkpoint_args(checkpoint_args) + + def test_checkpoint_args_raise_on_scaling_mismatch(self): + runtime_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + checkpoint_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=512, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + + with ( + patch('megatron.training.checkpointing.get_args', return_value=runtime_args), + patch('megatron.training.checkpointing.get_checkpoint_version', return_value=3.0), + ): + with pytest.raises(AssertionError, match='Resolved scaling context'): + check_checkpoint_args(checkpoint_args) + + def test_checkpoint_args_raise_on_legacy_mup_knob_mismatch(self): + runtime_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + mup_embedding_mult=1.2, + mup_output_mult=0.35, + mup_attn_scale_power=0.8, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + checkpoint_args = _build_transformer_namespace( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_embedding_mult=1.1, + mup_output_mult=0.35, + mup_attn_scale_power=0.8, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_dist_ckpt=False, + ) + + with ( + patch('megatron.training.checkpointing.get_args', return_value=runtime_args), + patch('megatron.training.checkpointing.get_checkpoint_version', return_value=3.0), + ): + with pytest.raises(AssertionError, match='Resolved scaling context'): + check_checkpoint_args(checkpoint_args) + + def test_use_checkpoint_args_restores_scaling_surface(self): + args = SimpleNamespace( + load='dummy', + iteration=0, + use_mp_args_from_checkpoint_args=False, + use_tokenizer_model_from_checkpoint_args=False, + use_mup=False, + scaling_recipe=None, + scaling_base_hidden_size=None, + scaling_base_num_layers=None, + scaling_base_head_dim=None, + scaling_residual_branch_depth_power=None, + scaling_hidden_lr_depth_power=None, + scaling_block_out_proj_init_depth_power=None, + mup_width_mult=1.0, + mup_base_hidden_size=None, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=None, + mup_attn_scale_power=1.0, + ) + checkpoint_args = SimpleNamespace( + use_mup=True, + mup_width_mult=4.0, + mup_base_hidden_size=256, + mup_embedding_mult=1.5, + mup_output_mult=0.3, + mup_base_head_dim=64, + mup_attn_scale_power=0.75, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=12, + scaling_base_head_dim=64, + scaling_residual_branch_depth_power=-0.5, + scaling_hidden_lr_depth_power=-0.25, + scaling_block_out_proj_init_depth_power=-0.5, + ) + state_dict = {'args': checkpoint_args, 'checkpoint_version': 3.0, 'iteration': 17} + + with patch( + 'megatron.training.checkpointing._load_base_checkpoint', + return_value=(state_dict, 'dummy', False, None), + ): + load_args_from_checkpoint(args) + + assert args.iteration == 17 + assert args.use_mup is True + assert args.mup_embedding_mult == pytest.approx(1.5) + assert args.mup_output_mult == pytest.approx(0.3) + assert args.mup_attn_scale_power == pytest.approx(0.75) + assert args.scaling_recipe == 'mup' + assert args.scaling_base_hidden_size == 256 + assert args.scaling_base_num_layers == 12 + assert args.scaling_base_head_dim == 64 + assert args.scaling_residual_branch_depth_power == pytest.approx(-0.5) + assert args.scaling_hidden_lr_depth_power == pytest.approx(-0.25) + assert args.scaling_block_out_proj_init_depth_power == pytest.approx(-0.5) + + def test_use_checkpoint_args_restores_depth_mup_surface(self): + args = SimpleNamespace( + load='dummy', + iteration=0, + use_mp_args_from_checkpoint_args=False, + use_tokenizer_model_from_checkpoint_args=False, + use_mup=False, + scaling_recipe=None, + scaling_base_hidden_size=None, + scaling_base_num_layers=None, + scaling_base_head_dim=None, + scaling_residual_branch_depth_power=None, + scaling_hidden_lr_depth_power=None, + scaling_block_out_proj_init_depth_power=None, + mup_width_mult=1.0, + mup_base_hidden_size=None, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=None, + mup_attn_scale_power=1.0, + ) + checkpoint_args = SimpleNamespace( + use_mup=False, + mup_width_mult=4.0, + mup_base_hidden_size=256, + mup_embedding_mult=1.0, + mup_output_mult=0.25, + mup_base_head_dim=64, + mup_attn_scale_power=1.0, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_base_head_dim=64, + scaling_residual_branch_depth_power=-1.0, + scaling_hidden_lr_depth_power=0.0, + scaling_block_out_proj_init_depth_power=0.5, + ) + state_dict = {'args': checkpoint_args, 'checkpoint_version': 3.0, 'iteration': 23} + + with patch( + 'megatron.training.checkpointing._load_base_checkpoint', + return_value=(state_dict, 'dummy', False, None), + ): + load_args_from_checkpoint(args) + + assert args.iteration == 23 + assert args.use_mup is False + assert args.scaling_recipe == 'depth_mup' + assert args.scaling_base_hidden_size == 256 + assert args.scaling_base_num_layers == 6 + assert args.scaling_base_head_dim == 64 + assert args.scaling_residual_branch_depth_power == pytest.approx(-1.0) + assert args.scaling_hidden_lr_depth_power == pytest.approx(0.0) + assert args.scaling_block_out_proj_init_depth_power == pytest.approx(0.5) + def test_mup_backward_compatible(self): """Default config unchanged when MuP disabled.""" config = TransformerConfig(hidden_size=512, num_layers=4, num_attention_heads=8) @@ -68,6 +1083,456 @@ def test_mup_backward_compatible(self): assert config.mup_width_mult == 1.0 assert config.mup_base_hidden_size is None + +class TestMuPModelPolicy: + @pytest.mark.parametrize('recipe', ['mup', 'depth_mup']) + def test_model_policy_scales_embedding_and_logits(self, recipe): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe=recipe, + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + mup_embedding_mult=1.25, + mup_output_mult=0.2, + ) + policy = build_resolved_model_policy(config) + embeddings = torch.ones(2, 3) + logits = torch.ones(2, 3) + + assert torch.equal(policy.scale_embedding_activations(embeddings), embeddings * 1.25) + assert torch.equal(policy.scale_output_logits(logits), logits * 0.2) + + @pytest.mark.parametrize('recipe', ['mup', 'depth_mup']) + def test_model_policy_uses_embedding_init_for_untied_readout(self, recipe): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe=recipe, + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + ) + policy = build_resolved_model_policy(config) + + assert ( + policy.output_layer_init_method( + share_embeddings_and_output_weights=False, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, + ) + is config.embedding_init_method + ) + assert ( + policy.output_layer_init_method( + share_embeddings_and_output_weights=True, + default_init_method=config.init_method, + embedding_init_method=config.embedding_init_method, + ) + is config.init_method + ) + + def test_dense_block_output_init_depth_scaling(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_block_out_proj_init_depth_power=-0.5, + ) + policy = build_resolved_model_policy(config) + init_fn = policy.dense_block_output_init_method( + default_init_method=config.output_layer_init_method, + init_method_std=config.init_method_std, + num_layers=config.num_layers, + is_hybrid_model=config.is_hybrid_model, + output_layer_init_method_is_user_provided=False, + ) + weights = torch.empty(200_000) + init_fn(weights) + + expected_std = ( + config.init_method_std + / (math.sqrt(2 * config.num_layers) * math.sqrt(config.mup_width_mult)) + * (policy.context.depth_mult**config.scaling_block_out_proj_init_depth_power) + ) + actual_std = weights.std().item() + assert abs(actual_std - expected_std) < expected_std * 0.05 + + def test_depth_mup_default_block_output_init_rebases_to_base_depth(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + ) + policy = build_resolved_model_policy(config) + init_fn = policy.dense_block_output_init_method( + default_init_method=config.output_layer_init_method, + init_method_std=config.init_method_std, + num_layers=config.num_layers, + is_hybrid_model=config.is_hybrid_model, + output_layer_init_method_is_user_provided=False, + ) + weights = torch.empty(200_000) + init_fn(weights) + + expected_std = config.init_method_std / ( + math.sqrt(2 * config.scaling_base_num_layers) * math.sqrt(policy.context.width_mult) + ) + actual_std = weights.std().item() + assert abs(actual_std - expected_std) < expected_std * 0.05 + + def test_dense_block_output_init_can_be_disabled_per_site(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_block_out_proj_init_depth_power=-0.5, + ) + policy = build_resolved_model_policy(config) + + assert ( + policy.dense_block_output_init_method( + default_init_method=config.output_layer_init_method, + init_method_std=config.init_method_std, + num_layers=config.num_layers, + is_hybrid_model=config.is_hybrid_model, + output_layer_init_method_is_user_provided=False, + apply_depth_hook=False, + ) + is config.output_layer_init_method + ) + + def test_plain_mlp_requires_explicit_block_output_init_scaling_opt_in(self): + class DummyLinear(torch.nn.Module): + def __init__(self, init_method): + super().__init__() + self.init_method = init_method + + def forward(self, hidden_states): + return hidden_states, None + + def backward_dw(self): + return None + + def fc1_builder(input_size, output_size, *, init_method, **kwargs): + return DummyLinear(init_method) + + def fc2_builder(input_size, output_size, *, init_method, **kwargs): + return DummyLinear(init_method) + + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_block_out_proj_init_depth_power=-0.5, + ) + submodules = MLPSubmodules(linear_fc1=fc1_builder, linear_fc2=fc2_builder) + + plain_mlp = MLP(config, submodules, apply_block_output_init_scaling=False) + scaled_mlp = MLP(config, submodules, apply_block_output_init_scaling=True) + + assert plain_mlp.linear_fc2.init_method is config.output_layer_init_method + assert scaled_mlp.linear_fc2.init_method is not config.output_layer_init_method + + def test_transformer_layer_residual_branch_scaling_helper(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=6, + scaling_residual_branch_depth_power=-0.5, + ) + layer = object.__new__(TransformerLayer) + layer.model_scaling_policy = build_resolved_model_policy(config) + + output = torch.ones(4, 8) + bias = torch.ones(8) + scaled_output, scaled_bias = layer._scale_dense_residual_branch_output( + (output, bias), branch_name='self attention', using_fused_tp_inference_kernel=False + ) + expected_mult = ( + config.num_layers / config.scaling_base_num_layers + ) ** config.scaling_residual_branch_depth_power + assert torch.equal(scaled_output, output * expected_mult) + assert torch.equal(scaled_bias, bias * expected_mult) + + with pytest.raises(NotImplementedError, match='Residual-branch scaling'): + layer._scale_dense_residual_branch_output( + (output, bias), branch_name='self attention', using_fused_tp_inference_kernel=True + ) + + def test_depth_mup_rejects_inference_even_at_base_depth(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=12, + ) + layer = object.__new__(TransformerLayer) + layer.model_scaling_policy = build_resolved_model_policy(config) + layer.training = False + + with pytest.raises(NotImplementedError, match='during inference'): + layer._scale_dense_residual_branch_output( + (torch.ones(2, 2), None), + branch_name='self attention', + using_fused_tp_inference_kernel=False, + ) + + def test_depth_mup_eval_context_allows_unfused_validation_scaling(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=24, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=12, + ) + layer = object.__new__(TransformerLayer) + layer.model_scaling_policy = build_resolved_model_policy(config) + layer.training = False + + output = torch.ones(2, 2) + bias = torch.full((2, 2), 3.0) + expected_mult = config.num_layers / config.scaling_base_num_layers + + with depth_mup_eval_context(True): + scaled_output, scaled_bias = layer._scale_dense_residual_branch_output( + (output, bias), branch_name='self attention', using_fused_tp_inference_kernel=False + ) + + assert torch.equal(scaled_output, output / expected_mult) + assert torch.equal(scaled_bias, bias / expected_mult) + + def test_depth_mup_rejects_fused_tp_inference_specifically(self): + config = TransformerConfig( + hidden_size=1024, + num_layers=12, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=12, + ) + layer = object.__new__(TransformerLayer) + layer.model_scaling_policy = build_resolved_model_policy(config) + layer.training = False + + with pytest.raises(NotImplementedError, match='fused TP inference'): + layer._scale_dense_residual_branch_output( + (torch.ones(2, 2), None), + branch_name='self attention', + using_fused_tp_inference_kernel=True, + ) + + def test_transformer_layer_rejects_cross_attention_for_depth_mup(self): + class DummyCrossAttention(torch.nn.Module): + def __init__(self, *args, **kwargs): + super().__init__() + + def forward(self, *args, **kwargs): + return torch.ones(1, 1, 1), None + + config = TransformerConfig( + hidden_size=16, + num_layers=12, + num_attention_heads=4, + scaling_recipe='depth_mup', + scaling_base_hidden_size=8, + scaling_base_num_layers=6, + ) + submodules = TransformerLayerSubmodules(cross_attention=DummyCrossAttention) + + with pytest.raises(NotImplementedError, match='Cross-attention is out of scope for v1'): + TransformerLayer(config=config, submodules=submodules) + + def test_depth_mup_rejects_bert_model(self): + from megatron.core.models.bert.bert_model import BertModel + + config = TransformerConfig( + hidden_size=16, + num_layers=12, + num_attention_heads=4, + scaling_recipe='depth_mup', + scaling_base_hidden_size=8, + scaling_base_num_layers=6, + ) + + with pytest.raises(NotImplementedError, match='BertModel is out of scope for v1'): + BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=None, + vocab_size=128, + max_sequence_length=16, + ) + + def test_depth_mup_rejects_t5_model(self): + from megatron.core.models.T5.t5_model import T5Model + + config = TransformerConfig( + hidden_size=16, + num_layers=12, + num_attention_heads=4, + scaling_recipe='depth_mup', + scaling_base_hidden_size=8, + scaling_base_num_layers=6, + ) + + with pytest.raises(NotImplementedError, match='T5Model is out of scope for v1'): + T5Model( + config=config, + encoder_config=config, + transformer_encoder_layer_spec=None, + transformer_decoder_layer_spec=None, + vocab_size=128, + max_sequence_length=16, + ) + + def test_depth_mup_rejects_mamba_model(self): + from megatron.core.models.mamba.mamba_model import MambaModel + + config = TransformerConfig( + hidden_size=16, + num_layers=12, + num_attention_heads=4, + scaling_recipe='depth_mup', + scaling_base_hidden_size=8, + scaling_base_num_layers=6, + ) + + with pytest.raises(NotImplementedError, match='MambaModel is out of scope for v1'): + MambaModel(config=config, mamba_stack_spec=None, vocab_size=128, max_sequence_length=16) + + def test_overlap_helper_routes_through_residual_branch_scaler(self): + class DummyCtx: + def __enter__(self): + return None + + def __exit__(self, exc_type, exc, tb): + return False + + class DummyConfig: + bias_dropout_fusion = False + + class DummyLayer: + training = True + hidden_dropout = 0.1 + is_moe_layer = False + config = DummyConfig() + + def __init__(self): + self.scaler_called = False + + def _scale_dense_residual_branch_output( + self, + output_with_bias, + *, + branch_name, + using_fused_tp_inference_kernel, + apply_depth_hook, + ): + self.scaler_called = True + assert branch_name == 'mlp' + assert using_fused_tp_inference_kernel is False + assert apply_depth_hook is True + output, _ = output_with_bias + return (output * 2.0, None) + + def bias_dropout_add_exec_handler(self): + return DummyCtx() + + def mlp_bda(self, training, bias_dropout_fusion): + assert training is True + assert bias_dropout_fusion is False + + def apply(output_with_bias, residual, hidden_dropout): + output, bias = output_with_bias + assert bias is None + return output + residual + hidden_dropout + + return apply + + layer = DummyLayer() + output = torch.ones(3, 4) + residual = torch.ones(3, 4) + hidden_states = _apply_mlp_bda_with_scaling(layer, output, residual) + + assert layer.scaler_called is True + assert torch.equal(hidden_states, output * 2.0 + residual + layer.hidden_dropout) + + def test_overlap_helper_disables_depth_hook_for_moe_layers(self): + class DummyCtx: + def __enter__(self): + return None + + def __exit__(self, exc_type, exc, tb): + return False + + class DummyConfig: + bias_dropout_fusion = False + + class DummyLayer: + training = True + hidden_dropout = 0.1 + is_moe_layer = True + config = DummyConfig() + + def __init__(self): + self.scaler_called = False + + def _scale_dense_residual_branch_output( + self, + output_with_bias, + *, + branch_name, + using_fused_tp_inference_kernel, + apply_depth_hook, + ): + self.scaler_called = True + assert branch_name == 'mlp' + assert using_fused_tp_inference_kernel is False + assert apply_depth_hook is False + return output_with_bias + + def bias_dropout_add_exec_handler(self): + return DummyCtx() + + def mlp_bda(self, training, bias_dropout_fusion): + assert training is True + assert bias_dropout_fusion is False + + def apply(output_with_bias, residual, hidden_dropout): + output, bias = output_with_bias + assert bias is None + return output + residual + hidden_dropout + + return apply + + layer = DummyLayer() + output = torch.ones(3, 4) + residual = torch.ones(3, 4) + hidden_states = _apply_mlp_bda_with_scaling(layer, output, residual) + + assert layer.scaler_called is True + assert torch.equal(hidden_states, output + residual + layer.hidden_dropout) + def test_mup_base_hidden_size_must_be_positive(self): """mup_base_hidden_size must be positive.""" with pytest.raises(AssertionError) as exc_info: @@ -154,6 +1619,19 @@ def test_standard_attention_scale_power_05(self): expected_scale = 1.0 / math.sqrt(kv_channels) # 1/8 = 0.125 assert abs(config.softmax_scale - expected_scale) < 1e-6 + def test_depth_mup_inherits_width_mup_attention_scaling(self): + config = TransformerConfig( + hidden_size=512, + num_layers=8, + num_attention_heads=8, + scaling_recipe='depth_mup', + scaling_base_hidden_size=128, + scaling_base_num_layers=4, + scaling_base_head_dim=64, + ) + expected_scale = (config.scaling_base_head_dim**0.5) / config.kv_channels + assert config.softmax_scale == expected_scale + def test_attention_scale_not_set_when_disabled(self): """softmax_scale should not be auto-set when use_mup=False.""" config = TransformerConfig( @@ -169,7 +1647,7 @@ class TestMuPWarnings: def test_mup_warns_with_custom_init_method(self): """Warn when MuP is enabled and init_method is user-provided.""" - with pytest.warns(UserWarning, match="use_mup is enabled"): + with pytest.warns(UserWarning, match="scaling recipe 'mup' is enabled"): TransformerConfig( hidden_size=512, num_layers=4, @@ -181,7 +1659,7 @@ def test_mup_warns_with_custom_init_method(self): def test_mup_warns_with_custom_output_layer_init_method(self): """Warn when MuP is enabled and output_layer_init_method is user-provided.""" - with pytest.warns(UserWarning, match="use_mup is enabled"): + with pytest.warns(UserWarning, match="scaling recipe 'mup' is enabled"): TransformerConfig( hidden_size=512, num_layers=4, @@ -191,6 +1669,20 @@ def test_mup_warns_with_custom_output_layer_init_method(self): output_layer_init_method=init_method_normal(0.01), ) + def test_mup_warns_with_custom_embedding_init_method(self): + """Warn when MuP is enabled and embedding_init_method is user-provided.""" + with pytest.warns( + UserWarning, match="scaling recipe 'mup' is enabled, but custom embedding_init_method" + ): + TransformerConfig( + hidden_size=512, + num_layers=4, + num_attention_heads=8, + use_mup=True, + mup_base_hidden_size=128, + embedding_init_method=init_method_normal(0.01), + ) + class TestMuPLRScaling: """Tests for MuP learning rate and Adam epsilon scaling.""" @@ -351,6 +1843,415 @@ def test_mup_with_decoupled_lr_scales_hidden_only_for_lr(self): assert shared_output_override['min_lr'] == pytest.approx(2e-6) assert 'eps' not in shared_output_override + def test_scaling_policy_matches_legacy_optimizer_overrides_for_width_only(self): + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=8, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + + policy_overrides = get_scaling_config_overrides(optimizer_config, scaling_policy) + legacy_overrides = get_mup_config_overrides( + optimizer_config, model_config.mup_width_mult, optimizer_type='adam' + ) + + hidden_param = torch.nn.Parameter(torch.zeros(10, 10)) + bias_param = torch.nn.Parameter(torch.zeros(10)) + embedding_param = torch.nn.Parameter(torch.zeros(10, 10)) + embedding_param.is_embedding_parameter = True + output_param = torch.nn.Parameter(torch.zeros(10, 10)) + output_param.is_embedding_parameter = True + output_param.is_embedding_or_output_parameter = True + + sample_params = [ + (hidden_param, 'decoder.layers.0.self_attention.linear_qkv.weight'), + (bias_param, 'decoder.layers.0.self_attention.linear_qkv.bias'), + (embedding_param, 'embedding.word_embeddings.weight'), + (output_param, 'output_layer.weight'), + ] + + for param, name in sample_params: + assert _combined_override_for_param( + policy_overrides, param, name + ) == _combined_override_for_param(legacy_overrides, param, name) + + def test_scaling_policy_applies_hidden_lr_depth_power_for_adam(self): + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + scaling_hidden_lr_depth_power=-0.5, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + overrides = get_scaling_config_overrides(optimizer_config, scaling_policy) + + hidden_param = torch.nn.Parameter(torch.zeros(10, 10)) + bias_param = torch.nn.Parameter(torch.zeros(10)) + hidden_override = _combined_override_for_param( + overrides, hidden_param, 'decoder.layers.0.self_attention.linear_qkv.weight' + ) + bias_override = _combined_override_for_param( + overrides, bias_param, 'decoder.layers.0.self_attention.linear_qkv.bias' + ) + + expected_lr_mult = (1.0 / model_config.mup_width_mult) * ( + scaling_policy.context.depth_mult**model_config.scaling_hidden_lr_depth_power + ) + assert hidden_override['max_lr'] == pytest.approx(optimizer_config.lr * expected_lr_mult) + assert hidden_override['min_lr'] == pytest.approx( + optimizer_config.min_lr * expected_lr_mult + ) + assert hidden_override['eps'] == pytest.approx( + optimizer_config.adam_eps / model_config.mup_width_mult + ) + assert 'max_lr' not in bias_override + assert 'min_lr' not in bias_override + assert 'eps' not in bias_override + + def test_depth_mup_adamw_matches_paper_table_5_role_specific_overrides(self): + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, weight_decay=0.1) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + standard_overrides = get_standard_config_overrides( + optimizer_config, scaling_policy=scaling_policy + ) + scaling_overrides = get_scaling_config_overrides(optimizer_config, scaling_policy) + overrides = {**standard_overrides, **scaling_overrides} + + hidden_param = torch.nn.Parameter(torch.zeros(10, 10)) + hidden_vector_param = torch.nn.Parameter(torch.zeros(10)) + embedding_param = torch.nn.Parameter(torch.zeros(10, 10)) + embedding_param.is_embedding_parameter = True + embedding_param.is_embedding_or_output_parameter = True + output_param = torch.nn.Parameter(torch.zeros(10, 10)) + output_param.is_embedding_parameter = True + output_param.is_embedding_or_output_parameter = True + hidden_override = _combined_override_for_param( + overrides, hidden_param, 'decoder.layers.0.self_attention.linear_qkv.weight' + ) + hidden_vector_override = _combined_override_for_param( + overrides, hidden_vector_param, 'decoder.layers.0.mlp.linear_fc1.bias' + ) + embedding_override = _combined_override_for_param( + overrides, embedding_param, 'embedding.word_embeddings.weight' + ) + output_override = _combined_override_for_param( + overrides, output_param, 'output_layer.weight' + ) + + expected_lr_mult = 1.0 / scaling_policy.context.width_mult + expected_eps_mult = (1.0 / scaling_policy.context.width_mult) * ( + 1.0 / scaling_policy.context.depth_mult + ) + assert scaling_policy.hidden_eps_depth_power == pytest.approx(-1.0) + assert hidden_override['max_lr'] == pytest.approx(optimizer_config.lr * expected_lr_mult) + assert hidden_override['min_lr'] == pytest.approx( + optimizer_config.min_lr * expected_lr_mult + ) + assert hidden_override['eps'] == pytest.approx( + optimizer_config.adam_eps * expected_eps_mult + ) + assert hidden_override['wd_mult'] == pytest.approx(scaling_policy.context.width_mult) + assert 'max_lr' not in hidden_vector_override + assert 'min_lr' not in hidden_vector_override + assert hidden_vector_override['eps'] == pytest.approx( + optimizer_config.adam_eps * expected_eps_mult + ) + assert 'wd_mult' not in hidden_vector_override + assert 'max_lr' not in embedding_override + assert 'min_lr' not in embedding_override + assert embedding_override['eps'] == pytest.approx( + optimizer_config.adam_eps / scaling_policy.context.width_mult + ) + assert 'wd_mult' not in embedding_override + assert 'max_lr' not in output_override + assert 'min_lr' not in output_override + assert output_override['eps'] == pytest.approx( + optimizer_config.adam_eps / scaling_policy.context.width_mult + ) + assert 'wd_mult' not in output_override + + def test_depth_mup_applies_spectral_adam_policy_defaults(self): + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + + assert scaling_policy.hidden_lr_multiplier == pytest.approx(1.0 / 4.0) + assert scaling_policy.hidden_eps_depth_power == pytest.approx(-1.0) + assert scaling_policy.hidden_eps_multiplier == pytest.approx(1.0 / 16.0) + assert scaling_policy.hidden_vector_eps_multiplier == pytest.approx(1.0 / 16.0) + assert scaling_policy.embedding_class_eps_multiplier == pytest.approx(1.0 / 4.0) + assert scaling_policy.hidden_matrix_wd_multiplier == pytest.approx(4.0) + + def test_depth_mup_residual_multiplier_exact_depth_factors(self): + base_depth_config = TransformerConfig( + hidden_size=512, + num_layers=12, + num_attention_heads=8, + scaling_recipe='depth_mup', + scaling_base_hidden_size=512, + scaling_base_num_layers=12, + ) + double_depth_config = TransformerConfig( + hidden_size=512, + num_layers=24, + num_attention_heads=8, + scaling_recipe='depth_mup', + scaling_base_hidden_size=512, + scaling_base_num_layers=12, + ) + + assert build_resolved_model_policy( + base_depth_config + ).residual_branch_multiplier == pytest.approx(1.0) + assert build_resolved_model_policy( + double_depth_config + ).residual_branch_multiplier == pytest.approx(0.5) + + def test_depth_mup_hidden_depth_factor_excludes_embedding_class_eps(self): + config = TransformerConfig( + hidden_size=512, + num_layers=24, + num_attention_heads=8, + scaling_recipe='depth_mup', + scaling_base_hidden_size=512, + scaling_base_num_layers=12, + ) + policy = build_resolved_training_policy(config, optimizer_type='adam') + + assert policy.hidden_eps_multiplier == pytest.approx(0.5) + assert policy.hidden_vector_eps_multiplier == pytest.approx(0.5) + assert policy.embedding_class_eps_multiplier == pytest.approx(1.0) + + def test_depth_mup_weight_decay_treats_hidden_bias_and_norm_vectors_explicitly(self): + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, weight_decay=0.1) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + standard_overrides = get_standard_config_overrides( + optimizer_config, scaling_policy=scaling_policy + ) + + bias_param = torch.nn.Parameter(torch.zeros(10)) + norm_scale_param = torch.nn.Parameter(torch.zeros(10)) + unknown_vector_param = torch.nn.Parameter(torch.zeros(10)) + assert ( + _combined_override_for_param( + standard_overrides, bias_param, 'decoder.layers.0.mlp.linear_fc1.bias' + ) + == {} + ) + assert _combined_override_for_param( + standard_overrides, norm_scale_param, 'decoder.layers.0.input_layernorm.weight' + ) == {'wd_mult': 0.0} + assert _combined_override_for_param( + standard_overrides, unknown_vector_param, 'decoder.layers.0.some_scalar' + ) == {'wd_mult': 0.0} + + def test_depth_mup_respects_apply_wd_to_qk_layernorm_for_qk_norm_vectors(self): + optimizer_config = OptimizerConfig( + lr=1e-3, min_lr=1e-5, weight_decay=0.1, apply_wd_to_qk_layernorm=True + ) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + standard_overrides = get_standard_config_overrides( + optimizer_config, scaling_policy=scaling_policy + ) + + q_norm_param = torch.nn.Parameter(torch.zeros(10)) + ordinary_norm_param = torch.nn.Parameter(torch.zeros(10)) + assert ( + _combined_override_for_param( + standard_overrides, + q_norm_param, + 'decoder.layers.0.self_attention.q_layernorm.weight', + ) + == {} + ) + assert _combined_override_for_param( + standard_overrides, ordinary_norm_param, 'decoder.layers.0.input_layernorm.weight' + ) == {'wd_mult': 0.0} + + def test_depth_mup_rejects_coupled_adam_nonzero_weight_decay(self): + optimizer_config = OptimizerConfig( + lr=1e-3, min_lr=1e-5, weight_decay=0.1, decoupled_weight_decay=False + ) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + + with pytest.raises(ValueError, match='requires decoupled_weight_decay=True'): + get_scaling_config_overrides(optimizer_config, scaling_policy) + + def test_depth_mup_allows_coupled_adam_when_weight_decay_is_zero(self): + optimizer_config = OptimizerConfig( + lr=1e-3, min_lr=1e-5, weight_decay=0.0, decoupled_weight_decay=False + ) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + + assert get_scaling_config_overrides(optimizer_config, scaling_policy) + + def test_depth_mup_cli_validation_rejects_coupled_adam_nonzero_weight_decay(self): + args = SimpleNamespace( + scaling_recipe='depth_mup', + optimizer='adam', + weight_decay=0.1, + decoupled_weight_decay=False, + ) + + with pytest.raises(ValueError, match='requires decoupled_weight_decay=True'): + validate_depth_mup_optimizer_support(args) + + def test_depth_mup_cli_validation_treats_missing_decoupled_weight_decay_as_default_adamw(self): + args = SimpleNamespace(scaling_recipe='depth_mup', optimizer='adam', weight_decay=0.1) + + validate_depth_mup_optimizer_support(args) + + def test_depth_mup_with_decoupled_lr_preserves_embedding_output_lr_and_scales_eps(self): + optimizer_config = OptimizerConfig( + lr=1e-3, min_lr=1e-5, decoupled_lr=2e-4, decoupled_min_lr=2e-6, weight_decay=0.1 + ) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='adam') + standard_overrides = get_standard_config_overrides( + optimizer_config, scaling_policy=scaling_policy + ) + scaling_overrides = get_scaling_config_overrides(optimizer_config, scaling_policy) + overrides = {**standard_overrides, **scaling_overrides} + + output_param = torch.nn.Parameter(torch.zeros(10, 10)) + output_param.is_embedding_parameter = True + output_param.is_embedding_or_output_parameter = True + output_override = _combined_override_for_param( + overrides, output_param, 'output_layer.weight' + ) + + assert output_override['max_lr'] == pytest.approx(2e-4) + assert output_override['min_lr'] == pytest.approx(2e-6) + assert output_override['eps'] == pytest.approx( + optimizer_config.adam_eps / scaling_policy.context.width_mult + ) + + def test_depth_mup_rejects_sgd(self): + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + + with pytest.raises(ValueError, match="supports optimizer='adam' only"): + build_resolved_training_policy(model_config, optimizer_type='sgd') + + def test_depth_mup_rejects_muon(self): + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='depth_mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + ) + + with pytest.raises(ValueError, match="supports optimizer='adam' only"): + build_resolved_training_policy(model_config, optimizer_type='muon') + + def test_scaling_policy_applies_hidden_lr_depth_power_for_sgd(self): + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) + model_config = TransformerConfig( + hidden_size=1024, + num_layers=16, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_num_layers=4, + scaling_hidden_lr_depth_power=-0.5, + ) + scaling_policy = build_resolved_training_policy(model_config, optimizer_type='sgd') + overrides = get_scaling_config_overrides(optimizer_config, scaling_policy) + + hidden_param = torch.nn.Parameter(torch.zeros(10, 10)) + bias_param = torch.nn.Parameter(torch.zeros(10)) + hidden_override = _combined_override_for_param( + overrides, hidden_param, 'decoder.layers.0.mlp.linear_fc1.weight' + ) + bias_override = _combined_override_for_param( + overrides, bias_param, 'decoder.layers.0.mlp.linear_fc1.bias' + ) + + expected_hidden_lr_mult = scaling_policy.context.depth_mult ** ( + model_config.scaling_hidden_lr_depth_power + ) + assert hidden_override['max_lr'] == pytest.approx( + optimizer_config.lr * expected_hidden_lr_mult + ) + assert hidden_override['min_lr'] == pytest.approx( + optimizer_config.min_lr * expected_hidden_lr_mult + ) + assert bias_override['max_lr'] == pytest.approx( + optimizer_config.lr * model_config.mup_width_mult + ) + assert bias_override['min_lr'] == pytest.approx( + optimizer_config.min_lr * model_config.mup_width_mult + ) + class TestMuPConfigIntegration: """Integration tests for MuP config with init methods.""" @@ -380,6 +2281,21 @@ def test_mup_output_layer_init(self): class TestMuPOptimizerTypeHandling: """Tests for MuP optimizer-specific override behavior.""" + def test_adam_with_decoupled_weight_decay_uses_standard_optimizer_path(self): + param = torch.nn.Parameter(torch.ones(1)) + optimizer_config = OptimizerConfig( + optimizer='adam', lr=1e-3, weight_decay=0.1, decoupled_weight_decay=True + ) + + raw_optimizer, _ = _get_megatron_optimizer_based_on_param_groups( + config=optimizer_config, + model_chunks=[], + param_groups=[{'params': [param]}], + skip_megatron_wrapping=True, + ) + + assert isinstance(raw_optimizer, optimizer_module.Adam) + def test_sgd_scales_vector_like_lr_only(self): """SGD scales vector-like params by width_mult; hidden params keep base LR.""" optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) @@ -534,6 +2450,8 @@ def test_muon_excludes_muon_managed_matrices_from_mup_overrides(self, optimizer_ muon_managed_param.is_embedding_or_output_parameter = False output_param = torch.nn.Parameter(torch.zeros(10, 10)) output_param.is_embedding_or_output_parameter = True + output_param.is_output_parameter = True + output_param.is_embedding_parameter = True bias_param = torch.nn.Parameter(torch.zeros(10)) muon_managed_matches = [ @@ -563,9 +2481,9 @@ def test_muon_excludes_muon_managed_matrices_from_mup_overrides(self, optimizer_ assert 'min_lr' not in muon_managed_override assert 'eps' not in muon_managed_override - # Output params remain in the MuP override path (handled by chained Adam optimizer). - assert output_override['max_lr'] == pytest.approx(1e-3 / width_mult) - assert output_override['min_lr'] == pytest.approx(1e-5 / width_mult) + # Output params remain on the scalar optimizer path at base LR/eps. + assert 'max_lr' not in output_override + assert 'min_lr' not in output_override assert 'eps' not in output_override # Vector-like params stay unscaled.