diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 8cfb22620bb..c32ed77d1d7 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -132,6 +132,8 @@ def get_mup_config_overrides( - Non-Adam optimizers: - hidden (matrix-like) lr = base_lr / width_mult - no eps override is applied. + - for Muon optimizers, matrix-like params managed by Muon itself are + excluded from these Adam-style MuP overrides. With decoupled_lr enabled, embedding/output params continue using decoupled LR and MuP will not override those explicit decoupled values. @@ -147,6 +149,7 @@ def get_mup_config_overrides( optimizer_type_lower = optimizer_type.lower() is_sgd_optimizer = optimizer_type_lower == 'sgd' is_adam_optimizer = 'adam' in optimizer_type_lower + is_muon_optimizer = 'muon' in optimizer_type_lower decoupled_lr_enabled = config.decoupled_lr is not None if decoupled_lr_enabled: @@ -159,6 +162,18 @@ def get_mup_config_overrides( message += " MuP Adam epsilon scaling remains applied to hidden matrix-like parameters." log_single_rank(logger, logging.WARNING, message) + if is_muon_optimizer: + muon_scale_mode = getattr(config, 'muon_scale_mode', 'spectral') + if muon_scale_mode == 'spectral': + log_single_rank( + logger, + logging.WARNING, + "Both MuP and muon_scale_mode=spectral are enabled. " + "Muon-managed matrix parameters will continue using spectral Muon scaling. " + "Set --muon-scale-mode unit_rms_norm to use unit_rms_norm scaling for " + "Muon-managed matrices with MuP.", + ) + if mup_width_mult == 1.0: # No scaling needed when width_mult is 1 return {} @@ -184,9 +199,16 @@ def is_vector_like_parameter(param: torch.nn.Parameter, param_name: str) -> bool return True return False + def is_muon_managed_matrix_parameter(param: torch.nn.Parameter, _: str) -> bool: + if not is_muon_optimizer: + return False + return param.dim() == 2 and not getattr(param, 'is_embedding_or_output_parameter', False) + def should_scale_lr_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: if decoupled_lr_enabled and getattr(param, 'is_embedding_or_output_parameter', False): return False + if is_muon_managed_matrix_parameter(param, param_name): + return False return not is_vector_like_parameter(param, param_name) def should_scale_vector_like_lr_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: @@ -197,6 +219,8 @@ def should_scale_vector_like_lr_with_mup(param: torch.nn.Parameter, param_name: def should_scale_eps_with_mup(param: torch.nn.Parameter, param_name: str) -> bool: if is_vector_like_parameter(param, param_name): return False + if is_muon_managed_matrix_parameter(param, param_name): + return False # MuP Appendix B.3: eps scales with fan_in when non-negligible. # This implementation follows the common denominator form: sqrt(v) + eps. return True diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 0d1391d8b97..b8c51b6651d 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2230,7 +2230,9 @@ def _add_regularization_args(parser): 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'], - help='Scale mode for Muon optimizer') + help='Scale mode for Muon optimizer. With MuP, set ' + '--muon-scale-mode unit_rms_norm to use unit_rms_norm scaling, ' + 'or set --muon-scale-mode spectral to keep spectral scaling.') group.add_argument('--muon-fp32-matmul-prec', type=str, default='medium', choices=['low', 'medium', 'high'], help='FP32 matmul precision for Newton-Schulz iteration') diff --git a/tests/unit_tests/transformer/test_mup.py b/tests/unit_tests/transformer/test_mup.py index 605a4e2607b..f1d99cad1e6 100644 --- a/tests/unit_tests/transformer/test_mup.py +++ b/tests/unit_tests/transformer/test_mup.py @@ -9,7 +9,9 @@ 4. LR override computation """ +import logging import math +from unittest.mock import patch import pytest import torch @@ -518,6 +520,108 @@ def test_non_adam_does_not_set_eps_override(self): assert override['min_lr'] == pytest.approx(1e-5 / width_mult) assert 'eps' not in override + @pytest.mark.parametrize('optimizer_type', ['muon', 'dist_muon']) + def test_muon_excludes_muon_managed_matrices_from_mup_overrides(self, optimizer_type): + """Muon-managed 2D params should use Muon scaling only, not MuP LR overrides.""" + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, muon_scale_mode='unit_rms_norm') + width_mult = 4.0 + + overrides = get_mup_config_overrides( + optimizer_config, width_mult, optimizer_type=optimizer_type + ) + + muon_managed_param = torch.nn.Parameter(torch.zeros(10, 10)) + muon_managed_param.is_embedding_or_output_parameter = False + output_param = torch.nn.Parameter(torch.zeros(10, 10)) + output_param.is_embedding_or_output_parameter = True + bias_param = torch.nn.Parameter(torch.zeros(10)) + + muon_managed_matches = [ + override + for param_key, override in overrides.items() + if param_key.matches( + muon_managed_param, 'decoder.layers.0.self_attention.linear_proj.weight' + ) + ] + output_matches = [ + override + for param_key, override in overrides.items() + if param_key.matches(output_param, 'output_layer.weight') + ] + bias_matches = [ + override + for param_key, override in overrides.items() + if param_key.matches(bias_param, 'decoder.layers.0.self_attention.linear_proj.bias') + ] + + muon_managed_override = combine_param_group_overrides(muon_managed_matches) + output_override = combine_param_group_overrides(output_matches) + bias_override = combine_param_group_overrides(bias_matches) + + # Muon-managed matrix params are excluded from Adam-style MuP LR overrides. + assert 'max_lr' not in muon_managed_override + assert 'min_lr' not in muon_managed_override + assert 'eps' not in muon_managed_override + + # Output params remain in the MuP override path (handled by chained Adam optimizer). + assert output_override['max_lr'] == pytest.approx(1e-3 / width_mult) + assert output_override['min_lr'] == pytest.approx(1e-5 / width_mult) + assert 'eps' not in output_override + + # Vector-like params stay unscaled. + assert 'max_lr' not in bias_override + assert 'min_lr' not in bias_override + assert 'eps' not in bias_override + + @pytest.mark.parametrize('optimizer_type', ['muon', 'dist_muon']) + def test_muon_warns_for_spectral_scale_mode(self, optimizer_type): + """Muon+MuP should warn when scale mode is spectral.""" + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, muon_scale_mode='spectral') + width_mult = 4.0 + + with patch('megatron.core.optimizer.log_single_rank') as mock_warn: + overrides = get_mup_config_overrides( + optimizer_config, width_mult, optimizer_type=optimizer_type + ) + + assert len(overrides) == 1 + mock_warn.assert_called_once() + _, level, message = mock_warn.call_args[0] + assert level == logging.WARNING + assert "Both MuP and muon_scale_mode=spectral are enabled." in message + assert "--muon-scale-mode unit_rms_norm" in message + + @pytest.mark.parametrize('optimizer_type', ['muon', 'dist_muon']) + def test_muon_unit_rms_norm_mode_has_no_warning(self, optimizer_type): + """Muon+MuP should not warn when scale mode is unit_rms_norm.""" + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, muon_scale_mode='unit_rms_norm') + width_mult = 4.0 + + with patch('megatron.core.optimizer.log_single_rank') as mock_warn: + overrides = get_mup_config_overrides( + optimizer_config, width_mult, optimizer_type=optimizer_type + ) + + assert len(overrides) == 1 + mock_warn.assert_not_called() + + @pytest.mark.parametrize('optimizer_type', ['muon', 'dist_muon']) + def test_muon_warns_for_spectral_mode_at_unity_width_mult(self, optimizer_type): + """Muon+MuP warning should still fire when width_mult==1.0.""" + optimizer_config = OptimizerConfig(lr=1e-3, min_lr=1e-5, muon_scale_mode='spectral') + width_mult = 1.0 + + with patch('megatron.core.optimizer.log_single_rank') as mock_warn: + overrides = get_mup_config_overrides( + optimizer_config, width_mult, optimizer_type=optimizer_type + ) + + assert len(overrides) == 0 + mock_warn.assert_called_once() + _, level, message = mock_warn.call_args[0] + assert level == logging.WARNING + assert "Both MuP and muon_scale_mode=spectral are enabled." in message + class TestMuPMTPLossScaling: """Tests for MuP scaling integration with MTP loss processing."""