diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index b4a2f1e5835..fc237ffcc56 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -2,7 +2,6 @@ import copy import logging import warnings -from collections import defaultdict from dataclasses import astuple from typing import Any, Callable, Dict, List, Optional, Tuple, Union @@ -61,13 +60,7 @@ from ..transformer.module import MegatronModule from ..utils import get_model_config, get_pg_rank, get_pg_size, is_te_min_version, log_single_rank from .distrib_optimizer import DistributedOptimizer -from .emerging_optimizers import ( - _EMERGING_OPTIMIZERS, - HAVE_EMERGING_OPTIMIZERS, - _create_emerging_optimizer, -) from .grad_scaler import ConstantGradScaler, DynamicGradScaler -from .layer_wise_optimizer import LayerWiseDistributedOptimizer from .optimizer import ( ChainedOptimizer, Float16OptimizerWithFloat16Params, @@ -75,8 +68,6 @@ MegatronOptimizer, param_group_identifier_keys, ) - -# Subclass aliases kept for backward compatibility; all are OptimizerConfig. from .optimizer_config import ( AdamOptimizerConfig, OptimizerConfig, @@ -325,6 +316,14 @@ def _get_param_groups( # Map (pg_overrides, is_expert_parallel) to params. params_map = {} + if config_overrides is None: + # TODO remove this default behavior eventually. + # This is only needed for backwards compatibility with the old config overrides API where + # the config_overrides argument by default lead to bias parameters and length 1 parameters. + # We assume that users of decoupled LR already provide config overrides so will adapt + # to the new API. + config_overrides = get_standard_config_overrides(config=config) + for model_chunk in model_chunks: for name, param in model_chunk.named_parameters(): if not param.requires_grad: @@ -459,8 +458,7 @@ def _get_megatron_optimizer_based_on_param_groups( intra_dist_opt_group: Optional[torch.distributed.ProcessGroup] = None, distributed_optimizer_instance_id: Optional[int] = 0, pg_collection: Optional[ProcessGroupCollection] = None, - skip_megatron_wrapping: bool = False, -) -> Union[MegatronOptimizer, Tuple[Optional[torch.optim.Optimizer], Optional[Callable]]]: +) -> MegatronOptimizer: """Get Megatron optimizer based on parameter groups. Args: @@ -476,24 +474,12 @@ def _get_megatron_optimizer_based_on_param_groups( optimizer. Defaults to None. distributed_optimizer_instance_id (int, optional): Distributed optimizer instance. Defaults 0. - skip_megatron_wrapping (bool): if True, return a - ``(optimizer, init_state_fn)`` tuple of the raw PyTorch optimizer - without any Megatron wrapping. Useful when the caller - (e.g. LayerWiseDistributedOptimizer) performs its own wrapping. Returns: - Instance of MegatronOptimizer, or ``(optimizer, init_state_fn)`` when - *skip_megatron_wrapping=True*. + Instance of MegatronOptimizer. """ - # All param_groups passed here must belong to the same optimizer type (adam / sgd). - # Callers are responsible for splitting by optimizer type before calling this function. - - if skip_megatron_wrapping and config.use_precision_aware_optimizer: - raise ValueError( - "skip_megatron_wrapping=True is incompatible with use_precision_aware_optimizer." - ) - if skip_megatron_wrapping and config.optimizer_cpu_offload: - raise ValueError("skip_megatron_wrapping=True is incompatible with optimizer_cpu_offload.") + # TODO: Logic needs to be updated to handle different optimizer types (i.e., param_groups + # passed into this function need to correspond to the same optimizer). # When freezing sub-models we may have no trainable parameters on a rank and # hence an empty param_groups. However, we still need to create an optimizer @@ -628,9 +614,6 @@ def init_state_fn(opt, config=None): optimizer = None init_state_fn = None - if skip_megatron_wrapping: - return optimizer, init_state_fn - # Mixed precision optimizer. # - Note: both the Float16Optimizer and the DistributedOptimizer inherit # from the MixedPrecisionOptimizer, which manages any optimizer where @@ -721,142 +704,6 @@ def check_config_overrides_consistency( return True -def _get_megatron_emerging_optimizer( - config: OptimizerConfig, - model_chunks: List[MegatronModule], - config_overrides: Optional[Dict[ParamKey, Any]] = None, - pg_collection: Optional[ProcessGroupCollection] = None, -) -> MegatronOptimizer: - """Build an emerging optimizer (e.g. Muon) for the given model chunks. - - Parameter separation (e.g., linear weights -> Muon, rest -> Adam) is expressed as a - config_override, the same mechanism used for weight-decay and learning-rate overrides. - Adam/SGD groups are delegated to _get_megatron_optimizer_based_on_param_groups so they - go through the exact same code path as the standard optimizer factory. - - When ``config.use_layer_wise_distributed_optimizer`` is True, the underlying optimizers - are wrapped with :class:`LayerWiseDistributedOptimizer`. - """ - eopt_name = config.optimizer - use_layer_wise = config.use_layer_wise_distributed_optimizer - - # Handle legacy "dist_*" optimizer names (e.g. "dist_muon" → "muon" + layer-wise). - if eopt_name.startswith('dist_'): - bare_name = eopt_name[len('dist_') :] - warnings.warn( - f"optimizer='{eopt_name}' is deprecated. " - f"Use optimizer='{bare_name}' with use_layer_wise_distributed_optimizer=True.", - DeprecationWarning, - stacklevel=3, - ) - eopt_name = bare_name - use_layer_wise = True - - if not HAVE_EMERGING_OPTIMIZERS: - raise ImportError( - f"emerging-optimizers package is required for optimizer='{eopt_name}'. " - "Install it with: pip install emerging-optimizers" - ) - if eopt_name not in _EMERGING_OPTIMIZERS: - raise ValueError(f"Unsupported emerging optimizer: {eopt_name}") - if config.fp16: - raise ValueError('emerging optimizer with fp16 is not supported.') - - if pg_collection is None: - pg_collection = ProcessGroupCollection.use_mpu_process_groups() - - log_single_rank(logger, logging.INFO, f'Setting up emerging optimizer with config {config}') - - # Tag parameters with optimizer-specific attributes (expert_tp, is_qkv). - for model_chunk in model_chunks: - for name, param in model_chunk.named_parameters(): - if not param.requires_grad: - continue - if 'experts' in name and 'shared' not in name: - param.expert_tp = True - # TODO(deyuf): support MLA - if 'linear_qkv.weight' in name and len(param.shape) == 2: - param.is_qkv = True - - # Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam). - config_overrides.update(_EMERGING_OPTIMIZERS[eopt_name].default_param_overrides) - - # Build param groups and bucket by (optimizer_name, is_expert_parallel). - # Layer-wise distributed optimizer handles expert params internally so we skip that split. - all_param_groups = _get_param_groups(model_chunks, config, config_overrides) - grouped_param_groups = defaultdict(list) - for group in all_param_groups: - opt_name = group.get('optimizer', eopt_name) - is_expert = group['is_expert_parallel'] and not use_layer_wise - grouped_param_groups[(opt_name, is_expert)].append(group) - - # Build an optimizer for each (optimizer_name, is_expert) bucket and combine. - results = [] - for (opt_name, is_expert), groups in grouped_param_groups.items(): - if not groups: - continue - - model_parallel_group = pg_collection.tp_ep_pp if is_expert else pg_collection.mp - - if opt_name in _EMERGING_OPTIMIZERS: - optimizer, init_state_fn = _create_emerging_optimizer( - config, groups, eopt_name, model_chunks, pg_collection - ) - if use_layer_wise: - result = (optimizer, init_state_fn) - else: - if config.bf16: - optimizer = Float16OptimizerWithFloat16Params( - optimizer, config, None, init_state_fn - ) - else: - optimizer = FP32Optimizer(optimizer, config, init_state_fn) - setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group) - if pg_collection is None or not hasattr(pg_collection, 'tp'): - tp_group = parallel_state.get_tensor_model_parallel_group() - else: - tp_group = pg_collection.tp - setattr(optimizer, 'tp_group', tp_group) - result = optimizer - else: - fallback_config = copy.copy(config) - fallback_config.optimizer = opt_name - fallback_config.use_distributed_optimizer = False - result = _get_megatron_optimizer_based_on_param_groups( - config=fallback_config, - model_chunks=model_chunks, - param_groups=groups, - model_parallel_group=model_parallel_group, - pg_collection=pg_collection, - skip_megatron_wrapping=use_layer_wise, - ) - # TODO(deyuf): ChainedOptimizer currently asserts all sub-optimizers - # share the same config. Revisit this design now that emerging - # optimizers mix different optimizer types (e.g. Muon + Adam). - # For now, reset to the top-level config so the assertion holds. - if not use_layer_wise and hasattr(result, 'config'): - result.config = config - results.append(result) - - if use_layer_wise: - base_optimizers, init_fns = (), () - if results: - base_optimizers, init_fns = zip(*results) - log_single_rank( - logger, logging.INFO, f'Using LayerWiseDistributedOptimizer for {eopt_name}' - ) - return LayerWiseDistributedOptimizer( - list(base_optimizers), - config, - pg_collection, - init_state_fn_list=list(init_fns), - model_chunks=model_chunks if config.overlap_param_gather else None, - async_allgather=config.overlap_param_gather, - ) - - return ChainedOptimizer(results) - - def get_megatron_optimizer( config: OptimizerConfig, model_chunks: List[MegatronModule], @@ -867,10 +714,7 @@ def get_megatron_optimizer( ) -> MegatronOptimizer: """Retrieve the Megatron optimizer for model chunks. - Handles both standard optimizers (Adam, SGD) and emerging optimizers (e.g. Muon). We use separate optimizers for expert parameters and non-expert parameters. - For emerging optimizers with ``config.use_layer_wise_distributed_optimizer=True``, - the optimizer is automatically wrapped with :class:`LayerWiseDistributedOptimizer`. Args: config (OptimizerConfig): optimizer configuration object. @@ -887,25 +731,10 @@ def get_megatron_optimizer( Instance of MegatronOptimizer. """ - # None → apply standard defaults. To extend defaults with custom overrides, - # start from get_standard_config_overrides(config) and merge yours in. - if config_overrides is None: - config_overrides = get_standard_config_overrides(config) + log_single_rank(logger, logging.INFO, f'Setting up optimizer with config {config}') check_config_overrides_consistency(config, config_overrides) - # TODO: the standard and emerging optimizer paths handle pg_collection differently; - # unify them so both use a single pg_collection-based flow. - if config.optimizer not in ('adam', 'sgd'): - return _get_megatron_emerging_optimizer( - config=config, - model_chunks=model_chunks, - config_overrides=config_overrides, - pg_collection=pg_collection, - ) - - log_single_rank(logger, logging.INFO, f'Setting up optimizer with config {config}') - # Separate out first model chunk if overlapping param AG with optimizer step. if config.overlap_param_gather_with_optimizer_step: all_dense_model_chunks = [[model_chunks[0]], model_chunks[1:]] diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py deleted file mode 100644 index cc218d6ba40..00000000000 --- a/megatron/core/optimizer/emerging_optimizers.py +++ /dev/null @@ -1,450 +0,0 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -"""Emerging optimizer registry. - -To add a new emerging optimizer: - 1. Define its optimizer class (or import it). - 2. Write its ``__init_state_fn`` and ``__config_to_kwargs``. - 3. Add an ``EmergingOptimizerEntry`` to ``_EMERGING_OPTIMIZERS`` at the bottom. -""" - -import inspect -import logging -from dataclasses import dataclass, field -from typing import Any, Callable, Dict, List, Literal, Optional, get_args - -import torch -from torch.optim.optimizer import ParamsT - -from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.utils import get_pg_size, log_single_rank - -from .optimizer_config import ParamKey, ParamPredicate - -try: - from emerging_optimizers import registry - from emerging_optimizers.orthogonalized_optimizers import ( - AdaptiveMuon, - OrthogonalizedOptimizer, - get_muon_scale_factor, - ) - from emerging_optimizers.orthogonalized_optimizers.muon_utils import NSCoeffT, newton_schulz_tp - - # It is necessary to import optimizers for the registry to work. - from emerging_optimizers.scalar_optimizers import Lion # pylint: disable=unused-import - from emerging_optimizers.soap import SOAP # pylint: disable=unused-import - - HAVE_EMERGING_OPTIMIZERS = True -except ImportError: - HAVE_EMERGING_OPTIMIZERS = False - OrthogonalizedOptimizer = object - AdaptiveMuon = object - - -logger = logging.getLogger(__name__) - - -def get_supported_coefficient_types() -> tuple[str, ...]: - """Return the coefficient types supported by the installed emerging_optimizers. - - Reads the members of the ``NSCoeffT`` Literal type so that new types - added upstream are automatically available without code changes here. - """ - assert ( - HAVE_EMERGING_OPTIMIZERS - ), "emerging_optimizers >= 0.2 is required for NSCoeffT. Please install or upgrade it." - return get_args(NSCoeffT) - - -def validate_coefficient_type(coefficient_type: str) -> None: - """Raise ``ValueError`` if *coefficient_type* is not supported.""" - supported = get_supported_coefficient_types() - if coefficient_type not in supported: - raise ValueError( - f"Unsupported muon coefficient type '{coefficient_type}'. " - f"Supported types: {supported}" - ) - - -# =========================================================================== -# Registry dataclass and public API -# =========================================================================== - - -def _eopt_init_state_fn(opt, config=None): - """Initialize emerging optimizer state for torch_dist checkpoint format.""" - for group in opt.param_groups: - # Checkpoint init needs state for all parameters, including those without grads yet. - opt._init_group(group, skip_non_grad_params=False) - - -def _default_param_overrides_factory() -> Dict[ParamKey, Dict[str, Any]]: - """Default param overrides: route non-linear/embedding params to Adam.""" - return { - ParamKey( - predicate=ParamPredicate(name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding) - ): {'optimizer': 'adam'} - } - - -@dataclass -class EmergingOptimizerEntry: - """Everything needed to create and configure an emerging optimizer. - - Attributes: - optimizer_cls: The torch optimizer class. - init_state_fn: Lazily initialises optimizer state (needed for checkpoint formats). - config_to_kwargs: ``(config, model_chunks, pg_collection) -> dict`` of constructor kwargs. - default_param_overrides: Per-parameter config overrides applied automatically - (e.g. route non-linear params to Adam). - """ - - optimizer_cls: type - init_state_fn: Callable = _eopt_init_state_fn - config_to_kwargs: Callable | None = None - default_param_overrides: Dict[ParamKey, Dict[str, Any]] = field( - default_factory=_default_param_overrides_factory - ) - - -def _create_emerging_optimizer(config, param_groups, eopt_name, model_chunks, pg_collection): - """Instantiate an emerging optimizer and return it with its init_state_fn.""" - entry = _EMERGING_OPTIMIZERS[eopt_name] - if entry.config_to_kwargs is not None: - eopt_kwargs = entry.config_to_kwargs(config, model_chunks, pg_collection) - else: - eopt_kwargs = _default_adam_based_eopt_config_to_kwargs( - eopt_name, config, model_chunks, pg_collection - ) - optimizer = entry.optimizer_cls(param_groups, **eopt_kwargs) - return optimizer, entry.init_state_fn - - -# =========================================================================== -# Shared helpers -# =========================================================================== - - -def _is_nonlinear_or_embedding(param): - """True for parameters that should NOT use the emerging optimizer.""" - return getattr(param, 'is_embedding_or_output_parameter', False) or len(param.shape) != 2 - - -def _get_qkv_split_shapes(model_cfg) -> List[int]: - """Compute QKV split shapes from model config.""" - return [ - model_cfg.num_attention_heads // model_cfg.num_query_groups * model_cfg.kv_channels, - model_cfg.kv_channels, - model_cfg.kv_channels, - ] - - -# =========================================================================== -# Registry – populated below only when emerging_optimizers is installed. -# =========================================================================== - -_EMERGING_OPTIMIZERS: Dict[str, EmergingOptimizerEntry] = {} - - -# =========================================================================== -# Muon -# =========================================================================== - - -class TensorParallelMuon(OrthogonalizedOptimizer): - """Tensor Parallel Muon optimizer.""" - - def __init__( - self, - params: ParamsT, - lr: float = 3e-4, - momentum: float = 0.95, - nesterov: bool = True, - weight_decay: float = 0.01, - use_decoupled_weight_decay: bool = True, - split_qkv: bool = False, - is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: tuple[int, int, int] | None = None, - fp32_matmul_prec: str = "medium", - coefficient_type: str = "quintic", - num_ns_steps: int = 5, - scale_mode: str = "spectral", - extra_scale_factor: float = 1.0, - pg_collection: Optional[ProcessGroupCollection] = None, - tp_mode: Literal["blockwise", "duplicated", "distributed"] = "duplicated", - ) -> None: - if num_ns_steps < 1: - raise ValueError(f"num_ns_steps must be at least 1, got {num_ns_steps}") - - def scaled_orthogonalize_fn( - grad: torch.Tensor, - tp_group: torch.distributed.ProcessGroup, - partition_dim: int | None = None, - ) -> torch.Tensor: - log_single_rank( - logger, - logging.DEBUG, - f'Orthogonalizing grad with {num_ns_steps} steps, ' - f'{coefficient_type} coefficient, ' - f'{scale_mode} scale mode, extra_scale_factor={extra_scale_factor}', - ) - size = [grad.size(-2), grad.size(-1)] - if partition_dim is not None: - size[partition_dim] *= get_pg_size(tp_group) - orth_grad = newton_schulz_tp( - grad, - steps=num_ns_steps, - coefficient_type=coefficient_type, - tp_group=tp_group, - partition_dim=partition_dim, - tp_mode="duplicated" if tp_mode == "blockwise" else tp_mode, - ) - scale_factor = get_muon_scale_factor(size[0], size[1], mode=scale_mode) - return orth_grad * scale_factor * extra_scale_factor - - self.pg_collection = pg_collection - self.tp_mode = tp_mode - self.split_qkv = split_qkv - self.is_qkv_fn = is_qkv_fn - self.qkv_split_shapes = qkv_split_shapes - - weight_decay_method = "decoupled" if use_decoupled_weight_decay else "l2" - # Use explicit class call instead of super() so that subclasses with - # multiple inheritance (e.g. TensorParallelAdaptiveMuon) don't route - # through an intermediate class that doesn't accept scaled_orthogonalize_fn. - OrthogonalizedOptimizer.__init__( - self, - params, - lr, - momentum, - nesterov=nesterov, - weight_decay=weight_decay, - weight_decay_method=weight_decay_method, - fp32_matmul_prec=fp32_matmul_prec, - scaled_orthogonalize_fn=scaled_orthogonalize_fn, - ) - - def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> torch.Tensor: - """Orthogonalize the momentum. - - Args: - p: The parameter tensor. i is necessary to pass param tensor in addition to - momentum because a lot of information is only available in the param tensor, - attributes for example. - grad: The momentum tensor. - - Returns: - The orthogonalized gradient tensor. - """ - # TODO(deyuf): switch to group - if self.pg_collection: - tp_group = ( - self.pg_collection.expt_tp - if getattr(p, 'expert_tp', False) - else self.pg_collection.tp - ) - else: - tp_group = None - partition_dim = None if self.tp_mode == "blockwise" else getattr(p, "partition_dim", None) - if partition_dim == -1: - partition_dim = None - - if self.split_qkv and self.is_qkv_fn(p): # type: ignore[misc] - grad_shape = grad.shape - log_single_rank( - logger, - logging.DEBUG, - f'qkv split grad shape {grad_shape}, ' f'split shapes {self.qkv_split_shapes}', - ) - num_query_groups = grad_shape[0] // sum(self.qkv_split_shapes) - qkv_grads = torch.split( - grad.view(num_query_groups, sum(self.qkv_split_shapes), -1), - self.qkv_split_shapes, - dim=1, - ) - qkv_grads = [g.reshape(-1, grad_shape[-1]) for g in qkv_grads] - - qkv_grads = [ - self.scaled_orthogonalize_fn(g, tp_group, partition_dim).view( - num_query_groups, -1, grad_shape[-1] - ) - for g in qkv_grads - ] - grad = torch.cat(qkv_grads, dim=1).view(grad_shape) - else: - grad = self.scaled_orthogonalize_fn(grad, tp_group, partition_dim) - return grad - - -class TensorParallelAdaptiveMuon(TensorParallelMuon, AdaptiveMuon): - """Tensor Parallel Adaptive Muon optimizer. - - This class extends Muon by adding AdamW-style or NorMuon-style second moment - accumulation after orthogonalization. This idea was first explored in D.E. Carlson, - E. Collins, Ya-Ping Hsieh, L. Carin, and V. Cevher. *Preconditioned spectral - descent for deep learning.* In Advances in neural information processing systems 28 (2015). - The step() method is overridden to include second moment normalization logic. - - Args: - params: Iterable of parameters to optimize or dicts defining parameter groups. - lr: Learning rate. - momentum: The exponential decay rate for momentum. - nesterov: Whether to use Nesterov momentum. - weight_decay: Weight decay coefficient. - use_decoupled_weight_decay: Whether to use decoupled weight decay. - split_qkv: Whether to split QKV weights for orthogonalization. - is_qkv_fn: Function to determine if a tensor is a QKV weight. - qkv_split_shapes: Shapes for splitting QKV weights. - fp32_matmul_prec: Precision for FP32 matrix multiplication. - coefficient_type: The type of coefficient set to use for the Newton-Schulz iteration. - num_ns_steps: The number of iteration steps to use in the Newton-Schulz iteration. - scale_mode: The type of scale factor to use for the update. - extra_scale_factor: The additional scale factor to use for the update. - pg_collection: Process group collection for distributed training. - tp_mode: Tensor parallel mode ("blockwise", "duplicated", or "distributed"). - moment2_method: Method for second moment accumulation ("adamuon" or "normuon"). - beta2: The exponential decay rate for second moment. - eps: Small constant for numerical stability. - """ - - def __init__( - self, - params: ParamsT, - lr: float = 3e-4, - momentum: float = 0.95, - nesterov: bool = True, - weight_decay: float = 0.01, - use_decoupled_weight_decay: bool = True, - split_qkv: bool = False, - is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: tuple[int, int, int] | None = None, - fp32_matmul_prec: str = "medium", - coefficient_type: str = "quintic", - num_ns_steps: int = 5, - scale_mode: str = "spectral", - extra_scale_factor: float = 1.0, - pg_collection: Optional[ProcessGroupCollection] = None, - tp_mode: Literal["blockwise", "duplicated", "distributed"] = "duplicated", - moment2_method: Literal["adamuon", "normuon"] = "adamuon", - beta2: float = 0.95, - eps: float = 1e-8, - ) -> None: - TensorParallelMuon.__init__( - self, - params, - lr=lr, - momentum=momentum, - nesterov=nesterov, - weight_decay=weight_decay, - use_decoupled_weight_decay=use_decoupled_weight_decay, - split_qkv=split_qkv, - is_qkv_fn=is_qkv_fn, - qkv_split_shapes=qkv_split_shapes, - fp32_matmul_prec=fp32_matmul_prec, - coefficient_type=coefficient_type, - num_ns_steps=num_ns_steps, - scale_mode=scale_mode, - extra_scale_factor=extra_scale_factor, - pg_collection=pg_collection, - tp_mode=tp_mode, - ) - self.moment2_method = moment2_method - - for group in self.param_groups: - group.setdefault("beta2", beta2) - group.setdefault("eps", eps) - - @torch.no_grad() # type: ignore[misc] - def step(self, closure: Optional[Callable] = None) -> Optional[float]: - """Step function""" - return AdaptiveMuon.step(self, closure) - - -def _kwargs_from_config(optimizer_cls: type, prefix: str, config) -> Dict[str, Any]: - """Match ``optimizer_cls.__init__`` parameters to config attributes. - - For each init parameter, looks for ``{prefix}_{name}`` on *config* first, - then falls back to ``{name}`` (unprefixed). ``self`` and ``params`` are - always skipped. - """ - skip_params = {"self", "params"} - sig = inspect.signature(optimizer_cls.__init__) - kwargs: Dict[str, Any] = {} - for name in sig.parameters: - if name in skip_params: - continue - prefixed = f"{prefix}_{name}" - if hasattr(config, prefixed): - kwargs[name] = getattr(config, prefixed) - elif hasattr(config, name): - kwargs[name] = getattr(config, name) - return kwargs - - -def _muon_config_to_kwargs(config, model_chunks, pg_collection) -> Dict[str, Any]: - """Convert OptimizerConfig to TensorParallelMuon constructor kwargs.""" - kwargs = _kwargs_from_config(TensorParallelMuon, "muon", config) - kwargs["is_qkv_fn"] = lambda p: getattr(p, "is_qkv", False) - kwargs["qkv_split_shapes"] = _get_qkv_split_shapes(model_chunks[0].config) - kwargs["pg_collection"] = pg_collection - return kwargs - - -def _adaptive_muon_config_to_kwargs(config, model_chunks, pg_collection) -> Dict[str, Any]: - """Convert OptimizerConfig to TensorParallelAdaptiveMuon constructor kwargs.""" - kwargs = _muon_config_to_kwargs(config, model_chunks, pg_collection) - kwargs.update(_kwargs_from_config(TensorParallelAdaptiveMuon, "adaptive_muon", config)) - return kwargs - - -def _default_adam_based_eopt_config_to_kwargs( - eopt_name, config, model_chunks, pg_collection -) -> Dict[str, Any]: - """Convert OptimizerConfig to default emerging optimizer constructor kwargs.""" - kwargs = _kwargs_from_config(registry.get_optimizer_cls(eopt_name), eopt_name, config) - kwargs["betas"] = (config.adam_beta1, config.adam_beta2) - return kwargs - - -# ----------------------------------------------------------------------- -# Register emerging optimizers -# ----------------------------------------------------------------------- -_EMERGING_OPTIMIZERS.update( - { - 'muon': EmergingOptimizerEntry( - optimizer_cls=TensorParallelMuon, - init_state_fn=_eopt_init_state_fn, - config_to_kwargs=_muon_config_to_kwargs, - default_param_overrides={ - ParamKey( - predicate=ParamPredicate( - name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding - ) - ): {'optimizer': 'adam'} - }, - ), - "adaptive_muon": EmergingOptimizerEntry( - optimizer_cls=TensorParallelAdaptiveMuon, - init_state_fn=_eopt_init_state_fn, - config_to_kwargs=_adaptive_muon_config_to_kwargs, - default_param_overrides={ - ParamKey( - predicate=ParamPredicate( - name="nonlinear_or_embedding", fn=_is_nonlinear_or_embedding - ) - ): {'optimizer': 'adam'} - }, - ), - } -) - -# Register soap with default config -# TODO(skyw): register all emerging optimizers. -if HAVE_EMERGING_OPTIMIZERS: - for eopt_name in registry.get_optimizer_name_list(): - if eopt_name in _EMERGING_OPTIMIZERS: - # skip already registered local versions, e.g. TensorParallel versions. - continue - _EMERGING_OPTIMIZERS[eopt_name] = EmergingOptimizerEntry( - optimizer_cls=registry.get_optimizer_cls(eopt_name) - ) diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index 6e59e03ae42..a9fdc7ba72f 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -76,17 +76,19 @@ def __init__( optimizers ), "init_state_fn_list must be the same length as optimizers if provided" - # Wrap base torch optimizers with Float16 for bf16 training. - # Callers pass base optimizers; wrapping happens here *after* - # shard_params so master weights are only created for the local shard. + # wrap optimizer after sharding to avoid unnecessary master weight creation + # for higher precision, optimizers are wrapped with megatron already if config.bf16: + # unwrap FP32 optimizer, possibly from reusing get_megatron_optimizer for adam for i in range(len(optimizers)): opt = optimizers[i] - if isinstance(opt, (Float16OptimizerWithFloat16Params, FP32Optimizer)): + if isinstance(opt, Float16OptimizerWithFloat16Params): raise TypeError( - 'LayerWiseDistributedOptimizer expects base torch optimizers, ' - f'got {type(opt).__name__}. Do not pre-wrap with Megatron optimizers.' + 'LayerWiseDistributedOptimizer received Float16 optimizer already.' ) + # unwrap FP32 optimizer from reusing get_megatron_optimizer for adam + if isinstance(opt, FP32Optimizer): + opt = opt.optimizer optimizers[i] = Float16OptimizerWithFloat16Params( opt, config, None, init_state_fn_list[i] if init_state_fn_list else None ) diff --git a/megatron/core/optimizer/muon.py b/megatron/core/optimizer/muon.py index a3f7506f941..5dfddd786f5 100644 --- a/megatron/core/optimizer/muon.py +++ b/megatron/core/optimizer/muon.py @@ -1,16 +1,394 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Backward-compatible shim — all code now lives in ``emerging_optimizers``.""" +"""Megatron muon optimizer wrapper to handle tensor-parallel.""" -from typing import Any +import logging +from typing import Any, Callable, Dict, List, Literal, Optional, get_args +import torch +from torch.optim.optimizer import ParamsT -def get_megatron_muon_optimizer(*args: Any, **kwargs: Any) -> Any: - """Backward compatible muon optimizer getter. +from megatron.core.optimizer_param_scheduler import ParamGroupOverride +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.module import MegatronModule +from megatron.core.utils import get_pg_size, log_single_rank - .. deprecated:: - Use :func:`megatron.core.optimizer.get_megatron_optimizer` instead. +from . import HAVE_EMERGING_OPTIMIZERS, _get_param_groups, get_megatron_optimizer +from .layer_wise_optimizer import LayerWiseDistributedOptimizer +from .optimizer import ( + ChainedOptimizer, + Float16OptimizerWithFloat16Params, + FP32Optimizer, + MegatronOptimizer, +) +from .optimizer_config import OptimizerConfig, ParamKey + +if HAVE_EMERGING_OPTIMIZERS: + from emerging_optimizers.orthogonalized_optimizers import ( + OrthogonalizedOptimizer, + get_muon_scale_factor, + ) + from emerging_optimizers.orthogonalized_optimizers.muon_utils import newton_schulz_tp +else: + OrthogonalizedOptimizer = object + +if HAVE_EMERGING_OPTIMIZERS: + from emerging_optimizers.orthogonalized_optimizers.muon_utils import NSCoeffT + + +logger = logging.getLogger(__name__) + + +def get_supported_coefficient_types() -> tuple[str, ...]: + """Return the coefficient types supported by the installed emerging_optimizers. + + Reads the members of the ``NSCoeffT`` Literal type so that new types + added upstream are automatically available without code changes here. + """ + assert ( + HAVE_EMERGING_OPTIMIZERS + ), "emerging_optimizers >= 0.2 is required for NSCoeffT. Please install or upgrade it." + return get_args(NSCoeffT) # pylint: disable=possibly-used-before-assignment + + +def validate_coefficient_type(coefficient_type: str) -> None: + """Raise ``ValueError`` if *coefficient_type* is not supported.""" + supported = get_supported_coefficient_types() + if coefficient_type not in supported: + raise ValueError( + f"Unsupported muon coefficient type '{coefficient_type}'. " + f"Supported types: {supported}" + ) + + +class TensorParallelMuon(OrthogonalizedOptimizer): + """Tensor Parallel Muon optimizer.""" + + def __init__( + self, + params: ParamsT, + lr: float = 3e-4, + momentum_beta: float = 0.95, + use_nesterov: bool = True, + weight_decay: float = 0.01, + use_decoupled_weight_decay: bool = True, + split_qkv: bool = False, + is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, + qkv_split_shapes: tuple[int, int, int] | None = None, + fp32_matmul_prec: str = "medium", + coefficient_type: str = "quintic", + num_ns_steps: int = 5, + scale_mode: str = "spectral", + extra_scale_factor: float = 1.0, + pg_collection: Optional[ProcessGroupCollection] = None, + mode: Literal["blockwise", "duplicated", "distributed"] = "duplicated", + ) -> None: + if num_ns_steps < 1: + raise ValueError(f"num_ns_steps must be at least 1, got {num_ns_steps}") + validate_coefficient_type(coefficient_type) + + def scaled_orthogonalize_fn( + grad: torch.Tensor, + tp_group: torch.distributed.ProcessGroup, + partition_dim: int | None = None, + ) -> torch.Tensor: + log_single_rank( + logger, + logging.DEBUG, + f'Orthogonalizing grad with {num_ns_steps} steps, {coefficient_type} coefficient, ' + f'{scale_mode} scale mode, extra_scale_factor={extra_scale_factor}', + ) + size = [grad.size(-2), grad.size(-1)] + if partition_dim is not None: + size[partition_dim] *= get_pg_size(tp_group) + mode_value = "duplicated" if mode == "blockwise" else mode + mode_kwarg = {"tp_mode": mode_value} + ns_kwargs = dict( + steps=num_ns_steps, tp_group=tp_group, partition_dim=partition_dim, **mode_kwarg + ) + ns_kwargs["coefficient_type"] = coefficient_type + # pylint: disable-next=possibly-used-before-assignment + orth_grad = newton_schulz_tp(grad, **ns_kwargs) + # pylint: disable-next=possibly-used-before-assignment + scale_factor = get_muon_scale_factor(size[0], size[1], mode=scale_mode) + return orth_grad * scale_factor * extra_scale_factor + + self.pg_collection = pg_collection + self.mode = mode + self.split_qkv = split_qkv + self.is_qkv_fn = is_qkv_fn + self.qkv_split_shapes = qkv_split_shapes + + weight_decay_method = "decoupled" if use_decoupled_weight_decay else "l2" + nesterov_kwarg = {"nesterov": use_nesterov} + super().__init__( + params, + lr, + momentum_beta, + **nesterov_kwarg, + weight_decay=weight_decay, + weight_decay_method=weight_decay_method, + fp32_matmul_prec=fp32_matmul_prec, + scaled_orthogonalize_fn=scaled_orthogonalize_fn, + ) + + def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> torch.Tensor: + """Orthogonalize the momentum. + + Args: + p: The parameter tensor. i is necessary to pass param tensor in addition to momentum + because a lot of information is only available in the param tensor, + attributes for example. + grad: The momentum tensor. + + Returns: + The orthogonalized gradient tensor. + """ + # TODO(deyuf): switch to group + if self.pg_collection: + tp_group = ( + self.pg_collection.expt_tp + if getattr(p, 'expert_tp', False) + else self.pg_collection.tp + ) + else: + tp_group = None + partition_dim = None if self.mode == "blockwise" else getattr(p, "partition_dim", None) + if partition_dim == -1: + # emerging-optimizers use None instead of -1 to indicate no tensor parallel + partition_dim = None + + if self.split_qkv and self.is_qkv_fn(p): # type: ignore[misc] + # split grouped attention parameters (e.g., QKV, GQA, etc.) + grad_shape = grad.shape + log_single_rank( + logger, + logging.DEBUG, + f'qkv split grad shape {grad_shape}, split shapes {self.qkv_split_shapes}', + ) + num_query_groups = grad_shape[0] // sum(self.qkv_split_shapes) + qkv_grads = torch.split( + grad.view(num_query_groups, sum(self.qkv_split_shapes), -1), + self.qkv_split_shapes, + dim=1, + ) + qkv_grads = [g.reshape(-1, grad_shape[-1]) for g in qkv_grads] + + # Apply Newton-Schulz and scales to each component, concat back + qkv_grads = [ + self.scaled_orthogonalize_fn(g, tp_group, partition_dim).view( + num_query_groups, -1, grad_shape[-1] + ) + for g in qkv_grads + ] + grad = torch.cat(qkv_grads, dim=1).view(grad_shape) + else: + grad = self.scaled_orthogonalize_fn(grad, tp_group, partition_dim) + return grad + + +def get_megatron_muon_optimizer( + config: OptimizerConfig, + model_chunks: List[MegatronModule], + config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]] = None, + use_gloo_process_groups: bool = True, + layer_wise_distributed_optimizer: bool = False, + pg_collection: Optional[ProcessGroupCollection] = None, +) -> MegatronOptimizer: + """This function is used to get the muon optimizer for the model chunks. + It is used to get the muon optimizer for the model chunks. + + Args: + config (OptimizerConfig): optimizer configuration object. + model_chunks (List[MegatronModule]): model chunks to get optimizer for. + use_gloo_process_groups (bool): if false, disable use of Gloo process groups + in underlying Megatron optimizers. + layer_wise_distributed_optimizer (bool): if true, use layer-wise distributed optimizer. + Defaults to False. """ - from . import get_megatron_optimizer + # TODO: Mutating config.optimizer is a side effect; clean up after + # https://github.com/NVIDIA/Megatron-LM/pull/3638 lands. + # Set the nonlinear optimizer for muon (used for embeddings, biases, norms). + config.optimizer = config.muon_scalar_optimizer + + assert HAVE_EMERGING_OPTIMIZERS, "Emerging Optimizers >= 0.2 is not installed." + + # Dist-opt is not supported due to strong coupling with how DDP init grad buffer + # In theory we can change DDP to enable use muon and dist-opt-adam together + if config.use_distributed_optimizer: + raise Exception('muon with dist optimizer is not supported.') + # only support bf16 w/o loss scale now + if config.fp16: + raise Exception('muon with fp16 is not supported.') + + # before this function receive properly created collection + if pg_collection is None: + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + + log_single_rank(logger, logging.INFO, f'Setting up emerging optimizer with config {config}') + + # Needed for torch_dist ckpt_format, unlike torch ckpt_format + # For other emerging optimizers, need to implement init_state_fn as well + # TODO(boxiangw): Improve usability after optimizer refactor + # TODO(boxiangw): support precision aware optimizer + def muon_init_state_fn(opt, config=None): + for group in opt.param_groups: + for p in group['params']: + if len(opt.state[p]) == 0: + opt.state[p]['momentum_buffer'] = torch.zeros_like(p.data) + + def adam_init_state_fn(opt, config=None): + for group in opt.param_groups: + for p in group['params']: + if len(opt.state[p]) == 0: + if config is None or not config.use_precision_aware_optimizer: + opt.state[p]['exp_avg'] = torch.zeros_like(p.data) + opt.state[p]['exp_avg_sq'] = torch.zeros_like(p.data) + else: + opt.initialize_state(p) + + def lion_init_state_fn(opt, config=None): + for group in opt.param_groups: + for p in group['params']: + if len(opt.state[p]) == 0: + opt.state[p]['exp_avg'] = torch.zeros_like(p.data) + + nonlinear_init_state_fn = ( + lion_init_state_fn if config.muon_scalar_optimizer == 'lion' else adam_init_state_fn + ) + + optimizers = [] + # record list of non/linear params + linear_params = [] + nonlinear_params = [] + for model_chunk in model_chunks: + # use config to determine qkv split shapes. + # no need to check tp since tp splits by head and this is per head(group) dimension + num_attention_heads = model_chunk.config.num_attention_heads + num_query_groups = model_chunk.config.num_query_groups + kv_channels = model_chunk.config.kv_channels + qkv_split_shapes = [ + num_attention_heads // num_query_groups * kv_channels, + kv_channels, + kv_channels, + ] + for name, param in model_chunk.named_parameters(): + if not param.requires_grad: + continue + # add flag for expert weight so optimizer can figure which tp group it uses + # alternatively, create new param group and save tp_group. this require more + # change in optimizer + if 'experts' in name and 'shared' not in name: + param.expert_tp = True + # add flag for qkv parameter + # TODO(deyuf): support MLA + if 'linear_qkv.weight' in name and len(param.shape) == 2: + param.is_qkv = True + # TODO(deyuf): currently only allow 2D non-embedding weight to avoid breaking + if ( + not getattr(param, 'is_embedding_or_output_parameter', False) + and len(param.shape) == 2 + ): + linear_params.append(param) + else: + nonlinear_params.append(param) + + muon_kwargs = { + "lr": config.lr, + "momentum_beta": config.muon_momentum, + "use_nesterov": config.muon_use_nesterov, + "weight_decay": config.weight_decay, + "fp32_matmul_prec": config.muon_fp32_matmul_prec, + "coefficient_type": config.muon_coefficient_type, + "num_ns_steps": config.muon_num_ns_steps, + "scale_mode": config.muon_scale_mode, + "split_qkv": config.muon_split_qkv, + "is_qkv_fn": lambda p: getattr(p, "is_qkv", False), + "qkv_split_shapes": qkv_split_shapes, + "extra_scale_factor": config.muon_extra_scale_factor, + "pg_collection": pg_collection, + "mode": config.muon_tp_mode, + } + + # freezing nonlinear params and get param groups for muon + for param in nonlinear_params: + param.requires_grad = False + + linear_param_groups = _get_param_groups(model_chunks, config, config_overrides) + # if layerwise distributed optimizer is not used, need to handle ep params separately + expert_param_groups = [] + if not layer_wise_distributed_optimizer: + for group in linear_param_groups: + if group['is_expert_parallel']: + expert_param_groups.append(group) + linear_param_groups.remove(group) + + optimizer = TensorParallelMuon(linear_param_groups, **muon_kwargs) + + reset_config_bf16 = False + if config.bf16: + if layer_wise_distributed_optimizer: + # creating master weight before layerwise sharding will lead to unnecessary master + # weight so here we delay master weight creation into layer_wise unset config.bf16 + # will also result in all optimizers below(adam) to also not be wrapped + config.bf16 = False + reset_config_bf16 = True + else: + # if not using layer_wise wrapper, just create master weight here is fine + optimizer = Float16OptimizerWithFloat16Params( + optimizer, config, None, muon_init_state_fn + ) + else: + optimizer = FP32Optimizer(optimizer, config, muon_init_state_fn) + + optimizers.append(optimizer) + + # expert optimizer exists meaning layerwise distributed optimizer is not used + if len(expert_param_groups) > 0: + expert_optimizer = TensorParallelMuon(expert_param_groups, **muon_kwargs) + if config.bf16: + expert_optimizer = Float16OptimizerWithFloat16Params( + expert_optimizer, config, None, muon_init_state_fn + ) + else: + expert_optimizer = FP32Optimizer(expert_optimizer, config, muon_init_state_fn) + setattr(expert_optimizer, 'grad_stats_parallel_group', pg_collection.tp_ep_pp) + optimizers.append(expert_optimizer) + + # done with muon, unfreeze nonlinear and freeze linear + for param in nonlinear_params: + param.requires_grad = True + for param in linear_params: + param.requires_grad = False + + # call original get. linear params will be skipped since they're freezed + chained_adam = get_megatron_optimizer( + config, + model_chunks, + config_overrides=config_overrides, + use_gloo_process_groups=use_gloo_process_groups, + ) + + # unfreeze everything + for param in linear_params: + param.requires_grad = True + + # chain everything together + init_fns = [muon_init_state_fn] + len(chained_adam.chained_optimizers) * [ + nonlinear_init_state_fn + ] + optimizers += chained_adam.chained_optimizers - return get_megatron_optimizer(*args, **kwargs) + if layer_wise_distributed_optimizer: + log_single_rank(logger, logging.INFO, 'Using LayerWiseDistributedOptimizer for Muon') + if reset_config_bf16: + config.bf16 = True + return LayerWiseDistributedOptimizer( + optimizers, + config, + pg_collection, + init_state_fn_list=init_fns, + model_chunks=model_chunks, + async_allgather=config.overlap_param_gather, + ) + return ChainedOptimizer(optimizers) diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 514d109ddc8..fc4056f417e 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -207,8 +207,7 @@ class OptimizerConfig: """dtype of exp_avg_sq when enabling precision-aware-optimizer""" optimizer: str = 'adam' - """Optimizer name (e.g., 'adam', 'sgd', 'muon'). Can be overridden per-parameter group - via config_overrides to use different optimizers for different parameters.""" + """Optimizer name. NOTE: Deprecated, use individual optimizer classes instead.""" ############### # Loss scaling @@ -231,7 +230,7 @@ class OptimizerConfig: """Hysteresis for dynamic loss scaling.""" ################################################################################### - # Optimizer-specific parameters. + # Optimizer (NOTE: Deprecated, use individual optimizer classes instead.). ################################################################################### # Adam. adam_beta1: float = 0.9 @@ -256,14 +255,15 @@ class OptimizerConfig: sgd_momentum: float = 0.9 """Momentum factor for SGD optimizer.""" - # emerging optimizers. + # Muon. + # TODO: move muon configs to it's own `MuonConfig`. muon_momentum: float = 0.95 - """The momentum used by the internal SGD in Muon optimizer.""" + """The momentum used by the internal SGD.""" muon_split_qkv: bool = True """Whether to split QKV parameters for Muon optimizer.""" - muon_nesterov: bool = False + muon_use_nesterov: bool = False """Whether to use Nesterov-style momentum in the internal SGD.""" muon_scale_mode: str = "spectral" @@ -297,36 +297,12 @@ class OptimizerConfig: """Second beta coefficient for Lion optimizer (used in momentum EMA update). Defaults to 0.98.""" - soap_shampoo_beta: float = 0.95 - """The beta parameter for the Shampoo preconditioner.""" - - soap_precondition_frequency: int = 1 - """The frequency of the Shampoo preconditioner.""" - - soap_use_kl_shampoo: bool = True - """Whether to use the KL-Shampoo preconditioner.""" - - adaptive_muon_moment2_method: str = "adamuon" - """The method to use for the moment2 update in Adaptive Muon optimizer.""" - - adaptive_muon_beta2: float = 0.95 - """The beta2 parameter for the Adaptive Muon optimizer.""" - - adaptive_muon_eps: float = 1e-8 - """The eps parameter for the Adaptive Muon optimizer.""" - ####################### # Distributed optimizer ####################### use_distributed_optimizer: bool = False """Distribute optimizer state over data-parallel replicas.""" - use_layer_wise_distributed_optimizer: bool = False - """Use :class:`LayerWiseDistributedOptimizer` for emerging optimizers (e.g. Muon). - When set via ``--use-distributed-optimizer`` with an emerging optimizer, the training - arguments layer sets this flag and resets ``use_distributed_optimizer`` to False so - that the standard distributed-optimizer path is not triggered.""" - overlap_param_gather: bool = False """If true, overlap param all-gather with forward compute. This argument is intended to have the same value as the "overlap_param_gather" argument @@ -469,6 +445,33 @@ def __post_init__(self): ), "exp_avg_sq_dtype can only be fp32 when not using precision-aware optimizer" -# Backward-compatible aliases (deprecated; use OptimizerConfig directly). -AdamOptimizerConfig = OptimizerConfig -SGDOptimizerConfig = OptimizerConfig +@dataclass +class AdamOptimizerConfig(OptimizerConfig): + """Adam optimizer configuration object.""" + + optimizer: str = 'adam' + """Optimizer name.""" + + adam_beta1: float = 0.9 + """First coefficient for computing running averages of gradient and its square in Adam + optimizer. + """ + + adam_beta2: float = 0.999 + """Second coefficient for computing running averages of gradient and its square in Adam + optimizer. + """ + + adam_eps: float = 1e-08 + """Term added to the denominator to improve numerical stability in Adam optimizer.""" + + +@dataclass +class SGDOptimizerConfig(OptimizerConfig): + """SGD optimizer configuration object.""" + + optimizer: str = 'sgd' + """Optimizer name.""" + + sgd_momentum: float = 0.9 + """Momentum factor for SGD optimizer.""" diff --git a/megatron/core/optimizer_param_scheduler.py b/megatron/core/optimizer_param_scheduler.py index 53a2c1a3951..d8a7703bad2 100644 --- a/megatron/core/optimizer_param_scheduler.py +++ b/megatron/core/optimizer_param_scheduler.py @@ -16,7 +16,7 @@ logger = logging.getLogger(__name__) -class ParamGroupOverride(TypedDict, total=False): +class ParamGroupOverride(TypedDict): """Override values for a parameter group. These values may be optimizer-state/scheduler related. These are the values you see later in param_group.get(...) calls in the @@ -25,7 +25,7 @@ class ParamGroupOverride(TypedDict, total=False): Example: >>> param_group_override = ParamGroupOverride(min_lr=1e-4, wd_mult=0.1) - >>> param_group_override == ParamGroupOverride(optimizer='muon') # per-param optimizer + >>> param_group_override == ParamGroupOverride(newvar=3) # this is ok too """ @@ -34,7 +34,6 @@ class ParamGroupOverride(TypedDict, total=False): start_wd: float end_wd: float wd_mult: float - optimizer: str def get_canonical_lr_for_logging(param_groups: list[dict]) -> float | None: diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 0bfe6142d01..af4ffd80192 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -1472,20 +1472,14 @@ def validate_args(args, defaults={}): '--no-load-optim with --skip-train --perform-rl-step skips the optimizer; ' \ '--rl-offload-optimizer-during-inference is incompatible (no optimizer to offload).' - # emerging optimizer check - if args.optimizer not in ('sgd', 'adam'): - if args.optimizer == 'dist_muon': - warn_rank_0( - "optimizer='dist_muon' is deprecated. " - "Use --optimizer muon --use-distributed-optimizer instead." - ) - args.optimizer = 'muon' - args.use_layer_wise_distributed_optimizer = True + # Muon optimizer check + if 'muon' in args.optimizer: - if args.use_distributed_optimizer: - args.use_layer_wise_distributed_optimizer = True - args.use_distributed_optimizer = False + if args.optimizer == 'muon': + assert not args.overlap_grad_reduce, "Muon optimizer does not support overlap grad reduce. Use dist_muon instead." + assert not args.overlap_param_gather, "Muon optimizer does not support overlap param gather. Use dist_muon instead." + assert not args.use_distributed_optimizer, "Muon optimizer does not support distributed optimizer for now." assert not args.use_torch_fsdp2, "Muon optimizer does not support Torch-FSDP2 for now." assert not args.use_megatron_fsdp, "Muon optimizer does not support Megatron-FSDP for now." assert args.ckpt_format in ["torch", "torch_dist"], "Muon optimizer supports torch and torch_dist checkpoint format." @@ -2234,7 +2228,7 @@ def _add_regularization_args(parser): group.add_argument('--muon-no-split-qkv', action='store_false', default=True, dest='muon_split_qkv', help='Whether to split QKV parameters for Muon optimizer') - group.add_argument('--muon-nesterov', action='store_true', + group.add_argument('--muon-use-nesterov', action='store_true', help='Whether to use Nesterov-style momentum in the internal SGD') group.add_argument('--muon-scale-mode', type=str, default='spectral', choices=['spectral', 'unit_rms_norm', 'shape_scaling'], @@ -2482,10 +2476,8 @@ def _add_training_args(parser): help='use FlashAttention implementation of attention. ' 'https://arxiv.org/abs/2205.14135') group.add_argument('--optimizer', type=str, default='adam', - choices=['adam', 'sgd', 'muon', 'dist_muon', 'lion', 'soap', 'adaptive_muon'], - help='Optimizer function. ' - 'Note: dist_muon is deprecated; use --optimizer muon ' - 'with --use-distributed-optimizer instead.') + choices=['adam', 'sgd', 'muon', 'dist_muon', 'lion'], + help='Optimizer function') group.add_argument('--optimizer-cpu-offload', action='store_true', help='Offload optimizer state to CPU') group.add_argument('--optimizer-cuda-graph', action='store_true', diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index caf32b2a7bc..d29b1e28869 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -581,7 +581,7 @@ def save_checkpoint(iteration, model, optimizer, opt_param_scheduler, num_floati optimizer.save_parameter_state(optim_checkpoint_name) # LayerWiseDistributedOptimizer save optimizer state to file on different ranks - if getattr(args, "use_layer_wise_distributed_optimizer", False) and args.ckpt_format == 'torch': + if getattr(args, "optimizer", "adam").startswith("dist_") and args.ckpt_format == 'torch': dp_rank = mpu.get_data_parallel_rank() optim_checkpoint_name = os.path.join(os.path.dirname(checkpoint_name), f"layer_wise_optimizer_{dp_rank}.pt") ensure_directory_exists(optim_checkpoint_name) @@ -1892,7 +1892,7 @@ def load_model_state_dict(module, state_dict, strict: bool): if not release and not args.finetune and not args.no_load_optim: try: # Load state dict. - if getattr(args, "use_layer_wise_distributed_optimizer", False) and args.ckpt_format == 'torch': + if getattr(args, "optimizer", "adam").startswith("dist_") and args.ckpt_format == 'torch': # LayerWiseDistributedOptimizer load optimizer state from file on different ranks dp_rank = mpu.get_data_parallel_rank() optim_checkpoint_name = os.path.join(os.path.dirname(checkpoint_name), f"layer_wise_optimizer_{dp_rank}.pt") diff --git a/megatron/training/training.py b/megatron/training/training.py index b8707041695..69e7f8e8c9f 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -176,11 +176,8 @@ def set_startup_timestamps(program_start=None, main_entry=None): from megatron.core.distributed import finalize_model_grads from megatron.core.enums import ModelType -from megatron.core.optimizer import ( - get_megatron_optimizer, - OptimizerConfig, - ParamKey, -) +from megatron.core.optimizer import get_megatron_optimizer, AdamOptimizerConfig, SGDOptimizerConfig, OptimizerConfig, ParamKey +from megatron.core.optimizer.muon import get_megatron_muon_optimizer from megatron.core.rerun_state_machine import ( get_rerun_state_machine, destroy_rerun_state_machine, @@ -1575,11 +1572,23 @@ def get_optimizer_param_scheduler(optimizer): def get_megatron_optimizer_config(args: Any) -> OptimizerConfig: """Return a Megatron optimizer config object from Megatron's arguments.""" - kwargs = {} - for f in dataclasses.fields(OptimizerConfig): - if hasattr(args, f.name): - kwargs[f.name] = getattr(args, f.name) - config = OptimizerConfig(**kwargs) + config = None + if args.optimizer == 'adam' or 'muon' in args.optimizer: + # TODO(deyuf): Muon needs both adam + muon but get() only receive one config + # So for now we keep using adam config that's back compat with old way + kwargs = {} + for f in dataclasses.fields(AdamOptimizerConfig): + if hasattr(args, f.name): + kwargs[f.name] = getattr(args, f.name) + config = AdamOptimizerConfig(**kwargs) + elif args.optimizer == 'sgd': + kwargs = {} + for f in dataclasses.fields(SGDOptimizerConfig): + if hasattr(args, f.name): + kwargs[f.name] = getattr(args, f.name) + config = SGDOptimizerConfig(**kwargs) + else: + raise ValueError("Invalid optimizer type!") # Construct the appropriate config_overrides object. This default handles many cases, but # can be added to as needed by the user, or replaced entirely with a custom override. @@ -1628,13 +1637,25 @@ def setup_model_and_optimizer( if mup_overrides: config_overrides = {**(config_overrides or {}), **mup_overrides} - optimizer = get_megatron_optimizer( - config, - model, - config_overrides=config_overrides, - use_gloo_process_groups=args.use_gloo_process_groups, - dump_param_to_param_group_map=args.dump_param_to_param_group_map, - ) + if 'muon' not in config.optimizer: + # If the user is asking for a non-zero embedding init std, skip weight decay for embeddings + # to avoid embeddings from shrinking to zero as recommended in https://arxiv.org/abs/2312.16903 + # default_skip_embedding_weight_decay=args.embedding_init_method_std is not None, + optimizer = get_megatron_optimizer( + config, + model, + config_overrides=config_overrides, + use_gloo_process_groups=args.use_gloo_process_groups, + dump_param_to_param_group_map=args.dump_param_to_param_group_map, + ) + else: + optimizer = get_megatron_muon_optimizer( + config, + model, + config_overrides=config_overrides, + use_gloo_process_groups=args.use_gloo_process_groups, + layer_wise_distributed_optimizer='dist' in config.optimizer, + ) opt_param_scheduler = get_optimizer_param_scheduler(optimizer) one_logger and one_logger.log_metrics({"app_build_optimzer_finish_time": one_logger_utils.get_timestamp_in_ms()}) diff --git a/tests/unit_tests/dist_checkpointing/utils.py b/tests/unit_tests/dist_checkpointing/utils.py index 0aadaee3b29..ec95602b020 100644 --- a/tests/unit_tests/dist_checkpointing/utils.py +++ b/tests/unit_tests/dist_checkpointing/utils.py @@ -15,7 +15,7 @@ get_gpt_layer_with_transformer_engine_spec, ) from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer -from megatron.core.optimizer.optimizer import ChainedOptimizer +from megatron.core.optimizer.muon import get_megatron_muon_optimizer from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from megatron.training.arguments import parse_args @@ -178,6 +178,11 @@ def init_checkpointing_mock_args(args, ckpt_dir, fully_parallel=False): def setup_model_and_optimizer( seed, tp, pp, initialize_fn=initialize_gpt_model, bf16=True, dist_opt=True, optimizer='adam' ): + if 'muon' in optimizer and dist_opt: + raise ValueError( + "Layer-wise distributed optimizer with Muon is not supported with distributed optimizer." + ) + mock_args = parse_args(ignore_unknown_args=True) with mock.patch('megatron.training.training.get_args', new=lambda: mock_args): init_basic_mock_args(mock_args, tp, pp, bf16=bf16) @@ -192,42 +197,37 @@ def setup_model_and_optimizer( ) ) - optimizer_type = optimizer - use_layer_wise = False - if optimizer_type == 'dist_muon': - optimizer = 'muon' - use_layer_wise = True - if optimizer_type in ('muon', 'dist_muon') and dist_opt: - use_layer_wise = True - dist_opt = False - config = OptimizerConfig( bf16=bf16, params_dtype=torch.bfloat16 if bf16 else torch.float, use_distributed_optimizer=dist_opt, - use_layer_wise_distributed_optimizer=use_layer_wise, optimizer=optimizer, ) - if optimizer_type in ('muon', 'dist_muon'): + if 'muon' in optimizer: + # Use layer-wise distributed optimizer with Muon + optimizer_type = optimizer + # default lr None feels wrong. only change muon lr to avoid breaking old tests config.lr = 0.0 - optimizer = get_megatron_optimizer(config, model) + optimizer = get_megatron_muon_optimizer( + config, model, layer_wise_distributed_optimizer='dist' in optimizer_type + ) + else: + optimizer_type = optimizer + optimizer = get_megatron_optimizer(config, model) torch.manual_seed(seed + 1) model_parallel_cuda_manual_seed(seed + 1) - if isinstance(optimizer, ChainedOptimizer): - for opt in optimizer.chained_optimizers: - if not hasattr(opt, 'optimizer'): - opt.init_state_fn(opt) - else: - opt.init_state_fn(opt.optimizer) - else: + if not 'muon' in optimizer_type: for group in optimizer.optimizer.param_groups: for p in group['params']: if len(optimizer.optimizer.state[p]) == 0: optimizer.optimizer.state[p]['exp_avg'] = torch.rand_like(p.data) optimizer.optimizer.state[p]['exp_avg_sq'] = torch.rand_like(p.data) + else: + for opt in optimizer.chained_optimizers: + opt.init_state_fn(opt) optimizer.reload_model_params() CachedMetadataFileSystemReader.clear_metadata_cache() @@ -272,6 +272,10 @@ def setup_moe_model_and_optimizer( use_glu=False, optimizer='adam', ): + if 'muon' in optimizer and dist_opt: + raise ValueError( + "Layer-wise distributed optimizer with Muon is not supported with distributed optimizer." + ) mock_args = parse_args(ignore_unknown_args=True) with mock.patch('megatron.training.training.get_args', new=lambda: mock_args): init_basic_mock_args(mock_args, tp, pp, bf16=bf16) @@ -291,43 +295,37 @@ def setup_moe_model_and_optimizer( ) ) - optimizer_type = optimizer - use_layer_wise = False - if optimizer_type == 'dist_muon': - optimizer = 'muon' - use_layer_wise = True - if optimizer_type in ('muon', 'dist_muon') and dist_opt: - use_layer_wise = True - dist_opt = False - config = OptimizerConfig( bf16=bf16, params_dtype=torch.bfloat16 if bf16 else torch.float, use_distributed_optimizer=dist_opt, - use_layer_wise_distributed_optimizer=use_layer_wise, optimizer=optimizer, ) - if optimizer_type in ('muon', 'dist_muon'): + if 'muon' in optimizer: + optimizer_type = optimizer + # default lr None feels wrong. only change muon lr to avoid breaking old tests config.lr = 0.0 - optimizer = get_megatron_optimizer(config, model) + optimizer = get_megatron_muon_optimizer( + config, model, layer_wise_distributed_optimizer='dist' in optimizer_type + ) + else: + optimizer_type = optimizer + optimizer = get_megatron_optimizer(config, model) torch.manual_seed(seed + 1) model_parallel_cuda_manual_seed(seed + 1) - if optimizer_type in ('muon', 'dist_muon'): - for opt in optimizer.chained_optimizers: - if not hasattr(opt, 'optimizer'): - opt.init_state_fn(opt) - else: - opt.init_state_fn(opt.optimizer) - else: + if not 'muon' in optimizer_type: for opt in optimizer.chained_optimizers: for group in opt.param_groups: for p in group['params']: if len(opt.state[p]) == 0: opt.state[p]['exp_avg'] = torch.rand_like(p.data) opt.state[p]['exp_avg_sq'] = torch.rand_like(p.data) + else: + for opt in optimizer.chained_optimizers: + opt.init_state_fn(opt) optimizer.reload_model_params() CachedMetadataFileSystemReader.clear_metadata_cache() diff --git a/tests/unit_tests/test_emerging_optimizers.py b/tests/unit_tests/test_emerging_optimizers.py deleted file mode 100644 index 939512d3e5b..00000000000 --- a/tests/unit_tests/test_emerging_optimizers.py +++ /dev/null @@ -1,1669 +0,0 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -import os - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F -from packaging.version import Version - -from megatron.core import parallel_state -from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig -from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer -from megatron.core.optimizer.emerging_optimizers import ( - HAVE_EMERGING_OPTIMIZERS, - TensorParallelAdaptiveMuon, - TensorParallelMuon, - get_supported_coefficient_types, - validate_coefficient_type, -) -from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.transformer import TransformerConfig -from tests.unit_tests.test_utilities import Utils - -if HAVE_EMERGING_OPTIMIZERS: - from emerging_optimizers.scalar_optimizers import Lion - from emerging_optimizers.soap import SOAP -else: - SOAP = None - Lion = None - -# Skip all tests in this file for LTS versions or when emerging_optimizers is missing -pytestmark = [ - pytest.mark.skipif( - Version(os.getenv('NVIDIA_PYTORCH_VERSION', "24.01")) <= Version("25.05"), - reason="Skip emerging optimizer tests for LTS test", - ), - pytest.mark.skipif( - not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package is not installed" - ), -] - - -class Net(nn.Module): - def __init__(self): - super().__init__() - self.fc1 = nn.Linear(80, 48) - self.fc2 = nn.Linear(48, 32) - self.fc3 = nn.Linear(32, 24) - self.fc4 = nn.Linear(24, 16) - self.fc5 = nn.Linear(16, 10) - - def forward(self, x): - x = F.relu(self.fc1(x)) - x = F.relu(self.fc2(x)) - x = F.relu(self.fc3(x)) - x = F.relu(self.fc4(x)) - x = self.fc5(x) - return x - - -# =========================================================================== -# Muon optimizer tests -# =========================================================================== - - -def test_muon_optimizer_smoke(): - """Smoke test for TensorParallelMuon optimizer.""" - # Create a simple linear model for testing - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - # Create TensorParallelMuon optimizer - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - nesterov=True, - weight_decay=0.01, - use_decoupled_weight_decay=True, - split_qkv=False, - fp32_matmul_prec="medium", - num_ns_steps=5, - scale_mode="spectral", - extra_scale_factor=1.0, - pg_collection=None, - tp_mode="duplicated", - ) - - # Test basic properties - assert optimizer is not None, "Optimizer should not be None" - assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" - assert len(optimizer.param_groups) > 0, "Optimizer should have at least one parameter group" - - # Test forward and backward pass - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - # Store original weight - original_weight = model.weight.data.clone() - - # Test optimizer step - optimizer.step() - - # Verify weight was updated - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated after optimizer step" - - # Test zero_grad - optimizer.zero_grad() - assert model.weight.grad is None or torch.all( - model.weight.grad == 0 - ), "Gradients should be zeroed" - - # Test state_dict and load_state_dict - state_dict = optimizer.state_dict() - assert 'state' in state_dict, "State dict should contain state" - assert 'param_groups' in state_dict, "State dict should contain param_groups" - - # Load state dict should not raise error - optimizer.load_state_dict(state_dict) - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestMuonOptimizerMultiRank: - """Test class for Muon optimizer with multi-rank setup.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - """Setup and teardown for each test.""" - Utils.initialize_model_parallel() - yield - Utils.destroy_model_parallel() - - def create_ddp_model(self, model): - """Wrap model in DDP. - - Args: - model: Model to wrap - - Returns: - DDP-wrapped model - """ - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) - return DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - def test_get_megatron_optimizer_smoke(self): - """Smoke test for get_megatron_optimizer function.""" - model = Net().bfloat16().cuda() - model.requires_grad_(True) - model = self.create_ddp_model(model) - - # Ensure all parameters require gradients - for param in model.parameters(): - assert param.requires_grad, "All parameters should require gradients" - - # Create optimizer config for Muon - optimizer_config = OptimizerConfig( - optimizer='muon', # This will be changed internally to 'adam' for non-linear params - lr=0.01, - weight_decay=0.01, - bf16=True, - use_distributed_optimizer=False, # Muon doesn't support distributed optimizer - muon_momentum=0.95, - muon_nesterov=True, - muon_fp32_matmul_prec="medium", - muon_num_ns_steps=5, - muon_scale_mode="spectral", - muon_tp_mode="duplicated", - ) - - # Test creating the optimizer - optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True - ) - - # Test basic properties - assert optimizer is not None, "Optimizer should not be None" - assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" - assert hasattr(optimizer, 'chained_optimizers'), "Should be a ChainedOptimizer" - assert len(optimizer.chained_optimizers) >= 1, "Should have at least one chained optimizer" - - # Test forward and backward pass - input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - # Store original parameters - original_params = {} - for name, param in model.named_parameters(): - original_params[name] = param.data.clone() - - # Test optimizer step - optimizer.step() - - # Verify at least some parameters were updated - params_updated = 0 - for name, param in model.named_parameters(): - if not torch.equal(param.data, original_params[name]): - params_updated += 1 - - assert params_updated > 0, "At least some parameters should be updated after optimizer step" - - # Test zero_grad - optimizer.zero_grad() - for param in model.parameters(): - assert param.grad is None or torch.all( - param.grad == 0 - ), f"Gradients should be zeroed for all parameters" - - # Test state_dict and load_state_dict - state_dict = optimizer.state_dict() - assert isinstance(state_dict, list), "State dict should be a list" - - # Load state dict should not raise error - optimizer.load_state_dict(state_dict) - - def test_get_megatron_optimizer_validation(self): - """Test validation logic for get_megatron_optimizer.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.bfloat16, device='cuda') - model.requires_grad_(True) - model = self.create_ddp_model(model) - - # Test 1: FP16 should raise exception - optimizer_config_fp16 = OptimizerConfig( - optimizer='muon', - lr=0.01, - fp16=True, # This should cause an exception - use_distributed_optimizer=False, - ) - - with pytest.raises(Exception, match='emerging optimizer with fp16 is not supported'): - get_megatron_optimizer(config=optimizer_config_fp16, model_chunks=[model]) - - # Test 3: Invalid num_ns_steps should raise exception - optimizer_config_invalid_ns = OptimizerConfig( - optimizer='muon', - lr=0.01, - bf16=True, - use_distributed_optimizer=False, - muon_num_ns_steps=0, # This should cause an exception - ) - - with pytest.raises(ValueError, match='num_ns_steps must be at least 1'): - get_megatron_optimizer(config=optimizer_config_invalid_ns, model_chunks=[model]) - - def test_get_megatron_optimizer_layer_wise(self): - """Test get_megatron_optimizer with layer-wise distributed optimizer.""" - model = Net().bfloat16().cuda() - model.requires_grad_(True) - model = self.create_ddp_model(model) - - optimizer_config = OptimizerConfig( - optimizer='muon', - lr=0.01, - weight_decay=0.01, - bf16=True, - use_layer_wise_distributed_optimizer=True, - muon_momentum=0.95, - muon_nesterov=True, - muon_fp32_matmul_prec="medium", - muon_num_ns_steps=5, - muon_scale_mode="spectral", - muon_tp_mode="duplicated", - ) - - # use_layer_wise_distributed_optimizer=True triggers LayerWiseDistributedOptimizer - optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True - ) - - # Verify it's a LayerWiseDistributedOptimizer - from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer - - assert isinstance( - optimizer, LayerWiseDistributedOptimizer - ), "Should return LayerWiseDistributedOptimizer" - - # Test forward and backward pass - input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - # Test optimizer step - update_successful, grad_norm, num_zeros = optimizer.step() - - assert update_successful, "Optimizer step should be successful" - assert grad_norm is not None or grad_norm is None, "Grad norm should be returned" - - -@pytest.mark.parametrize("mode", ["duplicated", "blockwise", "distributed"]) -def test_muon_optimizer_different_modes_single_rank(mode): - """Test TensorParallelMuon optimizer with different modes on single rank. - - When TP size is 1, all modes should produce the same result. - """ - # Set random seed for reproducibility - torch.manual_seed(42) - torch.cuda.manual_seed(42) - - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.normal_(0, 0.02) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.0, # Disable weight decay for deterministic comparison - num_ns_steps=5, - pg_collection=None, - tp_mode=mode, - ) - - # Use fixed input for deterministic results - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - # Verify weight was updated - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with mode={mode}" - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestMuonOptimizerMultiRankTP: - """Test class for Muon optimizer with multi-rank and tensor parallel setup.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - """Setup and teardown for each test with tensor parallel.""" - world = int(os.getenv('WORLD_SIZE', '1')) - Utils.initialize_model_parallel(tensor_model_parallel_size=min(world, 2)) - yield - Utils.destroy_model_parallel() - - def create_tp_model_and_optimizer(self, mode): - """Create model with TP and optimizer. - - Args: - mode: Muon optimizer mode - - Returns: - tuple: (model, optimizer, pg_collection) - """ - rank = int(os.getenv('RANK', '0')) - pg_collection = ProcessGroupCollection.use_mpu_process_groups() - - # Create model with partition_dim for TP - torch.manual_seed(42 + rank) - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.normal_(0, 0.02) - model.weight.partition_dim = 0 # Set partition dimension for TP - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.0, - num_ns_steps=5, - pg_collection=pg_collection, - tp_mode=mode, - ) - - return model, optimizer - - @pytest.mark.parametrize("mode", ["duplicated", "distributed"]) - def test_muon_optimizer_modes_multirank_same_result(self, mode): - """Test that duplicated and distributed modes produce same results with TP > 1.""" - model, optimizer = self.create_tp_model_and_optimizer(mode) - - # Use fixed input for deterministic results - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - # Verify weight was updated - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with mode={mode}" - - def test_muon_optimizer_blockwise_mode_different_result(self): - """Test that blockwise mode produces different results than duplicated/distributed with TP > 1.""" - model, optimizer = self.create_tp_model_and_optimizer("blockwise") - - # Use fixed input for deterministic results - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - # Verify weight was updated - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated with mode=blockwise" - - -# All non-custom coefficient types supported by emerging_optimizers. -_TESTABLE_COEFFICIENT_TYPES = ( - [t for t in get_supported_coefficient_types() if t != "custom"] - if HAVE_EMERGING_OPTIMIZERS - else [] -) - -# A reasonable default NS step count for testing; get_coefficient_iterator -# cycles/repeats coefficients so any step count works with any type. -_DEFAULT_NS_STEPS = 5 - - -@pytest.mark.parametrize("coefficient_type", _TESTABLE_COEFFICIENT_TYPES) -def test_muon_optimizer_coefficient_types(coefficient_type): - """Test TensorParallelMuon optimizer with different coefficient types.""" - model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - coefficient_type=coefficient_type, - num_ns_steps=_DEFAULT_NS_STEPS, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with coefficient_type={coefficient_type}" - - -@pytest.mark.parametrize("scale_mode", ["spectral", "unit_rms_norm", "shape_scaling"]) -def test_muon_optimizer_scale_modes(scale_mode): - """Test TensorParallelMuon optimizer with different scale modes.""" - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - scale_mode=scale_mode, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with scale_mode={scale_mode}" - - -@pytest.mark.parametrize("nesterov", [True, False]) -def test_muon_optimizer_nesterov(nesterov): - """Test TensorParallelMuon optimizer with and without Nesterov momentum.""" - model = torch.nn.Linear(50, 25, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - momentum=0.9, - nesterov=nesterov, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 50, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with nesterov={nesterov}" - - -def test_muon_optimizer_multiple_steps(): - """Test TensorParallelMuon optimizer across multiple optimization steps.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.01, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - weights_history = [model.weight.data.clone()] - - for i in range(3): - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer.step() - optimizer.zero_grad() - weights_history.append(model.weight.data.clone()) - - # Verify weights changed at each step - for i in range(len(weights_history) - 1): - assert not torch.equal( - weights_history[i], weights_history[i + 1] - ), f"Weight should change at step {i}" - - -def test_muon_optimizer_qkv_split(): - """Test TensorParallelMuon optimizer with QKV splitting.""" - # Create a model with QKV-like parameter - qkv_size = 3 * 64 * 16 # Combined Q, K, V dimensions, 16 heads x 64 per head - hidden_size = 1024 - model = torch.nn.Linear(hidden_size, qkv_size, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - # Mark parameter as QKV - model.weight.is_qkv = True - - # QKV split shapes: [Q_size, K_size, V_size] - qkv_split_shapes = (64, 64, 64) - - # Test with split_qkv=True - optimizer_split = TensorParallelMuon( - params=[model.weight], - lr=0.01, - split_qkv=True, - is_qkv_fn=lambda p: getattr(p, 'is_qkv', False), - qkv_split_shapes=qkv_split_shapes, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, hidden_size, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer_split.step() - weight_with_split = model.weight.data.clone() - - assert not torch.equal( - weight_with_split, original_weight - ), "QKV weight should be updated with split_qkv=True" - - # Reset model and test with split_qkv=False - model.weight.data.fill_(1.0) - optimizer_no_split = TensorParallelMuon( - params=[model.weight], - lr=0.01, - split_qkv=False, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer_no_split.step() - weight_without_split = model.weight.data.clone() - - assert not torch.equal( - weight_without_split, original_weight - ), "QKV weight should be updated with split_qkv=False" - - # Ensure the two results are different - assert not torch.equal( - weight_with_split, weight_without_split - ), "Weights should be different between split_qkv=True and split_qkv=False" - - -def test_muon_optimizer_extra_scale_factor(): - """Test TensorParallelMuon optimizer with different extra_scale_factor values.""" - model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - extra_scale_factor=2.0, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated with extra_scale_factor" - - -def test_get_supported_coefficient_types_returns_tuple(): - """Test that get_supported_coefficient_types returns a non-empty tuple of strings.""" - supported = get_supported_coefficient_types() - assert isinstance(supported, tuple) - assert len(supported) > 0 - for t in supported: - assert isinstance(t, str) - - -def test_get_supported_coefficient_types_contains_known_types(): - """Test that the known coefficient types are present in the supported set.""" - supported = get_supported_coefficient_types() - for expected in ("simple", "quintic", "polar_express"): - assert expected in supported, f"Expected '{expected}' in supported types {supported}" - - -def test_validate_coefficient_type_accepts_valid(): - """Test that validate_coefficient_type does not raise for valid types.""" - for t in get_supported_coefficient_types(): - validate_coefficient_type(t) # should not raise - - -def test_validate_coefficient_type_rejects_invalid(): - """Test that validate_coefficient_type raises ValueError for an invalid type.""" - with pytest.raises(ValueError, match="Unsupported muon coefficient type"): - validate_coefficient_type("nonexistent_type_xyz") - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestMuonCoefficientTypeMultiRank: - """Test coefficient_type integration through get_megatron_optimizer.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - Utils.initialize_model_parallel() - yield - Utils.destroy_model_parallel() - - def create_ddp_model(self, model): - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) - return DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - @pytest.mark.parametrize("coefficient_type", _TESTABLE_COEFFICIENT_TYPES) - def test_get_megatron_optimizer_coefficient_type(self, coefficient_type): - """Test that coefficient_type flows through get_megatron_optimizer.""" - model = Net().bfloat16().cuda() - model.requires_grad_(True) - model = self.create_ddp_model(model) - - optimizer_config = OptimizerConfig( - optimizer='muon', - lr=0.01, - weight_decay=0.01, - bf16=True, - use_distributed_optimizer=False, - muon_coefficient_type=coefficient_type, - muon_num_ns_steps=_DEFAULT_NS_STEPS, - muon_tp_mode="duplicated", - ) - - optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True - ) - - assert optimizer is not None - - input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer.step() - - -@pytest.mark.parametrize("num_ns_steps", [5, 15, 25]) -def test_muon_optimizer_num_ns_steps(num_ns_steps): - """Test TensorParallelMuon optimizer with different numbers of Newton-Schulz steps.""" - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelMuon( - params=[model.weight], - lr=0.01, - coefficient_type="quintic", - num_ns_steps=num_ns_steps, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with num_ns_steps={num_ns_steps}" - - -# =========================================================================== -# Adaptive Muon optimizer tests -# =========================================================================== - - -def test_adaptive_muon_optimizer_smoke(): - """Smoke test for TensorParallelAdaptiveMuon optimizer.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - nesterov=True, - weight_decay=0.01, - use_decoupled_weight_decay=True, - split_qkv=False, - fp32_matmul_prec="medium", - num_ns_steps=5, - scale_mode="spectral", - extra_scale_factor=1.0, - pg_collection=None, - tp_mode="duplicated", - moment2_method="adamuon", - beta2=0.95, - eps=1e-8, - ) - - assert optimizer is not None - assert hasattr(optimizer, 'param_groups') - assert len(optimizer.param_groups) > 0 - - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated after optimizer step" - - optimizer.zero_grad() - assert model.weight.grad is None or torch.all( - model.weight.grad == 0 - ), "Gradients should be zeroed" - - state_dict = optimizer.state_dict() - assert 'state' in state_dict - assert 'param_groups' in state_dict - optimizer.load_state_dict(state_dict) - - -@pytest.mark.parametrize("mode", ["duplicated", "blockwise", "distributed"]) -def test_adaptive_muon_optimizer_different_modes_single_rank(mode): - """Test TensorParallelAdaptiveMuon with different modes on single rank.""" - torch.manual_seed(42) - torch.cuda.manual_seed(42) - - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.normal_(0, 0.02) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.0, - num_ns_steps=5, - pg_collection=None, - tp_mode=mode, - ) - - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with mode={mode}" - - -@pytest.mark.parametrize("moment2_method", ["adamuon", "normuon"]) -def test_adaptive_muon_optimizer_moment2_methods(moment2_method): - """Test TensorParallelAdaptiveMuon with different moment2 methods.""" - model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - moment2_method=moment2_method, - ) - - input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with moment2_method={moment2_method}" - - -@pytest.mark.parametrize("beta2", [0.5, 0.95, 0.999]) -def test_adaptive_muon_optimizer_beta2(beta2): - """Test TensorParallelAdaptiveMuon with different beta2 values.""" - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - beta2=beta2, - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with beta2={beta2}" - - -def test_adaptive_muon_optimizer_multiple_steps(): - """Test TensorParallelAdaptiveMuon across multiple optimization steps.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.01, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - weights_history = [model.weight.data.clone()] - - for i in range(3): - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer.step() - optimizer.zero_grad() - weights_history.append(model.weight.data.clone()) - - for i in range(len(weights_history) - 1): - assert not torch.equal( - weights_history[i], weights_history[i + 1] - ), f"Weight should change at step {i}" - - -@pytest.mark.parametrize("nesterov", [True, False]) -def test_adaptive_muon_optimizer_nesterov(nesterov): - """Test TensorParallelAdaptiveMuon with and without Nesterov momentum.""" - model = torch.nn.Linear(50, 25, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - momentum=0.9, - nesterov=nesterov, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, 50, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with nesterov={nesterov}" - - -def test_adaptive_muon_optimizer_qkv_split(): - """Test TensorParallelAdaptiveMuon with QKV splitting.""" - qkv_size = 3 * 64 * 16 # Combined Q, K, V dimensions - hidden_size = 1024 - model = torch.nn.Linear(hidden_size, qkv_size, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - model.weight.is_qkv = True - qkv_split_shapes = (64, 64, 64) - - optimizer_split = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - split_qkv=True, - is_qkv_fn=lambda p: getattr(p, 'is_qkv', False), - qkv_split_shapes=qkv_split_shapes, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - input_tensor = torch.randn(16, hidden_size, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer_split.step() - weight_with_split = model.weight.data.clone() - - assert not torch.equal( - weight_with_split, original_weight - ), "QKV weight should be updated with split_qkv=True" - - model.weight.data.fill_(1.0) - optimizer_no_split = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - split_qkv=False, - num_ns_steps=5, - pg_collection=None, - tp_mode="duplicated", - ) - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer_no_split.step() - weight_without_split = model.weight.data.clone() - - assert not torch.equal( - weight_without_split, original_weight - ), "QKV weight should be updated with split_qkv=False" - - assert not torch.equal( - weight_with_split, weight_without_split - ), "Weights should be different between split_qkv=True and split_qkv=False" - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestAdaptiveMuonOptimizerMultiRank: - """Test class for Adaptive Muon optimizer with multi-rank setup.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - """Setup and teardown for each test.""" - Utils.initialize_model_parallel() - yield - Utils.destroy_model_parallel() - - def create_ddp_model(self, model): - """Wrap model in DDP.""" - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) - return DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - def test_get_megatron_optimizer_adaptive_muon_smoke(self): - """Smoke test for get_megatron_optimizer with adaptive_muon.""" - model = Net().bfloat16().cuda() - model.requires_grad_(True) - model = self.create_ddp_model(model) - - for param in model.parameters(): - assert param.requires_grad - - optimizer_config = OptimizerConfig( - optimizer='adaptive_muon', - lr=0.01, - weight_decay=0.01, - bf16=True, - use_distributed_optimizer=False, - muon_momentum=0.95, - muon_nesterov=True, - muon_fp32_matmul_prec="medium", - muon_num_ns_steps=5, - muon_scale_mode="spectral", - muon_tp_mode="duplicated", - adaptive_muon_moment2_method="adamuon", - adaptive_muon_beta2=0.95, - adaptive_muon_eps=1e-8, - ) - - optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True - ) - - assert optimizer is not None - assert hasattr(optimizer, 'param_groups') - assert hasattr(optimizer, 'chained_optimizers') - assert len(optimizer.chained_optimizers) >= 1 - - input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_params = {} - for name, param in model.named_parameters(): - original_params[name] = param.data.clone() - - optimizer.step() - - params_updated = 0 - for name, param in model.named_parameters(): - if not torch.equal(param.data, original_params[name]): - params_updated += 1 - - assert params_updated > 0, "At least some parameters should be updated after optimizer step" - - optimizer.zero_grad() - for param in model.parameters(): - assert param.grad is None or torch.all( - param.grad == 0 - ), "Gradients should be zeroed for all parameters" - - state_dict = optimizer.state_dict() - assert isinstance(state_dict, list) - optimizer.load_state_dict(state_dict) - - def test_get_megatron_optimizer_adaptive_muon_validation(self): - """Test validation logic for get_megatron_optimizer with adaptive_muon.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.bfloat16, device='cuda') - model.requires_grad_(True) - model = self.create_ddp_model(model) - - optimizer_config_fp16 = OptimizerConfig( - optimizer='adaptive_muon', lr=0.01, fp16=True, use_distributed_optimizer=False - ) - - with pytest.raises(Exception, match='emerging optimizer with fp16 is not supported'): - get_megatron_optimizer(config=optimizer_config_fp16, model_chunks=[model]) - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestAdaptiveMuonOptimizerMultiRankTP: - """Test class for Adaptive Muon optimizer with multi-rank and tensor parallel setup.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - """Setup and teardown for each test with tensor parallel.""" - world = int(os.getenv('WORLD_SIZE', '1')) - Utils.initialize_model_parallel(tensor_model_parallel_size=min(world, 2)) - yield - Utils.destroy_model_parallel() - - def create_tp_model_and_optimizer(self, mode): - """Create model with TP and optimizer.""" - rank = int(os.getenv('RANK', '0')) - pg_collection = ProcessGroupCollection.use_mpu_process_groups() - - torch.manual_seed(42 + rank) - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.normal_(0, 0.02) - model.weight.partition_dim = 0 - - optimizer = TensorParallelAdaptiveMuon( - params=[model.weight], - lr=0.01, - momentum=0.95, - weight_decay=0.0, - num_ns_steps=5, - pg_collection=pg_collection, - tp_mode=mode, - ) - - return model, optimizer - - @pytest.mark.parametrize("mode", ["duplicated", "distributed"]) - def test_adaptive_muon_optimizer_modes_multirank_same_result(self, mode): - """Test that duplicated and distributed modes produce same results with TP > 1.""" - model, optimizer = self.create_tp_model_and_optimizer(mode) - - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with mode={mode}" - - def test_adaptive_muon_optimizer_blockwise_mode(self): - """Test that blockwise mode works with TP > 1.""" - model, optimizer = self.create_tp_model_and_optimizer("blockwise") - - torch.manual_seed(42) - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated with mode=blockwise" - - -# =========================================================================== -# SOAP optimizer tests -# =========================================================================== - -skip_no_soap = pytest.mark.skipif( - not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package not installed" -) - - -@skip_no_soap -def test_soap_optimizer_smoke(): - """Smoke test for SOAP optimizer.""" - - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = SOAP( - params=[model.weight], - lr=0.01, - betas=(0.9, 0.999), - shampoo_beta=0.95, - weight_decay=0.01, - precondition_frequency=1, - ) - - # Test basic properties - assert optimizer is not None, "Optimizer should not be None" - assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" - assert len(optimizer.param_groups) > 0, "Optimizer should have at least one parameter group" - - # Test forward and backward pass - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - # Store original weight - original_weight = model.weight.data.clone() - - # Test optimizer step - optimizer.step() - - # Verify weight was updated - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated after optimizer step" - - # Test zero_grad - optimizer.zero_grad() - assert model.weight.grad is None or torch.all( - model.weight.grad == 0 - ), "Gradients should be zeroed" - - # Test state_dict and load_state_dict - state_dict = optimizer.state_dict() - assert 'state' in state_dict, "State dict should contain state" - assert 'param_groups' in state_dict, "State dict should contain param_groups" - - # Load state dict should not raise error - optimizer.load_state_dict(state_dict) - - -@skip_no_soap -def test_soap_optimizer_multiple_steps(): - """Test SOAP optimizer across multiple optimization steps.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = SOAP( - params=[model.weight], - lr=0.01, - betas=(0.9, 0.999), - shampoo_beta=0.95, - weight_decay=0.01, - precondition_frequency=1, - ) - - weights_history = [model.weight.data.clone()] - - for i in range(3): - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer.step() - optimizer.zero_grad() - weights_history.append(model.weight.data.clone()) - - # Verify weights changed at each step - for i in range(len(weights_history) - 1): - assert not torch.equal( - weights_history[i], weights_history[i + 1] - ), f"Weight should change at step {i}" - - -@skip_no_soap -@pytest.mark.parametrize("precondition_frequency", [1, 5, 10]) -def test_soap_optimizer_precondition_frequency(precondition_frequency): - """Test SOAP optimizer with different precondition frequencies.""" - - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = SOAP( - params=[model.weight], - lr=0.01, - betas=(0.9, 0.999), - shampoo_beta=0.95, - precondition_frequency=precondition_frequency, - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with precondition_frequency={precondition_frequency}" - - -@skip_no_soap -@pytest.mark.parametrize("use_kl_shampoo", [True, False]) -def test_soap_optimizer_kl_shampoo(use_kl_shampoo): - """Test SOAP optimizer with and without KL-Shampoo preconditioner.""" - - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = SOAP( - params=[model.weight], - lr=0.01, - betas=(0.9, 0.999), - shampoo_beta=0.95, - use_kl_shampoo=use_kl_shampoo, - precondition_frequency=1, - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with use_kl_shampoo={use_kl_shampoo}" - - -@skip_no_soap -@pytest.mark.parametrize("shampoo_beta", [0.5, 0.9, 0.99]) -def test_soap_optimizer_shampoo_beta(shampoo_beta): - """Test SOAP optimizer with different shampoo_beta values.""" - - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = SOAP( - params=[model.weight], - lr=0.01, - betas=(0.9, 0.999), - shampoo_beta=shampoo_beta, - precondition_frequency=1, - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with shampoo_beta={shampoo_beta}" - - -@pytest.mark.skipif( - int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" -) -class TestSoapOptimizerMultiRank: - """Test class for SOAP optimizer with multi-rank setup.""" - - @pytest.fixture(autouse=True) - def setup_and_teardown(self): - """Setup and teardown for each test.""" - Utils.initialize_model_parallel() - yield - Utils.destroy_model_parallel() - - def create_ddp_model(self, model): - """Wrap model in DDP.""" - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) - return DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - def test_get_megatron_optimizer_soap_smoke(self): - """Smoke test for get_megatron_optimizer with SOAP.""" - model = Net().bfloat16().cuda() - model.requires_grad_(True) - model = self.create_ddp_model(model) - - for param in model.parameters(): - assert param.requires_grad, "All parameters should require gradients" - - optimizer_config = OptimizerConfig( - optimizer='soap', - lr=0.01, - weight_decay=0.01, - bf16=True, - use_distributed_optimizer=False, - soap_shampoo_beta=0.95, - soap_precondition_frequency=1, - soap_use_kl_shampoo=True, - ) - - optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True - ) - - assert optimizer is not None, "Optimizer should not be None" - assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" - assert hasattr(optimizer, 'chained_optimizers'), "Should be a ChainedOptimizer" - assert len(optimizer.chained_optimizers) >= 1, "Should have at least one chained optimizer" - - # Test forward and backward pass - input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - # Store original parameters - original_params = {} - for name, param in model.named_parameters(): - original_params[name] = param.data.clone() - - # Test optimizer step - optimizer.step() - - # Verify at least some parameters were updated - params_updated = 0 - for name, param in model.named_parameters(): - if not torch.equal(param.data, original_params[name]): - params_updated += 1 - - assert params_updated > 0, "At least some parameters should be updated after optimizer step" - - # Test zero_grad - optimizer.zero_grad() - for param in model.parameters(): - assert param.grad is None or torch.all( - param.grad == 0 - ), "Gradients should be zeroed for all parameters" - - # Test state_dict and load_state_dict - state_dict = optimizer.state_dict() - assert isinstance(state_dict, list), "State dict should be a list" - optimizer.load_state_dict(state_dict) - - def test_get_megatron_optimizer_soap_validation(self): - """Test validation logic for get_megatron_optimizer with SOAP.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.bfloat16, device='cuda') - model.requires_grad_(True) - model = self.create_ddp_model(model) - - # FP16 should raise exception - optimizer_config_fp16 = OptimizerConfig( - optimizer='soap', lr=0.01, fp16=True, use_distributed_optimizer=False - ) - - with pytest.raises(Exception, match='emerging optimizer with fp16 is not supported'): - get_megatron_optimizer(config=optimizer_config_fp16, model_chunks=[model]) - - -# =========================================================================== -# Lion optimizer tests -# =========================================================================== - -skip_no_lion = pytest.mark.skipif( - not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package not installed" -) - - -@skip_no_lion -def test_lion_optimizer_smoke(): - """Smoke test for Lion optimizer.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = Lion(params=[model.weight], lr=1e-4, betas=(0.9, 0.99), weight_decay=0.01) - - assert optimizer is not None - assert hasattr(optimizer, 'param_groups') - assert len(optimizer.param_groups) > 0 - - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), "Weight should be updated after optimizer step" - - optimizer.zero_grad() - assert model.weight.grad is None or torch.all( - model.weight.grad == 0 - ), "Gradients should be zeroed" - - state_dict = optimizer.state_dict() - assert 'state' in state_dict - assert 'param_groups' in state_dict - optimizer.load_state_dict(state_dict) - - -@skip_no_lion -def test_lion_optimizer_multiple_steps(): - """Test Lion optimizer across multiple optimization steps.""" - model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = Lion(params=[model.weight], lr=1e-4, betas=(0.9, 0.99), weight_decay=0.01) - - weights_history = [model.weight.data.clone()] - - for i in range(3): - input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - optimizer.step() - optimizer.zero_grad() - weights_history.append(model.weight.data.clone()) - - for i in range(len(weights_history) - 1): - assert not torch.equal( - weights_history[i], weights_history[i + 1] - ), f"Weight should change at step {i}" - - -@skip_no_lion -@pytest.mark.parametrize("betas", [(0.9, 0.99), (0.95, 0.999), (0.5, 0.9)]) -def test_lion_optimizer_betas(betas): - """Test Lion optimizer with different beta values.""" - model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = Lion(params=[model.weight], lr=1e-4, betas=betas) - - input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with betas={betas}" - - -@skip_no_lion -@pytest.mark.parametrize("weight_decay", [0.0, 0.01, 0.1]) -def test_lion_optimizer_weight_decay(weight_decay): - """Test Lion optimizer with different weight decay values.""" - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = Lion(params=[model.weight], lr=1e-4, betas=(0.9, 0.99), weight_decay=weight_decay) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with weight_decay={weight_decay}" - - -@skip_no_lion -@pytest.mark.parametrize("weight_decay_method", ["decoupled", "l2"]) -def test_lion_optimizer_weight_decay_method(weight_decay_method): - """Test Lion optimizer with different weight decay methods.""" - model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') - model.requires_grad_(True) - model.weight.data.fill_(1.0) - - optimizer = Lion( - params=[model.weight], - lr=1e-4, - betas=(0.9, 0.99), - weight_decay=0.01, - weight_decay_method=weight_decay_method, - ) - - input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_weight = model.weight.data.clone() - optimizer.step() - - assert not torch.equal( - model.weight.data, original_weight - ), f"Weight should be updated with weight_decay_method={weight_decay_method}" - - -@skip_no_lion -def test_lion_optimizer_multi_layer_net(): - """Test Lion optimizer with the multi-layer Net model.""" - model = Net().cuda() - model.requires_grad_(True) - - optimizer = Lion(params=model.parameters(), lr=1e-4, betas=(0.9, 0.99), weight_decay=0.01) - - input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') - output = model(input_tensor) - loss = output.sum() - loss.backward() - - original_params = {name: p.data.clone() for name, p in model.named_parameters()} - optimizer.step() - - params_updated = 0 - for name, param in model.named_parameters(): - if not torch.equal(param.data, original_params[name]): - params_updated += 1 - - assert params_updated > 0, "At least some parameters should be updated after optimizer step" diff --git a/tests/unit_tests/test_layer_wise_optimizer.py b/tests/unit_tests/test_layer_wise_optimizer.py index 74837a30a44..c484ca104ee 100644 --- a/tests/unit_tests/test_layer_wise_optimizer.py +++ b/tests/unit_tests/test_layer_wise_optimizer.py @@ -10,8 +10,9 @@ from megatron.core import parallel_state from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig -from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer +from megatron.core.optimizer import OptimizerConfig from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer +from megatron.core.optimizer.muon import get_megatron_muon_optimizer from megatron.core.optimizer.optimizer import Float16OptimizerWithFloat16Params from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import TransformerConfig @@ -117,14 +118,19 @@ def create_model_and_optimizer( use_distributed_optimizer=False, clip_grad=clip_grad, muon_tp_mode="duplicated", - use_layer_wise_distributed_optimizer=use_layer_wise, ) pg_collection = ProcessGroupCollection.use_mpu_process_groups() pg_collection.dp_cp = parallel_state.get_data_parallel_group(with_context_parallel=True) pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group() - optimizer = get_megatron_optimizer(optimizer_config, [model], pg_collection=pg_collection) + optimizer = get_megatron_muon_optimizer( + config=optimizer_config, + model_chunks=[model], + use_gloo_process_groups=True, + layer_wise_distributed_optimizer=use_layer_wise, + pg_collection=pg_collection, + ) return model, optimizer, pg_collection def create_model_and_optimizer_with_overlap_param_gather( @@ -140,7 +146,7 @@ def create_model_and_optimizer_with_overlap_param_gather( """Create model, DDP wrapper, and optimizer with overlap-param-gather enabled. This variant sets overlap_param_gather=True in DDP config and uses - get_megatron_optimizer with layer_wise_distributed_optimizer=True, + get_megatron_muon_optimizer with layer_wise_distributed_optimizer=True, enabling the bucket-based async param gather path. Args: @@ -185,17 +191,17 @@ def create_model_and_optimizer_with_overlap_param_gather( clip_grad=clip_grad, overlap_param_gather=async_allgather, muon_tp_mode="duplicated", - use_layer_wise_distributed_optimizer=True, ) pg_collection = ProcessGroupCollection.use_mpu_process_groups() pg_collection.dp_cp = parallel_state.get_data_parallel_group(with_context_parallel=True) pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group() - optimizer = get_megatron_optimizer( + optimizer = get_megatron_muon_optimizer( config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True, + layer_wise_distributed_optimizer=True, pg_collection=pg_collection, ) return model, optimizer, pg_collection @@ -344,34 +350,8 @@ def test_multiple_optimizers(self): """ model, optimizer, pg_collection = self.create_model_and_optimizer() - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) - model = DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - optimizer_config = OptimizerConfig( - optimizer='adam', lr=0.01, bf16=True, use_distributed_optimizer=False - ) - - # Split parameters into two groups for testing multiple optimizers - params = list(model.parameters()) - mid_point = len(params) // 2 - param_groups_1 = [{'params': params[:mid_point]}] - param_groups_2 = [{'params': params[mid_point:]}] - - # Create two separate plain base optimizers (LayerWise wraps them itself) - base_optimizer_1 = torch.optim.Adam(param_groups_1, lr=optimizer_config.lr) - base_optimizer_2 = torch.optim.Adam(param_groups_2, lr=optimizer_config.lr) - - pg_collection = ProcessGroupCollection.use_mpu_process_groups() - pg_collection.dp_cp = parallel_state.get_data_parallel_group(with_context_parallel=True) - pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group() - - optimizer = LayerWiseDistributedOptimizer( - [base_optimizer_1, base_optimizer_2], optimizer_config, pg_collection - ) - - assert len(optimizer.chained_optimizers) == 2, "Should have two chained optimizers" + # get_megatron_muon_optimizer produces muon + adam chained optimizers + assert len(optimizer.chained_optimizers) >= 2, "Should have multiple chained optimizers" # Set gradients and test optimizer step - this will trigger allgather for param in model.parameters(): @@ -411,7 +391,7 @@ def test_bf16_error(self): pg_collection.dp_cp = parallel_state.get_data_parallel_group(with_context_parallel=True) pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group() - # Create optimizer (non-layer-wise) — produces Float16-wrapped chained optimizers + # Create muon optimizer (non-layer-wise) — produces Float16-wrapped chained optimizers optimizer_config = OptimizerConfig( optimizer='muon', lr=0.01, @@ -419,8 +399,12 @@ def test_bf16_error(self): use_distributed_optimizer=False, muon_tp_mode="duplicated", ) - muon_optimizer = get_megatron_optimizer( - config=optimizer_config, model_chunks=[model], use_gloo_process_groups=True + muon_optimizer = get_megatron_muon_optimizer( + config=optimizer_config, + model_chunks=[model], + use_gloo_process_groups=True, + layer_wise_distributed_optimizer=False, + pg_collection=pg_collection, ) # Extract a Float16-wrapped chained optimizer @@ -428,11 +412,12 @@ def test_bf16_error(self): assert isinstance(wrapped_optimizer, Float16OptimizerWithFloat16Params) # Should raise TypeError when receiving already-wrapped Float16 optimizer + # Use a fresh config since get_megatron_muon_optimizer mutates config.optimizer lw_config = OptimizerConfig( optimizer='muon', lr=0.01, bf16=True, use_distributed_optimizer=False ) with pytest.raises( - TypeError, match='LayerWiseDistributedOptimizer expects base torch optimizers' + TypeError, match='LayerWiseDistributedOptimizer received Float16 optimizer already' ): LayerWiseDistributedOptimizer([wrapped_optimizer], lw_config, pg_collection) diff --git a/tests/unit_tests/test_muon_optimizer.py b/tests/unit_tests/test_muon_optimizer.py new file mode 100644 index 00000000000..7a15674d7b2 --- /dev/null +++ b/tests/unit_tests/test_muon_optimizer.py @@ -0,0 +1,787 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import os + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F +from packaging.version import Version + +from megatron.core import parallel_state +from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig +from megatron.core.optimizer import HAVE_EMERGING_OPTIMIZERS, OptimizerConfig +from megatron.core.optimizer.muon import ( + TensorParallelMuon, + get_megatron_muon_optimizer, + get_supported_coefficient_types, + validate_coefficient_type, +) +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer import TransformerConfig +from tests.unit_tests.test_utilities import Utils + +# Skip all tests in this file for LTS versions or when emerging_optimizers is missing +pytestmark = [ + pytest.mark.skipif( + Version(os.getenv('NVIDIA_PYTORCH_VERSION', "24.01")) <= Version("25.05"), + reason="Skip muon optimizer for LTS test", + ), + pytest.mark.skipif( + not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package is not installed" + ), +] + + +class Net(nn.Module): + def __init__(self): + super().__init__() + self.fc1 = nn.Linear(80, 48) + self.fc2 = nn.Linear(48, 32) + self.fc3 = nn.Linear(32, 24) + self.fc4 = nn.Linear(24, 16) + self.fc5 = nn.Linear(16, 10) + + def forward(self, x): + x = F.relu(self.fc1(x)) + x = F.relu(self.fc2(x)) + x = F.relu(self.fc3(x)) + x = F.relu(self.fc4(x)) + x = self.fc5(x) + return x + + +def test_muon_optimizer_smoke(): + """Smoke test for TensorParallelMuon optimizer.""" + # Create a simple linear model for testing + model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + # Create TensorParallelMuon optimizer + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + momentum_beta=0.95, + use_nesterov=True, + weight_decay=0.01, + use_decoupled_weight_decay=True, + split_qkv=False, + fp32_matmul_prec="medium", + num_ns_steps=5, + scale_mode="spectral", + extra_scale_factor=1.0, + pg_collection=None, + mode="duplicated", + ) + + # Test basic properties + assert optimizer is not None, "Optimizer should not be None" + assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" + assert len(optimizer.param_groups) > 0, "Optimizer should have at least one parameter group" + + # Test forward and backward pass + input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + # Store original weight + original_weight = model.weight.data.clone() + + # Test optimizer step + optimizer.step() + + # Verify weight was updated + assert not torch.equal( + model.weight.data, original_weight + ), "Weight should be updated after optimizer step" + + # Test zero_grad + optimizer.zero_grad() + assert model.weight.grad is None or torch.all( + model.weight.grad == 0 + ), "Gradients should be zeroed" + + # Test state_dict and load_state_dict + state_dict = optimizer.state_dict() + assert 'state' in state_dict, "State dict should contain state" + assert 'param_groups' in state_dict, "State dict should contain param_groups" + + # Load state dict should not raise error + optimizer.load_state_dict(state_dict) + + +@pytest.mark.skipif( + int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" +) +class TestMuonOptimizerMultiRank: + """Test class for Muon optimizer with multi-rank setup.""" + + @pytest.fixture(autouse=True) + def setup_and_teardown(self): + """Setup and teardown for each test.""" + Utils.initialize_model_parallel() + yield + Utils.destroy_model_parallel() + + def create_ddp_model(self, model): + """Wrap model in DDP. + + Args: + model: Model to wrap + + Returns: + DDP-wrapped model + """ + ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) + return DistributedDataParallel( + TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model + ) + + def test_get_megatron_muon_optimizer_smoke(self): + """Smoke test for get_megatron_muon_optimizer function.""" + model = Net().bfloat16().cuda() + model.requires_grad_(True) + model = self.create_ddp_model(model) + + # Ensure all parameters require gradients + for param in model.parameters(): + assert param.requires_grad, "All parameters should require gradients" + + # Create optimizer config for Muon + optimizer_config = OptimizerConfig( + optimizer='muon', # This will be changed internally to 'adam' for non-linear params + lr=0.01, + weight_decay=0.01, + bf16=True, + use_distributed_optimizer=False, # Muon doesn't support distributed optimizer + muon_momentum=0.95, + muon_use_nesterov=True, + muon_fp32_matmul_prec="medium", + muon_num_ns_steps=5, + muon_scale_mode="spectral", + muon_tp_mode="duplicated", + ) + + # Test creating the optimizer + optimizer = get_megatron_muon_optimizer( + config=optimizer_config, + model_chunks=[model], + use_gloo_process_groups=True, + layer_wise_distributed_optimizer=False, + ) + + # Test basic properties + assert optimizer is not None, "Optimizer should not be None" + assert hasattr(optimizer, 'param_groups'), "Optimizer should have param_groups" + assert hasattr(optimizer, 'chained_optimizers'), "Should be a ChainedOptimizer" + assert len(optimizer.chained_optimizers) >= 1, "Should have at least one chained optimizer" + + # Test forward and backward pass + input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + # Store original parameters + original_params = {} + for name, param in model.named_parameters(): + original_params[name] = param.data.clone() + + # Test optimizer step + optimizer.step() + + # Verify at least some parameters were updated + params_updated = 0 + for name, param in model.named_parameters(): + if not torch.equal(param.data, original_params[name]): + params_updated += 1 + + assert params_updated > 0, "At least some parameters should be updated after optimizer step" + + # Test zero_grad + optimizer.zero_grad() + for param in model.parameters(): + assert param.grad is None or torch.all( + param.grad == 0 + ), f"Gradients should be zeroed for all parameters" + + # Test state_dict and load_state_dict + state_dict = optimizer.state_dict() + assert isinstance(state_dict, list), "State dict should be a list" + + # Load state dict should not raise error + optimizer.load_state_dict(state_dict) + + def test_get_megatron_muon_optimizer_validation(self): + """Test validation logic for get_megatron_muon_optimizer.""" + model = torch.nn.Linear(100, 50, bias=False, dtype=torch.bfloat16, device='cuda') + model.requires_grad_(True) + model = self.create_ddp_model(model) + + # Test 1: Distributed optimizer should raise exception + optimizer_config_dist = OptimizerConfig( + optimizer='muon', + lr=0.01, + bf16=True, + use_distributed_optimizer=True, # This should cause an exception + ) + + with pytest.raises(Exception, match='muon with dist optimizer is not supported'): + get_megatron_muon_optimizer(config=optimizer_config_dist, model_chunks=[model]) + + # Test 2: FP16 should raise exception + optimizer_config_fp16 = OptimizerConfig( + optimizer='muon', + lr=0.01, + fp16=True, # This should cause an exception + use_distributed_optimizer=False, + ) + + with pytest.raises(Exception, match='muon with fp16 is not supported'): + get_megatron_muon_optimizer(config=optimizer_config_fp16, model_chunks=[model]) + + # Test 3: Invalid num_ns_steps should raise exception + optimizer_config_invalid_ns = OptimizerConfig( + optimizer='muon', + lr=0.01, + bf16=True, + use_distributed_optimizer=False, + muon_num_ns_steps=0, # This should cause an exception + ) + + with pytest.raises(ValueError, match='num_ns_steps must be at least 1'): + get_megatron_muon_optimizer(config=optimizer_config_invalid_ns, model_chunks=[model]) + + def test_get_megatron_muon_optimizer_layer_wise(self): + """Test get_megatron_muon_optimizer with layer-wise distributed optimizer.""" + model = Net().bfloat16().cuda() + model.requires_grad_(True) + model = self.create_ddp_model(model) + + optimizer_config = OptimizerConfig( + optimizer='muon', + lr=0.01, + weight_decay=0.01, + bf16=True, + use_distributed_optimizer=False, + muon_momentum=0.95, + muon_use_nesterov=True, + muon_fp32_matmul_prec="medium", + muon_num_ns_steps=5, + muon_scale_mode="spectral", + muon_tp_mode="duplicated", + ) + + # Test with layer_wise_distributed_optimizer=True + optimizer = get_megatron_muon_optimizer( + config=optimizer_config, + model_chunks=[model], + use_gloo_process_groups=True, + layer_wise_distributed_optimizer=True, + ) + + # Verify it's a LayerWiseDistributedOptimizer + from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer + + assert isinstance( + optimizer, LayerWiseDistributedOptimizer + ), "Should return LayerWiseDistributedOptimizer" + + # Test forward and backward pass + input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + # Test optimizer step + update_successful, grad_norm, num_zeros = optimizer.step() + + assert update_successful, "Optimizer step should be successful" + assert grad_norm is not None or grad_norm is None, "Grad norm should be returned" + + +@pytest.mark.parametrize("mode", ["duplicated", "blockwise", "distributed"]) +def test_muon_optimizer_different_modes_single_rank(mode): + """Test TensorParallelMuon optimizer with different modes on single rank. + + When TP size is 1, all modes should produce the same result. + """ + # Set random seed for reproducibility + torch.manual_seed(42) + torch.cuda.manual_seed(42) + + model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.normal_(0, 0.02) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + momentum_beta=0.95, + weight_decay=0.0, # Disable weight decay for deterministic comparison + num_ns_steps=5, + pg_collection=None, + mode=mode, + ) + + # Use fixed input for deterministic results + torch.manual_seed(42) + input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') + + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + # Verify weight was updated + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with mode={mode}" + + +@pytest.mark.skipif( + int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" +) +class TestMuonOptimizerMultiRankTP: + """Test class for Muon optimizer with multi-rank and tensor parallel setup.""" + + @pytest.fixture(autouse=True) + def setup_and_teardown(self): + """Setup and teardown for each test with tensor parallel.""" + world = int(os.getenv('WORLD_SIZE', '1')) + Utils.initialize_model_parallel(tensor_model_parallel_size=min(world, 2)) + yield + Utils.destroy_model_parallel() + + def create_tp_model_and_optimizer(self, mode): + """Create model with TP and optimizer. + + Args: + mode: Muon optimizer mode + + Returns: + tuple: (model, optimizer, pg_collection) + """ + rank = int(os.getenv('RANK', '0')) + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + + # Create model with partition_dim for TP + torch.manual_seed(42 + rank) + model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.normal_(0, 0.02) + model.weight.partition_dim = 0 # Set partition dimension for TP + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + momentum_beta=0.95, + weight_decay=0.0, + num_ns_steps=5, + pg_collection=pg_collection, + mode=mode, + ) + + return model, optimizer + + @pytest.mark.parametrize("mode", ["duplicated", "distributed"]) + def test_muon_optimizer_modes_multirank_same_result(self, mode): + """Test that duplicated and distributed modes produce same results with TP > 1.""" + model, optimizer = self.create_tp_model_and_optimizer(mode) + + # Use fixed input for deterministic results + torch.manual_seed(42) + input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') + + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + # Verify weight was updated + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with mode={mode}" + + def test_muon_optimizer_blockwise_mode_different_result(self): + """Test that blockwise mode produces different results than duplicated/distributed with TP > 1.""" + model, optimizer = self.create_tp_model_and_optimizer("blockwise") + + # Use fixed input for deterministic results + torch.manual_seed(42) + input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') + + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + # Verify weight was updated + assert not torch.equal( + model.weight.data, original_weight + ), "Weight should be updated with mode=blockwise" + + +# All non-custom coefficient types supported by emerging_optimizers. +_TESTABLE_COEFFICIENT_TYPES = ( + [t for t in get_supported_coefficient_types() if t != "custom"] + if HAVE_EMERGING_OPTIMIZERS + else [] +) + +# A reasonable default NS step count for testing; get_coefficient_iterator +# cycles/repeats coefficients so any step count works with any type. +_DEFAULT_NS_STEPS = 5 + + +@pytest.mark.parametrize("coefficient_type", _TESTABLE_COEFFICIENT_TYPES) +def test_muon_optimizer_coefficient_types(coefficient_type): + """Test TensorParallelMuon optimizer with different coefficient types.""" + model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + coefficient_type=coefficient_type, + num_ns_steps=_DEFAULT_NS_STEPS, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with coefficient_type={coefficient_type}" + + +@pytest.mark.parametrize("scale_mode", ["spectral", "unit_rms_norm", "shape_scaling"]) +def test_muon_optimizer_scale_modes(scale_mode): + """Test TensorParallelMuon optimizer with different scale modes.""" + model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + scale_mode=scale_mode, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with scale_mode={scale_mode}" + + +@pytest.mark.parametrize("use_nesterov", [True, False]) +def test_muon_optimizer_nesterov(use_nesterov): + """Test TensorParallelMuon optimizer with and without Nesterov momentum.""" + model = torch.nn.Linear(50, 25, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + momentum_beta=0.9, + use_nesterov=use_nesterov, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, 50, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with use_nesterov={use_nesterov}" + + +def test_muon_optimizer_multiple_steps(): + """Test TensorParallelMuon optimizer across multiple optimization steps.""" + model = torch.nn.Linear(100, 50, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + momentum_beta=0.95, + weight_decay=0.01, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + weights_history = [model.weight.data.clone()] + + for i in range(3): + input_tensor = torch.randn(32, 100, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + optimizer.step() + optimizer.zero_grad() + weights_history.append(model.weight.data.clone()) + + # Verify weights changed at each step + for i in range(len(weights_history) - 1): + assert not torch.equal( + weights_history[i], weights_history[i + 1] + ), f"Weight should change at step {i}" + + +def test_muon_optimizer_qkv_split(): + """Test TensorParallelMuon optimizer with QKV splitting.""" + # Create a model with QKV-like parameter + qkv_size = 3 * 64 * 16 # Combined Q, K, V dimensions, 16 heads x 64 per head + hidden_size = 1024 + model = torch.nn.Linear(hidden_size, qkv_size, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + # Mark parameter as QKV + model.weight.is_qkv = True + + # QKV split shapes: [Q_size, K_size, V_size] + qkv_split_shapes = (64, 64, 64) + + # Test with split_qkv=True + optimizer_split = TensorParallelMuon( + params=[model.weight], + lr=0.01, + split_qkv=True, + is_qkv_fn=lambda p: getattr(p, 'is_qkv', False), + qkv_split_shapes=qkv_split_shapes, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, hidden_size, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer_split.step() + weight_with_split = model.weight.data.clone() + + assert not torch.equal( + weight_with_split, original_weight + ), "QKV weight should be updated with split_qkv=True" + + # Reset model and test with split_qkv=False + model.weight.data.fill_(1.0) + optimizer_no_split = TensorParallelMuon( + params=[model.weight], + lr=0.01, + split_qkv=False, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + output = model(input_tensor) + loss = output.sum() + loss.backward() + + optimizer_no_split.step() + weight_without_split = model.weight.data.clone() + + assert not torch.equal( + weight_without_split, original_weight + ), "QKV weight should be updated with split_qkv=False" + + # Ensure the two results are different + assert not torch.equal( + weight_with_split, weight_without_split + ), "Weights should be different between split_qkv=True and split_qkv=False" + + +def test_muon_optimizer_extra_scale_factor(): + """Test TensorParallelMuon optimizer with different extra_scale_factor values.""" + model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + extra_scale_factor=2.0, + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, 80, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + assert not torch.equal( + model.weight.data, original_weight + ), "Weight should be updated with extra_scale_factor" + + +def test_get_supported_coefficient_types_returns_tuple(): + """Test that get_supported_coefficient_types returns a non-empty tuple of strings.""" + supported = get_supported_coefficient_types() + assert isinstance(supported, tuple) + assert len(supported) > 0 + for t in supported: + assert isinstance(t, str) + + +def test_get_supported_coefficient_types_contains_known_types(): + """Test that the known coefficient types are present in the supported set.""" + supported = get_supported_coefficient_types() + for expected in ("simple", "quintic", "polar_express"): + assert expected in supported, f"Expected '{expected}' in supported types {supported}" + + +def test_validate_coefficient_type_accepts_valid(): + """Test that validate_coefficient_type does not raise for valid types.""" + for t in get_supported_coefficient_types(): + validate_coefficient_type(t) # should not raise + + +def test_validate_coefficient_type_rejects_invalid(): + """Test that validate_coefficient_type raises ValueError for an invalid type.""" + with pytest.raises(ValueError, match="Unsupported muon coefficient type"): + validate_coefficient_type("nonexistent_type_xyz") + + +def test_muon_optimizer_invalid_coefficient_type(): + """Test that TensorParallelMuon raises ValueError for an invalid coefficient_type.""" + model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + + with pytest.raises(ValueError, match="Unsupported muon coefficient type"): + TensorParallelMuon( + params=[model.weight], + lr=0.01, + coefficient_type="nonexistent_type_xyz", + num_ns_steps=5, + pg_collection=None, + mode="duplicated", + ) + + +@pytest.mark.skipif( + int(os.getenv('WORLD_SIZE', '1')) == 1, reason="Multi-rank test requires WORLD_SIZE > 1" +) +class TestMuonCoefficientTypeMultiRank: + """Test coefficient_type integration through get_megatron_muon_optimizer.""" + + @pytest.fixture(autouse=True) + def setup_and_teardown(self): + Utils.initialize_model_parallel() + yield + Utils.destroy_model_parallel() + + def create_ddp_model(self, model): + ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=False) + return DistributedDataParallel( + TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model + ) + + @pytest.mark.parametrize("coefficient_type", _TESTABLE_COEFFICIENT_TYPES) + def test_get_megatron_muon_optimizer_coefficient_type(self, coefficient_type): + """Test that coefficient_type flows through get_megatron_muon_optimizer.""" + model = Net().bfloat16().cuda() + model.requires_grad_(True) + model = self.create_ddp_model(model) + + optimizer_config = OptimizerConfig( + optimizer='muon', + lr=0.01, + weight_decay=0.01, + bf16=True, + use_distributed_optimizer=False, + muon_coefficient_type=coefficient_type, + muon_num_ns_steps=_DEFAULT_NS_STEPS, + muon_tp_mode="duplicated", + ) + + optimizer = get_megatron_muon_optimizer( + config=optimizer_config, + model_chunks=[model], + use_gloo_process_groups=True, + layer_wise_distributed_optimizer=False, + ) + + assert optimizer is not None + + input_tensor = torch.randn(16, 80, dtype=torch.bfloat16, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + optimizer.step() + + +@pytest.mark.parametrize("num_ns_steps", [5, 15, 25]) +def test_muon_optimizer_num_ns_steps(num_ns_steps): + """Test TensorParallelMuon optimizer with different numbers of Newton-Schulz steps.""" + model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda') + model.requires_grad_(True) + model.weight.data.fill_(1.0) + + optimizer = TensorParallelMuon( + params=[model.weight], + lr=0.01, + coefficient_type="quintic", + num_ns_steps=num_ns_steps, + pg_collection=None, + mode="duplicated", + ) + + input_tensor = torch.randn(16, 60, dtype=torch.float32, device='cuda') + output = model(input_tensor) + loss = output.sum() + loss.backward() + + original_weight = model.weight.data.clone() + optimizer.step() + + assert not torch.equal( + model.weight.data, original_weight + ), f"Weight should be updated with num_ns_steps={num_ns_steps}" diff --git a/tests/unit_tests/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 56af8545042..2488900ba72 100644 --- a/tests/unit_tests/test_optimizer.py +++ b/tests/unit_tests/test_optimizer.py @@ -106,10 +106,10 @@ def test_get_param_groups_no_overrides(mock_get_world_size): def test_get_param_groups_default_overrides(mock_get_world_size): """Test that the default overrides are applied to the parameter groups.""" net = Net() + # NOTE: to get legacy default overrides, supply None. opt_config = OptimizerConfig(optimizer='adam', lr=0.01) - config_overrides = get_standard_config_overrides(opt_config) - check_config_overrides_consistency(opt_config, config_overrides) - param_groups = _get_param_groups([net], opt_config, config_overrides) + check_config_overrides_consistency(opt_config, None) + param_groups = _get_param_groups([net], opt_config, None) assert len(param_groups) == 2 pg0, pg1 = param_groups wd_mults = {pg0['wd_mult'], pg1['wd_mult']}