Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
243 changes: 100 additions & 143 deletions megatron/core/optimizer/__init__.py

Large diffs are not rendered by default.

116 changes: 1 addition & 115 deletions megatron/core/optimizer/optimizer_config.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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
Expand All @@ -69,72 +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:
"""Base optimizer configuration object."""
Expand Down
69 changes: 3 additions & 66 deletions megatron/core/optimizer_param_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -101,7 +38,7 @@ class OptimizerParamScheduler:

def __init__(
self,
optimizer: "MegatronOptimizer",
optimizer: MegatronOptimizer,
init_lr: float,
max_lr: float,
min_lr: float,
Expand Down
18 changes: 12 additions & 6 deletions megatron/training/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -1441,9 +1439,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

Expand Down
1 change: 0 additions & 1 deletion tests/unit_tests/optimizer/__init__.py

This file was deleted.

38 changes: 0 additions & 38 deletions tests/unit_tests/optimizer/test_optimizer_config.py

This file was deleted.

Loading
Loading