Skip to content
Merged
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
29 changes: 28 additions & 1 deletion megatron/core/optimizer/emerging_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,13 @@
from torch.optim.optimizer import ParamsT

from megatron.core.process_groups_config import ProcessGroupCollection
from megatron.core.utils import get_pg_rank, get_pg_size, log_single_rank
from megatron.core.utils import (
get_emerging_optimizers_version,
get_pg_rank,
get_pg_size,
is_emerging_optimizers_min_version,
log_single_rank,
)

from .optimizer_config import ParamKey, ParamPredicate

Expand All @@ -43,6 +49,12 @@

logger = logging.getLogger(__name__)

# newton_schulz_tp() gained the use_syrk kwarg in emerging_optimizers 0.4.0. Earlier releases
# expose use_syrk on the non-TP newton_schulz() only, so 0.3.x still rejects it here. Spelled
# ".dev0" so pre-release builds of that line are accepted too, matching how the TE minimums
# elsewhere in the tree are written.
_SYRK_MIN_EO_VERSION = "0.4.0.dev0"


def get_supported_coefficient_types() -> tuple[str, ...]:
"""Return the coefficient types supported by the installed emerging_optimizers.
Expand Down Expand Up @@ -178,9 +190,16 @@ def __init__(
extra_scale_factor: float = 1.0,
pg_collection: Optional[ProcessGroupCollection] = None,
tp_mode: Literal["blockwise", "duplicated", "distributed"] = "duplicated",
use_syrk: bool = False,
) -> None:
if num_ns_steps < 1:
raise ValueError(f"num_ns_steps must be at least 1, got {num_ns_steps}")
if use_syrk and not is_emerging_optimizers_min_version(_SYRK_MIN_EO_VERSION):
raise ValueError(
f"use_syrk requires emerging_optimizers >= {_SYRK_MIN_EO_VERSION}, but "
f"{get_emerging_optimizers_version()} is installed. Upgrade "
"emerging_optimizers or drop --muon-use-syrk."
)

def scaled_orthogonalize_fn(
grad: torch.Tensor,
Expand All @@ -197,13 +216,17 @@ def scaled_orthogonalize_fn(
size = [grad.size(-2), grad.size(-1)]
if partition_dim is not None:
size[partition_dim] *= get_pg_size(tp_group)
# Only forward the kwarg when enabled; older emerging_optimizers do not
# accept it at all, and __init__ has already rejected use_syrk on those.
ns_kwargs = {"use_syrk": True} if use_syrk else {}
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,
**ns_kwargs,
)
scale_factor = get_muon_scale_factor(size[0], size[1], mode=scale_mode)
return orth_grad * scale_factor * extra_scale_factor
Expand Down Expand Up @@ -353,6 +376,8 @@ class TensorParallelAdaptiveMuon(TensorParallelMuon, AdaptiveMuon):
extra_scale_factor: The additional scale factor to use for the update.
pg_collection: Process group collection for distributed training.
tp_mode: Tensor parallel mode ("blockwise", "duplicated", or "distributed").
use_syrk: Whether to use the Triton SYRK kernel for the Gram matrix in
Newton-Schulz. Requires emerging_optimizers >= 0.4.0.
moment2_method: Method for second moment accumulation ("adamuon" or "normuon").
beta2: The exponential decay rate for second moment.
eps: Small constant for numerical stability.
Expand All @@ -376,6 +401,7 @@ def __init__(
extra_scale_factor: float = 1.0,
pg_collection: Optional[ProcessGroupCollection] = None,
tp_mode: Literal["blockwise", "duplicated", "distributed"] = "duplicated",
use_syrk: bool = False,
moment2_method: Literal["adamuon", "normuon"] = "adamuon",
beta2: float = 0.95,
eps: float = 1e-8,
Expand All @@ -398,6 +424,7 @@ def __init__(
extra_scale_factor=extra_scale_factor,
pg_collection=pg_collection,
tp_mode=tp_mode,
use_syrk=use_syrk,
)
self.moment2_method = moment2_method

Expand Down
3 changes: 3 additions & 0 deletions megatron/core/optimizer/optimizer_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -282,6 +282,9 @@ class OptimizerConfig:
muon_tp_mode: str = "blockwise"
"""How to perform NS calculation for tensor parallel weights. Defaults to "blockwise"."""

muon_use_syrk: bool = False
"""Use the Triton SYRK kernel for the Gram matrix in Newton-Schulz iteration."""

muon_extra_scale_factor: float = 1.0
"""Additional scale factor for the muon update."""

Expand Down
34 changes: 34 additions & 0 deletions megatron/core/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@
_flashinfer_version = None
_mamba_ssm_version = None
_causal_conv1d_version = None
_emerging_optimizers_version = None


_Wrapped = TypeVar('_Wrapped', bound=Callable)
Expand Down Expand Up @@ -487,6 +488,39 @@ def is_flashinfer_min_version(version, check_equality=True):
return flashinfer_version > PkgVersion(version)


def get_emerging_optimizers_version():
"""Get emerging_optimizers version from __version__; if not available use pip's. Use caching."""
if not HAVE_PACKAGING:
raise ImportError(
"packaging is not installed. Please install it with `pip install packaging`."
)

def get_emerging_optimizers_version_str():
import emerging_optimizers

if hasattr(emerging_optimizers, "__version__"):
return str(emerging_optimizers.__version__)
else:
# The distribution name is hyphenated even though the module is not.
return version("emerging-optimizers")

global _emerging_optimizers_version
if _emerging_optimizers_version is None:
_emerging_optimizers_version = PkgVersion(get_emerging_optimizers_version_str())
return _emerging_optimizers_version


def is_emerging_optimizers_min_version(version, check_equality=True):
"""Check if minimum version of `emerging_optimizers` is installed."""
if not HAVE_PACKAGING:
raise ImportError(
"packaging is not installed. Please install it with `pip install packaging`."
)
if check_equality:
return get_emerging_optimizers_version() >= PkgVersion(version)
return get_emerging_optimizers_version() > PkgVersion(version)


_VALID_DSA_KERNEL_BACKENDS = ("none", "tilelang", "cudnn")


Expand Down
3 changes: 3 additions & 0 deletions megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -2537,6 +2537,9 @@ def _add_regularization_args(parser):
group.add_argument('--muon-tp-mode', type=str, default='blockwise',
choices=['blockwise', 'duplicated', 'distributed'],
help='How to perform NS calculation for tensor model parallel weights')
group.add_argument('--muon-use-syrk', action='store_true',
help='Use the Triton SYRK kernel for the Gram matrix '
'in Newton-Schulz iteration.')
group.add_argument('--muon-extra-scale-factor', type=float, default=1.0,
help='Additional scale factor for the muon update')
group.add_argument('--muon-scalar-optimizer', type=str, default='adam',
Expand Down
38 changes: 38 additions & 0 deletions tests/unit_tests/test_emerging_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -1758,3 +1758,41 @@ def test_lion_optimizer_multi_layer_net():
params_updated += 1

assert params_updated > 0, "At least some parameters should be updated after optimizer step"


# ===========================================================================
# use_syrk version gate
# ===========================================================================


@pytest.mark.parametrize("optimizer_cls", [TensorParallelMuon, TensorParallelAdaptiveMuon])
def test_muon_use_syrk_rejected_on_old_emerging_optimizers(monkeypatch, optimizer_cls):
"""use_syrk must raise on emerging_optimizers < 0.4.0 rather than silently falling back.

Covers TensorParallelAdaptiveMuon too, since it forwards use_syrk through
TensorParallelMuon.__init__ and that forwarding is what applies the gate to both.
"""
import megatron.core.optimizer.emerging_optimizers as eo_module

monkeypatch.setattr(eo_module, "is_emerging_optimizers_min_version", lambda _version: False)
monkeypatch.setattr(eo_module, "get_emerging_optimizers_version", lambda: "0.2.0")

model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda')
with pytest.raises(ValueError, match="use_syrk requires emerging_optimizers"):
optimizer_cls(
params=[model.weight], lr=0.01, pg_collection=None, tp_mode="duplicated", use_syrk=True
)


@pytest.mark.parametrize("optimizer_cls", [TensorParallelMuon, TensorParallelAdaptiveMuon])
def test_muon_use_syrk_default_off_ignores_version(monkeypatch, optimizer_cls):
"""The gate only fires when use_syrk is requested; the default path stays version-agnostic."""
import megatron.core.optimizer.emerging_optimizers as eo_module

monkeypatch.setattr(eo_module, "is_emerging_optimizers_min_version", lambda _version: False)

model = torch.nn.Linear(60, 30, bias=False, dtype=torch.float32, device='cuda')
optimizer = optimizer_cls(
params=[model.weight], lr=0.01, pg_collection=None, tp_mode="duplicated"
)
assert optimizer is not None
Loading