diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index 53ac956b35c..99d8605fba2 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -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 @@ -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. @@ -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, @@ -197,6 +216,9 @@ 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, @@ -204,6 +226,7 @@ def scaled_orthogonalize_fn( 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 @@ -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. @@ -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, @@ -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 diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 24f9a032c47..20045108c89 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -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.""" diff --git a/megatron/core/utils.py b/megatron/core/utils.py index 72373e9ac3b..c29700eb709 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -73,6 +73,7 @@ _flashinfer_version = None _mamba_ssm_version = None _causal_conv1d_version = None +_emerging_optimizers_version = None _Wrapped = TypeVar('_Wrapped', bound=Callable) @@ -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") diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index d4f9eb9c0de..5bdce87c0ed 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -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', diff --git a/tests/unit_tests/test_emerging_optimizers.py b/tests/unit_tests/test_emerging_optimizers.py index e3b9f666fb2..bb8919f4c72 100644 --- a/tests/unit_tests/test_emerging_optimizers.py +++ b/tests/unit_tests/test_emerging_optimizers.py @@ -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