From 927f0dcdf020981e6dbbfa4de07322abaaf1032f Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:22:08 -0700 Subject: [PATCH 01/10] Merge optimizer changes from dev to main Bring over optimizer improvements from the dev branch including: - Emerging optimizers registry and refactoring (Lion, AdaptiveMuon, etc.) - Optimizer state and master weight CPU offloading - Layer-wise optimizer improvements with --overlap-param-gather support - Muon optimizer cleanup and consolidation into emerging optimizers framework - Updated unit tests (new: emerging_optimizers, lion, state_offloading; removed: muon_optimizer) - Updated functional test golden values Co-Authored-By: Claude Opus 4.6 (1M context) --- megatron/core/optimizer/__init__.py | 219 ++- .../optimizer_state_offloader.py | 315 ++++ megatron/core/optimizer/distrib_optimizer.py | 138 +- .../core/optimizer/emerging_optimizers.py | 378 ++++ .../core/optimizer/layer_wise_optimizer.py | 18 +- megatron/core/optimizer/muon.py | 410 +---- megatron/core/optimizer/optimizer.py | 22 +- megatron/core/optimizer/optimizer_config.py | 80 +- .../golden_values_dev_dgx_h100.json | 784 ++++---- tests/unit_tests/dist_checkpointing/utils.py | 76 +- tests/unit_tests/test_emerging_optimizers.py | 1574 +++++++++++++++++ tests/unit_tests/test_layer_wise_optimizer.py | 2 +- tests/unit_tests/test_lion_optimizer.py | 10 +- tests/unit_tests/test_muon_optimizer.py | 792 --------- tests/unit_tests/test_optimizer.py | 6 +- .../test_optimizer_state_offloading.py | 337 ++++ 16 files changed, 3442 insertions(+), 1719 deletions(-) create mode 100644 megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py create mode 100644 megatron/core/optimizer/emerging_optimizers.py create mode 100644 tests/unit_tests/test_emerging_optimizers.py delete mode 100644 tests/unit_tests/test_muon_optimizer.py create mode 100644 tests/unit_tests/test_optimizer_state_offloading.py diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index ef23ea22244..b64c871104d 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -2,6 +2,7 @@ import copy import logging import warnings +from collections import defaultdict from dataclasses import astuple from typing import Any, Callable, Dict, List, Optional, Tuple, Union @@ -34,19 +35,12 @@ USING_PYTORCH_OPTIMIZER = True try: - from importlib.metadata import PackageNotFoundError - from importlib.metadata import version as _pkg_version - - _eo_ver = tuple(int(x) for x in _pkg_version('emerging-optimizers').split('.')[:2]) -except (ImportError, PackageNotFoundError): - _eo_ver = (0, 0) - -HAVE_EMERGING_OPTIMIZERS = _eo_ver >= (0, 1) -HAVE_EO_V02 = _eo_ver >= (0, 2) - -if HAVE_EO_V02: from emerging_optimizers.scalar_optimizers import Lion + HAVE_LION = True +except ImportError: + HAVE_LION = False + from megatron.core import parallel_state from megatron.core.optimizer.cpu_offloading.hybrid_optimizer import HybridDeviceOptimizer from megatron.core.optimizer_param_scheduler import ( @@ -61,7 +55,13 @@ 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, @@ -69,6 +69,8 @@ MegatronOptimizer, param_group_identifier_keys, ) + +# Subclass aliases kept for backward compatibility; all are OptimizerConfig. from .optimizer_config import ( AdamOptimizerConfig, OptimizerConfig, @@ -317,14 +319,6 @@ 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,7 +453,8 @@ 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, -) -> MegatronOptimizer: + skip_megatron_wrapping: bool = False, +) -> Union[MegatronOptimizer, Tuple[Optional[torch.optim.Optimizer], Optional[Callable]]]: """Get Megatron optimizer based on parameter groups. Args: @@ -475,12 +470,24 @@ 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. + Instance of MegatronOptimizer, or ``(optimizer, init_state_fn)`` when + *skip_megatron_wrapping=True*. """ - # 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). + # 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.") # 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 @@ -582,12 +589,12 @@ def init_state_fn(opt, config=None): opt.initialize_state(p) elif config.optimizer == 'lion': - if not HAVE_EO_V02: + if not HAVE_LION: raise ImportError( - "Lion optimizer requires emerging_optimizers >= 0.2. " - "Please install or upgrade it to use --optimizer lion." + "Lion optimizer requires the 'emerging_optimizers' package. " + "Please install it to use --optimizer lion." ) - optimizer = Lion( # pylint: disable=possibly-used-before-assignment + optimizer = Lion( param_groups, lr=config.lr, betas=(config.lion_beta1, config.lion_beta2), @@ -614,6 +621,9 @@ 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 @@ -704,6 +714,141 @@ 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, + ) + + return ChainedOptimizer(results) + + def get_megatron_optimizer( config: OptimizerConfig, model_chunks: List[MegatronModule], @@ -714,7 +859,10 @@ 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. @@ -731,10 +879,25 @@ def get_megatron_optimizer( Instance of MegatronOptimizer. """ - log_single_rank(logger, logging.INFO, f'Setting up optimizer with config {config}') + # 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) 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/cpu_offloading/optimizer_state_offloader.py b/megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py new file mode 100644 index 00000000000..81fd116c8ba --- /dev/null +++ b/megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py @@ -0,0 +1,315 @@ +# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. + +"""Optimizer state offloading class.""" + +from typing import TYPE_CHECKING, Dict, List, Tuple + +import torch + +if TYPE_CHECKING: + from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer + + +class OptimizerStateOffloader: + """ + Manages offloading of optimizer states and master weights to CPU. + Used with DistributedOptimizer to reduce GPU memory usage. + + Supports overlapped D2H/H2D transfers using CUDA streams. + + Master weights can be stored in two locations: + - In adam optimizer state (when use_precision_aware_optimizer_no_fp8_or_ds_fp8 is True) + - In mcore's shard_fp32_from_float16_groups + """ + + OPTIMIZER_STATE_KEYS = ('exp_avg', 'exp_avg_sq') + MASTER_WEIGHT_KEY = 'master_param' + + def __init__(self, distrib_optimizer: "DistributedOptimizer"): + """ + Args: + distrib_optimizer: The DistributedOptimizer to offload states and master weights from. + """ + self.dist_optimizer = distrib_optimizer + self.adam_optimizer = distrib_optimizer.optimizer + + # Only support TE FusedAdam optimizer for now. + try: + from transformer_engine.pytorch.optimizers import FusedAdam + + assert isinstance(self.adam_optimizer, FusedAdam), ( + f"OptimizerStateOffloader requires TE FusedAdam optimizer, " + f"but got {type(self.adam_optimizer).__name__}" + ) + except ImportError: + raise ImportError( + "OptimizerStateOffloader requires transformer_engine.pytorch.optimizers.FusedAdam" + ) + + # Check if master weights are stored in adam optimizer state + self.optimizer_contains_master_weights = self.adam_optimizer.master_weights + + # CUDA streams for async transfers + self._d2h_stream = torch.cuda.Stream() + self._h2d_stream = torch.cuda.Stream() + + # CPU buffers for optimizer states: {param: {key: cpu_tensor}} + self._opt_state_cpu_buffers: Dict[torch.Tensor, Dict[str, torch.Tensor]] = {} + + # CPU buffers for mcore master weights, matching the structure of source groups + # List[List[cpu_tensor]] + self._shard_fp32_from_float16_cpu_buffers: List[List[torch.Tensor]] = [] + + # State tracking + self._offloaded = False + self._offloaded_state_keys: Tuple[str, ...] = () + self._offloaded_mcore_master_weights = False + + # Track whether optimizer states (exp_avg, exp_avg_sq) have been initialized. + # These are lazily initialized by FusedAdam during the first optimizer.step(). + # Master weights (shard_fp32_from_float16_groups) are available from the start. + self._optimizer_states_initialized = False + + def mark_optimizer_states_initialized(self): + """ + Mark that optimizer states (exp_avg, exp_avg_sq) are now available. + Should be called after the first optimizer.step() completes. + """ + self._optimizer_states_initialized = True + + def _get_state_keys_to_offload( + self, offload_optimizer_states: bool, offload_master_weights: bool + ) -> Tuple[str, ...]: + """Get the state keys in FusedAdam to offload based on configuration.""" + keys = [] + # Skip optimizer states offloading if they haven't been initialized yet. + # Optimizer states are lazily initialized by FusedAdam during the first optimizer.step(). + if self._optimizer_states_initialized: + if offload_optimizer_states: + keys.extend(self.OPTIMIZER_STATE_KEYS) + if offload_master_weights and self.optimizer_contains_master_weights: + keys.append(self.MASTER_WEIGHT_KEY) + return tuple(keys) + + def _ensure_state_cpu_buffer( + self, param: torch.Tensor, state_key: str, gpu_tensor: torch.Tensor, pin_memory: bool = True + ) -> torch.Tensor: + """Get or create a CPU buffer for a state tensor.""" + if param not in self._opt_state_cpu_buffers: + self._opt_state_cpu_buffers[param] = {} + + if state_key not in self._opt_state_cpu_buffers[param]: + cpu_buffer = torch.empty( + gpu_tensor.size(), + dtype=gpu_tensor.dtype, + layout=gpu_tensor.layout, + device='cpu', + pin_memory=pin_memory, + ) + self._opt_state_cpu_buffers[param][state_key] = cpu_buffer + + return self._opt_state_cpu_buffers[param][state_key] + + def _offload_shard_groups( + self, + shard_groups: List[List[torch.Tensor]], + cpu_buffers: List[List[torch.Tensor]], + pin_memory: bool = True, + ): + """Offload a shard group to CPU buffers.""" + # Initialize CPU buffers on first call + if len(cpu_buffers) == 0: + for group in shard_groups: + group_buffers = [] + for gpu_tensor in group: + cpu_buffer = torch.empty( + gpu_tensor.size(), + dtype=gpu_tensor.dtype, + layout=gpu_tensor.layout, + device='cpu', + pin_memory=pin_memory, + ) + group_buffers.append(cpu_buffer) + cpu_buffers.append(group_buffers) + + # Copy D2H + for group_idx, group in enumerate(shard_groups): + for param_idx, gpu_tensor in enumerate(group): + cpu_buffer = cpu_buffers[group_idx][param_idx] + cpu_buffer.copy_(gpu_tensor, non_blocking=pin_memory) + gpu_tensor.record_stream(self._d2h_stream) + + def _offload_states( + self, + offload_optimizer_states: bool, + offload_master_weights: bool, + use_pin_memory: bool = True, + ): + """Offload optimizer states and/or master weights to CPU.""" + # Offload states from adam optimizer + self._offloaded_state_keys = self._get_state_keys_to_offload( + offload_optimizer_states, offload_master_weights + ) + states = self.adam_optimizer.state + + for param, param_state in states.items(): + for state_key in self._offloaded_state_keys: + if state_key not in param_state: + continue + + gpu_tensor = param_state[state_key] + if not isinstance(gpu_tensor, torch.Tensor) or not gpu_tensor.is_cuda: + continue + + cpu_buffer = self._ensure_state_cpu_buffer( + param, state_key, gpu_tensor, use_pin_memory + ) + cpu_buffer.copy_(gpu_tensor, non_blocking=use_pin_memory) + gpu_tensor.record_stream(self._d2h_stream) + + # Offload mcore master weights if not in optimizer state + if offload_master_weights and not self.optimizer_contains_master_weights: + self._offload_shard_groups( + self.dist_optimizer.shard_fp32_from_float16_groups, + self._shard_fp32_from_float16_cpu_buffers, + use_pin_memory, + ) + self._offloaded_mcore_master_weights = True + + def _release_states(self): + """Replace optimizer state GPU tensors with CPU tensors to free GPU memory.""" + states = self.adam_optimizer.state + + for param, param_state in states.items(): + if param not in self._opt_state_cpu_buffers: + continue + + for state_key in self._offloaded_state_keys: + if state_key not in self._opt_state_cpu_buffers[param]: + continue + + param_state[state_key].untyped_storage().resize_(0) + + if self._offloaded_mcore_master_weights: + for group in self.dist_optimizer.shard_fp32_from_float16_groups: + for gpu_tensor in group: + gpu_tensor.untyped_storage().resize_(0) + + def _reload_shard_groups( + self, + shard_groups: List[List[torch.Tensor]], + cpu_buffers: List[List[torch.Tensor]], + is_allocate_stage: bool, + ): + """Reload shard groups from CPU to GPU.""" + for group_idx, group in enumerate(shard_groups): + for param_idx, _ in enumerate(group): + cpu_buffer = cpu_buffers[group_idx][param_idx] + if is_allocate_stage: + shard_groups[group_idx][param_idx].untyped_storage().resize_( + cpu_buffer.untyped_storage().size() + ) + else: + shard_groups[group_idx][param_idx].copy_( + cpu_buffer, non_blocking=cpu_buffer.is_pinned() + ) + + def _reload_states(self, is_allocate_stage: bool): + """ + Reload optimizer states and/or master weights from CPU to GPU. + + If is_allocate_stage is True, only allocate GPU memory for the states and master weights, + but do not copy the data from CPU to GPU. Otherwise, copy the data from CPU to GPU. + The two processes are separated to make sure that the GPU memory is allocated on the + default stream to avoid fragmentation. + """ + # Reload states to adam optimizer + states = self.adam_optimizer.state + + for param, param_state in states.items(): + if param not in self._opt_state_cpu_buffers: + continue + + for state_key in self._offloaded_state_keys: + if state_key not in self._opt_state_cpu_buffers[param]: + continue + + cpu_buffer = self._opt_state_cpu_buffers[param][state_key] + if is_allocate_stage: + param_state[state_key].untyped_storage().resize_( + cpu_buffer.untyped_storage().size() + ) + else: + param_state[state_key].copy_(cpu_buffer, non_blocking=cpu_buffer.is_pinned()) + + # Reload mcore master weights if not in optimizer state + if self._offloaded_mcore_master_weights: + self._reload_shard_groups( + self.dist_optimizer.shard_fp32_from_float16_groups, + self._shard_fp32_from_float16_cpu_buffers, + is_allocate_stage, + ) + + def offload(self, offload_optimizer_states: bool = True, offload_master_weights: bool = True): + """ + Offload optimizer states and/or master weights to CPU. + Starts async D2H transfer that can overlap with other operations. + + Args: + offload_optimizer_states: Whether to offload exp_avg, exp_avg_sq. + offload_master_weights: Whether to offload master weights. + """ + if not offload_optimizer_states and not offload_master_weights: + return + + # Wait for current stream finishing updating the optimizer states. + self._d2h_stream.wait_stream(torch.cuda.current_stream()) + + with torch.cuda.stream(self._d2h_stream): + self._offload_states(offload_optimizer_states, offload_master_weights) + + self._offloaded = True + + def release_gpu_memory(self): + """ + Release GPU memory for optimizer states and master weights after D2H copy completes. + + This is separated from offload() to allow delayed GPU memory release, + which is needed for mxfp8 + overlap_param_gather case where master weights + must remain on GPU until after _copy_main_params_to_param_buffer() is called. + """ + if not self._offloaded: + return + + self._release_states() + + def reload(self): + """ + Reload optimizer states and/or master weights from CPU to GPU. + Call before optimizer.step() to ensure states are on GPU. + """ + if not self._offloaded: + return + + # Allocate GPU memory on the current stream to avoid fragmentation. + self._reload_states(is_allocate_stage=True) + + self._h2d_stream.wait_stream(self._d2h_stream) + self._h2d_stream.wait_stream(torch.cuda.current_stream()) + + # Reload states on the h2d stream to overlap with other operations. + with torch.cuda.stream(self._h2d_stream): + self._reload_states(is_allocate_stage=False) + + self._offloaded_state_keys = () + self._offloaded_mcore_master_weights = False + self._offloaded = False + + def sync_before_step(self): + """ + Wait for H2D reload to complete before optimizer.step(). + Must be called to ensure states are on GPU before optimizer uses them. + + This is separated from reload() to make it possible to move the reload ahead of time. + """ + torch.cuda.current_stream().wait_stream(self._h2d_stream) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index eeda383a75d..beb00391759 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -52,6 +52,7 @@ from ..fp8_utils import dequantize_fp8_tensor, is_float8tensor, quantize_param_shard from ..transformer.fsdp_dtensor_checkpoint import handle_experts_in_state_dict from ..transformer.module import MegatronModule +from .cpu_offloading.optimizer_state_offloader import OptimizerStateOffloader from .grad_scaler import MegatronGradScaler from .optimizer import MixedPrecisionOptimizer, _zero_grad_group_helper, param_group_identifier_keys from .optimizer_config import OptimizerConfig @@ -361,7 +362,10 @@ def _build_model_and_main_param_groups( if model_param.type() in ['torch.cuda.HalfTensor', 'torch.cuda.BFloat16Tensor']: # Generate sharded model param. - if is_float8tensor(model_param) and config.fp8_recipe != "delayed": + if ( + cls._is_distopt_quantized_param(model_param) + and config.fp8_recipe != "delayed" + ): # MXFP8Tensor and BlockwiseQTensor don't support view(-1) shard_model_param = None else: @@ -381,7 +385,7 @@ def _build_model_and_main_param_groups( # precision at the beginning of training (this problem will not occur if the # training is long enough or if the main params are loaded from a # checkpoint). - if is_float8tensor(model_param): + if cls._is_distopt_quantized_param(model_param): if hasattr(model_param, 'get_high_precision_init_val'): shard_main_param = ( model_param.get_high_precision_init_val() @@ -519,6 +523,8 @@ def __init__( "due to checkpointing requirements." ) + self._state_offloader: Optional[OptimizerStateOffloader] = None + # when freezing sub-models we have no real optimizer # but still need a stub DistributedOptimizer class if optimizer is None: @@ -607,6 +613,9 @@ def __init__( self.optimizer.param_groups = [g["orig_group"] for g in self.opt_group_ranges] self.optimizer.load_state_dict(self.optimizer.state_dict()) + if self.config.offload_optimizer_states: + self._state_offloader = OptimizerStateOffloader(self) + def _get_model_param_range_map(self, param: torch.nn.Parameter): """ Given a model param, get the index sub-range of the param that this @@ -913,6 +922,70 @@ def _get_main_param_and_optimizer_states(self, model_param): tensors[k] = v return tensors + @staticmethod + def _is_grouped_quantized_tensor(tensor: torch.Tensor) -> bool: + """Check if tensor is a TE GroupedTensor using quantized storage.""" + return ( + hasattr(tensor, "split_into_quantized_tensors") + and callable(tensor.split_into_quantized_tensors) + and getattr(tensor, "quantizer", None) is not None + ) + + @classmethod + def _is_distopt_quantized_param(cls, tensor: torch.Tensor) -> bool: + """Check if tensor should follow quantized parameter path in dist optimizer.""" + return is_float8tensor(tensor) or cls._is_grouped_quantized_tensor(tensor) + + def _expand_quantized_param_shard_for_cast( + self, + model_param: torch.Tensor, + shard_main_param: Optional[torch.Tensor], + start_offset: Optional[int], + ): + """Expand one quantized model param to cast-ready entries. + + For grouped quantized tensors, split into member quantized tensors and map the sharded + master slice to per-member offset ranges, while preserving deterministic ordering across + DP ranks. + """ + if not self._is_grouped_quantized_tensor(model_param): + return [model_param], [shard_main_param], [start_offset] + + quantized_members = model_param.quantized_tensors + if quantized_members is None: + quantized_members = model_param.split_into_quantized_tensors() + + shard_start = 0 if start_offset is None else start_offset + shard_size = 0 if shard_main_param is None else shard_main_param.numel() + shard_end = shard_start + shard_size + shard_flat = None if shard_main_param is None else shard_main_param.view(-1) + + expanded_model_params = [] + expanded_shard_main_params = [] + expanded_start_offsets = [] + member_offset = 0 + for member in quantized_members: + member_numel = member.numel() + member_start = member_offset + member_end = member_start + member_numel + overlap_start = max(member_start, shard_start) + overlap_end = min(member_end, shard_end) + + member_master = None + member_start_offset = None + if overlap_start < overlap_end: + local_start = overlap_start - shard_start + local_end = overlap_end - shard_start + member_master = shard_flat[local_start:local_end] + member_start_offset = overlap_start - member_start + + expanded_model_params.append(member) + expanded_shard_main_params.append(member_master) + expanded_start_offsets.append(member_start_offset) + member_offset = member_end + + return expanded_model_params, expanded_shard_main_params, expanded_start_offsets + def _set_main_param_and_optimizer_states(self, model_param, tensors): """Set the main param and optimizer states corresponding to the input model_param. @@ -2145,7 +2218,7 @@ def split_state_dict_if_needed(self, state_dict): fp8_gbuf_indices = [] for gbuf_idx, gbuf_range_maps in enumerate(self.gbuf_ranges): for dtype, _ in gbuf_range_maps.items(): - if is_float8tensor(self.buffers[gbuf_idx].params[0]): + if self._is_distopt_quantized_param(self.buffers[gbuf_idx].params[0]): fp8_gbuf_indices.append(gbuf_idx) if len(fp8_gbuf_indices) == 0: return @@ -2167,7 +2240,7 @@ def split_state_dict_if_needed(self, state_dict): new_state_dict = {'buckets_coalesced': state_dict['buckets_coalesced']} for gbuf_idx, gbuf_range_maps in enumerate(self.gbuf_ranges): for dtype, _ in gbuf_range_maps.items(): - if not is_float8tensor(self.buffers[gbuf_idx].params[0]): + if not self._is_distopt_quantized_param(self.buffers[gbuf_idx].params[0]): new_state_dict[gbuf_idx] = state_dict[dtype_to_gbuf_idx[dtype]] for fp8_gbuf_idx in fp8_gbuf_indices: @@ -2367,7 +2440,7 @@ def _get_fp8_params_and_shard_fp32_from_fp8(self): idx = 0 for buffer in buffers: for param in buffer.params: - if is_float8tensor(param): + if self._is_distopt_quantized_param(param): fp8_params.append(param) shard_fp32_from_fp8.append(None) shard_offsets_in_fp8.append(None) @@ -2382,7 +2455,7 @@ def get_shard_fp32_from_fp8(shard_main_groups, model_groups): """ for shard_main_group, model_group in zip(shard_main_groups, model_groups): for shard_main_param, model_param in zip(shard_main_group, model_group): - if is_float8tensor(model_param): + if self._is_distopt_quantized_param(model_param): param_range_map = self._get_model_param_range_map(model_param) param_range = param_range_map["param"] assert param_range.size == shard_main_param.nelement() @@ -2459,8 +2532,29 @@ def _copy_main_params_to_model_params(self): if self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8: return + fp8_params, shard_fp32_from_fp8, shard_offsets_in_fp8 = ( + self._get_fp8_params_and_shard_fp32_from_fp8() + ) + expanded_fp8_params = [] + expanded_shard_fp32_from_fp8 = [] + expanded_shard_offsets_in_fp8 = [] + for model_param, shard_main_param, start_offset in zip( + fp8_params, shard_fp32_from_fp8, shard_offsets_in_fp8 + ): + sub_model_params, sub_shard_main_params, sub_start_offsets = ( + self._expand_quantized_param_shard_for_cast( + model_param, shard_main_param, start_offset + ) + ) + expanded_fp8_params.extend(sub_model_params) + expanded_shard_fp32_from_fp8.extend(sub_shard_main_params) + expanded_shard_offsets_in_fp8.extend(sub_start_offsets) + quantize_param_shard( - *self._get_fp8_params_and_shard_fp32_from_fp8(), self.data_parallel_group + expanded_fp8_params, + expanded_shard_fp32_from_fp8, + expanded_shard_offsets_in_fp8, + self.data_parallel_group, ) # Utility method for copying group params. @@ -2480,7 +2574,7 @@ def copy_group_params(shard_main_groups, model_groups): world_range.start : world_range.end ] - if is_float8tensor(model_param): + if self._is_distopt_quantized_param(model_param): # FP8 params are quantized in the above "quantize_param_shard" function. continue else: @@ -2592,8 +2686,12 @@ def copy_group_params(model_groups, shard_main_groups): # Use param from state_dict to initialize main_param model_param = model_param_to_state_dict_param_map[model_param] - if is_float8tensor(model_param): - shard_model_param = dequantize_fp8_tensor(model_param).view(-1)[ + if self._is_distopt_quantized_param(model_param): + if self._is_grouped_quantized_tensor(model_param): + dequantized_model_param = model_param.float() + else: + dequantized_model_param = dequantize_fp8_tensor(model_param) + shard_model_param = dequantized_model_param.view(-1)[ param_range.start : param_range.end ] else: @@ -2612,6 +2710,8 @@ def step_with_ready_grads(self) -> bool: Under the hood, either launch synchronous param all-gathers or get ready to launch asynchorous all-gathers that get overlapped with the next forward pass. """ + if self._state_offloader is not None: + self._state_offloader.sync_before_step() update_successful = super().step_with_ready_grads() timers = self.config.timers @@ -2632,4 +2732,22 @@ def step_with_ready_grads(self) -> bool: if timers is not None: timers('params-all-gather').stop() + if self._state_offloader is not None: + self._state_offloader.mark_optimizer_states_initialized() + return update_successful + + def offload_states(self): + """Offload states to CPU.""" + if self._state_offloader is not None: + self._state_offloader.offload() + + def reload_offloaded_states(self): + """Start async reload of offloaded states.""" + if self._state_offloader is not None: + self._state_offloader.reload() + + def release_offloaded_gpu_states(self): + """Release GPU memory after D2H completes. For delayed release case.""" + if self._state_offloader is not None: + self._state_offloader.release_gpu_memory() diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py new file mode 100644 index 00000000000..74a0d0204f3 --- /dev/null +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -0,0 +1,378 @@ +# 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 + +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 newton_schulz_tp + from emerging_optimizers.scalar_optimizers import Lion # pylint: disable=unused-import + + # It is necessary to import optimizers for the registry to work. + 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__) + + +# =========================================================================== +# 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. +# =========================================================================== + + +# =========================================================================== +# 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.""" + + 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 = { + 'muon': EmergingOptimizerEntry( + optimizer_cls=TensorParallelMuon, config_to_kwargs=_muon_config_to_kwargs + ), + "adaptive_muon": EmergingOptimizerEntry( + optimizer_cls=TensorParallelAdaptiveMuon, config_to_kwargs=_adaptive_muon_config_to_kwargs + ), +} + +# 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 a9fdc7ba72f..6e0f32ab357 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -46,7 +46,6 @@ def __init__( pg_collection: Optional[ProcessGroupCollection] = None, init_state_fn_list: Optional[List[Callable]] = None, model_chunks: Optional[List] = None, - async_allgather: bool = False, ) -> None: """ Initialize LayerWiseDistributedOptimizer. @@ -57,14 +56,13 @@ def __init__( pg_collection: ProcessGroupCollection. init_state_fn_list: List of init state functions. model_chunks: DDP-wrapped model chunks (needed for async_allgather). - async_allgather: If True, defer param all-gather to forward pre-hooks. """ self.pg_collection = pg_collection self.shard_params(optimizers) # Set up async all-gather using DDP bucket infrastructure. - self.async_allgather = async_allgather + self.async_allgather = config.overlap_param_gather if self.async_allgather: assert ( model_chunks is not None @@ -76,19 +74,17 @@ def __init__( optimizers ), "init_state_fn_list must be the same length as optimizers if provided" - # wrap optimizer after sharding to avoid unnecessary master weight creation - # for higher precision, optimizers are wrapped with megatron already + # 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. 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): + if isinstance(opt, (Float16OptimizerWithFloat16Params, FP32Optimizer)): raise TypeError( - 'LayerWiseDistributedOptimizer received Float16 optimizer already.' + 'LayerWiseDistributedOptimizer expects base torch optimizers, ' + f'got {type(opt).__name__}. Do not pre-wrap with Megatron optimizers.' ) - # 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 046be78ad10..329ce60dd1f 100644 --- a/megatron/core/optimizer/muon.py +++ b/megatron/core/optimizer/muon.py @@ -1,402 +1,30 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -"""Megatron muon optimizer wrapper to handle tensor-parallel.""" +"""Backward-compatible shim — all code now lives in ``emerging_optimizers``.""" -import logging -from typing import Any, Callable, Dict, List, Literal, Optional, get_args +from typing import Any -import torch -from torch.optim.optimizer import ParamsT +# TODO: Remove this separate try/except once the next version of emerging_optimizers +# (which includes Lion) is released. Then Lion can be imported in the block above. +try: + from emerging_optimizers.scalar_optimizers import Lion # pylint: disable=unused-import -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 + HAVE_LION = True +except ImportError: + HAVE_LION = False -from . import HAVE_EMERGING_OPTIMIZERS, HAVE_EO_V02, _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 +def get_megatron_muon_optimizer(*args: Any, **kwargs: Any) -> Any: + """Backward compatible muon optimizer getter. -if HAVE_EO_V02: - 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. + .. deprecated:: + Use :func:`megatron.core.optimizer.get_megatron_optimizer` instead. """ - assert ( - HAVE_EO_V02 - ), "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 HAVE_EO_V02 else ("quintic",) - 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} if HAVE_EO_V02 else {"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} if HAVE_EO_V02 else {"use_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. - """ - # 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 - - if config.muon_scalar_optimizer == 'lion': - assert HAVE_EO_V02, ( - "Lion optimizer requires emerging_optimizers >= 0.2. " - "Please upgrade to use --muon-scalar-optimizer lion." - ) - else: - assert HAVE_EMERGING_OPTIMIZERS, "Emerging Optimizers 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 + from . import get_megatron_optimizer - # chain everything together - init_fns = [muon_init_state_fn] + len(chained_adam.chained_optimizers) * [ - nonlinear_init_state_fn - ] - optimizers += chained_adam.chained_optimizers + if kwargs.pop('layer_wise_distributed_optimizer', False): + config = args[0] if args else kwargs.get('config') + if config is not None: + config.use_layer_wise_distributed_optimizer = True - 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) + return get_megatron_optimizer(*args, **kwargs) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index df8ec8ef613..f5d66b8db4f 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -1161,20 +1161,26 @@ def _split_state_dict(self, state_dict): state_dicts = [None] * len(self.chained_optimizers) if state_dict is not None: if len(self.model_chunks) == 1: - state_dicts[0] = state_dict + # When there is only one global model chunk, all sub-optimizers + # (e.g., dense and MoE parts) use the same model state dict. + state_dicts = [state_dict] * len(self.chained_optimizers) else: - # Split state_dict if needed + # Split state_dict by model chunk object. prefix = "model" if "model0" in state_dict.keys() else "model_" - offset = 0 + chunk_to_global_idx = {chunk: idx for idx, chunk in enumerate(self.model_chunks)} for optimizer_idx, optimizer in enumerate(self.chained_optimizers): if hasattr(optimizer, "model_chunks"): d = {} - for chunk_idx in range(len(optimizer.model_chunks)): + for chunk_idx, model_chunk in enumerate(optimizer.model_chunks): + assert model_chunk in chunk_to_global_idx, ( + "Sub-optimizer model chunk was not found in " + "chained optimizer model chunks" + ) + global_idx = chunk_to_global_idx[model_chunk] assert ( - f"{prefix}{offset}" in state_dict - ), f"Wrong state_dict format, cannot find '{prefix}{offset}'" - d[f"{prefix}{chunk_idx}"] = state_dict[f"{prefix}{offset}"] - offset += 1 + f"{prefix}{global_idx}" in state_dict + ), f"Wrong state_dict format, cannot find '{prefix}{global_idx}'" + d[f"{prefix}{chunk_idx}"] = state_dict[f"{prefix}{global_idx}"] if len(d) > 0: state_dicts[optimizer_idx] = d return state_dicts diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 9e6375b978c..d425a56be71 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -142,7 +142,6 @@ class OptimizerConfig: ############## # General ############## - lr: Optional[float] = None """Initial learning rate. Depending on decay style and initial warmup, the learning rate at each iteration would be different. @@ -207,7 +206,8 @@ class OptimizerConfig: """dtype of exp_avg_sq when enabling precision-aware-optimizer""" optimizer: str = 'adam' - """Optimizer name. NOTE: Deprecated, use individual optimizer classes instead.""" + """Optimizer name (e.g., 'adam', 'sgd', 'muon'). Can be overridden per-parameter group + via config_overrides to use different optimizers for different parameters.""" ############### # Loss scaling @@ -230,7 +230,7 @@ class OptimizerConfig: """Hysteresis for dynamic loss scaling.""" ################################################################################### - # Optimizer (NOTE: Deprecated, use individual optimizer classes instead.). + # Optimizer-specific parameters. ################################################################################### # Adam. adam_beta1: float = 0.9 @@ -255,15 +255,14 @@ class OptimizerConfig: sgd_momentum: float = 0.9 """Momentum factor for SGD optimizer.""" - # Muon. - # TODO: move muon configs to it's own `MuonConfig`. + # emerging optimizers. muon_momentum: float = 0.95 - """The momentum used by the internal SGD.""" + """The momentum used by the internal SGD in Muon optimizer.""" muon_split_qkv: bool = True """Whether to split QKV parameters for Muon optimizer.""" - muon_use_nesterov: bool = False + muon_nesterov: bool = False """Whether to use Nesterov-style momentum in the internal SGD.""" muon_scale_mode: str = "spectral" @@ -272,10 +271,6 @@ class OptimizerConfig: muon_fp32_matmul_prec: str = "medium" """The precision to use for the fp32 matmul. Defaults to "medium".""" - muon_coefficient_type: str = "quintic" - """Newton-Schulz coefficient type for the Muon optimizer. Valid types are discovered - dynamically from the installed ``emerging_optimizers`` package. Defaults to "quintic".""" - muon_num_ns_steps: int = 5 """The number of iteration steps to use in the Newton-Schulz iteration.""" @@ -285,6 +280,24 @@ class OptimizerConfig: muon_extra_scale_factor: float = 1.0 """Additional scale factor for the muon update.""" + 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.""" + muon_scalar_optimizer: str = 'adam' """Optimizer for nonlinear parameters (embeddings, biases, norms) when using muon. One of 'adam' or 'lion'. Defaults to 'adam'.""" @@ -303,6 +316,12 @@ class OptimizerConfig: 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 @@ -341,6 +360,12 @@ class OptimizerConfig: pin_cpu_params: bool = True """If True, pin the optimizer parameters to CPU memory.""" + offload_optimizer_states: bool = False + """ + If True, offload optimizer states to CPU after each optimizer step and + reload them before the next optimizer step. + """ + ################ # Miscellaneous ################ @@ -442,33 +467,6 @@ def __post_init__(self): ), "exp_avg_sq_dtype can only be fp32 when not using precision-aware optimizer" -@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.""" +# Backward-compatible aliases (deprecated; use OptimizerConfig directly). +AdamOptimizerConfig = OptimizerConfig +SGDOptimizerConfig = OptimizerConfig diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json index f529a646a7e..9533c3e29a1 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json @@ -8,102 +8,102 @@ "2": 10.91072, "3": 10.91895, "4": 10.91763, - "5": 10.90484, - "6": 10.90203, - "7": 10.89753, - "8": 10.91294, - "9": 10.91701, - "10": 10.91028, - "11": 10.90124, - "12": 10.89698, - "13": 10.88788, - "14": 10.89478, - "15": 10.87488, - "16": 10.87022, - "17": 10.86892, - "18": 10.85196, - "19": 10.87008, - "20": 10.7881, - "21": 10.77222, - "22": 10.7669, - "23": 10.75865, - "24": 10.71955, - "25": 10.71987, - "26": 10.71249, - "27": 10.68554, - "28": 10.61292, - "29": 10.58664, - "30": 10.56554, - "31": 10.55749, - "32": 10.54875, - "33": 10.50948, - "34": 10.48165, - "35": 10.46995, - "36": 10.45309, - "37": 10.42791, - "38": 10.43268, - "39": 10.40324, - "40": 10.3773, - "41": 10.36856, - "42": 10.33125, - "43": 10.31537, - "44": 10.29014, - "45": 10.30253, - "46": 10.26536, - "47": 10.25557, - "48": 10.20689, - "49": 10.21031, - "50": 10.2105, - "51": 10.21191, - "52": 10.16277, - "53": 10.16315, - "54": 10.13391, - "55": 10.10867, - "56": 10.13455, + "5": 10.90462, + "6": 10.90222, + "7": 10.89756, + "8": 10.91282, + "9": 10.91678, + "10": 10.9104, + "11": 10.9015, + "12": 10.89781, + "13": 10.8883, + "14": 10.89516, + "15": 10.87477, + "16": 10.87004, + "17": 10.86866, + "18": 10.85186, + "19": 10.87023, + "20": 10.78833, + "21": 10.7724, + "22": 10.76686, + "23": 10.75821, + "24": 10.71892, + "25": 10.72027, + "26": 10.71214, + "27": 10.68529, + "28": 10.61314, + "29": 10.58641, + "30": 10.56586, + "31": 10.5575, + "32": 10.5488, + "33": 10.50937, + "34": 10.48155, + "35": 10.47006, + "36": 10.45297, + "37": 10.42758, + "38": 10.43258, + "39": 10.40282, + "40": 10.37727, + "41": 10.36865, + "42": 10.33123, + "43": 10.31512, + "44": 10.29023, + "45": 10.30268, + "46": 10.26547, + "47": 10.25564, + "48": 10.20686, + "49": 10.21056, + "50": 10.21037, + "51": 10.21194, + "52": 10.16248, + "53": 10.16319, + "54": 10.13395, + "55": 10.10854, + "56": 10.13474, "57": 10.13262, - "58": 10.12407, - "59": 10.06503, - "60": 10.09528, - "61": 10.04743, - "62": 10.01537, - "63": 10.08286, - "64": 10.03273, - "65": 9.99833, - "66": 10.03902, - "67": 10.01293, - "68": 9.97751, - "69": 9.99331, - "70": 9.97079, - "71": 9.99817, - "72": 9.97548, - "73": 9.95979, - "74": 9.95289, - "75": 9.91425, - "76": 9.9499, - "77": 9.94212, - "78": 9.89883, - "79": 9.89693, - "80": 9.91029, - "81": 9.93356, - "82": 9.88352, - "83": 9.83982, - "84": 9.78195, - "85": 9.76266, - "86": 9.87794, - "87": 9.90072, - "88": 9.87398, - "89": 9.82485, - "90": 9.81362, - "91": 9.8199, - "92": 9.81611, - "93": 9.74343, - "94": 9.82156, - "95": 9.8122, - "96": 9.79476, - "97": 9.74624, - "98": 9.76879, - "99": 9.81836, - "100": 9.7074 + "58": 10.124, + "59": 10.06483, + "60": 10.09511, + "61": 10.04736, + "62": 10.01513, + "63": 10.08268, + "64": 10.03239, + "65": 9.99804, + "66": 10.03859, + "67": 10.01247, + "68": 9.97703, + "69": 9.9927, + "70": 9.97031, + "71": 9.99747, + "72": 9.97476, + "73": 9.95896, + "74": 9.95212, + "75": 9.9133, + "76": 9.94908, + "77": 9.94119, + "78": 9.89795, + "79": 9.89601, + "80": 9.90926, + "81": 9.93266, + "82": 9.8826, + "83": 9.83875, + "84": 9.78078, + "85": 9.76158, + "86": 9.87689, + "87": 9.89972, + "88": 9.87298, + "89": 9.82372, + "90": 9.81265, + "91": 9.81889, + "92": 9.81491, + "93": 9.74217, + "94": 9.82042, + "95": 9.81103, + "96": 9.79363, + "97": 9.74488, + "98": 9.76721, + "99": 9.81701, + "100": 9.70593 } }, "num-zeros": { @@ -111,106 +111,106 @@ "end_step": 100, "step_interval": 1, "values": { - "1": 2589.0, - "2": 2610.0, - "3": 2532.0, - "4": 2530.0, - "5": 2535.0, - "6": 2504.0, - "7": 2664.0, - "8": 2529.0, - "9": 2641.0, - "10": 2550.0, - "11": 2654.0, - "12": 2438.0, - "13": 2617.0, - "14": 2645.0, - "15": 2328.0, - "16": 2493.0, - "17": 2550.0, - "18": 2599.0, - "19": 2441.0, - "20": 2491.0, - "21": 2583.0, - "22": 2562.0, - "23": 2470.0, - "24": 2588.0, - "25": 2439.0, - "26": 2535.0, - "27": 2589.0, - "28": 2534.0, - "29": 2637.0, - "30": 2716.0, - "31": 2705.0, - "32": 2812.0, - "33": 2835.0, - "34": 2727.0, - "35": 2870.0, - "36": 2698.0, - "37": 2921.0, - "38": 2783.0, - "39": 2848.0, - "40": 3037.0, - "41": 3154.0, - "42": 2864.0, - "43": 3103.0, - "44": 3123.0, - "45": 3271.0, - "46": 3208.0, - "47": 3206.0, - "48": 3309.0, - "49": 3457.0, - "50": 3466.0, - "51": 3276.0, - "52": 3448.0, - "53": 3254.0, - "54": 3504.0, - "55": 3230.0, - "56": 3568.0, - "57": 2933.0, - "58": 4052.0, - "59": 3626.0, - "60": 3510.0, - "61": 3371.0, - "62": 3642.0, - "63": 4019.0, - "64": 4041.0, - "65": 3371.0, - "66": 3826.0, - "67": 4156.0, - "68": 3811.0, - "69": 3545.0, - "70": 3831.0, - "71": 3834.0, - "72": 3593.0, - "73": 4098.0, - "74": 3711.0, - "75": 3649.0, - "76": 3907.0, - "77": 4118.0, - "78": 4212.0, - "79": 4428.0, - "80": 33291.0, - "81": 8226.0, - "82": 528724.0, - "83": 3499.0, - "84": 31529.0, - "85": 528713.0, - "86": 529264.0, - "87": 581775.0, - "88": 529230.0, - "89": 529270.0, - "90": 529149.0, - "91": 528757.0, - "92": 529091.0, - "93": 549748.0, - "94": 529131.0, - "95": 553058.0, - "96": 560607.0, - "97": 529708.0, - "98": 529488.0, - "99": 529121.0, - "100": 529245.0 + "1": 6427.0, + "2": 6618.0, + "3": 6705.0, + "4": 6626.0, + "5": 6454.0, + "6": 6215.0, + "7": 6854.0, + "8": 6253.0, + "9": 6519.0, + "10": 6579.0, + "11": 6610.0, + "12": 6245.0, + "13": 6667.0, + "14": 6918.0, + "15": 6294.0, + "16": 6413.0, + "17": 6473.0, + "18": 6473.0, + "19": 6481.0, + "20": 6284.0, + "21": 6610.0, + "22": 6553.0, + "23": 6354.0, + "24": 6699.0, + "25": 6464.0, + "26": 6614.0, + "27": 6724.0, + "28": 6671.0, + "29": 7037.0, + "30": 6976.0, + "31": 7135.0, + "32": 7146.0, + "33": 7088.0, + "34": 7123.0, + "35": 7319.0, + "36": 7225.0, + "37": 7638.0, + "38": 7696.0, + "39": 7778.0, + "40": 7985.0, + "41": 8138.0, + "42": 7526.0, + "43": 8067.0, + "44": 7962.0, + "45": 8660.0, + "46": 8468.0, + "47": 8513.0, + "48": 8547.0, + "49": 8878.0, + "50": 8823.0, + "51": 8750.0, + "52": 8942.0, + "53": 8470.0, + "54": 9274.0, + "55": 8387.0, + "56": 9552.0, + "57": 7729.0, + "58": 10444.0, + "59": 9320.0, + "60": 9455.0, + "61": 8934.0, + "62": 9447.0, + "63": 10085.0, + "64": 10049.0, + "65": 8632.0, + "66": 9644.0, + "67": 10241.0, + "68": 9905.0, + "69": 8978.0, + "70": 9730.0, + "71": 9629.0, + "72": 9249.0, + "73": 10081.0, + "74": 14397.0, + "75": 8917.0, + "76": 10143.0, + "77": 10427.0, + "78": 10760.0, + "79": 68696.0, + "80": 132664.0, + "81": 80159.0, + "82": 1117640.0, + "83": 67014.0, + "84": 1112297.0, + "85": 2106479.0, + "86": 2108092.0, + "87": 1279087.0, + "88": 2107686.0, + "89": 2111718.0, + "90": 1059710.0, + "91": 2106808.0, + "92": 2106945.0, + "93": 3155405.0, + "94": 2107876.0, + "95": 2155420.0, + "96": 2170260.0, + "97": 2108441.0, + "98": 2107668.0, + "99": 2107336.0, + "100": 2107900.0 } }, "mem-allocated-bytes": { @@ -327,104 +327,104 @@ "values": { "1": 974333952.0, "2": 1142500864.0, - "3": 1142675968.0, - "4": 1147437056.0, - "5": 1147925504.0, - "6": 1147925504.0, - "7": 1148942336.0, - "8": 1148942336.0, - "9": 1148942336.0, - "10": 1148942336.0, - "11": 1148942336.0, - "12": 1148942336.0, - "13": 1148942336.0, - "14": 1148942336.0, - "15": 1148942336.0, - "16": 1148942336.0, - "17": 1148942336.0, - "18": 1148942336.0, - "19": 1148942336.0, - "20": 1148942336.0, - "21": 1148942336.0, - "22": 1148942336.0, - "23": 1148942336.0, - "24": 1148942336.0, - "25": 1148942336.0, - "26": 1149713920.0, - "27": 1149713920.0, - "28": 1149713920.0, - "29": 1149713920.0, - "30": 1149713920.0, - "31": 1149713920.0, - "32": 1149713920.0, - "33": 1149713920.0, - "34": 1149713920.0, - "35": 1149713920.0, - "36": 1149713920.0, - "37": 1149713920.0, - "38": 1149713920.0, - "39": 1149713920.0, - "40": 1149713920.0, - "41": 1149713920.0, - "42": 1149713920.0, - "43": 1149713920.0, - "44": 1149713920.0, - "45": 1149713920.0, - "46": 1149713920.0, - "47": 1149713920.0, - "48": 1149713920.0, - "49": 1149713920.0, - "50": 1149713920.0, - "51": 1149713920.0, - "52": 1149713920.0, - "53": 1149713920.0, - "54": 1149713920.0, - "55": 1149713920.0, - "56": 1149713920.0, - "57": 1149713920.0, - "58": 1149713920.0, - "59": 1149713920.0, - "60": 1149713920.0, - "61": 1149713920.0, - "62": 1149713920.0, - "63": 1149713920.0, - "64": 1149713920.0, - "65": 1149713920.0, - "66": 1149713920.0, - "67": 1149713920.0, - "68": 1149713920.0, - "69": 1149713920.0, - "70": 1149713920.0, - "71": 1149713920.0, - "72": 1149713920.0, - "73": 1149713920.0, - "74": 1149713920.0, - "75": 1149713920.0, - "76": 1149713920.0, - "77": 1149713920.0, - "78": 1149713920.0, - "79": 1149713920.0, - "80": 1149713920.0, - "81": 1149713920.0, - "82": 1149713920.0, - "83": 1149713920.0, - "84": 1149713920.0, - "85": 1149713920.0, - "86": 1149713920.0, - "87": 1149713920.0, - "88": 1149713920.0, - "89": 1149713920.0, - "90": 1149713920.0, - "91": 1149713920.0, - "92": 1149713920.0, - "93": 1149713920.0, - "94": 1149713920.0, - "95": 1149713920.0, - "96": 1149713920.0, - "97": 1149713920.0, - "98": 1149713920.0, - "99": 1149713920.0, - "100": 1149713920.0 + "3": 1142671872.0, + "4": 1147373568.0, + "5": 1147845632.0, + "6": 1147845632.0, + "7": 1148584448.0, + "8": 1148584448.0, + "9": 1148584448.0, + "10": 1148584448.0, + "11": 1148584448.0, + "12": 1148584448.0, + "13": 1148584448.0, + "14": 1148584448.0, + "15": 1148584448.0, + "16": 1148584448.0, + "17": 1148584448.0, + "18": 1148584448.0, + "19": 1148584448.0, + "20": 1148584448.0, + "21": 1148584448.0, + "22": 1148584448.0, + "23": 1148584448.0, + "24": 1148584448.0, + "25": 1148584448.0, + "26": 1148584448.0, + "27": 1148584448.0, + "28": 1148584448.0, + "29": 1148584448.0, + "30": 1148584448.0, + "31": 1148584448.0, + "32": 1148584448.0, + "33": 1148584448.0, + "34": 1148584448.0, + "35": 1148595200.0, + "36": 1148595200.0, + "37": 1148595200.0, + "38": 1148595200.0, + "39": 1148595200.0, + "40": 1148595200.0, + "41": 1148595200.0, + "42": 1148595200.0, + "43": 1148595200.0, + "44": 1148595200.0, + "45": 1148595200.0, + "46": 1148595200.0, + "47": 1148595200.0, + "48": 1148595200.0, + "49": 1148595200.0, + "50": 1148595200.0, + "51": 1148595200.0, + "52": 1148595200.0, + "53": 1148595200.0, + "54": 1148595200.0, + "55": 1148595200.0, + "56": 1148595200.0, + "57": 1148595200.0, + "58": 1148595200.0, + "59": 1148595200.0, + "60": 1148595200.0, + "61": 1148595200.0, + "62": 1148595200.0, + "63": 1148595200.0, + "64": 1148595200.0, + "65": 1148595200.0, + "66": 1148595200.0, + "67": 1148595200.0, + "68": 1148595200.0, + "69": 1148595200.0, + "70": 1148595200.0, + "71": 1148595200.0, + "72": 1148595200.0, + "73": 1148595200.0, + "74": 1148595200.0, + "75": 1148595200.0, + "76": 1148595200.0, + "77": 1148595200.0, + "78": 1148595200.0, + "79": 1148595200.0, + "80": 1148595200.0, + "81": 1148595200.0, + "82": 1148595200.0, + "83": 1148595200.0, + "84": 1148595200.0, + "85": 1148595200.0, + "86": 1148595200.0, + "87": 1148595200.0, + "88": 1148595200.0, + "89": 1148595200.0, + "90": 1148595200.0, + "91": 1148595200.0, + "92": 1148595200.0, + "93": 1148595200.0, + "94": 1148595200.0, + "95": 1148595200.0, + "96": 1148595200.0, + "97": 1148595200.0, + "98": 1148595200.0, + "99": 1148595200.0, + "100": 1148595200.0 } }, "iteration-time": { @@ -433,105 +433,105 @@ "step_interval": 1, "values": { "1": "nan", - "2": 11.7836, - "3": 0.58975, - "4": 0.56544, - "5": 0.5504, - "6": 0.56842, - "7": 0.5491, - "8": 0.54138, - "9": 0.53371, - "10": 0.5342, - "11": 0.53224, - "12": 0.52891, - "13": 0.52976, - "14": 0.53162, - "15": 0.52297, - "16": 0.52336, - "17": 0.52793, - "18": 0.52225, - "19": 0.52121, - "20": 0.52937, - "21": 0.53168, - "22": 0.52349, - "23": 0.52045, - "24": 0.53318, - "25": 0.52745, - "26": 0.51972, - "27": 0.52474, - "28": 0.53885, - "29": 0.54406, - "30": 0.52979, - "31": 0.52273, - "32": 0.52354, - "33": 0.52179, - "34": 0.52809, - "35": 0.52207, - "36": 0.52789, - "37": 0.51996, - "38": 0.53223, - "39": 0.52549, - "40": 0.53308, - "41": 0.53147, - "42": 0.53153, - "43": 0.5292, - "44": 0.52056, - "45": 0.52578, - "46": 0.51549, - "47": 0.51842, - "48": 0.51917, - "49": 0.52488, - "50": 0.52255, - "51": 0.64477, - "52": 0.51979, - "53": 0.52383, - "54": 0.52192, - "55": 0.51931, - "56": 0.51907, - "57": 0.52009, - "58": 0.51807, - "59": 0.51736, - "60": 0.51892, - "61": 0.51809, - "62": 0.52089, - "63": 0.52315, - "64": 0.51504, - "65": 0.51491, - "66": 0.51739, - "67": 0.51455, - "68": 0.51564, - "69": 1.04071, - "70": 0.5162, - "71": 0.51607, - "72": 0.5156, - "73": 0.51835, - "74": 0.51882, - "75": 0.52265, - "76": 0.51863, - "77": 0.51483, - "78": 0.51774, - "79": 0.52634, - "80": 0.52171, - "81": 0.52135, - "82": 0.52168, - "83": 0.53375, - "84": 0.51785, - "85": 0.52358, - "86": 0.51614, - "87": 0.52652, - "88": 0.51691, - "89": 0.51638, - "90": 0.52191, - "91": 0.51655, - "92": 0.51846, - "93": 0.51379, - "94": 0.51835, - "95": 0.91609, - "96": 0.51869, - "97": 0.51813, - "98": 0.5255, - "99": 0.52418, - "100": 0.53762 + "2": 8.7306, + "3": 0.82541, + "4": 0.79111, + "5": 0.78772, + "6": 0.78491, + "7": 0.77321, + "8": 0.80845, + "9": 0.76281, + "10": 0.76741, + "11": 0.76405, + "12": 0.7464, + "13": 0.74032, + "14": 0.74249, + "15": 0.7361, + "16": 0.73487, + "17": 0.72656, + "18": 0.73602, + "19": 0.72939, + "20": 0.72896, + "21": 0.7316, + "22": 0.73357, + "23": 0.72972, + "24": 0.73707, + "25": 0.73966, + "26": 0.719, + "27": 0.72924, + "28": 0.74616, + "29": 0.75162, + "30": 0.75031, + "31": 0.74663, + "32": 0.73337, + "33": 0.73723, + "34": 0.73465, + "35": 0.73771, + "36": 0.7385, + "37": 0.73536, + "38": 0.74515, + "39": 0.73575, + "40": 0.74509, + "41": 0.73501, + "42": 0.74091, + "43": 0.74268, + "44": 0.73316, + "45": 0.7359, + "46": 0.72733, + "47": 0.73408, + "48": 0.73042, + "49": 0.73455, + "50": 0.72958, + "51": 0.8591, + "52": 0.81718, + "53": 0.74131, + "54": 0.74839, + "55": 0.74974, + "56": 0.75244, + "57": 0.74244, + "58": 0.73823, + "59": 0.74268, + "60": 0.74576, + "61": 0.74499, + "62": 0.74408, + "63": 0.74442, + "64": 0.74569, + "65": 0.73634, + "66": 0.74134, + "67": 1.30864, + "68": 0.74506, + "69": 0.7469, + "70": 0.73887, + "71": 0.74595, + "72": 0.73832, + "73": 0.73662, + "74": 0.74627, + "75": 0.75627, + "76": 0.74451, + "77": 0.73734, + "78": 0.73831, + "79": 0.74279, + "80": 0.74483, + "81": 0.74523, + "82": 0.7475, + "83": 0.75273, + "84": 0.74267, + "85": 0.73974, + "86": 0.73832, + "87": 0.74642, + "88": 0.73886, + "89": 0.73962, + "90": 0.82905, + "91": 0.73775, + "92": 0.7538, + "93": 0.75623, + "94": 0.74641, + "95": 0.74354, + "96": 0.73224, + "97": 0.73277, + "98": 0.73692, + "99": 0.73794, + "100": 0.73356 } } } \ No newline at end of file diff --git a/tests/unit_tests/dist_checkpointing/utils.py b/tests/unit_tests/dist_checkpointing/utils.py index ec95602b020..0aadaee3b29 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.muon import get_megatron_muon_optimizer +from megatron.core.optimizer.optimizer import ChainedOptimizer 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,11 +178,6 @@ 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) @@ -197,37 +192,42 @@ 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 '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 + if optimizer_type in ('muon', 'dist_muon'): config.lr = 0.0 - 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) + optimizer = get_megatron_optimizer(config, model) torch.manual_seed(seed + 1) model_parallel_cuda_manual_seed(seed + 1) - if not 'muon' in optimizer_type: + 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: 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,10 +272,6 @@ 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) @@ -295,37 +291,43 @@ 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 'muon' in optimizer: - optimizer_type = optimizer - # default lr None feels wrong. only change muon lr to avoid breaking old tests + if optimizer_type in ('muon', 'dist_muon'): config.lr = 0.0 - 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) + optimizer = get_megatron_optimizer(config, model) torch.manual_seed(seed + 1) model_parallel_cuda_manual_seed(seed + 1) - if not 'muon' in optimizer_type: + 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: 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 new file mode 100644 index 00000000000..53d780fd832 --- /dev/null +++ b/tests/unit_tests/test_emerging_optimizers.py @@ -0,0 +1,1574 @@ +# 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, +) +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 +pytestmark = pytest.mark.skipif( + Version(os.getenv('NVIDIA_PYTORCH_VERSION', "24.01")) <= Version("25.05"), + reason="Skip emerging optimizer tests for LTS test", +) + + +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" + + +@pytest.mark.parametrize( + "coefficient_type_and_steps", [("simple", 3), ("quintic", 5), ("polar_express", 8)] +) +def test_muon_optimizer_coefficient_types(coefficient_type_and_steps): + """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_and_steps[0], + num_ns_steps=coefficient_type_and_steps[1], + 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_and_steps[0]} and num_ns_steps={coefficient_type_and_steps[1]}" + + +@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" + + +@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 c484ca104ee..d8b0e97b524 100644 --- a/tests/unit_tests/test_layer_wise_optimizer.py +++ b/tests/unit_tests/test_layer_wise_optimizer.py @@ -417,7 +417,7 @@ def test_bf16_error(self): optimizer='muon', lr=0.01, bf16=True, use_distributed_optimizer=False ) with pytest.raises( - TypeError, match='LayerWiseDistributedOptimizer received Float16 optimizer already' + TypeError, match='LayerWiseDistributedOptimizer expects base torch optimizers' ): LayerWiseDistributedOptimizer([wrapped_optimizer], lw_config, pg_collection) diff --git a/tests/unit_tests/test_lion_optimizer.py b/tests/unit_tests/test_lion_optimizer.py index 589ed82764c..5cd479e655a 100644 --- a/tests/unit_tests/test_lion_optimizer.py +++ b/tests/unit_tests/test_lion_optimizer.py @@ -14,7 +14,7 @@ import torch.nn as nn from megatron.core.optimizer import ( - HAVE_EO_V02, + HAVE_LION, OptimizerConfig, _get_megatron_optimizer_based_on_param_groups, _get_param_groups, @@ -22,7 +22,7 @@ from megatron.core.optimizer.optimizer import FP32Optimizer requires_emerging_optimizers = pytest.mark.skipif( - not HAVE_EO_V02, reason="emerging_optimizers package not installed" + not HAVE_LION, reason="emerging_optimizers package not installed" ) @@ -97,9 +97,9 @@ def test_lion_import_error_without_package(self): """Should raise ImportError with helpful message if emerging_optimizers not installed.""" import megatron.core.optimizer as opt_module - original_have_lion = opt_module.HAVE_EO_V02 + original_have_lion = opt_module.HAVE_LION try: - opt_module.HAVE_EO_V02 = False + opt_module.HAVE_LION = False model = SimpleModel() config = OptimizerConfig(optimizer="lion", lr=1e-4) @@ -107,7 +107,7 @@ def test_lion_import_error_without_package(self): with pytest.raises(ImportError, match="emerging_optimizers"): _create_lion_optimizer(model, config) finally: - opt_module.HAVE_EO_V02 = original_have_lion + opt_module.HAVE_LION = original_have_lion @requires_emerging_optimizers diff --git a/tests/unit_tests/test_muon_optimizer.py b/tests/unit_tests/test_muon_optimizer.py deleted file mode 100644 index 0f0a90c91ed..00000000000 --- a/tests/unit_tests/test_muon_optimizer.py +++ /dev/null @@ -1,792 +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 HAVE_EMERGING_OPTIMIZERS, HAVE_EO_V02, 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" - ), -] - -requires_eo_v02 = pytest.mark.skipif( - not HAVE_EO_V02, reason="emerging_optimizers >= 0.2 is required" -) - - -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_EO_V02 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" - - -@requires_eo_v02 -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) - - -@requires_eo_v02 -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}" - - -@requires_eo_v02 -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 2488900ba72..56af8545042 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) - check_config_overrides_consistency(opt_config, None) - param_groups = _get_param_groups([net], opt_config, None) + 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) assert len(param_groups) == 2 pg0, pg1 = param_groups wd_mults = {pg0['wd_mult'], pg1['wd_mult']} diff --git a/tests/unit_tests/test_optimizer_state_offloading.py b/tests/unit_tests/test_optimizer_state_offloading.py new file mode 100644 index 00000000000..baaab355182 --- /dev/null +++ b/tests/unit_tests/test_optimizer_state_offloading.py @@ -0,0 +1,337 @@ +# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. + +"""Unit tests for OptimizerStateOffloader.""" + +import pytest +import torch +import torch.nn as nn + +from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig +from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer +from megatron.core.transformer import TransformerConfig +from tests.unit_tests.test_utilities import Utils + +try: + from transformer_engine.pytorch.optimizers import FusedAdam # noqa: F401 + + TE_FUSED_ADAM_AVAILABLE = True +except ImportError: + TE_FUSED_ADAM_AVAILABLE = False + + +class SimpleModel(nn.Module): + """Simple model for testing.""" + + def __init__(self, hidden_size=256): + super().__init__() + self.fc1 = nn.Linear(hidden_size, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + + def forward(self, x): + return self.fc2(torch.relu(self.fc1(x))) + + +def create_model_and_optimizer(hidden_size=256, offload_optimizer_states=True, **optimizer_kwargs): + """Helper to create model and optimizer for tests.""" + model = SimpleModel(hidden_size=hidden_size).bfloat16().cuda() + ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=True) + model = DistributedDataParallel( + TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model + ) + + default_config = dict( + optimizer='adam', + bf16=True, + lr=0.001, + use_distributed_optimizer=True, + offload_optimizer_states=offload_optimizer_states, + ) + default_config.update(optimizer_kwargs) + + optimizer_config = OptimizerConfig(**default_config) + optim = get_megatron_optimizer(optimizer_config, [model]) + return model, optim + + +def run_forward_backward_step(model, optim, hidden_size=256): + """Run a single forward-backward-step cycle.""" + input_tensor = torch.randn(8, hidden_size, dtype=torch.bfloat16, device='cuda') + output = model(input_tensor) + output.sum().backward() + optim.step() + optim.zero_grad() + + +# ============================================================================= +# Test 1: Basic OptimizerStateOffloader Initialization +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +def test_offloader_initialization(): + """Test that OptimizerStateOffloader initializes correctly.""" + Utils.initialize_model_parallel() + model, optim = create_model_and_optimizer() + dist_optim = optim.chained_optimizers[0] + + # Offloader is created in __init__ when offload_optimizer_states=True + assert dist_optim._state_offloader is not None + offloader = dist_optim._state_offloader + + # Verify offloader properties + assert offloader.adam_optimizer is not None + assert offloader._d2h_stream is not None + assert offloader._h2d_stream is not None + assert offloader._offloaded is False + + # Before first step, optimizer states are not initialized yet + assert offloader._optimizer_states_initialized is False + + # Run one step to initialize optimizer states + run_forward_backward_step(model, optim) + + # After first step, optimizer states should be marked as initialized + assert offloader._optimizer_states_initialized is True + Utils.destroy_model_parallel() + + +# ============================================================================= +# Test 2: Early Master Weight Offloading Before First Step +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +def test_early_master_weight_offloading(): + """Test that master weights can be offloaded before the first optimizer step.""" + Utils.initialize_model_parallel() + model, optim = create_model_and_optimizer() + dist_optim = optim.chained_optimizers[0] + + # Offloader is created in __init__ + assert dist_optim._state_offloader is not None + offloader = dist_optim._state_offloader + + # Before first step, optimizer states are not initialized + assert offloader._optimizer_states_initialized is False + + # Capture original master weights before offload + original_master_weights = [] + for group in dist_optim.shard_fp32_from_float16_groups: + group_weights = [tensor.clone() for tensor in group] + original_master_weights.append(group_weights) + + # Offload before first step - should only offload master weights + offloader.offload() + offloader.release_gpu_memory() + torch.cuda.synchronize() + + # Verify master weights were offloaded (storage resized to 0) + for group in dist_optim.shard_fp32_from_float16_groups: + for tensor in group: + assert tensor.untyped_storage().size() == 0, "Master weight should be offloaded" + + # Reload master weights + offloader.reload() + offloader.sync_before_step() + + # Verify master weights match after reload + for group_idx, group in enumerate(dist_optim.shard_fp32_from_float16_groups): + for param_idx, tensor in enumerate(group): + original = original_master_weights[group_idx][param_idx] + torch.testing.assert_close( + tensor, + original, + msg=f"Master weight [{group_idx}][{param_idx}] mismatch after offload/reload", + ) + + # Now run a step and verify optimizer states can be offloaded after + run_forward_backward_step(model, optim) + assert offloader._optimizer_states_initialized is True + + Utils.destroy_model_parallel() + + +# ============================================================================= +# Test 3: Offload and Reload Correctness +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +@pytest.mark.parametrize("offload_optimizer_states", [True, False]) +@pytest.mark.parametrize("offload_master_weights", [True, False]) +def test_offload_reload_correctness(offload_optimizer_states, offload_master_weights): + """Test that offload/reload preserves optimizer state values.""" + if not offload_optimizer_states and not offload_master_weights: + pytest.skip("At least one offload type required") + + Utils.initialize_model_parallel() + model, optim = create_model_and_optimizer() + dist_optim = optim.chained_optimizers[0] + + # Run steps to build up optimizer state + for _ in range(3): + run_forward_backward_step(model, optim) + + offloader = dist_optim._state_offloader + + # Capture original states before offload + original_states = {} + for param, state in offloader.adam_optimizer.state.items(): + original_states[param] = { + k: v.clone() for k, v in state.items() if isinstance(v, torch.Tensor) + } + + # Offload + offloader.offload( + offload_optimizer_states=offload_optimizer_states, + offload_master_weights=offload_master_weights, + ) + + # Release GPU memory + offloader.release_gpu_memory() + torch.cuda.synchronize() + + # Reload + offloader.reload() + offloader.sync_before_step() + + # Verify states match after reload + for param, state in offloader.adam_optimizer.state.items(): + if param in original_states: + for key, original_tensor in original_states[param].items(): + if key in state and isinstance(state[key], torch.Tensor): + reloaded_tensor = state[key] + assert reloaded_tensor.device.type == 'cuda', f"State {key} should be on GPU" + torch.testing.assert_close( + reloaded_tensor, + original_tensor, + msg=f"State {key} mismatch after offload/reload", + ) + Utils.destroy_model_parallel() + + +# ============================================================================= +# Test 4: GPU Memory Release Verification +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +def test_gpu_memory_release(): + """Test that GPU memory is actually freed after release_gpu_memory().""" + Utils.initialize_model_parallel() + # Use larger model for measurable memory impact + model, optim = create_model_and_optimizer(hidden_size=1024) + dist_optim = optim.chained_optimizers[0] + + # Initialize optimizer states + run_forward_backward_step(model, optim, hidden_size=1024) + + offloader = dist_optim._state_offloader + + # Measure memory before offload + torch.cuda.synchronize() + torch.cuda.empty_cache() + memory_before = torch.cuda.memory_allocated() + + # Offload and release + offloader.offload() + offloader.release_gpu_memory() + + # Wait for async operations + torch.cuda.synchronize() + torch.cuda.empty_cache() + memory_after = torch.cuda.memory_allocated() + + # Memory should decrease + memory_freed = memory_before - memory_after + assert memory_freed > 0, f"Expected memory to be freed, but got {memory_freed} bytes difference" + Utils.destroy_model_parallel() + + +# ============================================================================= +# Test 5: Multiple Offload/Reload Cycles +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +def test_multiple_offload_reload_cycles(): + """Test that multiple offload/reload cycles work correctly.""" + Utils.initialize_model_parallel() + model, optim = create_model_and_optimizer() + dist_optim = optim.chained_optimizers[0] + + # Initialize + run_forward_backward_step(model, optim) + + offloader = dist_optim._state_offloader + + # Run multiple cycles + for cycle in range(5): + # Offload + offloader.offload() + offloader.release_gpu_memory() + + # Reload + offloader.reload() + offloader.sync_before_step() + + # Run optimizer step + run_forward_backward_step(model, optim) + + # Verify model can still produce valid outputs + input_tensor = torch.randn(8, 256, dtype=torch.bfloat16, device='cuda') + output = model(input_tensor) + assert not output.isnan().any(), "Model output contains NaN after multiple cycles" + Utils.destroy_model_parallel() + + +# ============================================================================= +# Test 6: Training Correctness with Offloading +# ============================================================================= +@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") +def test_training_correctness_with_offloading(): + """Test that training with offloading produces same results as without.""" + Utils.initialize_model_parallel() + torch.manual_seed(42) + + # Model 1: with offloading + model1, optim1 = create_model_and_optimizer(offload_optimizer_states=True, lr=0.01) + + # Model 2: without offloading (reference) + torch.manual_seed(42) + model2, optim2 = create_model_and_optimizer(offload_optimizer_states=False, lr=0.01) + + # Train both models + n_steps = 10 + torch.manual_seed(123) + dist_optim1 = optim1.chained_optimizers[0] + + # Offloader is created in __init__ when offload_optimizer_states=True + assert dist_optim1._state_offloader is not None + offloader = dist_optim1._state_offloader + + for step in range(n_steps): + input_tensor = torch.randn(8, 256, dtype=torch.bfloat16, device='cuda') + + # Model 1 with offloading + # Offload states (master weights can be offloaded from the start, + # optimizer states will be skipped until after first step) + offloader.offload() + offloader.release_gpu_memory() + + output1 = model1(input_tensor) + loss1 = output1.sum() + loss1.backward() + + offloader.reload() + offloader.sync_before_step() + optim1.step() + optim1.zero_grad() + + # Model 2 without offloading + output2 = model2(input_tensor) + loss2 = output2.sum() + loss2.backward() + optim2.step() + optim2.zero_grad() + + # Compare final model weights + for (n1, p1), (n2, p2) in zip(model1.named_parameters(), model2.named_parameters()): + torch.testing.assert_close( + p1.data, + p2.data, + atol=1e-5, + rtol=1e-4, + msg=f"Parameter {n1} mismatch between offloaded and non-offloaded training", + ) + Utils.destroy_model_parallel() From b5cd5653a6f7d61b81f2fa6c2bb249f25b73b87b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:34:10 -0700 Subject: [PATCH 02/10] Add missing optimizer-related changes in arguments.py and param_and_grad_buffer.py - arguments.py: Add CLI args for new optimizers (soap, adaptive_muon), --offload-optimizer-states flag, --muon-nesterov rename, dist_muon deprecation handling, use_layer_wise_distributed_optimizer logic, and emerging optimizer validation checks - param_and_grad_buffer.py: Add grad mode tracking for cached param views to avoid stale views when grad mode changes Co-Authored-By: Claude Opus 4.6 (1M context) --- .../core/distributed/param_and_grad_buffer.py | 4 ++ megatron/training/arguments.py | 40 ++++++++++++++++--- 2 files changed, 38 insertions(+), 6 deletions(-) diff --git a/megatron/core/distributed/param_and_grad_buffer.py b/megatron/core/distributed/param_and_grad_buffer.py index 6c36f119e19..074ca7a6bb6 100644 --- a/megatron/core/distributed/param_and_grad_buffer.py +++ b/megatron/core/distributed/param_and_grad_buffer.py @@ -243,6 +243,9 @@ def __init__( # or bucket.grad_data. self.cached_param_buffer_shard_list = [None] * len(self.buckets) self.cached_grad_buffer_shard_list = [None] * len(self.buckets) + # Track grad mode used to create cached param views. Rebuild if mode changes to avoid + # mixing no_grad-created views with in-place updates in grad-enabled mode. + self._cached_param_buffer_shards_grad_enabled = None def reset(self): """ @@ -399,6 +402,7 @@ def start_param_sync(self, force_sync: bool = False): bucket.layerwise_gather_list = None bucket._layerwise_src_buffer = None self.param_gather_handle = None + else: # Standard distributed optimizer path: use _coalescing_manager. # all_gather_into_tensor writes directly into a contiguous output buffer and diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index d3026d94b68..3ec1eb5ce87 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -1472,17 +1472,30 @@ 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).' - # Muon optimizer check - if 'muon' in args.optimizer: + # emerging optimizer check + if not hasattr(args, 'use_layer_wise_distributed_optimizer'): + args.use_layer_wise_distributed_optimizer = False + 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 + + 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." + assert args.experimental_attention_variant is None, "Muon optimizer does not support attention variant for now." # Optimizer CPU offload check if args.optimizer_cpu_offload: @@ -1495,6 +1508,11 @@ def validate_args(args, defaults={}): "must be used in conjunction with `--fp8-recipe delayed`." ) + if args.offload_optimizer_states: + assert args.use_distributed_optimizer, "offload_optimizer_states is only supported with distributed optimizer" + assert args.optimizer == 'adam', "offload_optimizer_states is only supported with adam optimizer" + assert not args.use_megatron_fsdp, "offload_optimizer_states does not support Megatron-FSDP for now." + if args.non_persistent_ckpt_type == "local": assert args.non_persistent_local_ckpt_dir is not None, "Tried to use local checkpointing without specifying --local-ckpt-dir!" if args.replication: @@ -2228,7 +2246,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-use-nesterov', action='store_true', + group.add_argument('--muon-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'], @@ -2476,8 +2494,10 @@ 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'], - help='Optimizer function') + choices=['adam', 'sgd', 'muon', 'dist_muon', 'soap', "adaptive_muon", "lion"], + help='Optimizer function. ' + 'Note: dist_muon is deprecated; use --optimizer muon ' + 'with --use-distributed-optimizer instead.') group.add_argument('--optimizer-cpu-offload', action='store_true', help='Offload optimizer state to CPU') group.add_argument('--optimizer-offload-fraction', type=float, default=1.0, @@ -2494,6 +2514,14 @@ def _add_training_args(parser): help='Disable pinning of CPU memory for gradients.') group.add_argument('--no-pin-cpu-params', action='store_false', dest='pin_cpu_params', help='Disable pinning of CPU memory for parameters.') + group.add_argument('--offload-optimizer-states', + action='store_true', + dest='offload_optimizer_states', + help='Offload optimizer states to CPU after each optimizer step and ' + 'reload them before the next optimizer step. ' + 'Only support TE FusedAdam optimizer.' + 'Note that this still uses pure GPU optimizer instead of ' + 'HybridDeviceOptimizer for --optimizer-cpu-offload.') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], help='Single pass vs multiple pass data loader') From b1be5d1814791baf8d227cf6d9188d71ec1d5b6c Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:35:31 -0700 Subject: [PATCH 03/10] Add muon_coefficient_type config field back to OptimizerConfig Co-Authored-By: Claude Opus 4.6 (1M context) --- megatron/core/optimizer/optimizer_config.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index d425a56be71..c87d7c7274d 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -271,6 +271,10 @@ class OptimizerConfig: muon_fp32_matmul_prec: str = "medium" """The precision to use for the fp32 matmul. Defaults to "medium".""" + muon_coefficient_type: str = "quintic" + """Newton-Schulz coefficient type for the Muon optimizer. Valid types are discovered + dynamically from the installed ``emerging_optimizers`` package. Defaults to "quintic".""" + muon_num_ns_steps: int = 5 """The number of iteration steps to use in the Newton-Schulz iteration.""" From d33584332b054b118fc2f4bfe44dc2c65d74b63c Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:43:41 -0700 Subject: [PATCH 04/10] Remove optimizer state offloading changes, keep only emerging-optimizers Reverts: OptimizerStateOffloader, distrib_optimizer quantized tensor changes, offload_optimizer_states config/CLI arg, and offloading test. These will be merged separately. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../optimizer_state_offloader.py | 315 ---------------- megatron/core/optimizer/distrib_optimizer.py | 138 +------ megatron/core/optimizer/optimizer_config.py | 6 - megatron/training/arguments.py | 13 - .../test_optimizer_state_offloading.py | 337 ------------------ 5 files changed, 10 insertions(+), 799 deletions(-) delete mode 100644 megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py delete mode 100644 tests/unit_tests/test_optimizer_state_offloading.py diff --git a/megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py b/megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py deleted file mode 100644 index 81fd116c8ba..00000000000 --- a/megatron/core/optimizer/cpu_offloading/optimizer_state_offloader.py +++ /dev/null @@ -1,315 +0,0 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. - -"""Optimizer state offloading class.""" - -from typing import TYPE_CHECKING, Dict, List, Tuple - -import torch - -if TYPE_CHECKING: - from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer - - -class OptimizerStateOffloader: - """ - Manages offloading of optimizer states and master weights to CPU. - Used with DistributedOptimizer to reduce GPU memory usage. - - Supports overlapped D2H/H2D transfers using CUDA streams. - - Master weights can be stored in two locations: - - In adam optimizer state (when use_precision_aware_optimizer_no_fp8_or_ds_fp8 is True) - - In mcore's shard_fp32_from_float16_groups - """ - - OPTIMIZER_STATE_KEYS = ('exp_avg', 'exp_avg_sq') - MASTER_WEIGHT_KEY = 'master_param' - - def __init__(self, distrib_optimizer: "DistributedOptimizer"): - """ - Args: - distrib_optimizer: The DistributedOptimizer to offload states and master weights from. - """ - self.dist_optimizer = distrib_optimizer - self.adam_optimizer = distrib_optimizer.optimizer - - # Only support TE FusedAdam optimizer for now. - try: - from transformer_engine.pytorch.optimizers import FusedAdam - - assert isinstance(self.adam_optimizer, FusedAdam), ( - f"OptimizerStateOffloader requires TE FusedAdam optimizer, " - f"but got {type(self.adam_optimizer).__name__}" - ) - except ImportError: - raise ImportError( - "OptimizerStateOffloader requires transformer_engine.pytorch.optimizers.FusedAdam" - ) - - # Check if master weights are stored in adam optimizer state - self.optimizer_contains_master_weights = self.adam_optimizer.master_weights - - # CUDA streams for async transfers - self._d2h_stream = torch.cuda.Stream() - self._h2d_stream = torch.cuda.Stream() - - # CPU buffers for optimizer states: {param: {key: cpu_tensor}} - self._opt_state_cpu_buffers: Dict[torch.Tensor, Dict[str, torch.Tensor]] = {} - - # CPU buffers for mcore master weights, matching the structure of source groups - # List[List[cpu_tensor]] - self._shard_fp32_from_float16_cpu_buffers: List[List[torch.Tensor]] = [] - - # State tracking - self._offloaded = False - self._offloaded_state_keys: Tuple[str, ...] = () - self._offloaded_mcore_master_weights = False - - # Track whether optimizer states (exp_avg, exp_avg_sq) have been initialized. - # These are lazily initialized by FusedAdam during the first optimizer.step(). - # Master weights (shard_fp32_from_float16_groups) are available from the start. - self._optimizer_states_initialized = False - - def mark_optimizer_states_initialized(self): - """ - Mark that optimizer states (exp_avg, exp_avg_sq) are now available. - Should be called after the first optimizer.step() completes. - """ - self._optimizer_states_initialized = True - - def _get_state_keys_to_offload( - self, offload_optimizer_states: bool, offload_master_weights: bool - ) -> Tuple[str, ...]: - """Get the state keys in FusedAdam to offload based on configuration.""" - keys = [] - # Skip optimizer states offloading if they haven't been initialized yet. - # Optimizer states are lazily initialized by FusedAdam during the first optimizer.step(). - if self._optimizer_states_initialized: - if offload_optimizer_states: - keys.extend(self.OPTIMIZER_STATE_KEYS) - if offload_master_weights and self.optimizer_contains_master_weights: - keys.append(self.MASTER_WEIGHT_KEY) - return tuple(keys) - - def _ensure_state_cpu_buffer( - self, param: torch.Tensor, state_key: str, gpu_tensor: torch.Tensor, pin_memory: bool = True - ) -> torch.Tensor: - """Get or create a CPU buffer for a state tensor.""" - if param not in self._opt_state_cpu_buffers: - self._opt_state_cpu_buffers[param] = {} - - if state_key not in self._opt_state_cpu_buffers[param]: - cpu_buffer = torch.empty( - gpu_tensor.size(), - dtype=gpu_tensor.dtype, - layout=gpu_tensor.layout, - device='cpu', - pin_memory=pin_memory, - ) - self._opt_state_cpu_buffers[param][state_key] = cpu_buffer - - return self._opt_state_cpu_buffers[param][state_key] - - def _offload_shard_groups( - self, - shard_groups: List[List[torch.Tensor]], - cpu_buffers: List[List[torch.Tensor]], - pin_memory: bool = True, - ): - """Offload a shard group to CPU buffers.""" - # Initialize CPU buffers on first call - if len(cpu_buffers) == 0: - for group in shard_groups: - group_buffers = [] - for gpu_tensor in group: - cpu_buffer = torch.empty( - gpu_tensor.size(), - dtype=gpu_tensor.dtype, - layout=gpu_tensor.layout, - device='cpu', - pin_memory=pin_memory, - ) - group_buffers.append(cpu_buffer) - cpu_buffers.append(group_buffers) - - # Copy D2H - for group_idx, group in enumerate(shard_groups): - for param_idx, gpu_tensor in enumerate(group): - cpu_buffer = cpu_buffers[group_idx][param_idx] - cpu_buffer.copy_(gpu_tensor, non_blocking=pin_memory) - gpu_tensor.record_stream(self._d2h_stream) - - def _offload_states( - self, - offload_optimizer_states: bool, - offload_master_weights: bool, - use_pin_memory: bool = True, - ): - """Offload optimizer states and/or master weights to CPU.""" - # Offload states from adam optimizer - self._offloaded_state_keys = self._get_state_keys_to_offload( - offload_optimizer_states, offload_master_weights - ) - states = self.adam_optimizer.state - - for param, param_state in states.items(): - for state_key in self._offloaded_state_keys: - if state_key not in param_state: - continue - - gpu_tensor = param_state[state_key] - if not isinstance(gpu_tensor, torch.Tensor) or not gpu_tensor.is_cuda: - continue - - cpu_buffer = self._ensure_state_cpu_buffer( - param, state_key, gpu_tensor, use_pin_memory - ) - cpu_buffer.copy_(gpu_tensor, non_blocking=use_pin_memory) - gpu_tensor.record_stream(self._d2h_stream) - - # Offload mcore master weights if not in optimizer state - if offload_master_weights and not self.optimizer_contains_master_weights: - self._offload_shard_groups( - self.dist_optimizer.shard_fp32_from_float16_groups, - self._shard_fp32_from_float16_cpu_buffers, - use_pin_memory, - ) - self._offloaded_mcore_master_weights = True - - def _release_states(self): - """Replace optimizer state GPU tensors with CPU tensors to free GPU memory.""" - states = self.adam_optimizer.state - - for param, param_state in states.items(): - if param not in self._opt_state_cpu_buffers: - continue - - for state_key in self._offloaded_state_keys: - if state_key not in self._opt_state_cpu_buffers[param]: - continue - - param_state[state_key].untyped_storage().resize_(0) - - if self._offloaded_mcore_master_weights: - for group in self.dist_optimizer.shard_fp32_from_float16_groups: - for gpu_tensor in group: - gpu_tensor.untyped_storage().resize_(0) - - def _reload_shard_groups( - self, - shard_groups: List[List[torch.Tensor]], - cpu_buffers: List[List[torch.Tensor]], - is_allocate_stage: bool, - ): - """Reload shard groups from CPU to GPU.""" - for group_idx, group in enumerate(shard_groups): - for param_idx, _ in enumerate(group): - cpu_buffer = cpu_buffers[group_idx][param_idx] - if is_allocate_stage: - shard_groups[group_idx][param_idx].untyped_storage().resize_( - cpu_buffer.untyped_storage().size() - ) - else: - shard_groups[group_idx][param_idx].copy_( - cpu_buffer, non_blocking=cpu_buffer.is_pinned() - ) - - def _reload_states(self, is_allocate_stage: bool): - """ - Reload optimizer states and/or master weights from CPU to GPU. - - If is_allocate_stage is True, only allocate GPU memory for the states and master weights, - but do not copy the data from CPU to GPU. Otherwise, copy the data from CPU to GPU. - The two processes are separated to make sure that the GPU memory is allocated on the - default stream to avoid fragmentation. - """ - # Reload states to adam optimizer - states = self.adam_optimizer.state - - for param, param_state in states.items(): - if param not in self._opt_state_cpu_buffers: - continue - - for state_key in self._offloaded_state_keys: - if state_key not in self._opt_state_cpu_buffers[param]: - continue - - cpu_buffer = self._opt_state_cpu_buffers[param][state_key] - if is_allocate_stage: - param_state[state_key].untyped_storage().resize_( - cpu_buffer.untyped_storage().size() - ) - else: - param_state[state_key].copy_(cpu_buffer, non_blocking=cpu_buffer.is_pinned()) - - # Reload mcore master weights if not in optimizer state - if self._offloaded_mcore_master_weights: - self._reload_shard_groups( - self.dist_optimizer.shard_fp32_from_float16_groups, - self._shard_fp32_from_float16_cpu_buffers, - is_allocate_stage, - ) - - def offload(self, offload_optimizer_states: bool = True, offload_master_weights: bool = True): - """ - Offload optimizer states and/or master weights to CPU. - Starts async D2H transfer that can overlap with other operations. - - Args: - offload_optimizer_states: Whether to offload exp_avg, exp_avg_sq. - offload_master_weights: Whether to offload master weights. - """ - if not offload_optimizer_states and not offload_master_weights: - return - - # Wait for current stream finishing updating the optimizer states. - self._d2h_stream.wait_stream(torch.cuda.current_stream()) - - with torch.cuda.stream(self._d2h_stream): - self._offload_states(offload_optimizer_states, offload_master_weights) - - self._offloaded = True - - def release_gpu_memory(self): - """ - Release GPU memory for optimizer states and master weights after D2H copy completes. - - This is separated from offload() to allow delayed GPU memory release, - which is needed for mxfp8 + overlap_param_gather case where master weights - must remain on GPU until after _copy_main_params_to_param_buffer() is called. - """ - if not self._offloaded: - return - - self._release_states() - - def reload(self): - """ - Reload optimizer states and/or master weights from CPU to GPU. - Call before optimizer.step() to ensure states are on GPU. - """ - if not self._offloaded: - return - - # Allocate GPU memory on the current stream to avoid fragmentation. - self._reload_states(is_allocate_stage=True) - - self._h2d_stream.wait_stream(self._d2h_stream) - self._h2d_stream.wait_stream(torch.cuda.current_stream()) - - # Reload states on the h2d stream to overlap with other operations. - with torch.cuda.stream(self._h2d_stream): - self._reload_states(is_allocate_stage=False) - - self._offloaded_state_keys = () - self._offloaded_mcore_master_weights = False - self._offloaded = False - - def sync_before_step(self): - """ - Wait for H2D reload to complete before optimizer.step(). - Must be called to ensure states are on GPU before optimizer uses them. - - This is separated from reload() to make it possible to move the reload ahead of time. - """ - torch.cuda.current_stream().wait_stream(self._h2d_stream) diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index beb00391759..eeda383a75d 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -52,7 +52,6 @@ from ..fp8_utils import dequantize_fp8_tensor, is_float8tensor, quantize_param_shard from ..transformer.fsdp_dtensor_checkpoint import handle_experts_in_state_dict from ..transformer.module import MegatronModule -from .cpu_offloading.optimizer_state_offloader import OptimizerStateOffloader from .grad_scaler import MegatronGradScaler from .optimizer import MixedPrecisionOptimizer, _zero_grad_group_helper, param_group_identifier_keys from .optimizer_config import OptimizerConfig @@ -362,10 +361,7 @@ def _build_model_and_main_param_groups( if model_param.type() in ['torch.cuda.HalfTensor', 'torch.cuda.BFloat16Tensor']: # Generate sharded model param. - if ( - cls._is_distopt_quantized_param(model_param) - and config.fp8_recipe != "delayed" - ): + if is_float8tensor(model_param) and config.fp8_recipe != "delayed": # MXFP8Tensor and BlockwiseQTensor don't support view(-1) shard_model_param = None else: @@ -385,7 +381,7 @@ def _build_model_and_main_param_groups( # precision at the beginning of training (this problem will not occur if the # training is long enough or if the main params are loaded from a # checkpoint). - if cls._is_distopt_quantized_param(model_param): + if is_float8tensor(model_param): if hasattr(model_param, 'get_high_precision_init_val'): shard_main_param = ( model_param.get_high_precision_init_val() @@ -523,8 +519,6 @@ def __init__( "due to checkpointing requirements." ) - self._state_offloader: Optional[OptimizerStateOffloader] = None - # when freezing sub-models we have no real optimizer # but still need a stub DistributedOptimizer class if optimizer is None: @@ -613,9 +607,6 @@ def __init__( self.optimizer.param_groups = [g["orig_group"] for g in self.opt_group_ranges] self.optimizer.load_state_dict(self.optimizer.state_dict()) - if self.config.offload_optimizer_states: - self._state_offloader = OptimizerStateOffloader(self) - def _get_model_param_range_map(self, param: torch.nn.Parameter): """ Given a model param, get the index sub-range of the param that this @@ -922,70 +913,6 @@ def _get_main_param_and_optimizer_states(self, model_param): tensors[k] = v return tensors - @staticmethod - def _is_grouped_quantized_tensor(tensor: torch.Tensor) -> bool: - """Check if tensor is a TE GroupedTensor using quantized storage.""" - return ( - hasattr(tensor, "split_into_quantized_tensors") - and callable(tensor.split_into_quantized_tensors) - and getattr(tensor, "quantizer", None) is not None - ) - - @classmethod - def _is_distopt_quantized_param(cls, tensor: torch.Tensor) -> bool: - """Check if tensor should follow quantized parameter path in dist optimizer.""" - return is_float8tensor(tensor) or cls._is_grouped_quantized_tensor(tensor) - - def _expand_quantized_param_shard_for_cast( - self, - model_param: torch.Tensor, - shard_main_param: Optional[torch.Tensor], - start_offset: Optional[int], - ): - """Expand one quantized model param to cast-ready entries. - - For grouped quantized tensors, split into member quantized tensors and map the sharded - master slice to per-member offset ranges, while preserving deterministic ordering across - DP ranks. - """ - if not self._is_grouped_quantized_tensor(model_param): - return [model_param], [shard_main_param], [start_offset] - - quantized_members = model_param.quantized_tensors - if quantized_members is None: - quantized_members = model_param.split_into_quantized_tensors() - - shard_start = 0 if start_offset is None else start_offset - shard_size = 0 if shard_main_param is None else shard_main_param.numel() - shard_end = shard_start + shard_size - shard_flat = None if shard_main_param is None else shard_main_param.view(-1) - - expanded_model_params = [] - expanded_shard_main_params = [] - expanded_start_offsets = [] - member_offset = 0 - for member in quantized_members: - member_numel = member.numel() - member_start = member_offset - member_end = member_start + member_numel - overlap_start = max(member_start, shard_start) - overlap_end = min(member_end, shard_end) - - member_master = None - member_start_offset = None - if overlap_start < overlap_end: - local_start = overlap_start - shard_start - local_end = overlap_end - shard_start - member_master = shard_flat[local_start:local_end] - member_start_offset = overlap_start - member_start - - expanded_model_params.append(member) - expanded_shard_main_params.append(member_master) - expanded_start_offsets.append(member_start_offset) - member_offset = member_end - - return expanded_model_params, expanded_shard_main_params, expanded_start_offsets - def _set_main_param_and_optimizer_states(self, model_param, tensors): """Set the main param and optimizer states corresponding to the input model_param. @@ -2218,7 +2145,7 @@ def split_state_dict_if_needed(self, state_dict): fp8_gbuf_indices = [] for gbuf_idx, gbuf_range_maps in enumerate(self.gbuf_ranges): for dtype, _ in gbuf_range_maps.items(): - if self._is_distopt_quantized_param(self.buffers[gbuf_idx].params[0]): + if is_float8tensor(self.buffers[gbuf_idx].params[0]): fp8_gbuf_indices.append(gbuf_idx) if len(fp8_gbuf_indices) == 0: return @@ -2240,7 +2167,7 @@ def split_state_dict_if_needed(self, state_dict): new_state_dict = {'buckets_coalesced': state_dict['buckets_coalesced']} for gbuf_idx, gbuf_range_maps in enumerate(self.gbuf_ranges): for dtype, _ in gbuf_range_maps.items(): - if not self._is_distopt_quantized_param(self.buffers[gbuf_idx].params[0]): + if not is_float8tensor(self.buffers[gbuf_idx].params[0]): new_state_dict[gbuf_idx] = state_dict[dtype_to_gbuf_idx[dtype]] for fp8_gbuf_idx in fp8_gbuf_indices: @@ -2440,7 +2367,7 @@ def _get_fp8_params_and_shard_fp32_from_fp8(self): idx = 0 for buffer in buffers: for param in buffer.params: - if self._is_distopt_quantized_param(param): + if is_float8tensor(param): fp8_params.append(param) shard_fp32_from_fp8.append(None) shard_offsets_in_fp8.append(None) @@ -2455,7 +2382,7 @@ def get_shard_fp32_from_fp8(shard_main_groups, model_groups): """ for shard_main_group, model_group in zip(shard_main_groups, model_groups): for shard_main_param, model_param in zip(shard_main_group, model_group): - if self._is_distopt_quantized_param(model_param): + if is_float8tensor(model_param): param_range_map = self._get_model_param_range_map(model_param) param_range = param_range_map["param"] assert param_range.size == shard_main_param.nelement() @@ -2532,29 +2459,8 @@ def _copy_main_params_to_model_params(self): if self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8: return - fp8_params, shard_fp32_from_fp8, shard_offsets_in_fp8 = ( - self._get_fp8_params_and_shard_fp32_from_fp8() - ) - expanded_fp8_params = [] - expanded_shard_fp32_from_fp8 = [] - expanded_shard_offsets_in_fp8 = [] - for model_param, shard_main_param, start_offset in zip( - fp8_params, shard_fp32_from_fp8, shard_offsets_in_fp8 - ): - sub_model_params, sub_shard_main_params, sub_start_offsets = ( - self._expand_quantized_param_shard_for_cast( - model_param, shard_main_param, start_offset - ) - ) - expanded_fp8_params.extend(sub_model_params) - expanded_shard_fp32_from_fp8.extend(sub_shard_main_params) - expanded_shard_offsets_in_fp8.extend(sub_start_offsets) - quantize_param_shard( - expanded_fp8_params, - expanded_shard_fp32_from_fp8, - expanded_shard_offsets_in_fp8, - self.data_parallel_group, + *self._get_fp8_params_and_shard_fp32_from_fp8(), self.data_parallel_group ) # Utility method for copying group params. @@ -2574,7 +2480,7 @@ def copy_group_params(shard_main_groups, model_groups): world_range.start : world_range.end ] - if self._is_distopt_quantized_param(model_param): + if is_float8tensor(model_param): # FP8 params are quantized in the above "quantize_param_shard" function. continue else: @@ -2686,12 +2592,8 @@ def copy_group_params(model_groups, shard_main_groups): # Use param from state_dict to initialize main_param model_param = model_param_to_state_dict_param_map[model_param] - if self._is_distopt_quantized_param(model_param): - if self._is_grouped_quantized_tensor(model_param): - dequantized_model_param = model_param.float() - else: - dequantized_model_param = dequantize_fp8_tensor(model_param) - shard_model_param = dequantized_model_param.view(-1)[ + if is_float8tensor(model_param): + shard_model_param = dequantize_fp8_tensor(model_param).view(-1)[ param_range.start : param_range.end ] else: @@ -2710,8 +2612,6 @@ def step_with_ready_grads(self) -> bool: Under the hood, either launch synchronous param all-gathers or get ready to launch asynchorous all-gathers that get overlapped with the next forward pass. """ - if self._state_offloader is not None: - self._state_offloader.sync_before_step() update_successful = super().step_with_ready_grads() timers = self.config.timers @@ -2732,22 +2632,4 @@ def step_with_ready_grads(self) -> bool: if timers is not None: timers('params-all-gather').stop() - if self._state_offloader is not None: - self._state_offloader.mark_optimizer_states_initialized() - return update_successful - - def offload_states(self): - """Offload states to CPU.""" - if self._state_offloader is not None: - self._state_offloader.offload() - - def reload_offloaded_states(self): - """Start async reload of offloaded states.""" - if self._state_offloader is not None: - self._state_offloader.reload() - - def release_offloaded_gpu_states(self): - """Release GPU memory after D2H completes. For delayed release case.""" - if self._state_offloader is not None: - self._state_offloader.release_gpu_memory() diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index c87d7c7274d..df8c8249b2a 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -364,12 +364,6 @@ class OptimizerConfig: pin_cpu_params: bool = True """If True, pin the optimizer parameters to CPU memory.""" - offload_optimizer_states: bool = False - """ - If True, offload optimizer states to CPU after each optimizer step and - reload them before the next optimizer step. - """ - ################ # Miscellaneous ################ diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 3ec1eb5ce87..ad2a05f0241 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -1508,11 +1508,6 @@ def validate_args(args, defaults={}): "must be used in conjunction with `--fp8-recipe delayed`." ) - if args.offload_optimizer_states: - assert args.use_distributed_optimizer, "offload_optimizer_states is only supported with distributed optimizer" - assert args.optimizer == 'adam', "offload_optimizer_states is only supported with adam optimizer" - assert not args.use_megatron_fsdp, "offload_optimizer_states does not support Megatron-FSDP for now." - if args.non_persistent_ckpt_type == "local": assert args.non_persistent_local_ckpt_dir is not None, "Tried to use local checkpointing without specifying --local-ckpt-dir!" if args.replication: @@ -2514,14 +2509,6 @@ def _add_training_args(parser): help='Disable pinning of CPU memory for gradients.') group.add_argument('--no-pin-cpu-params', action='store_false', dest='pin_cpu_params', help='Disable pinning of CPU memory for parameters.') - group.add_argument('--offload-optimizer-states', - action='store_true', - dest='offload_optimizer_states', - help='Offload optimizer states to CPU after each optimizer step and ' - 'reload them before the next optimizer step. ' - 'Only support TE FusedAdam optimizer.' - 'Note that this still uses pure GPU optimizer instead of ' - 'HybridDeviceOptimizer for --optimizer-cpu-offload.') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], help='Single pass vs multiple pass data loader') diff --git a/tests/unit_tests/test_optimizer_state_offloading.py b/tests/unit_tests/test_optimizer_state_offloading.py deleted file mode 100644 index baaab355182..00000000000 --- a/tests/unit_tests/test_optimizer_state_offloading.py +++ /dev/null @@ -1,337 +0,0 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. - -"""Unit tests for OptimizerStateOffloader.""" - -import pytest -import torch -import torch.nn as nn - -from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig -from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer -from megatron.core.transformer import TransformerConfig -from tests.unit_tests.test_utilities import Utils - -try: - from transformer_engine.pytorch.optimizers import FusedAdam # noqa: F401 - - TE_FUSED_ADAM_AVAILABLE = True -except ImportError: - TE_FUSED_ADAM_AVAILABLE = False - - -class SimpleModel(nn.Module): - """Simple model for testing.""" - - def __init__(self, hidden_size=256): - super().__init__() - self.fc1 = nn.Linear(hidden_size, hidden_size) - self.fc2 = nn.Linear(hidden_size, hidden_size) - - def forward(self, x): - return self.fc2(torch.relu(self.fc1(x))) - - -def create_model_and_optimizer(hidden_size=256, offload_optimizer_states=True, **optimizer_kwargs): - """Helper to create model and optimizer for tests.""" - model = SimpleModel(hidden_size=hidden_size).bfloat16().cuda() - ddp_config = DistributedDataParallelConfig(use_distributed_optimizer=True) - model = DistributedDataParallel( - TransformerConfig(num_attention_heads=1, num_layers=1), ddp_config, model - ) - - default_config = dict( - optimizer='adam', - bf16=True, - lr=0.001, - use_distributed_optimizer=True, - offload_optimizer_states=offload_optimizer_states, - ) - default_config.update(optimizer_kwargs) - - optimizer_config = OptimizerConfig(**default_config) - optim = get_megatron_optimizer(optimizer_config, [model]) - return model, optim - - -def run_forward_backward_step(model, optim, hidden_size=256): - """Run a single forward-backward-step cycle.""" - input_tensor = torch.randn(8, hidden_size, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - output.sum().backward() - optim.step() - optim.zero_grad() - - -# ============================================================================= -# Test 1: Basic OptimizerStateOffloader Initialization -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -def test_offloader_initialization(): - """Test that OptimizerStateOffloader initializes correctly.""" - Utils.initialize_model_parallel() - model, optim = create_model_and_optimizer() - dist_optim = optim.chained_optimizers[0] - - # Offloader is created in __init__ when offload_optimizer_states=True - assert dist_optim._state_offloader is not None - offloader = dist_optim._state_offloader - - # Verify offloader properties - assert offloader.adam_optimizer is not None - assert offloader._d2h_stream is not None - assert offloader._h2d_stream is not None - assert offloader._offloaded is False - - # Before first step, optimizer states are not initialized yet - assert offloader._optimizer_states_initialized is False - - # Run one step to initialize optimizer states - run_forward_backward_step(model, optim) - - # After first step, optimizer states should be marked as initialized - assert offloader._optimizer_states_initialized is True - Utils.destroy_model_parallel() - - -# ============================================================================= -# Test 2: Early Master Weight Offloading Before First Step -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -def test_early_master_weight_offloading(): - """Test that master weights can be offloaded before the first optimizer step.""" - Utils.initialize_model_parallel() - model, optim = create_model_and_optimizer() - dist_optim = optim.chained_optimizers[0] - - # Offloader is created in __init__ - assert dist_optim._state_offloader is not None - offloader = dist_optim._state_offloader - - # Before first step, optimizer states are not initialized - assert offloader._optimizer_states_initialized is False - - # Capture original master weights before offload - original_master_weights = [] - for group in dist_optim.shard_fp32_from_float16_groups: - group_weights = [tensor.clone() for tensor in group] - original_master_weights.append(group_weights) - - # Offload before first step - should only offload master weights - offloader.offload() - offloader.release_gpu_memory() - torch.cuda.synchronize() - - # Verify master weights were offloaded (storage resized to 0) - for group in dist_optim.shard_fp32_from_float16_groups: - for tensor in group: - assert tensor.untyped_storage().size() == 0, "Master weight should be offloaded" - - # Reload master weights - offloader.reload() - offloader.sync_before_step() - - # Verify master weights match after reload - for group_idx, group in enumerate(dist_optim.shard_fp32_from_float16_groups): - for param_idx, tensor in enumerate(group): - original = original_master_weights[group_idx][param_idx] - torch.testing.assert_close( - tensor, - original, - msg=f"Master weight [{group_idx}][{param_idx}] mismatch after offload/reload", - ) - - # Now run a step and verify optimizer states can be offloaded after - run_forward_backward_step(model, optim) - assert offloader._optimizer_states_initialized is True - - Utils.destroy_model_parallel() - - -# ============================================================================= -# Test 3: Offload and Reload Correctness -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -@pytest.mark.parametrize("offload_optimizer_states", [True, False]) -@pytest.mark.parametrize("offload_master_weights", [True, False]) -def test_offload_reload_correctness(offload_optimizer_states, offload_master_weights): - """Test that offload/reload preserves optimizer state values.""" - if not offload_optimizer_states and not offload_master_weights: - pytest.skip("At least one offload type required") - - Utils.initialize_model_parallel() - model, optim = create_model_and_optimizer() - dist_optim = optim.chained_optimizers[0] - - # Run steps to build up optimizer state - for _ in range(3): - run_forward_backward_step(model, optim) - - offloader = dist_optim._state_offloader - - # Capture original states before offload - original_states = {} - for param, state in offloader.adam_optimizer.state.items(): - original_states[param] = { - k: v.clone() for k, v in state.items() if isinstance(v, torch.Tensor) - } - - # Offload - offloader.offload( - offload_optimizer_states=offload_optimizer_states, - offload_master_weights=offload_master_weights, - ) - - # Release GPU memory - offloader.release_gpu_memory() - torch.cuda.synchronize() - - # Reload - offloader.reload() - offloader.sync_before_step() - - # Verify states match after reload - for param, state in offloader.adam_optimizer.state.items(): - if param in original_states: - for key, original_tensor in original_states[param].items(): - if key in state and isinstance(state[key], torch.Tensor): - reloaded_tensor = state[key] - assert reloaded_tensor.device.type == 'cuda', f"State {key} should be on GPU" - torch.testing.assert_close( - reloaded_tensor, - original_tensor, - msg=f"State {key} mismatch after offload/reload", - ) - Utils.destroy_model_parallel() - - -# ============================================================================= -# Test 4: GPU Memory Release Verification -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -def test_gpu_memory_release(): - """Test that GPU memory is actually freed after release_gpu_memory().""" - Utils.initialize_model_parallel() - # Use larger model for measurable memory impact - model, optim = create_model_and_optimizer(hidden_size=1024) - dist_optim = optim.chained_optimizers[0] - - # Initialize optimizer states - run_forward_backward_step(model, optim, hidden_size=1024) - - offloader = dist_optim._state_offloader - - # Measure memory before offload - torch.cuda.synchronize() - torch.cuda.empty_cache() - memory_before = torch.cuda.memory_allocated() - - # Offload and release - offloader.offload() - offloader.release_gpu_memory() - - # Wait for async operations - torch.cuda.synchronize() - torch.cuda.empty_cache() - memory_after = torch.cuda.memory_allocated() - - # Memory should decrease - memory_freed = memory_before - memory_after - assert memory_freed > 0, f"Expected memory to be freed, but got {memory_freed} bytes difference" - Utils.destroy_model_parallel() - - -# ============================================================================= -# Test 5: Multiple Offload/Reload Cycles -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -def test_multiple_offload_reload_cycles(): - """Test that multiple offload/reload cycles work correctly.""" - Utils.initialize_model_parallel() - model, optim = create_model_and_optimizer() - dist_optim = optim.chained_optimizers[0] - - # Initialize - run_forward_backward_step(model, optim) - - offloader = dist_optim._state_offloader - - # Run multiple cycles - for cycle in range(5): - # Offload - offloader.offload() - offloader.release_gpu_memory() - - # Reload - offloader.reload() - offloader.sync_before_step() - - # Run optimizer step - run_forward_backward_step(model, optim) - - # Verify model can still produce valid outputs - input_tensor = torch.randn(8, 256, dtype=torch.bfloat16, device='cuda') - output = model(input_tensor) - assert not output.isnan().any(), "Model output contains NaN after multiple cycles" - Utils.destroy_model_parallel() - - -# ============================================================================= -# Test 6: Training Correctness with Offloading -# ============================================================================= -@pytest.mark.skipif(not TE_FUSED_ADAM_AVAILABLE, reason="Requires TE FusedAdam") -def test_training_correctness_with_offloading(): - """Test that training with offloading produces same results as without.""" - Utils.initialize_model_parallel() - torch.manual_seed(42) - - # Model 1: with offloading - model1, optim1 = create_model_and_optimizer(offload_optimizer_states=True, lr=0.01) - - # Model 2: without offloading (reference) - torch.manual_seed(42) - model2, optim2 = create_model_and_optimizer(offload_optimizer_states=False, lr=0.01) - - # Train both models - n_steps = 10 - torch.manual_seed(123) - dist_optim1 = optim1.chained_optimizers[0] - - # Offloader is created in __init__ when offload_optimizer_states=True - assert dist_optim1._state_offloader is not None - offloader = dist_optim1._state_offloader - - for step in range(n_steps): - input_tensor = torch.randn(8, 256, dtype=torch.bfloat16, device='cuda') - - # Model 1 with offloading - # Offload states (master weights can be offloaded from the start, - # optimizer states will be skipped until after first step) - offloader.offload() - offloader.release_gpu_memory() - - output1 = model1(input_tensor) - loss1 = output1.sum() - loss1.backward() - - offloader.reload() - offloader.sync_before_step() - optim1.step() - optim1.zero_grad() - - # Model 2 without offloading - output2 = model2(input_tensor) - loss2 = output2.sum() - loss2.backward() - optim2.step() - optim2.zero_grad() - - # Compare final model weights - for (n1, p1), (n2, p2) in zip(model1.named_parameters(), model2.named_parameters()): - torch.testing.assert_close( - p1.data, - p2.data, - atol=1e-5, - rtol=1e-4, - msg=f"Parameter {n1} mismatch between offloaded and non-offloaded training", - ) - Utils.destroy_model_parallel() From 605e6151f55b0b4d5b6a767a46e4187c5a4d35fd Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:52:26 -0700 Subject: [PATCH 05/10] Remove HAVE_LION guard, use HAVE_EMERGING_OPTIMIZERS instead Lion is part of the emerging_optimizers package, so a single HAVE_EMERGING_OPTIMIZERS check is sufficient. Co-Authored-By: Claude Opus 4.6 (1M context) --- megatron/core/optimizer/__init__.py | 11 +++-------- megatron/core/optimizer/muon.py | 9 --------- tests/unit_tests/test_lion_optimizer.py | 13 ++++++++----- 3 files changed, 11 insertions(+), 22 deletions(-) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index b64c871104d..318790d88d5 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -34,13 +34,6 @@ USING_PYTORCH_OPTIMIZER = True -try: - from emerging_optimizers.scalar_optimizers import Lion - - HAVE_LION = True -except ImportError: - HAVE_LION = False - from megatron.core import parallel_state from megatron.core.optimizer.cpu_offloading.hybrid_optimizer import HybridDeviceOptimizer from megatron.core.optimizer_param_scheduler import ( @@ -589,11 +582,13 @@ def init_state_fn(opt, config=None): opt.initialize_state(p) elif config.optimizer == 'lion': - if not HAVE_LION: + if not HAVE_EMERGING_OPTIMIZERS: raise ImportError( "Lion optimizer requires the 'emerging_optimizers' package. " "Please install it to use --optimizer lion." ) + from emerging_optimizers.scalar_optimizers import Lion + optimizer = Lion( param_groups, lr=config.lr, diff --git a/megatron/core/optimizer/muon.py b/megatron/core/optimizer/muon.py index 329ce60dd1f..af9b4cd019b 100644 --- a/megatron/core/optimizer/muon.py +++ b/megatron/core/optimizer/muon.py @@ -4,15 +4,6 @@ from typing import Any -# TODO: Remove this separate try/except once the next version of emerging_optimizers -# (which includes Lion) is released. Then Lion can be imported in the block above. -try: - from emerging_optimizers.scalar_optimizers import Lion # pylint: disable=unused-import - - HAVE_LION = True -except ImportError: - HAVE_LION = False - def get_megatron_muon_optimizer(*args: Any, **kwargs: Any) -> Any: """Backward compatible muon optimizer getter. diff --git a/tests/unit_tests/test_lion_optimizer.py b/tests/unit_tests/test_lion_optimizer.py index 5cd479e655a..e22dfab3100 100644 --- a/tests/unit_tests/test_lion_optimizer.py +++ b/tests/unit_tests/test_lion_optimizer.py @@ -14,15 +14,15 @@ import torch.nn as nn from megatron.core.optimizer import ( - HAVE_LION, OptimizerConfig, _get_megatron_optimizer_based_on_param_groups, _get_param_groups, ) +from megatron.core.optimizer.emerging_optimizers import HAVE_EMERGING_OPTIMIZERS from megatron.core.optimizer.optimizer import FP32Optimizer requires_emerging_optimizers = pytest.mark.skipif( - not HAVE_LION, reason="emerging_optimizers package not installed" + not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package not installed" ) @@ -96,10 +96,12 @@ def test_lion_param_groups_via_get_param_groups(self, mock_world_size): def test_lion_import_error_without_package(self): """Should raise ImportError with helpful message if emerging_optimizers not installed.""" import megatron.core.optimizer as opt_module + import megatron.core.optimizer.emerging_optimizers as eo_module - original_have_lion = opt_module.HAVE_LION + original = eo_module.HAVE_EMERGING_OPTIMIZERS try: - opt_module.HAVE_LION = False + eo_module.HAVE_EMERGING_OPTIMIZERS = False + opt_module.HAVE_EMERGING_OPTIMIZERS = False model = SimpleModel() config = OptimizerConfig(optimizer="lion", lr=1e-4) @@ -107,7 +109,8 @@ def test_lion_import_error_without_package(self): with pytest.raises(ImportError, match="emerging_optimizers"): _create_lion_optimizer(model, config) finally: - opt_module.HAVE_LION = original_have_lion + eo_module.HAVE_EMERGING_OPTIMIZERS = original + opt_module.HAVE_EMERGING_OPTIMIZERS = original @requires_emerging_optimizers From 04e0d35d6e33de6d1ff8b32274620889afed2710 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 09:53:40 -0700 Subject: [PATCH 06/10] update lion logic Signed-off-by: Hao Wu --- megatron/core/optimizer/__init__.py | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 318790d88d5..91ba924766d 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -51,6 +51,7 @@ from .emerging_optimizers import ( _EMERGING_OPTIMIZERS, HAVE_EMERGING_OPTIMIZERS, + Lion, _create_emerging_optimizer, ) from .grad_scaler import ConstantGradScaler, DynamicGradScaler @@ -587,7 +588,6 @@ def init_state_fn(opt, config=None): "Lion optimizer requires the 'emerging_optimizers' package. " "Please install it to use --optimizer lion." ) - from emerging_optimizers.scalar_optimizers import Lion optimizer = Lion( param_groups, diff --git a/pyproject.toml b/pyproject.toml index 9ece74b0c66..57877875f94 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -204,7 +204,7 @@ flash_mla = [ ] transformer-engine = { git = "https://github.com/NVIDIA/TransformerEngine.git", rev = "71bbefbf153418f943640df0f7373625dc93fa46" } nemo-run = { git = "https://github.com/NVIDIA-NeMo/Run.git", rev = "01a9a8ba360f7b2908728ad0516e0ad9d936966d" } -emerging_optimizers = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git", rev = "v0.1.0" } +emerging_optimizers = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git", rev = "v0.2.0" } nvidia-resiliency-ext = { git = "https://github.com/NVIDIA/nvidia-resiliency-ext.git", rev = "v0.5.0" } [tool.isort] From dedf3fbd143278bf46afd450a4886ec162724318 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 10:19:07 -0700 Subject: [PATCH 07/10] fix lion import Signed-off-by: Hao Wu --- megatron/core/optimizer/emerging_optimizers.py | 1 + 1 file changed, 1 insertion(+) diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index 74a0d0204f3..28d8a392f6f 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -39,6 +39,7 @@ HAVE_EMERGING_OPTIMIZERS = False OrthogonalizedOptimizer = object AdaptiveMuon = object + Lion = None logger = logging.getLogger(__name__) From 3c9dd05d97466744f22b358cc732524613b6cdd6 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 18:36:44 -0700 Subject: [PATCH 08/10] add safe global for loading torch ckpt Signed-off-by: Hao Wu --- megatron/core/safe_globals.py | 1 + 1 file changed, 1 insertion(+) diff --git a/megatron/core/safe_globals.py b/megatron/core/safe_globals.py index 8bcfe788f60..bd5ec5fb303 100755 --- a/megatron/core/safe_globals.py +++ b/megatron/core/safe_globals.py @@ -33,6 +33,7 @@ RerunState, BytesIO, Signals, + torch._C.Generator, ] From 2a21bbad0fe1a2bd4f5ee83de91b1c8b62c7353b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 18:38:49 -0700 Subject: [PATCH 09/10] remove overlap assert for muon Signed-off-by: Hao Wu --- megatron/training/arguments.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index ad2a05f0241..34213d61f51 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -854,9 +854,8 @@ def validate_args(args, defaults={}): ) if args.overlap_param_gather: - assert args.use_distributed_optimizer or args.use_megatron_fsdp \ - or args.optimizer == 'dist_muon', \ - '--overlap-param-gather only supported with distributed optimizer, megatron fsdp, or dist_muon' + assert args.use_distributed_optimizer or args.use_megatron_fsdp, \ + '--overlap-param-gather only supported with distributed optimizer, megatron fsdp' assert args.overlap_grad_reduce, \ 'Must use --overlap-param-gather with --overlap-grad-reduce' assert not args.use_legacy_models, \ @@ -1488,10 +1487,6 @@ def validate_args(args, defaults={}): 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_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." From ce7174ede7290bcbf0412edff2deda5d3487d906 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Mon, 30 Mar 2026 19:01:39 -0700 Subject: [PATCH 10/10] update uv lock Signed-off-by: Hao Wu --- uv.lock | 98 ++++++++++++++++++++++++++++----------------------------- 1 file changed, 48 insertions(+), 50 deletions(-) diff --git a/uv.lock b/uv.lock index 463d963242f..c1bffde4d37 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.12" resolution-markers = [ "python_full_version >= '3.14' and platform_machine != 's390x' and sys_platform == 'win32' and extra != 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts'", @@ -247,7 +247,7 @@ version = "1.4.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "frozenlist" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/61/62/06741b579156360248d1ec624842ad0edf697050bbaf7c3e46394e106ad1/aiosignal-1.4.0.tar.gz", hash = "sha256:f47eecd9468083c2029cc99945502cb7708b082c232f9aca65da147157b251c7", size = 25007, upload-time = "2025-07-03T22:54:43.528Z" } wheels = [ @@ -303,7 +303,7 @@ source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "idna" }, { name = "sniffio" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/95/7d/4c1bd541d4dffa1b52bd83fb8527089e097a106fc90b467a7313b105f840/anyio-4.9.0.tar.gz", hash = "sha256:673c0c244e15788651a4ff38710fea9675823028a6f08a5eda409e0c9840a028", size = 190949, upload-time = "2025-03-17T00:02:54.77Z" } wheels = [ @@ -631,7 +631,7 @@ name = "cffi" version = "2.0.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pycparser", marker = "implementation_name != 'PyPy'" }, + { name = "pycparser", marker = "implementation_name != 'PyPy' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/eb/56/b1ba7935a17738ae8453301356628e8147c79dbb825bcbc73dc7401f9846/cffi-2.0.0.tar.gz", hash = "sha256:44d1b5909021139fe36001ae048dbdde8214afa20200eda0f64c068cac5d5529", size = 523588, upload-time = "2025-09-08T23:24:04.541Z" } wheels = [ @@ -761,7 +761,7 @@ name = "click" version = "8.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "colorama", marker = "sys_platform == 'win32' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/3d/fa/656b739db8587d7b5dfa22e22ed02566950fbfbcdc20311993483657a5c0/click-8.3.1.tar.gz", hash = "sha256:12ff4785d337a1bb490bb7e9c2b1ee5da3112e94a8622f26a6c77f5d2fc6842a", size = 295065, upload-time = "2025-11-15T20:45:42.706Z" } wheels = [ @@ -930,7 +930,7 @@ name = "cuda-bindings" version = "13.2.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "cuda-pathfinder" }, + { name = "cuda-pathfinder", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/52/c8/b2589d68acf7e3d63e2be330b84bc25712e97ed799affbca7edd7eae25d6/cuda_bindings-13.2.0-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e865447abfb83d6a98ad5130ed3c70b1fc295ae3eeee39fd07b4ddb0671b6788", size = 5722404, upload-time = "2026-03-11T00:12:44.041Z" }, @@ -977,37 +977,37 @@ wheels = [ [package.optional-dependencies] cublas = [ - { name = "nvidia-cublas", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cublas", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cudart = [ - { name = "nvidia-cuda-runtime", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cuda-runtime", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cufft = [ - { name = "nvidia-cufft", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cufft", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cufile = [ - { name = "nvidia-cufile", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cufile", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cupti = [ - { name = "nvidia-cuda-cupti", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cuda-cupti", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] curand = [ - { name = "nvidia-curand", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-curand", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cusolver = [ - { name = "nvidia-cusolver", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cusolver", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] cusparse = [ - { name = "nvidia-cusparse", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cusparse", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] nvjitlink = [ - { name = "nvidia-nvjitlink", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-nvjitlink", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] nvrtc = [ - { name = "nvidia-cuda-nvrtc", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cuda-nvrtc", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] nvtx = [ - { name = "nvidia-nvtx", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-nvtx", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] [[package]] @@ -1174,12 +1174,11 @@ wheels = [ [[package]] name = "emerging-optimizers" -version = "0.1.0" -source = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.1.0#d5363b4a418128cd8111983b191c4b8869a9766b" } +version = "0.2.0" +source = { git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.2.0#1effa026ff096b7fa1063ca2fba19d98be6e6cdf" } dependencies = [ { name = "absl-py" }, - { name = "torch", marker = "sys_platform == 'never'" }, - { name = "typing-extensions" }, + { name = "torch", marker = "sys_platform == 'never' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] [[package]] @@ -1768,7 +1767,7 @@ dependencies = [ { name = "filelock" }, { name = "fsspec", version = "2026.2.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14' or sys_platform != 'win32' or extra == 'extra-13-megatron-core-dev' or extra == 'extra-13-megatron-core-lts'" }, { name = "fsspec", version = "2026.3.0", source = { registry = "https://pypi.org/simple" }, marker = "(python_full_version >= '3.14' and sys_platform == 'win32' and extra != 'extra-13-megatron-core-dev' and extra != 'extra-13-megatron-core-lts') or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, - { name = "hf-xet", marker = "platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64'" }, + { name = "hf-xet", marker = "platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, { name = "packaging" }, { name = "pyyaml" }, { name = "requests" }, @@ -2384,8 +2383,8 @@ requires-dist = [ { name = "datasets", marker = "extra == 'lts'" }, { name = "einops", marker = "extra == 'dev'", specifier = "~=0.8" }, { name = "einops", marker = "extra == 'lts'", specifier = "~=0.8" }, - { name = "emerging-optimizers", marker = "extra == 'dev'", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.1.0" }, - { name = "emerging-optimizers", marker = "extra == 'lts'", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.1.0" }, + { name = "emerging-optimizers", marker = "extra == 'dev'", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.2.0" }, + { name = "emerging-optimizers", marker = "extra == 'lts'", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.2.0" }, { name = "fastapi", marker = "extra == 'dev'", specifier = "~=0.50" }, { name = "fastapi", marker = "extra == 'lts'", specifier = "~=0.50" }, { name = "flash-linear-attention", marker = "extra == 'dev'", specifier = "~=0.4.0" }, @@ -2463,7 +2462,7 @@ linting = [ { name = "ruff", specifier = "~=0.9.0" }, ] no-pypi-wheels = [ - { name = "emerging-optimizers", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.1.0" }, + { name = "emerging-optimizers", git = "https://github.com/NVIDIA-NeMo/Emerging-Optimizers.git?rev=v0.2.0" }, { name = "flash-mla", git = "https://github.com/deepseek-ai/FlashMLA?rev=9edee0c022cd0938148a18e334203b0aab43aa19" }, ] test = [ @@ -3001,7 +3000,7 @@ name = "nvidia-cudnn-cu13" version = "9.19.0.56" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas" }, + { name = "nvidia-cublas", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/f1/84/26025437c1e6b61a707442184fa0c03d083b661adf3a3eecfd6d21677740/nvidia_cudnn_cu13-9.19.0.56-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:6ed29ffaee1176c612daf442e4dd6cfeb6a0caa43ddcbeb59da94953030b1be4", size = 433781201, upload-time = "2026-02-03T20:40:53.805Z" }, @@ -3030,7 +3029,7 @@ name = "nvidia-cufft" version = "12.0.0.61" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink" }, + { name = "nvidia-nvjitlink", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/8b/ae/f417a75c0259e85c1d2f83ca4e960289a5f814ed0cea74d18c353d3e989d/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:2708c852ef8cd89d1d2068bdbece0aa188813a0c934db3779b9b1faa8442e5f5", size = 214053554, upload-time = "2025-09-04T08:31:38.196Z" }, @@ -3062,9 +3061,9 @@ name = "nvidia-cusolver" version = "12.0.4.66" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas" }, - { name = "nvidia-cusparse" }, - { name = "nvidia-nvjitlink" }, + { name = "nvidia-cublas", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-cusparse", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-nvjitlink", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/c8/c3/b30c9e935fc01e3da443ec0116ed1b2a009bb867f5324d3f2d7e533e776b/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:02c2457eaa9e39de20f880f4bd8820e6a1cfb9f9a34f820eb12a155aa5bc92d2", size = 223467760, upload-time = "2025-09-04T08:33:04.222Z" }, @@ -3077,7 +3076,7 @@ name = "nvidia-cusparse" version = "12.6.3.3" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink" }, + { name = "nvidia-nvjitlink", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/f8/94/5c26f33738ae35276672f12615a64bd008ed5be6d1ebcb23579285d960a9/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:80bcc4662f23f1054ee334a15c72b8940402975e0eab63178fc7e670aa59472c", size = 162155568, upload-time = "2025-09-04T08:33:42.864Z" }, @@ -3697,7 +3696,7 @@ source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "numpy" }, { name = "python-dateutil" }, - { name = "tzdata", marker = "sys_platform == 'emscripten' or sys_platform == 'win32'" }, + { name = "tzdata", marker = "sys_platform == 'emscripten' or sys_platform == 'win32' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/2e/0c/b28ed414f080ee0ad153f848586d61d1878f91689950f037f976ce15f6c8/pandas-3.0.1.tar.gz", hash = "sha256:4186a699674af418f655dbd420ed87f50d56b4cd6603784279d9eef6627823c8", size = 4641901, upload-time = "2026-02-17T22:20:16.434Z" } wheels = [ @@ -4668,7 +4667,7 @@ source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "attrs" }, { name = "rpds-py" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/22/f5/df4e9027acead3ecc63e50fe1e36aca1523e1719559c499951bb4b53188f/referencing-0.37.0.tar.gz", hash = "sha256:44aefc3142c5b842538163acb373e24cce6632bd54bdb01b21ad5863489f50d8", size = 78036, upload-time = "2025-10-13T15:30:48.871Z" } wheels = [ @@ -5318,7 +5317,7 @@ version = "0.52.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions", marker = "python_full_version < '3.13' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/c4/68/79977123bb7be889ad680d79a40f339082c1978b5cfcf62c2d8d196873ac/starlette-0.52.1.tar.gz", hash = "sha256:834edd1b0a23167694292e94f597773bc3f89f362be6effee198165a35d62933", size = 2653702, upload-time = "2026-01-18T13:34:11.062Z" } wheels = [ @@ -5339,7 +5338,7 @@ name = "sympy" version = "1.14.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "mpmath" }, + { name = "mpmath", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/83/d3/803453b36afefb7c2bb238361cd4ae6125a569b4db67cd9e79846ba2d68c/sympy-1.14.0.tar.gz", hash = "sha256:d3d3fe8df1e5a0b42f0e7bdf50541697dbe7d23746e894990c030e2b05e72517", size = 7793921, upload-time = "2025-04-27T18:05:01.611Z" } wheels = [ @@ -5550,21 +5549,20 @@ name = "torch" version = "2.11.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "cuda-bindings", marker = "sys_platform == 'linux'" }, + { name = "cuda-bindings", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, { name = "cuda-toolkit", extra = ["cublas", "cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, - { name = "filelock" }, - { name = "fsspec", version = "2026.2.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14' or sys_platform != 'win32' or extra == 'extra-13-megatron-core-dev' or extra == 'extra-13-megatron-core-lts'" }, - { name = "fsspec", version = "2026.3.0", source = { registry = "https://pypi.org/simple" }, marker = "(python_full_version >= '3.14' and sys_platform == 'win32' and extra != 'extra-13-megatron-core-dev' and extra != 'extra-13-megatron-core-lts') or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, - { name = "jinja2" }, - { name = "networkx" }, - { name = "nvidia-cudnn-cu13", marker = "sys_platform == 'linux'" }, - { name = "nvidia-cusparselt-cu13", marker = "sys_platform == 'linux'" }, - { name = "nvidia-nccl-cu13", marker = "sys_platform == 'linux'" }, - { name = "nvidia-nvshmem-cu13", marker = "sys_platform == 'linux'" }, - { name = "setuptools" }, - { name = "sympy" }, - { name = "triton", marker = "sys_platform == 'never'" }, - { name = "typing-extensions" }, + { name = "filelock", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "fsspec", version = "2026.2.0", source = { registry = "https://pypi.org/simple" }, marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "jinja2", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "networkx", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-cudnn-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-cusparselt-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-nccl-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "nvidia-nvshmem-cu13", marker = "sys_platform == 'linux' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "setuptools", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "sympy", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "triton", marker = "sys_platform == 'never' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, + { name = "typing-extensions", marker = "(sys_platform != 'emscripten' and sys_platform != 'win32') or (sys_platform == 'emscripten' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts') or (sys_platform == 'win32' and extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/6f/8b/69e3008d78e5cee2b30183340cc425081b78afc5eff3d080daab0adda9aa/torch-2.11.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:4b5866312ee6e52ea625cd211dcb97d6a2cdc1131a5f15cc0d87eec948f6dd34", size = 80606338, upload-time = "2026-03-23T18:11:34.781Z" }, @@ -5615,7 +5613,7 @@ name = "tqdm" version = "4.67.3" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "colorama", marker = "sys_platform == 'win32' or (extra == 'extra-13-megatron-core-dev' and extra == 'extra-13-megatron-core-lts')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/09/a9/6ba95a270c6f1fbcd8dac228323f2777d886cb206987444e4bce66338dd4/tqdm-4.67.3.tar.gz", hash = "sha256:7d825f03f89244ef73f1d4ce193cb1774a8179fd96f31d7e1dcde62092b960bb", size = 169598, upload-time = "2026-02-03T17:35:53.048Z" } wheels = [