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
24 changes: 24 additions & 0 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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:
Expand All @@ -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 {}
Expand All @@ -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:
Expand All @@ -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
Expand Down
4 changes: 3 additions & 1 deletion megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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')
Expand Down
104 changes: 104 additions & 0 deletions tests/unit_tests/transformer/test_mup.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
4. LR override computation
"""

import logging
import math
from unittest.mock import patch

import pytest
import torch
Expand Down Expand Up @@ -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."""
Expand Down
Loading