diff --git a/docs/api-guide/internal/index.md b/docs/api-guide/internal/index.md index 312081ce70b..0f71bac266c 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_policy_infrastructure ``` diff --git a/docs/api-guide/internal/scaling_policy_infrastructure.md b/docs/api-guide/internal/scaling_policy_infrastructure.md new file mode 100644 index 00000000000..3a5f06d70a2 --- /dev/null +++ b/docs/api-guide/internal/scaling_policy_infrastructure.md @@ -0,0 +1,65 @@ + + +# Scaling Policy Infrastructure + +This internal policy layer centralizes Megatron's parameterization hooks behind a +scaling context. + +The current public recipes are `none` and `mup`. The policy resolver also accepts +legacy MuP aliases, then syncs them to the canonical recipe fields so model, +optimizer, YAML, and checkpoint paths see the same effective scaling context. +Standard Megatron behavior is represented as the identity policy, so code paths +can call the same hooks whether or not MuP is active. + +## Model Policy + +Model code should route scaling-sensitive decisions through the model scaling +policy instead of reading `use_mup` at each call site. The policy currently +covers: + +- hidden-weight initialization; +- output-projection initialization; +- attention softmax scale; +- embedding activation scaling; +- output logit scaling; +- residual branch output hooks. + +For non-MuP configs, every hook returns the current Megatron default. + +## Training Policy + +Optimizer code should route per-parameter hyperparameter multipliers through the +training scaling policy. The policy currently preserves the existing MuP rules: + +- Adam-family hidden matrix parameters use `lr / mup_width_mult`; +- Adam-family hidden matrix parameters use `eps / mup_width_mult`; +- SGD vector-like parameters use `lr * mup_width_mult`; +- Muon-managed matrices stay on Muon scaling rather than Adam-style MuP LR + overrides. + +The public compatibility function `get_mup_config_overrides` remains available +and delegates to the policy implementation. + +## Parameter Metadata + +Model construction may attach explicit parameterization metadata to parameters. +Optimizer grouping should prefer this metadata and keep existing name/shape +fallbacks only for compatibility with unannotated parameters. + +FSDP and other parameter-rewriting paths must preserve the metadata attributes so +optimizer grouping remains stable after wrapping or sharding. + +## Checkpoint Resume + +Distributed optimizer resume uses a shared helper for the existing stable +parameter-group identity: `wd_mult`, `lr_mult`, `is_expert_parallel`, and +`is_decoupled_lr`. The helper tolerates NeMo-style `pre_` field names and missing +legacy fields without adding newer optional fields such as `eps`, `max_lr`, +`min_lr`, or per-group `optimizer` to the stable resume identity. 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..7f1fc045d30 --- /dev/null +++ b/docs/user-guide/scaling-recipes.md @@ -0,0 +1,76 @@ + + +# Scaling Recipes + +Scaling recipes choose the parameterization used to transfer hyperparameters +between model sizes. The canonical flag is `--scaling-recipe`. + +Megatron currently exposes two recipes: + +| Recipe | Behavior | +| --- | --- | +| `none` | Standard Megatron parameterization. This is the default. | +| `mup` | Width MuP for hidden-size transfer. | + +## Standard Parameterization + +Use `--scaling-recipe none`, or omit `--scaling-recipe`, to keep the standard +parameterization. Scaling-specific fields such as `--scaling-base-hidden-size` +are rejected unless a scaling recipe is selected. + +## Width MuP + +Use `--scaling-recipe mup` when transferring hyperparameters from a base width to +a target width. + +```bash +--scaling-recipe mup \ +--scaling-base-hidden-size 1024 \ +--scaling-base-head-dim 64 +``` + +For MuP, Megatron derives the width multiplier internally: + +```text +width_mult = hidden_size / scaling_base_hidden_size +``` + +This derived value controls MuP initialization, attention scale, output-logit +scale, and optimizer multipliers. `--mup-width-mult` is no longer an independent +input. If it is provided on the CLI for compatibility, it must match the derived +value. + +## Legacy MuP Flags + +The following flags are accepted for checkpoint and script compatibility, but are +deprecated as user-facing inputs: + +| Deprecated flag | Canonical replacement | +| --- | --- | +| `--use-mup` | `--scaling-recipe mup` | +| `--mup-base-hidden-size` | `--scaling-base-hidden-size` | +| `--mup-base-head-dim` | `--scaling-base-head-dim` | +| `--mup-width-mult` | derived from `hidden_size / scaling_base_hidden_size` | + +`--mup-embedding-mult`, `--mup-output-mult`, and `--mup-attn-scale-power` remain +MuP-specific tuning knobs. When `--mup-output-mult` is left at `1.0`, Megatron +sets it to `1 / width_mult` for non-base widths. + +## Checkpoints and YAML + +Megatron stores and compares the resolved scaling recipe, not just the raw flag +spelling. A checkpoint created with legacy MuP aliases is compatible with the +canonical spelling when both resolve to the same effective recipe and base size. + +YAML configs use the same effective resolution rules as CLI configs. Existing +YAML files that omit the new canonical scaling fields default to +`--scaling-recipe none`. For compatibility with full legacy YAML files that +materialized old defaults, `mup_width_mult: 1.0` is treated as an omitted default; +non-`1.0` YAML values are still validated against the derived width multiplier. 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 266da6b74c4..777eb188f41 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 @@ -2891,6 +2891,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/models/T5/t5_model.py b/megatron/core/models/T5/t5_model.py index b2feb974643..a85854cbca7 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 build_model_scaling_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 @@ -56,10 +57,10 @@ def __init__( 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=build_model_scaling_policy(config).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, diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 3fd1e01f4a1..75875f9f7d3 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -135,10 +135,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, diff --git a/megatron/core/models/common/embeddings/language_model_embedding.py b/megatron/core/models/common/embeddings/language_model_embedding.py index 7e49ec6c02d..ffdcdb6c651 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_model_scaling_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 @@ -128,8 +129,7 @@ def forward(self, input_ids: Tensor, position_ids: Tensor, tokentype_ids: int = 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 = build_model_scaling_policy(self.config).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 34e3f6b1ba4..2650f480e37 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_model_scaling_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_model_scaling_policy(config) self._set_attention_backend() if pg_collection is None: pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -204,27 +212,55 @@ 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 # 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'): - for param in self.embedding.parameters(): - param.is_embedding_parameter = True if ( - self.config.use_mup + self.model_scaling_policy.enabled + and (self.pre_process or mtp_process) + and hasattr(self, 'embedding') + ): + self.model_scaling_policy.mark_embedding_class_parameters(self.embedding.parameters()) + if ( + 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 +300,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 + set_parameterization_metadata( + weight, role=ROLE_SHARED_EMBEDDING_OUTPUT, shared_group='lm_embedding_output' + ) # Keep optimizer grouping consistent for tied embedding/output copies. - if self.config.use_mup: - weight.is_embedding_parameter = True + 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 @@ -312,11 +353,7 @@ def _scale_logits(self, logits: Tensor) -> Tensor: Tensor: Scaled logits if MuP is enabled and mup_output_mult != 1.0, 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 build_model_scaling_policy(self.config).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/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index 9f8d9da4a10..48d58cf2fe5 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -252,10 +252,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, @@ -676,7 +676,7 @@ 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.config.use_mup 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 511b24673b0..f0df0f60a6e 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -296,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, @@ -543,7 +543,7 @@ 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.config.use_mup else None), ) sequence_parallel_override = False if ( diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 1598f6ed95c..070b3bd7d33 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -54,6 +54,15 @@ combine_param_group_overrides, param_group_override_to_tuple, ) +from megatron.core.parameterization import ( + TrainingScalingPolicy, + build_legacy_mup_training_policy, + is_embedding_class_parameter, + is_embedding_or_output_parameter, + is_hidden_matrix_parameter, + is_muon_managed_matrix_parameter, + is_vector_like_parameter, +) from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.fsdp_dtensor_checkpoint import get_global_unique_param_name @@ -73,7 +82,6 @@ Float16OptimizerWithFloat16Params, FP32Optimizer, MegatronOptimizer, - param_group_identifier_keys, ) # Subclass aliases kept for backward compatibility; all are OptimizerConfig. @@ -130,41 +138,23 @@ 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. - - 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 - - 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. - - With decoupled_lr enabled, embedding/output params continue using decoupled LR - and MuP will not override those explicit decoupled values. + """Compatibility wrapper for the existing MuP optimizer override surface.""" + 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) - 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. - Returns: - Dict[ParamKey, ParamGroupOverride]: MuP optimizer overrides. +def get_scaling_config_overrides( + config: OptimizerConfig, scaling_policy: TrainingScalingPolicy +) -> Dict[ParamKey, ParamGroupOverride]: + """Get optimizer overrides from an internal training scaling policy. + + This compatibility path keeps the public MuP surface unchanged and factors the + existing per-parameter MuP LR/epsilon behavior through this policy seam. """ - 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 {} decoupled_lr_enabled = config.decoupled_lr is not None if decoupled_lr_enabled: @@ -173,11 +163,11 @@ def get_mup_config_overrides( "absolute LR for embedding+output params, and MuP LR scaling will not " "override those parameters." ) - if is_adam_optimizer: + if scaling_policy.is_adam_optimizer: message += " MuP Adam epsilon scaling remains applied to hidden matrix-like parameters." 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,66 +179,43 @@ def get_mup_config_overrides( "Muon-managed matrices with MuP.", ) - if mup_width_mult == 1.0: - # No scaling needed when width_mult is 1 + if scaling_policy.context.width_mult == 1.0: 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 is_managed_by_layer_wise_optimizer(param) - - 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): + 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_embedding_class_parameter(param, param_name): + return False if is_vector_like_parameter(param, param_name): 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 - # MuP Appendix B.3: eps scales with fan_in when non-negligible. - # This implementation follows the common denominator form: sqrt(v) + eps. return True mup_overrides: Dict[ParamKey, ParamGroupOverride] = {} - if is_sgd_optimizer: - vector_like_lr_mult = mup_width_mult + if scaling_policy.is_sgd_optimizer: vector_like_lr_override: ParamGroupOverride = {} if base_lr is not None: - vector_like_lr_override["max_lr"] = base_lr * vector_like_lr_mult + vector_like_lr_override["max_lr"] = base_lr * scaling_policy.hidden_vector_lr_multiplier if base_min_lr is not None: - vector_like_lr_override["min_lr"] = base_min_lr * vector_like_lr_mult + vector_like_lr_override["min_lr"] = ( + base_min_lr * scaling_policy.hidden_vector_lr_multiplier + ) if vector_like_lr_override: vector_like_predicate = ParamWithNamePredicate( @@ -263,18 +230,18 @@ def should_scale_eps_with_mup(param: torch.nn.Parameter, param_name: str) -> boo lr_override: ParamGroupOverride = {} if base_lr is not None: - lr_override["max_lr"] = base_lr * hidden_lr_mult + lr_override["max_lr"] = base_lr * scaling_policy.hidden_lr_multiplier if base_min_lr is not None: - lr_override["min_lr"] = base_min_lr * hidden_lr_mult + lr_override["min_lr"] = base_min_lr * scaling_policy.hidden_lr_multiplier 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: + eps_override["eps"] = config.adam_eps * scaling_policy.hidden_eps_multiplier 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 diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index b388161a610..622603d0151 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 @@ -888,21 +892,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/optimizer.py b/megatron/core/optimizer/optimizer.py index ddc3dd8620e..7b98be7b79c 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -95,6 +95,25 @@ def _multi_tensor_copy_this_to_that( param_group_identifier_keys = ('wd_mult', 'lr_mult', 'is_expert_parallel', 'is_decoupled_lr') +param_group_identifier_defaults = { + 'wd_mult': 1.0, + 'lr_mult': 1.0, + 'is_expert_parallel': False, + 'is_decoupled_lr': False, +} + + +def get_param_group_identifier_tuple(param_group: Dict) -> tuple: + """Return the stable legacy 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(param_group_identifier_defaults[key]) + return tuple(values) class MegatronOptimizer(ABC): @@ -426,22 +445,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/parameterization/__init__.py b/megatron/core/parameterization/__init__.py new file mode 100644 index 00000000000..f3a31914eed --- /dev/null +++ b/megatron/core/parameterization/__init__.py @@ -0,0 +1,67 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from .model_policy import ModelScalingPolicy, build_model_scaling_policy +from .roles import ( + IS_OUTPUT_PARAMETER_ATTR, + PARAMETERIZATION_ROLE_ATTR, + PARAMETERIZATION_SHARED_GROUP_ATTR, + PARAMETERIZATION_TAGS_ATTR, + ROLE_EMBEDDING, + ROLE_HIDDEN_MATRIX, + ROLE_HIDDEN_VECTOR, + ROLE_MUON_MANAGED_MATRIX, + ROLE_OUTPUT, + ROLE_SHARED_EMBEDDING_OUTPUT, + get_parameterization_role, + is_embedding_class_parameter, + is_embedding_or_output_parameter, + is_hidden_matrix_parameter, + is_muon_managed_matrix_parameter, + is_output_parameter, + is_vector_like_parameter, + set_parameterization_metadata, +) +from .spec import ( + SCALING_RECIPE_MUP, + SCALING_RECIPE_NONE, + SCALING_RECIPE_VALUES, + ScalingContext, + ScalingUserConfig, + build_scaling_context, + build_scaling_user_config, + sync_legacy_mup_fields, +) +from .training_policy import TrainingScalingPolicy, build_legacy_mup_training_policy + +__all__ = [ + 'IS_OUTPUT_PARAMETER_ATTR', + 'PARAMETERIZATION_ROLE_ATTR', + 'PARAMETERIZATION_SHARED_GROUP_ATTR', + 'PARAMETERIZATION_TAGS_ATTR', + 'ROLE_EMBEDDING', + 'ROLE_HIDDEN_MATRIX', + 'ROLE_HIDDEN_VECTOR', + 'ROLE_MUON_MANAGED_MATRIX', + 'ROLE_OUTPUT', + 'ROLE_SHARED_EMBEDDING_OUTPUT', + 'ModelScalingPolicy', + 'SCALING_RECIPE_MUP', + 'SCALING_RECIPE_NONE', + 'SCALING_RECIPE_VALUES', + 'ScalingContext', + 'ScalingUserConfig', + 'TrainingScalingPolicy', + 'build_legacy_mup_training_policy', + 'build_model_scaling_policy', + 'build_scaling_context', + 'build_scaling_user_config', + 'get_parameterization_role', + 'is_embedding_class_parameter', + 'is_embedding_or_output_parameter', + 'is_hidden_matrix_parameter', + 'is_muon_managed_matrix_parameter', + 'is_output_parameter', + 'is_vector_like_parameter', + 'set_parameterization_metadata', + 'sync_legacy_mup_fields', +] diff --git a/megatron/core/parameterization/model_policy.py b/megatron/core/parameterization/model_policy.py new file mode 100644 index 00000000000..a3a84dee519 --- /dev/null +++ b/megatron/core/parameterization/model_policy.py @@ -0,0 +1,116 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +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 ScalingContext, build_scaling_context + + +@dataclass(frozen=True) +class ModelScalingPolicy: + """Model-side policy for existing Megatron scaling behavior.""" + + context: ScalingContext + + @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 1.0 + + @property + def dense_block_out_proj_init_multiplier(self) -> float: + return 1.0 + + 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.base_head_dim is None else self.context.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, + ): + del init_method_std, num_layers, is_hybrid_model + if output_layer_init_method_is_user_provided: + return default_init_method + return default_init_method + + 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]: + return output_with_bias + + +def build_model_scaling_policy(config) -> ModelScalingPolicy: + return ModelScalingPolicy(build_scaling_context(config)) diff --git a/megatron/core/parameterization/roles.py b/megatron/core/parameterization/roles.py new file mode 100644 index 00000000000..3d5513f9f46 --- /dev/null +++ b/megatron/core/parameterization/roles.py @@ -0,0 +1,74 @@ +# 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_HIDDEN_MATRIX = 'hidden_matrix' +ROLE_HIDDEN_VECTOR = 'hidden_vector' +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)) + + +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 + 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 is_hidden_matrix_parameter(param: Any, param_name: Optional[str] = None) -> bool: + if is_embedding_class_parameter(param, param_name): + return False + return param.dim() > 1 + + +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..9f730f2cd6e --- /dev/null +++ b/megatron/core/parameterization/spec.py @@ -0,0 +1,206 @@ +# 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_VALUES = (SCALING_RECIPE_NONE, SCALING_RECIPE_MUP) + + +@dataclass(frozen=True) +class ScalingUserConfig: + recipe: Optional[Literal['none', 'mup']] = None + base_hidden_size: Optional[int] = None + base_head_dim: Optional[float] = None + use_mup_alias: bool = False + mup_width_mult: Optional[float] = None + mup_width_mult_explicit: bool = False + 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 ScalingContext: + """Internal scaling context for standard and width-MuP parameterization.""" + + recipe: Literal['none', 'mup'] + width_mult: float = 1.0 + embedding_mult: float = 1.0 + output_mult: float = 1.0 + base_hidden_size: Optional[int] = None + base_head_dim: Optional[float] = None + attention_scale_power: float = 1.0 + + @property + def enabled(self) -> bool: + return self.recipe != SCALING_RECIPE_NONE + + @property + def uses_width_mup(self) -> bool: + return self.recipe == SCALING_RECIPE_MUP + + @property + def use_mup(self) -> bool: + return self.uses_width_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) -> list[str]: + candidates: dict[str, object] = { + 'scaling_base_hidden_size': user_config.base_hidden_size, + 'scaling_base_head_dim': user_config.base_head_dim, + 'mup_base_hidden_size': user_config.mup_base_hidden_size, + '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_explicit: + candidates['mup_width_mult'] = user_config.mup_width_mult + return [name for name, value in candidates.items() if value is not None] + + +def build_scaling_user_config(config: Any) -> ScalingUserConfig: + raw_mup_width_mult = getattr(config, 'mup_width_mult', None) + marker_present = hasattr(config, '_mup_width_mult_explicit') + mup_width_mult_explicit = bool(getattr(config, '_mup_width_mult_explicit', False)) + if ( + not marker_present + and not mup_width_mult_explicit + and raw_mup_width_mult not in (None, 1.0) + ): + # Direct TransformerConfig construction has no argparse provenance marker. + mup_width_mult_explicit = True + + return ScalingUserConfig( + recipe=getattr(config, 'scaling_recipe', None), + base_hidden_size=getattr(config, 'scaling_base_hidden_size', None), + base_head_dim=getattr(config, 'scaling_base_head_dim', None), + use_mup_alias=bool(getattr(config, 'use_mup', False)), + mup_width_mult=raw_mup_width_mult if mup_width_mult_explicit else None, + mup_width_mult_explicit=mup_width_mult_explicit, + 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 build_scaling_context(config: Any) -> ScalingContext: + user_config = build_scaling_user_config(config) + 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 and 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) + if non_default_fields: + raise ValueError( + "Scaling overrides require a non-'none' scaling recipe. Non-default fields: " + + ", ".join(non_default_fields) + ) + return ScalingContext(recipe=SCALING_RECIPE_NONE) + + 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.') + if base_head_dim is not None and base_head_dim <= 0: + raise AssertionError('--scaling-base-head-dim must be positive.') + + width_mult = config.hidden_size / base_hidden_size + if ( + user_config.mup_width_mult_explicit + and user_config.mup_width_mult is not None + 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}." + ) + + output_mult = user_config.mup_output_mult + if output_mult == 1.0 and width_mult != 1.0: + output_mult = 1.0 / width_mult + + return ScalingContext( + recipe=SCALING_RECIPE_MUP, + width_mult=width_mult, + embedding_mult=user_config.mup_embedding_mult, + output_mult=output_mult, + base_hidden_size=base_hidden_size, + base_head_dim=base_head_dim, + attention_scale_power=user_config.mup_attn_scale_power, + ) + + +def sync_legacy_mup_fields(config: Any, context: ScalingContext) -> None: + config.scaling_recipe = context.recipe + config.use_mup = context.recipe == SCALING_RECIPE_MUP + config.mup_width_mult = context.width_mult + config._mup_width_mult_explicit = False + config.mup_embedding_mult = context.embedding_mult + config.mup_output_mult = context.output_mult + config.mup_attn_scale_power = context.attention_scale_power + if context.recipe == SCALING_RECIPE_NONE: + config.scaling_base_hidden_size = None + config.scaling_base_head_dim = None + config.mup_base_hidden_size = None + config.mup_base_head_dim = None + return + + config.scaling_base_hidden_size = context.base_hidden_size + config.scaling_base_head_dim = context.base_head_dim + config.mup_base_hidden_size = context.base_hidden_size + config.mup_base_head_dim = context.base_head_dim diff --git a/megatron/core/parameterization/training_policy.py b/megatron/core/parameterization/training_policy.py new file mode 100644 index 00000000000..01e98851084 --- /dev/null +++ b/megatron/core/parameterization/training_policy.py @@ -0,0 +1,72 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from __future__ import annotations + +from dataclasses import dataclass + +from .spec import SCALING_RECIPE_MUP, ScalingContext + + +@dataclass(frozen=True) +class TrainingScalingPolicy: + """Optimizer-side policy for existing Megatron MuP hyperparameter multipliers.""" + + context: ScalingContext + optimizer_type: str = 'adam' + + @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 'adam' in self.optimizer_type_lower + + @property + def is_muon_optimizer(self) -> bool: + return 'muon' in self.optimizer_type_lower + + @property + def hidden_lr_multiplier(self) -> float: + if not self.enabled or not self.uses_width_mup: + return 1.0 + if self.is_sgd_optimizer: + return 1.0 + return 1.0 / self.context.width_mult + + @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_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 + + @property + def vector_like_lr_multiplier(self) -> float: + return self.hidden_vector_lr_multiplier + + +def build_legacy_mup_training_policy( + *, mup_width_mult: float, optimizer_type: str = 'adam' +) -> TrainingScalingPolicy: + return TrainingScalingPolicy( + context=ScalingContext(recipe=SCALING_RECIPE_MUP, width_mult=mup_width_mult), + optimizer_type=optimizer_type, + ) diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index b27f90c53d0..55da3a1321b 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -28,6 +28,7 @@ get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) +from megatron.core.parameterization import build_model_scaling_policy from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, ) @@ -391,11 +392,18 @@ def __init__( ) # Output. + model_scaling_policy = build_model_scaling_policy(self.config) self.linear_proj = submodules.linear_proj( self.query_projection_size, self.config.hidden_size, config=self.config, - init_method=not_none(self.config.output_layer_init_method), + 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=False, + ), bias=self.config.add_bias_linear, input_is_parallel=True, skip_bias_add=True, diff --git a/megatron/core/transformer/mlp.py b/megatron/core/transformer/mlp.py index 1a578151f1e..1b7bac4dba7 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_model_scaling_policy from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_config import TransformerConfig @@ -233,13 +234,20 @@ def __init__( else: self.activation_func = self.config.activation_func + model_scaling_policy = build_model_scaling_policy(self.config) self.linear_fc2 = submodules.linear_fc2( not_none(self.config.ffn_hidden_size), not_none( 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=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=False, + ), bias=self.config.add_bias_linear, input_is_parallel=True, skip_bias_add=True, diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 3e91a2b8042..52bf50b6123 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,11 @@ from megatron.core.enums import Fp4Recipe, Fp8Recipe from megatron.core.inference.moe import InferenceGroupedGemmBackend +from megatron.core.parameterization import ( + build_model_scaling_policy, + build_scaling_context, + sync_legacy_mup_fields, +) from megatron.core.quantization.quant_config import RecipeConfig from megatron.core.transformer.cuda_graph_config import ( ALLOWED_INFERENCE_SCOPES, @@ -30,14 +34,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__) @@ -365,25 +362,42 @@ class TransformerConfig(ModelParallelConfig): #################### # MuP (Maximal Update Parameterization) #################### + scaling_recipe: Optional[Literal['none', 'mup']] = None + """ + Canonical scaling recipe. ``none`` preserves standard parameterization, and ``mup`` + enables width MuP. If unset, legacy ``use_mup`` selects ``mup`` for backward + compatibility. + """ + + scaling_base_hidden_size: Optional[int] = None + """ + Canonical base hidden size for width scaling. For MuP, the width multiplier is + derived as hidden_size / scaling_base_hidden_size. + """ + + scaling_base_head_dim: Optional[float] = None + """ + Canonical base attention head dimension for MuP attention scaling. This aliases the + deprecated mup_base_head_dim field. + """ + 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 alias for scaling_recipe='mup'. Kept for checkpoint and script + compatibility. """ 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 derived MuP width multiplier. The canonical value is computed as + hidden_size / scaling_base_hidden_size. If this legacy input is non-default, it + must match the derived value. """ mup_base_hidden_size: Optional[int] = None """ - Base hidden size for MuP width scaling. This is the reference width from which - scaling factors are computed. Defaults to hidden_size if not specified (base model - case where width_mult=1.0). Set this to your base/proxy model's hidden size when - scaling up. + Deprecated alias for scaling_base_hidden_size. Set scaling_recipe='mup' and + scaling_base_hidden_size for new configs. """ mup_embedding_mult: float = 1.0 @@ -394,15 +408,15 @@ class TransformerConfig(ModelParallelConfig): mup_output_mult: float = 1.0 """ - Multiplier for output logits before softmax. When MuP is enabled and this is left - at 1.0, it is auto-set to 1/mup_width_mult to keep output variance stable across - widths. Override to customize output scaling. + Multiplier for output logits before softmax. When scaling_recipe='mup' and this is + left at 1.0, it is auto-set to 1/mup_width_mult to keep output variance stable + across widths. Override to customize output scaling. Default: 1.0. """ mup_base_head_dim: Optional[float] = None """ - Base head dimension for MuP attention scaling. When set, + Deprecated alias for scaling_base_head_dim. When set, softmax_scale = sqrt(mup_base_head_dim) / (kv_channels ** mup_attn_scale_power). Set to base model's d_head (e.g., 64) to match standard 1/sqrt(d_head) scaling at the base width, ensuring non-MuP compatibility for that specific value. @@ -412,7 +426,8 @@ 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 (MuP scaling when scaling_recipe='mup'). Set to 0.5 for standard + scaling. """ #################### @@ -1956,27 +1971,11 @@ 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_scaling_context(self) + sync_legacy_mup_fields(self, scaling_context) + # MuP (Maximal Update Parameterization) configuration + if scaling_context.uses_width_mup: overridden_init_methods = [] if self.init_method is not None: overridden_init_methods.append("init_method") @@ -1986,12 +1985,17 @@ def __post_init__(self): 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 " + "scaling recipe 'mup' is enabled, but custom " + overridden_init_methods_text + f" {verb} set. This may break MuP initialization assumptions.", UserWarning, ) + model_scaling_policy = build_model_scaling_policy(self) + 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), @@ -2015,29 +2019,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 ddd1e7d34cd..4243ec0898d 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -17,6 +17,7 @@ from megatron.core.dist_checkpointing.utils import apply_prefix_mapping from megatron.core.inference.utils import InferenceMode from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.parameterization import build_model_scaling_policy from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.cuda_graphs import is_graph_capturing from megatron.core.transformer.enums import CudaGraphModule, InferenceCudaGraphScope, LayerType @@ -335,6 +336,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_model_scaling_policy(config) # [Module 1: Input Layernorm] Optional Layernorm on the input data # TODO: add pytorch only layernorm @@ -680,6 +682,9 @@ def _forward_attention( hidden_states = attention_output_with_bias[0] else: with self.bias_dropout_add_exec_handler(): + attention_output_with_bias = self.model_scaling_policy.scale_residual_branch_output( + attention_output_with_bias + ) hidden_states = self.self_attn_bda(self.training, self.config.bias_dropout_fusion)( attention_output_with_bias, residual, self.hidden_dropout ) @@ -929,6 +934,9 @@ def _forward_post_mlp( hidden_states = mlp_output_with_bias[0] else: with self.bias_dropout_add_exec_handler(): + mlp_output_with_bias = self.model_scaling_policy.scale_residual_branch_output( + mlp_output_with_bias + ) hidden_states = self.mlp_bda(self.training, self.config.bias_dropout_fusion)( mlp_output_with_bias, residual, self.hidden_dropout ) diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index cd3ce44c3a4..f0ea0b71b73 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -38,6 +38,12 @@ ) from megatron.core.activations import squared_relu from megatron.core.fusions.fused_bias_geglu import quick_gelu +from megatron.core.parameterization import ( + SCALING_RECIPE_MUP, + SCALING_RECIPE_VALUES, + build_scaling_context, + sync_legacy_mup_fields, +) from megatron.training.global_vars import set_global_variables from megatron.training.utils import ( get_device_arch_version, @@ -59,6 +65,7 @@ def add_megatron_arguments(parser: argparse.ArgumentParser): # 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) @@ -252,6 +259,35 @@ def validate_model_config_args_from_heterogeneous_config(args): f"Arguments differ from heterogeneous config: {incompatible_args_str}" ) + +def warn_deprecated_mup_aliases(args): + """Warn when users select the legacy MuP flag surface instead of scaling recipes.""" + + deprecated_aliases = [] + if getattr(args, '_use_mup_explicit', getattr(args, 'use_mup', False)): + deprecated_aliases.append('--use-mup') + if getattr( + args, '_mup_base_hidden_size_explicit', getattr(args, 'mup_base_hidden_size', None) is not None + ): + deprecated_aliases.append('--mup-base-hidden-size') + if getattr( + args, '_mup_base_head_dim_explicit', getattr(args, 'mup_base_head_dim', None) is not None + ): + deprecated_aliases.append('--mup-base-head-dim') + if getattr(args, '_mup_width_mult_explicit', False): + deprecated_aliases.append('--mup-width-mult') + + if deprecated_aliases: + warn_rank_0( + "Deprecated MuP argument(s) " + + ", ".join(deprecated_aliases) + + " were provided. Use --scaling-recipe mup with " + "--scaling-base-hidden-size and --scaling-base-head-dim instead. " + "--mup-width-mult is derived from hidden_size / scaling_base_hidden_size.", + getattr(args, 'rank', 0), + ) + + def _eval_pattern(pattern): """ Validate and evaluate a string containing a Python list expression """ assert isinstance(pattern, str) @@ -1761,6 +1797,9 @@ 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." + warn_deprecated_mup_aliases(args) + sync_legacy_mup_fields(args, build_scaling_context(args)) + # Print arguments. _print_args("arguments", args) @@ -2087,6 +2126,17 @@ def _add_network_size_args(parser): "persist_layer_norm", "bias_dropout_fusion", "apply_rope_fusion", + # generated by the explicit scaling argument group + "scaling_recipe", + "scaling_base_hidden_size", + "scaling_base_head_dim", + "use_mup", + "mup_width_mult", + "mup_base_hidden_size", + "mup_embedding_mult", + "mup_output_mult", + "mup_base_head_dim", + "mup_attn_scale_power", ] transformer_factory = ArgumentGroupFactory(TransformerConfig, exclude=exclude) transformer_group = transformer_factory.build_group(parser, "transformer configuration") @@ -2647,6 +2697,86 @@ def _add_learning_rate_args(parser): return parser +def _add_scaling_args(parser): + group = parser.add_argument_group(title='scaling') + + class _StoreMupWidthMult(argparse.Action): + def __call__(self, parser, namespace, values, option_string=None): + setattr(namespace, self.dest, values) + setattr(namespace, '_mup_width_mult_explicit', True) + + group.add_argument( + '--scaling-recipe', + choices=SCALING_RECIPE_VALUES, + default=None, + help=( + "Canonical parameterization recipe. Use 'none' for standard parameterization " + "or 'mup' for width MuP." + ), + ) + group.add_argument( + '--scaling-base-hidden-size', + type=int, + default=None, + help=( + "Base hidden size for scaling recipes. For MuP, width multiplier is derived as " + "hidden_size / scaling_base_hidden_size." + ), + ) + group.add_argument( + '--scaling-base-head-dim', + type=float, + default=None, + help="Base attention head dimension for MuP attention scaling.", + ) + group.add_argument( + '--use-mup', + action='store_true', + help=f"Deprecated alias for --scaling-recipe {SCALING_RECIPE_MUP}.", + ) + group.add_argument( + '--mup-width-mult', + type=float, + default=1.0, + action=_StoreMupWidthMult, + help=( + "Deprecated derived MuP width multiplier. If supplied, it must match " + "hidden_size / scaling_base_hidden_size." + ), + ) + group.add_argument( + '--mup-base-hidden-size', + type=int, + default=None, + help="Deprecated alias for --scaling-base-hidden-size.", + ) + group.add_argument( + '--mup-embedding-mult', + type=float, + default=1.0, + help="MuP embedding activation multiplier.", + ) + group.add_argument( + '--mup-output-mult', + type=float, + default=1.0, + help="MuP output logit multiplier. Defaults to 1 / width_mult when left at 1.0.", + ) + group.add_argument( + '--mup-base-head-dim', + type=float, + default=None, + help="Deprecated alias for --scaling-base-head-dim.", + ) + group.add_argument( + '--mup-attn-scale-power', + type=float, + default=1.0, + help="MuP attention scale power. The default uses 1 / d_head.", + ) + return parser + + def _add_checkpointing_args(parser): from megatron.training.config import CheckpointConfig diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index 2363b7ae164..08146d66839 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -38,6 +38,10 @@ from megatron.core.msc_utils import MultiStorageClientFeature, open_file from megatron.core.num_microbatches_calculator import update_num_microbatches from megatron.core.optimizer import DistributedOptimizer +from megatron.core.parameterization import ( + build_scaling_context, + sync_legacy_mup_fields, +) from megatron.core.rerun_state_machine import get_rerun_state_machine from megatron.core.utils import get_pg_rank, get_pg_size, unwrap_model @@ -175,6 +179,39 @@ def _compare(arg_name, old_arg_name=None, default=None): _compare('tensor_model_parallel_size') _compare('pipeline_model_parallel_size') + checkpoint_scaling_context = build_scaling_context(checkpoint_args) + args_scaling_context = build_scaling_context(args) + assert checkpoint_scaling_context == args_scaling_context, ( + f"Scaling recipe from checkpoint ({checkpoint_scaling_context}) is not equal to " + f"the input argument value ({args_scaling_context})." + ) + + +_CHECKPOINT_SCALING_ARG_DEFAULTS = { + 'scaling_recipe': 'none', + 'scaling_base_hidden_size': None, + 'scaling_base_head_dim': None, + 'use_mup': False, + 'mup_width_mult': 1.0, + '_mup_width_mult_explicit': False, + '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, +} + + +def _sync_checkpoint_scaling_args(checkpoint_args): + """Populate canonical scaling fields on checkpoint args before force-copying them.""" + + sync_legacy_mup_fields( + checkpoint_args, build_scaling_context(checkpoint_args) + ) + for arg_name, default_value in _CHECKPOINT_SCALING_ARG_DEFAULTS.items(): + if not hasattr(checkpoint_args, arg_name): + setattr(checkpoint_args, arg_name, default_value) + def isfile(filename) -> bool: if MultiStorageClientFeature.is_enabled(): @@ -1502,7 +1539,9 @@ def load_args_from_checkpoint( if hasattr(checkpoint_args, 'num_layers'): setattr(checkpoint_args, 'num_layers', None) - def _set_arg(arg_name, old_arg_name=None, force=False): + _sync_checkpoint_scaling_args(checkpoint_args) + + def _set_arg(arg_name, old_arg_name=None, force=False, allow_none=False): if not force and getattr(args, arg_name, None) is not None: return @@ -1511,7 +1550,7 @@ def _set_arg(arg_name, old_arg_name=None, force=False): else: checkpoint_value = getattr(checkpoint_args, arg_name, None) - if checkpoint_value is not None: + if checkpoint_value is not None or allow_none: print_rank_0(f"Setting {arg_name} to {checkpoint_value} from checkpoint") setattr(args, arg_name, checkpoint_value) else: @@ -1543,6 +1582,24 @@ 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('scaling_recipe', force=True) + _set_arg('scaling_base_hidden_size', force=True, allow_none=True) + _set_arg('scaling_base_head_dim', force=True, allow_none=True) + _set_arg('use_mup', force=True) + setattr(args, '_use_mup_explicit', False) + _set_arg('mup_width_mult', force=True) + setattr( + args, + '_mup_width_mult_explicit', + getattr(checkpoint_args, '_mup_width_mult_explicit', False), + ) + _set_arg('mup_base_hidden_size', force=True, allow_none=True) + setattr(args, '_mup_base_hidden_size_explicit', False) + _set_arg('mup_embedding_mult', force=True) + _set_arg('mup_output_mult', force=True) + _set_arg('mup_base_head_dim', force=True, allow_none=True) + setattr(args, '_mup_base_head_dim_explicit', False) + _set_arg('mup_attn_scale_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 f69e4f30f6a..11c8128792e 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -163,7 +163,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 from megatron.core.utils import ( @@ -1009,7 +1009,7 @@ 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] + key_fn = get_param_group_identifier_tuple param_groups.sort(key=key_fn) inner_optimizer["param_groups"] = param_groups diff --git a/megatron/training/yaml_arguments.py b/megatron/training/yaml_arguments.py index d44f4d31822..3d26f122dfb 100644 --- a/megatron/training/yaml_arguments.py +++ b/megatron/training/yaml_arguments.py @@ -16,8 +16,10 @@ import torch.nn.functional as F +from megatron.core.parameterization import build_scaling_context, sync_legacy_mup_fields from megatron.core.transformer import TransformerConfig, MLATransformerConfig from megatron.core.utils import get_torch_version, is_torch_min_version +from megatron.training.arguments import 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 @@ -38,6 +40,20 @@ def env_constructor(loader, node): "bfloat16" : torch.bfloat16 } +DEFAULTABLE_SCALING_FIELDS = { + 'scaling_recipe', + 'scaling_base_hidden_size', + 'scaling_base_head_dim', + 'use_mup', + 'mup_width_mult', + 'mup_base_hidden_size', + 'mup_embedding_mult', + 'mup_output_mult', + 'mup_base_head_dim', + 'mup_attn_scale_power', +} + + def validate_yaml(args, defaults={}): # This is for legacy script env var setting @@ -246,6 +262,13 @@ def validate_yaml(args, defaults={}): assert args.language_model.hidden_size % args.language_model.num_attention_heads == 0 args.language_model.kv_channels = args.language_model.hidden_size // args.language_model.num_attention_heads + if getattr(args.language_model, 'mup_width_mult', 1.0) != 1.0: + args.language_model._mup_width_mult_explicit = True + warn_deprecated_mup_aliases(args.language_model) + sync_legacy_mup_fields( + args.language_model, build_scaling_context(args.language_model) + ) + #TODO: Implement arguments for encoder-decoder if args.seq_length is not None: assert args.encoder_seq_length is None @@ -380,6 +403,13 @@ def core_config_from_args(args, dataclass=TransformerConfig): 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: + if f.default is not dataclasses.MISSING: + kw_args[f.name] = f.default + elif 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") else: raise Exception(f"Missing argument {f.name} for {str(dataclass)} config") return kw_args @@ -438,4 +468,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/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 56af8545042..ee35b77b701 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 @@ -69,6 +70,87 @@ def forward(self, x): return x +def test_param_group_identifier_tuple_tolerates_missing_optional_keys(): + group = {"wd_mult": 1.0, "lr_mult": 1.0, "is_expert_parallel": False, "is_decoupled_lr": False} + + ident = get_param_group_identifier_tuple(group) + + assert ident == (1.0, 1.0, False, False) + + +def test_param_group_identifier_tuple_reads_pre_keys_and_optional_fields(): + group = { + "pre_wd_mult": 0.0, + "pre_lr_mult": 0.5, + "pre_is_expert_parallel": False, + "pre_is_decoupled_lr": True, + } + + ident = get_param_group_identifier_tuple(group) + + assert ident == (0.0, 0.5, False, True) + + +def test_param_group_identifier_tuple_defaults_missing_legacy_fields(): + group = {} + + ident = get_param_group_identifier_tuple(group) + + assert ident == (1.0, 1.0, False, False) + + +def test_param_group_matching_ignores_mutable_scheduler_values_on_resume(): + current_group = { + "wd_mult": 1.0, + "lr_mult": 1.0, + "is_expert_parallel": False, + "is_decoupled_lr": False, + "max_lr": 2e-4, + "min_lr": 2e-6, + "params": [0], + } + checkpoint_group = { + "wd_mult": 1.0, + "lr_mult": 1.0, + "is_expert_parallel": False, + "is_decoupled_lr": False, + "max_lr": 1e-4, + "min_lr": 1e-6, + "params": [7], + } + + assert get_param_group_identifier_tuple(current_group) == get_param_group_identifier_tuple( + checkpoint_group + ) + + reordered_groups = MegatronOptimizer._filter_and_reorder_param_groups( + [current_group], [checkpoint_group] + ) + + assert reordered_groups == [checkpoint_group] + + +def test_param_group_matching_normalizes_legacy_missing_identifier_fields(): + current_group = { + "wd_mult": 1.0, + "lr_mult": 1.0, + "is_expert_parallel": False, + "is_decoupled_lr": False, + "params": [0], + } + legacy_checkpoint_group = {"params": [7]} + + assert get_param_group_identifier_tuple(current_group) == get_param_group_identifier_tuple( + legacy_checkpoint_group + ) + + reordered_groups = MegatronOptimizer._filter_and_reorder_param_groups( + [current_group], [legacy_checkpoint_group] + ) + + assert reordered_groups == [legacy_checkpoint_group] + + @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/transformer/test_mup.py b/tests/unit_tests/transformer/test_mup.py index f1d99cad1e6..c430324edd4 100644 --- a/tests/unit_tests/transformer/test_mup.py +++ b/tests/unit_tests/transformer/test_mup.py @@ -9,19 +9,34 @@ 4. LR override computation """ +import argparse +import dataclasses import logging import math +import warnings +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 +from megatron.core.optimizer import ( + 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_legacy_mup_training_policy, + build_model_scaling_policy, + build_scaling_context, +) from megatron.core.transformer.multi_token_prediction import process_mtp_loss from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import init_method_normal, mup_scaled_init_method_normal +from megatron.training.arguments import add_megatron_arguments +from megatron.training.yaml_arguments import core_config_from_args class TestMuPConfigValidation: @@ -38,6 +53,8 @@ def test_mup_defaults_base_hidden_size(self): ) assert config.mup_base_hidden_size == 512 assert config.mup_width_mult == 1.0 + assert config.scaling_recipe == 'mup' + assert config.scaling_base_hidden_size == 512 def test_mup_width_mult_calculation(self): """width_mult = hidden_size / base_hidden_size.""" @@ -49,6 +66,8 @@ def test_mup_width_mult_calculation(self): mup_base_hidden_size=256, ) assert config.mup_width_mult == 4.0 + assert config.scaling_recipe == 'mup' + assert config.scaling_base_hidden_size == 256 def test_mup_width_mult_fractional(self): """width_mult can be fractional (smaller than base).""" @@ -67,6 +86,8 @@ def test_mup_backward_compatible(self): assert config.use_mup is False assert config.mup_width_mult == 1.0 assert config.mup_base_hidden_size is None + assert config.scaling_recipe == 'none' + assert config.scaling_base_hidden_size is None def test_mup_base_hidden_size_must_be_positive(self): """mup_base_hidden_size must be positive.""" @@ -80,6 +101,477 @@ def test_mup_base_hidden_size_must_be_positive(self): ) assert "positive" in str(exc_info.value).lower() + def test_scaling_recipe_mup_sets_legacy_fields(self): + """Canonical MuP fields populate the legacy fields used by existing call sites.""" + config = TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + ) + + assert config.use_mup is True + assert config.mup_base_hidden_size == 256 + assert config.mup_base_head_dim == 64 + assert config.mup_width_mult == pytest.approx(4.0) + assert config.scaling_base_hidden_size == 256 + assert config.scaling_base_head_dim == 64 + + def test_legacy_mup_fields_resolve_to_canonical_recipe(self): + """Legacy MuP flags remain compatible but are not separate state.""" + config = TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_base_head_dim=64, + ) + + assert config.scaling_recipe == 'mup' + assert config.scaling_base_hidden_size == 256 + assert config.scaling_base_head_dim == 64 + assert build_scaling_context(config).width_mult == pytest.approx(4.0) + + def test_scaling_recipe_none_rejects_scaling_overrides(self): + """Scaling fields cannot silently affect standard parameterization.""" + with pytest.raises(ValueError, match="Scaling overrides"): + TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='none', + scaling_base_hidden_size=256, + ) + + def test_use_mup_conflicts_with_scaling_recipe_none(self): + """The deprecated MuP boolean cannot override an explicit canonical recipe.""" + with pytest.raises(ValueError, match="conflicts"): + TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='none', + use_mup=True, + ) + + def test_canonical_and_legacy_base_hidden_must_match(self): + """Canonical and deprecated base hidden-size fields are aliases.""" + with pytest.raises(ValueError, match="conflicts"): + TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + mup_base_hidden_size=512, + ) + + def test_deprecated_width_mult_must_match_derived_value(self): + """mup_width_mult is accepted only when it matches the derived width.""" + config = TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + mup_width_mult=4.0, + ) + assert config.mup_width_mult == pytest.approx(4.0) + + with pytest.raises(ValueError, match="must match the derived"): + TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_recipe='mup', + scaling_base_hidden_size=256, + mup_width_mult=2.0, + ) + + def test_scaling_override_without_recipe_is_rejected(self): + """Base scaling fields do not implicitly enable MuP.""" + with pytest.raises(ValueError, match="Scaling overrides"): + TransformerConfig( + hidden_size=1024, + num_layers=4, + num_attention_heads=16, + scaling_base_hidden_size=256, + ) + + +class TestScalingRecipeSurfaces: + """Tests for public config surfaces that feed the scaling context.""" + + SCALING_FIELD_NAMES = { + 'scaling_recipe', + 'scaling_base_hidden_size', + 'scaling_base_head_dim', + 'use_mup', + 'mup_width_mult', + 'mup_base_hidden_size', + 'mup_embedding_mult', + 'mup_output_mult', + 'mup_base_head_dim', + 'mup_attn_scale_power', + } + + def test_cli_parser_accepts_canonical_scaling_args(self): + """The explicit scaling arg group owns canonical and legacy MuP flags.""" + parser = argparse.ArgumentParser(allow_abbrev=False) + add_megatron_arguments(parser) + + args, _ = parser.parse_known_args( + [ + '--scaling-recipe', + 'mup', + '--scaling-base-hidden-size', + '256', + '--mup-base-head-dim', + '64', + ] + ) + + assert args.scaling_recipe == 'mup' + assert args.scaling_base_hidden_size == 256 + assert args.mup_base_head_dim == 64 + assert args.mup_width_mult == 1.0 + + def test_cli_explicit_mup_width_mult_one_is_validated(self): + """Explicit legacy width multiplier must match the derived value, even at 1.0.""" + parser = argparse.ArgumentParser(allow_abbrev=False) + add_megatron_arguments(parser) + + args, _ = parser.parse_known_args( + [ + '--scaling-recipe', + 'mup', + '--scaling-base-hidden-size', + '256', + '--mup-width-mult', + '1.0', + ] + ) + args.hidden_size = 1024 + + with pytest.raises(ValueError, match="must match the derived"): + build_scaling_context(args) + + def test_yaml_core_config_defaults_missing_scaling_fields(self): + """Existing YAML files may omit the new scaling fields.""" + values = {} + for field in dataclasses.fields(TransformerConfig): + if field.name in self.SCALING_FIELD_NAMES: + continue + 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() + elif field.type is int: + values[field.name] = 1 + else: + values[field.name] = None + values['hidden_size'] = 512 + values['num_layers'] = 2 + values['num_attention_heads'] = 8 + + kwargs = core_config_from_args(SimpleNamespace(**values), TransformerConfig) + + assert kwargs['scaling_recipe'] is None + assert kwargs['scaling_base_hidden_size'] is None + assert kwargs['mup_width_mult'] == 1.0 + + def test_yaml_default_width_mult_is_not_treated_as_explicit(self): + """Full legacy YAML files may materialize the old default width multiplier.""" + yaml_args = SimpleNamespace( + hidden_size=1024, + scaling_recipe=None, + scaling_base_hidden_size=None, + scaling_base_head_dim=None, + use_mup=True, + mup_width_mult=1.0, + mup_base_hidden_size=256, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=None, + mup_attn_scale_power=1.0, + ) + + context = build_scaling_context(yaml_args) + + assert context.width_mult == pytest.approx(4.0) + + def test_scaling_context_matches_legacy_checkpoint_and_canonical_args(self): + """Checkpoint compatibility compares effective scaling, not flag spelling.""" + legacy_checkpoint_args = SimpleNamespace( + hidden_size=1024, + use_mup=True, + mup_base_hidden_size=256, + mup_width_mult=1.0, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=64, + mup_attn_scale_power=1.0, + ) + canonical_args = SimpleNamespace( + hidden_size=1024, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + use_mup=False, + 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, + ) + + assert build_scaling_context( + legacy_checkpoint_args + ) == build_scaling_context(canonical_args) + + def test_checkpoint_scaling_sync_populates_canonical_fields(self): + """Old checkpoints with only legacy MuP fields become canonical before copy.""" + from megatron.training.checkpointing import _sync_checkpoint_scaling_args + + legacy_checkpoint_args = SimpleNamespace( + hidden_size=1024, + use_mup=True, + mup_base_hidden_size=256, + mup_width_mult=1.0, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=64, + mup_attn_scale_power=1.0, + ) + + _sync_checkpoint_scaling_args(legacy_checkpoint_args) + + assert legacy_checkpoint_args.scaling_recipe == 'mup' + assert legacy_checkpoint_args.scaling_base_hidden_size == 256 + assert legacy_checkpoint_args.scaling_base_head_dim == 64 + assert legacy_checkpoint_args.mup_width_mult == pytest.approx(4.0) + + def test_check_checkpoint_args_compares_effective_scaling_context(self, monkeypatch): + """The real checkpoint check accepts legacy and canonical spellings if equivalent.""" + from megatron.training import checkpointing + + runtime_args = SimpleNamespace( + num_layers=2, + hidden_size=1024, + num_attention_heads=16, + add_position_embedding=True, + vocab_file=None, + data_parallel_random_init=False, + phase_transition_iterations=None, + use_dist_ckpt=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + use_mup=False, + 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( + num_layers=2, + hidden_size=1024, + num_attention_heads=16, + add_position_embedding=True, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + use_mup=True, + mup_base_hidden_size=256, + mup_width_mult=1.0, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_base_head_dim=64, + mup_attn_scale_power=1.0, + ) + monkeypatch.setattr(checkpointing, 'get_args', lambda: runtime_args) + monkeypatch.setattr(checkpointing, 'get_checkpoint_version', lambda: 3.0) + + checkpointing.check_checkpoint_args(checkpoint_args) + + def test_load_checkpoint_args_clears_optional_scaling_fields(self, monkeypatch): + """use-checkpoint-args must clear stale optional canonical fields.""" + from megatron.training import checkpointing + + checkpoint_args = SimpleNamespace( + num_layers=2, + hidden_size=1024, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_width_mult=1.0, + mup_embedding_mult=1.0, + mup_output_mult=1.0, + mup_attn_scale_power=1.0, + ) + state_dict = {'args': checkpoint_args, 'iteration': 7} + monkeypatch.setattr( + checkpointing, + '_load_base_checkpoint', + lambda *args, **kwargs: (state_dict, 'model_optim_rng.pt', False, None), + ) + runtime_args = SimpleNamespace( + load='dummy-checkpoint', + iteration=0, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + mup_base_head_dim=64, + use_tokenizer_model_from_checkpoint_args=False, + use_mp_args_from_checkpoint_args=False, + ) + + checkpointing.load_args_from_checkpoint(runtime_args) + + assert runtime_args.iteration == 7 + assert runtime_args.scaling_recipe == 'mup' + assert runtime_args.scaling_base_hidden_size == 256 + assert runtime_args.scaling_base_head_dim is None + assert runtime_args.mup_base_head_dim is None + assert runtime_args.mup_width_mult == pytest.approx(4.0) + + def test_load_non_mup_checkpoint_clears_width_mult_provenance(self, monkeypatch): + """Old no-scaling checkpoints must clear stale CLI scaling state.""" + from megatron.training import checkpointing + + checkpoint_args = SimpleNamespace(hidden_size=1024) + state_dict = {'args': checkpoint_args, 'iteration': 3} + monkeypatch.setattr( + checkpointing, + '_load_base_checkpoint', + lambda *args, **kwargs: (state_dict, 'model_optim_rng.pt', False, None), + ) + runtime_args = SimpleNamespace( + load='dummy-checkpoint', + iteration=0, + scaling_recipe='mup', + scaling_base_hidden_size=256, + scaling_base_head_dim=64, + use_mup=True, + mup_width_mult=4.0, + _mup_width_mult_explicit=True, + mup_base_hidden_size=256, + mup_embedding_mult=2.0, + mup_output_mult=0.25, + mup_base_head_dim=64, + mup_attn_scale_power=-0.5, + use_tokenizer_model_from_checkpoint_args=False, + use_mp_args_from_checkpoint_args=False, + ) + + checkpointing.load_args_from_checkpoint(runtime_args) + + assert runtime_args.iteration == 3 + assert runtime_args.scaling_recipe == 'none' + assert runtime_args.scaling_base_hidden_size is None + assert runtime_args.scaling_base_head_dim is None + assert runtime_args.use_mup is False + assert runtime_args.mup_width_mult == 1.0 + assert runtime_args._mup_width_mult_explicit is False + assert runtime_args.mup_base_hidden_size is None + assert runtime_args.mup_embedding_mult == 1.0 + assert runtime_args.mup_output_mult == 1.0 + assert runtime_args.mup_base_head_dim is None + assert runtime_args.mup_attn_scale_power == 1.0 + assert build_scaling_context(runtime_args).recipe == 'none' + + def test_checkpoint_derived_width_mult_does_not_warn_as_deprecated_cli(self): + """Normalized checkpoint state should not look like user-provided --mup-width-mult.""" + from megatron.training.arguments import warn_deprecated_mup_aliases + + checkpoint_derived_args = SimpleNamespace( + rank=0, + use_mup=False, + mup_base_hidden_size=None, + mup_base_head_dim=None, + mup_width_mult=4.0, + _mup_width_mult_explicit=False, + ) + + with warnings.catch_warnings(record=True) as caught_warnings: + warnings.simplefilter("always") + warn_deprecated_mup_aliases(checkpoint_derived_args) + + assert len(caught_warnings) == 0 + + def test_false_width_mult_marker_overrides_stale_non_default_value(self): + """Marker-present False means checkpoint/internal provenance, not explicit user input.""" + checkpoint_derived_args = SimpleNamespace( + hidden_size=1024, + scaling_recipe='none', + scaling_base_hidden_size=None, + scaling_base_head_dim=None, + use_mup=False, + mup_width_mult=4.0, + _mup_width_mult_explicit=False, + 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, + ) + + assert build_scaling_context(checkpoint_derived_args).recipe == 'none' + + def test_checkpoint_derived_legacy_aliases_do_not_warn_as_deprecated_cli(self): + """Checkpoint-synced legacy fields should not masquerade as user CLI aliases.""" + from megatron.training.arguments import warn_deprecated_mup_aliases + + checkpoint_derived_args = SimpleNamespace( + rank=0, + use_mup=True, + _use_mup_explicit=False, + mup_base_hidden_size=256, + _mup_base_hidden_size_explicit=False, + mup_base_head_dim=64, + _mup_base_head_dim_explicit=False, + mup_width_mult=4.0, + _mup_width_mult_explicit=False, + ) + + with warnings.catch_warnings(record=True) as caught_warnings: + warnings.simplefilter("always") + warn_deprecated_mup_aliases(checkpoint_derived_args) + + assert len(caught_warnings) == 0 + + def test_user_provided_legacy_aliases_warn_as_deprecated_cli(self): + """Real user-provided legacy aliases should still produce a deprecation warning.""" + from megatron.training.arguments import warn_deprecated_mup_aliases + + user_args = SimpleNamespace( + rank=0, + use_mup=True, + _use_mup_explicit=True, + mup_base_hidden_size=256, + _mup_base_hidden_size_explicit=True, + mup_base_head_dim=64, + _mup_base_head_dim_explicit=True, + mup_width_mult=1.0, + _mup_width_mult_explicit=True, + ) + + with pytest.warns(UserWarning) as caught_warnings: + warn_deprecated_mup_aliases(user_args) + + warning_text = str(caught_warnings[0].message) + assert '--use-mup' in warning_text + assert '--mup-base-hidden-size' in warning_text + assert '--mup-base-head-dim' in warning_text + assert '--mup-width-mult' in warning_text + class TestMuPInitMethods: """Tests for MuP initialization methods.""" @@ -169,7 +661,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 +673,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, @@ -195,6 +687,20 @@ def test_mup_warns_with_custom_output_layer_init_method(self): class TestMuPLRScaling: """Tests for MuP learning rate and Adam epsilon scaling.""" + def test_mup_overrides_route_through_training_scaling_policy(self): + """The new policy seam preserves the legacy MuP override surface.""" + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) + width_mult = 4.0 + + legacy_overrides = get_mup_config_overrides(optimizer_config, width_mult) + policy_overrides = get_scaling_config_overrides( + optimizer_config, + build_legacy_mup_training_policy(mup_width_mult=width_mult, optimizer_type='adam'), + ) + + assert legacy_overrides.keys() == policy_overrides.keys() + assert list(legacy_overrides.values()) == list(policy_overrides.values()) + def test_mup_lr_override_computation(self): """Hidden LR and Adam eps scale as 1/width_mult.""" optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5) @@ -355,6 +861,47 @@ def test_mup_with_decoupled_lr_scales_hidden_only_for_lr(self): class TestMuPConfigIntegration: """Integration tests for MuP config with init methods.""" + def test_model_scaling_policy_matches_legacy_config_fields(self): + """Model policy exposes the same effective MuP values as TransformerConfig.""" + config = TransformerConfig( + hidden_size=1024, + num_layers=8, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + mup_embedding_mult=3.0, + ) + policy = build_model_scaling_policy(config) + + assert policy.enabled is True + assert policy.context.width_mult == pytest.approx(config.mup_width_mult) + assert policy.context.output_mult == pytest.approx(config.mup_output_mult) + assert policy.context.embedding_mult == pytest.approx(config.mup_embedding_mult) + assert config.softmax_scale == pytest.approx( + policy.resolve_attention_softmax_scale( + softmax_scale=None, kv_channels=config.kv_channels + ) + ) + + def test_model_scaling_policy_tracks_post_init_multiplier_mutations(self): + """Policy resolution should preserve legacy live reads of mutable config fields.""" + config = TransformerConfig( + hidden_size=1024, + num_layers=8, + num_attention_heads=16, + use_mup=True, + mup_base_hidden_size=256, + ) + logits = torch.ones(2, 4) + embeddings = torch.ones(2, 4) + + config.mup_output_mult = 0.25 + config.mup_embedding_mult = 3.0 + policy = build_model_scaling_policy(config) + + assert torch.equal(policy.scale_output_logits(logits), logits * 0.25) + assert torch.equal(policy.scale_embedding_activations(embeddings), embeddings * 3.0) + def test_mup_output_layer_init(self): """Output layer init should also scale with MuP.""" config = TransformerConfig(