diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 11aa6c49585..234bee274be 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -3,7 +3,7 @@ import logging import warnings from dataclasses import astuple -from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from typing import Callable, Dict, List, Optional, Tuple, Union import torch from torch.optim import SGD as CPUSGD @@ -35,11 +35,6 @@ from megatron.core import parallel_state from megatron.core.optimizer.cpu_offloading.hybrid_optimizer import HybridDeviceOptimizer -from megatron.core.optimizer_param_scheduler import ( - ParamGroupOverride, - combine_param_group_overrides, - param_group_override_to_tuple, -) from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.fsdp_dtensor_checkpoint import get_global_unique_param_name @@ -55,92 +50,66 @@ MegatronOptimizer, param_group_identifier_keys, ) -from .optimizer_config import ( - AdamOptimizerConfig, - OptimizerConfig, - ParamKey, - ParamPredicate, - ParamWithNamePredicate, - SGDOptimizerConfig, -) +from .optimizer_config import AdamOptimizerConfig, OptimizerConfig, ParamKey, SGDOptimizerConfig logger = logging.getLogger(__name__) -def get_standard_config_overrides(config: OptimizerConfig) -> Dict[ParamKey, ParamGroupOverride]: - """Get standard config overrides for the optimizer, handling decoupled LR and common wd skips. +def _matches(param: torch.nn.Parameter, param_name: str, param_key: ParamKey) -> bool: + """Returns true if passed-in parameter (with name) matches `param_key`. Args: - config (OptimizerConfig): optimizer configuration object. + param (torch.nn.Parameter): Handle to parameter object. + param_name (str): Name of parameter in underlying PyTorch module. + param_key (ParamKey): ParamKey object. Returns: - Dict[ParamKey, ParamGroupOverride]: standard config overrides. + bool: True if parameter matches passed-in param_key. """ - config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]] = {} - # First, figure out how we are going to do wd skipping. The two main approaches are: - # 1. The classic megatron approach of skipping all len 1 and bias parameters. - # 2. The Qwen3-Next approach of doing 1, other than qk layernorm parameters. - if config.apply_wd_to_qk_layernorm: - shape_1_not_qkln_param = ParamWithNamePredicate( - name="s1_not_qkln", - fn=lambda param, name: (len(param.shape) == 1 or name.endswith(".bias")) - and not ("q_layernorm." in name or "k_layernorm." in name), - ) - param_wd_mult_key = ParamKey(with_name_predicate=shape_1_not_qkln_param) - else: - param_length_1_match = ParamPredicate( - name="param_len_1", fn=lambda param: len(param.shape) == 1 - ) - param_wd_mult_key = ParamKey(name="*.bias", predicate=param_length_1_match) - config_overrides[param_wd_mult_key] = ParamGroupOverride(wd_mult=0.0) - - if config.decoupled_lr is not None: - decoupled_lr_config: ParamGroupOverride = {"max_lr": config.decoupled_lr} - decoupled_param_key = ParamKey(attr="is_embedding_or_output_parameter") - if config.decoupled_min_lr is not None: - decoupled_lr_config["min_lr"] = config.decoupled_min_lr - config_overrides[decoupled_param_key] = decoupled_lr_config + # Check if name matches. + if isinstance(param_key.name, str): + target_names = [param_key.name] + else: + target_names = list(param_key.name) + for target_name in target_names: + if param_name in target_name: + return True + + # Check if attribute matches. + if isinstance(param_key.attr, str): + target_attrs = [param_key.attr] + else: + target_attrs = list(param_key.attr) + for target_attr in target_attrs: + if getattr(param, target_attr, False): + return True - return config_overrides + return False def _get_param_groups( model_chunks: List[MegatronModule], config: OptimizerConfig, - config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]], + config_overrides: Optional[Dict[ParamKey, OptimizerConfig]], ) -> List[Dict]: """Create parameter groups for optimizer. Creates parameter groups from provided optimizer config object. - NOTE There can be more than one match between a ParamKey and a parameter. - What we do is merge all of the matching ParamKey overrides into a single ParamGroupOverride - for that parameter and use that as the key for that parameter. Any parameters that get - the same set of merged overrides will be mapped into the same parameter group. - Args: model_chunks (List[MegatronModule]): model chunks to create parameter groups for. config (OptimizerConfig): optimizer configuration object. - config_overrides (Optional[Dict[ParamKey, ParamGroupOverride]): optimizer overrides, - specified on a per-layer basis. NOTE: if you want to skip applying weight decay on bias - and length 1 parameters, and also do not want to do any other overrides, set this to an - empty dictionary rather than the default value of None. + config_overrides (Optional[Dict[LayerKey, OptimizerConfig]): optimizer overrides, + specified on a per-layer basis. Returns: List of parameter groups. """ - # Map (pg_overrides, is_expert_parallel) to params. + # Map (wd_mult, is_expert_parallel, param_group_hyperparameters_config) 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) + configs_map = {} for model_chunk in model_chunks: for name, param in model_chunk.named_parameters(): @@ -148,31 +117,47 @@ def _get_param_groups( continue uses_default_config = False - # Get optimizer config overrides for this parameter. - param_overrides_list: list[ParamGroupOverride] = [] - if config_overrides is not None: - for param_key, param_override in config_overrides.items(): - if param_key.matches(param, name): - param_overrides_list.append(param_override) - - if param_overrides_list: - param_override: ParamGroupOverride | None = combine_param_group_overrides( - param_overrides_list - ) + # Get optimizer config for this parameter. + if config_overrides is None: + config_for_param = config + uses_default_config = True else: - param_override = None + config_for_param = None + for param_key in config_overrides: + if _matches(param, name, param_key): + config_for_param = config_overrides[param_key] + break + # Fall back to default config. + if config_for_param is None: + config_for_param = config + uses_default_config = True is_expert_parallel = not getattr(param, 'allreduce', True) - # Create config_tuple that is hash-able, and has a consistent ordering of the keys. - param_override_tuple: tuple[tuple[str, Any], ...] | None = ( - param_group_override_to_tuple(param_override) - ) - key = (param_override_tuple, is_expert_parallel) + # TODO: Make sure there is a way to support old no_weight_decay_func functionality + # and default_skip_embedding_weight_decay: + # or (default_skip_embedding_weight_decay and "embedding" in name) + no_wd = name.endswith(".bias") or len(param.shape) == 1 + if not no_wd: + wd_mult = 1.0 + else: + wd_mult = 0.0 + + # Create config_tuple that is hash-able. Remove timers object before + # creating config_tuple. + config_for_param_copy = copy.deepcopy(config_for_param) + config_for_param_copy.timers = None + config_tuple = astuple(config_for_param_copy) + key = (wd_mult, is_expert_parallel, config_tuple) if key not in params_map: params_map[key] = [] params_map[key].append(param) + if key in configs_map: + assert (config_for_param, uses_default_config) == configs_map[key] + else: + configs_map[key] = (config_for_param, uses_default_config) + # Distributed checkpoint requires all ranks to have the same param groups, # so we need to align the param groups across ranks, otherwise we may have # runtime error when loading the checkpoint or numerical error when resuming training. @@ -183,47 +168,34 @@ def _get_param_groups( for key in keys: if key not in params_key: params_key.append(key) - # Need to pick one of the param_override_tuples to use for the param group. + param_groups = [] - # Sort keys, None first. - for key in sorted(params_key, key=lambda x: (x[0] is not None, x[0])): - param_override_tuple, is_expert_parallel = key + for key in params_key: + wd_mult, is_expert_parallel, _ = key params = params_map[key] if key in params_map else [] - if param_override_tuple is None: - param_override: ParamGroupOverride = {} + config, uses_default_config = None, True + if key not in configs_map: + assert params == [] else: - param_override: ParamGroupOverride = {k: v for (k, v) in param_override_tuple} - - # False if param_group_override is None or empty tuple or if we do not modify the - # LR schedule. - # NOTE: "default_config" is used for logging the learning rate in training.py. - # so set to True if we do not modify the learning rate. - # if param_group['default_config']: - # learning_rate = param_group['lr'] - uses_default_lr_schedule: bool = (not bool(param_override_tuple)) or not any( - ["lr" in k for k in param_override] - ) + config, uses_default_config = configs_map[key] + assert config is not None # TODO: Remove "backwards compatible" fields below eventually. - default_config: ParamGroupOverride = { - 'wd_mult': 1.0, - 'lr_mult': 1.0, - 'is_decoupled_lr': False, - # The following two fields may be important to keep even when we remove the - # above "backwards compatible" fields. - "max_lr": config.lr, # user may override this in param_override - "min_lr": config.min_lr, # user may override this in param_override - } - assert ( - "params" not in param_override - ), "'params' should not be in param_override, this is a protected key" param_group = { 'params': params, + 'wd_mult': wd_mult, # For backwards compatibility. + 'lr_mult': 1.0, # For backwards compatibility. 'is_expert_parallel': is_expert_parallel, - 'default_config': uses_default_lr_schedule, - **default_config, - **param_override, # keep **param_override last so that users can override other fields. + 'is_decoupled_lr': False, # For backwards compatibility. + 'default_config': uses_default_config, } + + # Stick relevant fields into param_group from config object. + if config is not None: + param_group['max_lr'] = config.lr + param_group['min_lr'] = config.min_lr + # TODO: Add other relevant arguments (e.g., weight decay, optimizer) + # here as well. param_groups.append(param_group) return param_groups @@ -233,7 +205,7 @@ def _get_param_groups_and_buffers( model_chunks: List[MegatronModule], model_chunk_offset: int, config: OptimizerConfig, - config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]], + config_overrides: Optional[Dict[ParamKey, OptimizerConfig]], filter_fn: Callable, buffer_name: str, ) -> Tuple[List[Dict], Dict[int, List[_ParamAndGradBuffer]]]: @@ -244,8 +216,8 @@ def _get_param_groups_and_buffers( groups for. model_chunk_offset (int): offset of model_chunks in global model_chunks list. config (OptimizerConfig): optimizer configuration object. - config_overrides (Optional[Dict[ParamKey, ParamGroupOverride]): optimizer/scheduler - overrides, specified on the basis of ParamKey matches with each parameter. + config_overrides (Optional[Dict[LayerKey, OptimizerConfig]): optimizer overrides, + specified on a per-layer basis. lr (float): learning rate. min_lr (float): minimum learning rate. filter_fn (callable): filtering function for param_groups. @@ -475,37 +447,10 @@ def init_state_fn(opt, config=None): return optimizer -def check_config_overrides_consistency( - config: OptimizerConfig, config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]] -): - """Check if the config overrides are consistent with the config.""" - - # TODO: Remove `optimizer` from this eventually (e.g., if we use Muon for some layers and - # Adam for other layers). This would need some more refactoring to work though (param_groups - # filtered by optimizer passed into _get_megatron_optimizer_based_on_param_groups). - if config_overrides is not None: - fields_to_check_for_consistency = [ - 'overlap_param_gather_with_optimizer_step', - 'optimizer', - 'optimizer_cpu_offload', - ] - for field_name in fields_to_check_for_consistency: - base_field = getattr(config, field_name, None) - all_config_overrides = list(config_overrides.values()) - for config_override in all_config_overrides: - if field_name in config_override: - field = config_override[field_name] - if field != base_field: - raise ValueError( - f"Field {field_name} should not be overriden in a config override." - ) - return True - - def get_megatron_optimizer( config: OptimizerConfig, model_chunks: List[MegatronModule], - config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]] = None, + config_overrides: Optional[Dict[ParamKey, OptimizerConfig]] = None, use_gloo_process_groups: bool = True, pg_collection: Optional[ProcessGroupCollection] = None, dump_param_to_param_group_map: Optional[str] = None, @@ -531,7 +476,19 @@ def get_megatron_optimizer( log_single_rank(logger, logging.INFO, f'Setting up optimizer with config {config}') - check_config_overrides_consistency(config, config_overrides) + # TODO: Remove `optimizer` from this eventually (e.g., if we use Muon for some layers and + # Adam for other layers). This would need some more refactoring to work though (param_groups + # filtered by optimizer passed into _get_megatron_optimizer_based_on_param_groups). + fields_to_check_for_consistency = [ + 'overlap_param_gather_with_optimizer_step', + 'optimizer', + 'optimizer_cpu_offload', + ] + for field_name in fields_to_check_for_consistency: + field = getattr(config, field_name, None) + if config_overrides is not None: + all_configs = list(config_overrides.values()) + assert all([getattr(x, field_name, None) == field for x in all_configs]) # Separate out first model chunk if overlapping param AG with optimizer step. if config.overlap_param_gather_with_optimizer_step: diff --git a/megatron/core/optimizer/muon.py b/megatron/core/optimizer/muon.py index 57eb1e94478..0623c1e0585 100644 --- a/megatron/core/optimizer/muon.py +++ b/megatron/core/optimizer/muon.py @@ -8,7 +8,6 @@ import torch from torch.optim.optimizer import ParamsT -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 @@ -165,7 +164,7 @@ def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> t def get_megatron_muon_optimizer( config: OptimizerConfig, model_chunks: List[MegatronModule], - config_overrides: Optional[Dict[ParamKey, ParamGroupOverride]] = None, + config_overrides: Optional[Dict] = None, use_gloo_process_groups: bool = True, layer_wise_distributed_optimizer: bool = False, pg_collection: Optional[ProcessGroupCollection] = None, diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 2d3e3ca08e0..6778f23baa2 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -1,6 +1,5 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -import fnmatch from dataclasses import dataclass, field from typing import Callable, Optional, Tuple, Union @@ -9,58 +8,6 @@ from ..utils import is_te_min_version -@dataclass(frozen=True) -class ParamPredicate: - """Wraps a matching function to make it hashable for ParamKey. - Example: - >>> shape_1_param = ParamPredicate(name="s1", fn=lambda param: len(param.shape) == 1) - >>> shape_1_param(torch.empty(10)) - True - >>> shape_1_param_copy = ParamPredicate(name="s1", fn=lambda param: len(param.shape) == 1) - >>> shape_1_param == shape_1_param_copy # name is used to match - True - >>> {shape_1_param, shape_1_param_copy} == {shape_1_param} # set hashing works properly - - NOTE: - __hash__ and __eq__ are automatically generated by @dataclass(frozen=True) - based solely on 'name' because we set compare=False/hash=False on 'fn'. - """ - - name: str - fn: Callable[[torch.nn.Parameter], bool] = field(compare=False, hash=False) - - def __call__(self, param: torch.nn.Parameter) -> bool: - return self.fn(param) - - -@dataclass(frozen=True) -class ParamWithNamePredicate: - """Wraps a matching function to make it hashable for ParamKey. - Example: - >>> shape_1_not_qkln_param = ParamWithNamePredicate( - name="s1_not_qkln", - fn=lambda param, name: ( - len(param.shape) == 1 or name.endswith(".bias") - and not ("q_layernorm." in name or "k_layernorm." in name) - ) - ) - >>> shape_1_not_qkln_param(torch.empty(10), "interesting.bias") - True - >>> shape_1_not_qkln_param(torch.empty(10), "interesting.q_layernorm.bias") - False - - NOTE: - __hash__ and __eq__ are automatically generated by @dataclass(frozen=True) - based solely on 'name' because we set compare=False/hash=False on 'fn'. - """ - - name: str - fn: Callable[[torch.nn.Parameter, str], bool] = field(compare=False, hash=False) - - def __call__(self, param: torch.nn.Parameter, name: str) -> bool: - return self.fn(param, name) - - @dataclass(frozen=True, slots=True) class ParamKey: """Key to group parameters by. All such grouped parameters can share an @@ -69,71 +16,11 @@ class ParamKey: # TODO: Can add layer_id here later. name: Union[str, Tuple[str]] = field(default_factory=tuple) - """Parameter name(s), will use unix filesystem path syntax for matching.""" + """Parameter name(s).""" attr: Union[str, Tuple[str]] = field(default_factory=tuple) """Parameter attribute(s).""" - predicate: Union[ParamPredicate, Tuple[ParamPredicate]] = field(default_factory=tuple) - """Predicate(s) to match parameters by. If multiple predicates are provided, any must match.""" - - with_name_predicate: Union[ParamWithNamePredicate, Tuple[ParamWithNamePredicate]] = field( - default_factory=tuple - ) - """ - Predicate(s) to match parameters with their name. If multiple predicates are provided, - any must match. This is useful if you need to filter out some parameters from an otherwise - positive match by their name. - """ - - def matches(self, param: torch.nn.Parameter, param_name: str) -> bool: - """Returns true if passed-in parameter (with name) matches `param_key`. - - Args: - param (torch.nn.Parameter): Handle to parameter object. - param_name (str): Name of parameter in underlying PyTorch module. - - Returns: - bool: True if parameter matches passed-in param_key. - """ - - # Check if name matches. - if isinstance(self.name, str): - target_names = [self.name] - else: - target_names = list(self.name) - for target_name in target_names: - if fnmatch.fnmatch(param_name, target_name): - return True - - # Check if attribute matches. - if isinstance(self.attr, str): - target_attrs = [self.attr] - else: - target_attrs = list(self.attr) - for target_attr in target_attrs: - if getattr(param, target_attr, False): - return True - - # Check if predicate matches. - if isinstance(self.predicate, ParamPredicate): - if self.predicate(param): - return True - else: - for predicate in self.predicate: - if predicate(param): - return True - - # Check if with_name_predicate matches. - if isinstance(self.with_name_predicate, ParamWithNamePredicate): - if self.with_name_predicate(param, param_name): - return True - else: - for predicate in self.with_name_predicate: - if predicate(param, param_name): - return True - return False - @dataclass class OptimizerConfig: diff --git a/megatron/core/optimizer_param_scheduler.py b/megatron/core/optimizer_param_scheduler.py index 7ff6fee35a7..9f771c612e8 100644 --- a/megatron/core/optimizer_param_scheduler.py +++ b/megatron/core/optimizer_param_scheduler.py @@ -3,77 +3,14 @@ """Learning rate decay and weight decay incr functions.""" import logging import math -from typing import TYPE_CHECKING, Any, Optional, TypedDict +from typing import Optional +from megatron.core.optimizer import MegatronOptimizer from megatron.core.utils import log_single_rank -if TYPE_CHECKING: - # Avoid circular import. - from megatron.core.optimizer import MegatronOptimizer - logger = logging.getLogger(__name__) -class ParamGroupOverride(TypedDict): - """Override values for a parameter group. These values may be optimizer-state/scheduler related. - - These are the values you see later in param_group.get(...) calls in the - OptimizerParamScheduler.get_lr and get_wd methods. If you use a custom optimizer - or scheduler, you could override those variables instead. - - Example: - >>> param_group_override = ParamGroupOverride(min_lr=1e-4, wd_mult=0.1) - >>> param_group_override == ParamGroupOverride(newvar=3) # this is ok too - - """ - - max_lr: float - min_lr: float - start_wd: float - end_wd: float - wd_mult: float - - -def param_group_override_to_tuple( - param_group_override: ParamGroupOverride | None, -) -> tuple[tuple[str, Any], ...] | None: - """Convert a param group override to a tuple for use as a key in a dictionary. - - The tuple is sorted by the keys of the param group override to handle different orderings of - the keys in different override dictionaries which still mean the same thing. - """ - if param_group_override is None: - return None - return tuple(sorted(param_group_override.items())) - - -def combine_param_group_overrides( - param_group_overrides: list[ParamGroupOverride | None], -) -> ParamGroupOverride: - """Combine a list of param group overrides into a single param group override. - - This function ensures that the overrides are not conflicting as well. - - Args: - param_group_overrides (list[ParamGroupOverride]): list of param group overrides to combine - - Returns: - ParamGroupOverride: combined param group override - """ - combined_override = ParamGroupOverride() - for override in param_group_overrides: - if override is None: - continue - for key, value in override.items(): - if key in combined_override: - if combined_override[key] != value: - raise ValueError( - f"Conflicting overrides for {key}: {combined_override[key]} and {value}" - ) - combined_override[key] = value - return combined_override - - class OptimizerParamScheduler: """Anneals learning rate and weight decay @@ -101,7 +38,7 @@ class OptimizerParamScheduler: def __init__( self, - optimizer: "MegatronOptimizer", + optimizer: MegatronOptimizer, init_lr: float, max_lr: float, min_lr: float, diff --git a/megatron/training/training.py b/megatron/training/training.py index 563c228367f..33ddf6c20a3 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -42,8 +42,7 @@ def set_startup_timestamps(program_start=None, main_entry=None): import math import os import sys -from contextlib import nullcontext -from typing import Any, Optional, Dict +from typing import Any, Optional import torch.distributed @@ -99,7 +98,6 @@ def set_startup_timestamps(program_start=None, main_entry=None): is_vp_first_stage, is_vp_last_stage, ) -from megatron.core.optimizer import get_standard_config_overrides from megatron.training.checkpointing import load_checkpoint from megatron.training.checkpointing import save_checkpoint, save_grads from megatron.training.checkpointing import checkpoint_exists @@ -1468,9 +1466,17 @@ def get_megatron_optimizer_config(args: Any) -> OptimizerConfig: else: raise ValueError("Invalid optimizer type!") - # Construct the appropriate config_overrides object. This default handles many cases, but - # can be added to as needed by the user, or replaced entirely with a custom override. - config_overrides = get_standard_config_overrides(config=config) + # Construct the appropriate config_overrides object. + # TODO: add more logic here as needed down the road. + if args.decoupled_lr is not None: + decoupled_param_key = ParamKey(attr="is_embedding_or_output_parameter") + decoupled_optimizer_config = copy.deepcopy(config) + decoupled_optimizer_config.lr = args.decoupled_lr + if args.decoupled_min_lr is not None: + decoupled_optimizer_config.min_lr = args.decoupled_min_lr + config_overrides = {decoupled_param_key: decoupled_optimizer_config} + else: + config_overrides = None return config, config_overrides diff --git a/tests/unit_tests/optimizer/__init__.py b/tests/unit_tests/optimizer/__init__.py deleted file mode 100644 index b5dff7b5663..00000000000 --- a/tests/unit_tests/optimizer/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. diff --git a/tests/unit_tests/optimizer/test_optimizer_config.py b/tests/unit_tests/optimizer/test_optimizer_config.py deleted file mode 100644 index 0ecb877ed27..00000000000 --- a/tests/unit_tests/optimizer/test_optimizer_config.py +++ /dev/null @@ -1,38 +0,0 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -import torch - -from megatron.core.optimizer.optimizer_config import ParamKey, ParamPredicate - - -def test_paramkey_matches(): - len_1_predicate = ParamPredicate(name="param_len_1", fn=lambda param: len(param.shape) == 1) - endswith_bias = ParamKey(name="*.bias") - has_dotbias = ParamKey(name="*.bias*") - len_1_param = ParamKey(predicate=len_1_predicate) - has_bias_or_len1_param = ParamKey(name="*.bias", predicate=len_1_predicate) - has_attr = ParamKey(attr="is_embedding_or_output_parameter") - - assert endswith_bias.matches(torch.nn.Parameter(torch.empty(10, 10)), "interesting.bias") - assert not endswith_bias.matches( - torch.nn.Parameter(torch.empty(10, 10)), "something.bias.other" - ) - assert has_dotbias.matches(torch.nn.Parameter(torch.empty(10)), "random.biasstuff") - assert not has_dotbias.matches(torch.nn.Parameter(torch.empty(10, 10)), "random_bias_name") - assert len_1_param.matches(torch.nn.Parameter(torch.empty(10)), "interesting.bias") - assert not len_1_param.matches(torch.nn.Parameter(torch.empty(10, 10)), "interesting_bias") - assert has_bias_or_len1_param.matches( - torch.nn.Parameter(torch.empty(10, 10)), "interesting.bias" - ) - assert has_bias_or_len1_param.matches(torch.nn.Parameter(torch.empty(10)), "interesting_bias") - assert not has_bias_or_len1_param.matches( - torch.nn.Parameter(torch.empty(10, 10)), "random_bias_name" - ) - p_with_attr = torch.nn.Parameter(torch.empty(10, 10)) - setattr(p_with_attr, "is_embedding_or_output_parameter", True) - assert has_attr.matches(p_with_attr, "interesting.bias") - assert not has_attr.matches(torch.nn.Parameter(torch.empty(10, 10)), "interesting.bias") - - # We expect that if the return of the attribute is False, it should not match even if - # it has the attribute. - setattr(p_with_attr, "is_embedding_or_output_parameter", False) - assert not has_attr.matches(p_with_attr, "interesting.bias") diff --git a/tests/unit_tests/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 6b1da8c4e3f..63841dc2bde 100644 --- a/tests/unit_tests/test_optimizer.py +++ b/tests/unit_tests/test_optimizer.py @@ -1,7 +1,6 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import os -from unittest.mock import patch import pytest import torch @@ -13,17 +12,7 @@ from transformer_engine.pytorch.fp8 import fp8_autocast from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig -from megatron.core.optimizer import ( - ChainedOptimizer, - OptimizerConfig, - ParamKey, - ParamPredicate, - _get_param_groups, - check_config_overrides_consistency, - get_megatron_optimizer, - get_standard_config_overrides, -) -from megatron.core.optimizer_param_scheduler import ParamGroupOverride +from megatron.core.optimizer import ChainedOptimizer, OptimizerConfig, get_megatron_optimizer from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import TransformerConfig from megatron.core.utils import is_te_min_version, is_torch_min_version @@ -35,7 +24,7 @@ from transformer_engine.pytorch.fp8 import check_fp8_block_scaling_support fp8_block_scaling_available, reason_for_no_fp8_block_scaling = check_fp8_block_scaling_support() - from transformer_engine.common.recipe import DelayedScaling, Float8BlockScaling, Format + from transformer_engine.common.recipe import Float8BlockScaling, Format except: fp8_block_scaling_available = False reason_for_no_fp8_block_scaling = "FP8 block scaled GEMM requires Hopper and CUDA >= 12.9." @@ -69,223 +58,6 @@ def forward(self, x): return x -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_no_overrides(mock_get_world_size): - net = Net() - # NOTE: to get no overrides, supply an empty dictionary rather than None. - param_groups = _get_param_groups([net], OptimizerConfig(optimizer='adam', lr=0.01), {}) - assert len(param_groups) == 1 - pg0 = param_groups[0] - assert pg0.keys() == { - 'params', - 'is_expert_parallel', - 'default_config', - 'wd_mult', - 'lr_mult', - 'is_decoupled_lr', - 'max_lr', - 'min_lr', - } - assert pg0['params'] == list(net.parameters()) - assert pg0['is_expert_parallel'] == False - assert pg0['default_config'] == True - assert pg0['wd_mult'] == 1.0 - assert pg0['lr_mult'] == 1.0 - assert pg0['is_decoupled_lr'] == False - assert pg0['max_lr'] == 0.01 # from the optimizer config default for lr - assert pg0['min_lr'] is None # from the optimizer config default. - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -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) - assert len(param_groups) == 2 - pg0, pg1 = param_groups - wd_mults = {pg0['wd_mult'], pg1['wd_mult']} - assert wd_mults == {1.0, 0.0} - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_with_overrides(mock_get_world_size): - net = Net() - config_overrides = { - ParamKey( - name="*.bias", - predicate=ParamPredicate(name="param_len_1", fn=lambda param: len(param.shape) == 1), - ): ParamGroupOverride(wd_mult=0.0) - } - opt_config = OptimizerConfig(optimizer='adam', lr=0.01) - check_config_overrides_consistency(opt_config, config_overrides) - param_groups = _get_param_groups([net], opt_config, config_overrides) - assert len(param_groups) == 2 - p_set = set(net.parameters()) - - assert p_set == set(param_groups[0]['params']) | set(param_groups[1]['params']) - assert len(p_set) == len(param_groups[0]['params']) + len(param_groups[1]['params']) - assert param_groups[0]['wd_mult'] == 0.0 or param_groups[1]['wd_mult'] == 0.0 - assert param_groups[0]['wd_mult'] == 1.0 or param_groups[1]['wd_mult'] == 1.0 - assert len(param_groups[0]['params']) > 0 and len(param_groups[1]['params']) > 0 - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_multiple_matches(mock_get_world_size): - net = Net() - - param_groups = _get_param_groups( - [net], - OptimizerConfig(optimizer='adam', lr=0.01), - { - ParamKey(name="*.bias"): ParamGroupOverride(min_lr=1e-4, wd_mult=0.0), - ParamKey( - predicate=ParamPredicate(name="param_len_1", fn=lambda param: len(param.shape) == 1) - ): ParamGroupOverride(wd_mult=0.0, min_lr=1e-4), - }, - ) - config_overrides = { - ParamKey( - name="*.bias", - predicate=ParamPredicate(name="param_len_1", fn=lambda param: len(param.shape) == 1), - ): ParamGroupOverride(min_lr=1e-4, wd_mult=0.0) - } - opt_config = OptimizerConfig(optimizer='adam', lr=0.01) - check_config_overrides_consistency(opt_config, config_overrides) - param_groups2 = _get_param_groups([net], opt_config, config_overrides) - assert len(param_groups) == 2 - assert param_groups == param_groups2 - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_overlapping_matches(mock_get_world_size): - """In this test, we see if we can have two matches that create three param groups.""" - net = Net() - # We expect that all convolution parameters will have wd_mult=0.0 - # However the conv1 related parameters will additionally have a different LR schedule. - # this should create three param groups (no match, conv1 (both wd_mult=0.0 and LR schedule), conv2 (only wd_mult=0.0)) - config_overrides = { - ParamKey(name="*conv*"): ParamGroupOverride(wd_mult=0.0), - ParamKey(name="*conv1*"): ParamGroupOverride(min_lr=10, max_lr=20), - } - opt_config = OptimizerConfig(optimizer='adam', lr=0.01) - check_config_overrides_consistency(opt_config, config_overrides) - param_groups = _get_param_groups([net], opt_config, config_overrides) - assert len(param_groups) == 3 - p_set = set(net.parameters()) - assert p_set == set(param_groups[0]['params']) | set(param_groups[1]['params']) | set( - param_groups[2]['params'] - ) - assert len(p_set) == len(param_groups[0]['params']) + len(param_groups[1]['params']) + len( - param_groups[2]['params'] - ) - assert ( - param_groups[0]['wd_mult'] == 1.0 - ), "We expect the first param group to be the None one, which should have wd_mult=1.0" - assert ( - param_groups[1]['wd_mult'] == 0.0 - ), "We expect the second param group to be the conv1 one, which should have wd_mult=0.0" - assert ( - param_groups[2]['wd_mult'] == 0.0 - ), "We expect the third param group to be the conv2 one, which should have wd_mult=0.0" - assert param_groups[1]['min_lr'] == 10 - assert param_groups[1]['max_lr'] == 20 - assert param_groups[2]['min_lr'] is None - assert param_groups[2]['max_lr'] == 0.01 - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_with_standard_config_overrides(apply_wd_to_qk_layernorm: bool): - """In this test, we see if the standard config overrides are applied correctly.""" - - # Initialize the model with layernorm - net = Net() - - config = OptimizerConfig(optimizer='adam', lr=0.01) - config_overrides = get_standard_config_overrides(config=config) - param_groups = _get_param_groups([net], config, config_overrides) - - assert len(param_groups) == 2 - p_set = set(net.parameters()) - - assert p_set == set(param_groups[0]['params']) | set(param_groups[1]['params']) - assert len(p_set) == len(param_groups[0]['params']) + len(param_groups[1]['params']) - assert param_groups[0]['wd_mult'] == 0.0 or param_groups[1]['wd_mult'] == 0.0 - assert param_groups[0]['wd_mult'] == 1.0 or param_groups[1]['wd_mult'] == 1.0 - assert len(param_groups[0]['params']) > 0 and len(param_groups[1]['params']) > 0 - - # Both param groups should have 5 parameters. - # Param group A (wd_mult=1.0): conv1.weight, conv2.weight, fc1.weight, fc2.weight, fc3.weight - # Param group B (wd_mult=0.0): conv1.bias, conv2.bias, fc1.bias, fc2.bias, fc3.bias - assert len(param_groups[0]['params']) == 5, ( - f"Expected 5 parameters in the first param group, " - f"but got {len(param_groups[0]['params'])}" - ) - assert len(param_groups[1]['params']) == 5, ( - f"Expected 5 parameters in the second param group, " - f"but got {len(param_groups[1]['params'])}" - ) - - -@patch('torch.distributed.get_world_size', return_value=1) -@patch( - 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) -) -def test_get_param_groups_appling_wd_to_qk_layernorm(apply_wd_to_qk_layernorm: bool): - """In this test, we see if the `apply_wd_to_qk_layernorm` config is applied correctly.""" - - # Initialize the model with layernorm - net = Net(add_layernorm=True) - - config = OptimizerConfig( - optimizer='adam', lr=0.01, apply_wd_to_qk_layernorm=apply_wd_to_qk_layernorm - ) - config_overrides = get_standard_config_overrides(config=config) - param_groups = _get_param_groups([net], config, config_overrides) - - assert len(param_groups) == 2 - p_set = set(net.parameters()) - - assert p_set == set(param_groups[0]['params']) | set(param_groups[1]['params']) - assert len(p_set) == len(param_groups[0]['params']) + len(param_groups[1]['params']) - assert param_groups[0]['wd_mult'] == 1.0 - assert param_groups[1]['wd_mult'] == 0.0 - - # There are two param groups, having 7, and 6 parameters respectively. - # Param group A (wd_mult=1.0): conv1.weight, conv2.weight, fc1.weight, fc2.weight, fc3.weight, - # q_layernorm.weight, k_layernorm.weight - # Param group B (wd_mult=0.0): conv1.bias, conv2.bias, fc1.bias, fc2.bias, fc3.bias, - # layernorm.weight - assert len(param_groups[0]['params']) == 7, ( - f"Expected 5 parameters in the first param group, " - f"but got {len(param_groups[0]['params'])}" - ) - assert len(param_groups[1]['params']) == 6, ( - f"Expected 6 parameters in the second param group, " - f"but got {len(param_groups[1]['params'])}" - ) - - def test_chained_optimizer(): net = Net() optimizer_1 = Adam(list(net.parameters())[:2], lr=0.01) diff --git a/tests/unit_tests/test_utilities.py b/tests/unit_tests/test_utilities.py index 39c78efb2b9..f16f88f7865 100644 --- a/tests/unit_tests/test_utilities.py +++ b/tests/unit_tests/test_utilities.py @@ -1,4 +1,3 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import os from datetime import timedelta @@ -28,8 +27,8 @@ def __init__( class Utils: - world_size = int(os.environ.get('WORLD_SIZE', '1')) - rank = int(os.environ.get('LOCAL_RANK', '0')) + world_size = int(os.environ['WORLD_SIZE']) + rank = int(os.environ['LOCAL_RANK']) inited = False store = None