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
58 changes: 31 additions & 27 deletions src/maxtext/optimizers/muon.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,7 @@ def _base_ns_iterator(
}


class MuonDimensionNumbers(NamedTuple):
class ShardedMuonDimensionNumbers(NamedTuple):
"""Specification for which weight axes participate in matrix projection.

Muon defines an orthogonalization for 2D matrix weights for matrix-vector
Expand Down Expand Up @@ -243,7 +243,10 @@ class MuonDimensionNumbers(NamedTuple):
sharding: jax.sharding.NamedSharding | None = None


def _normalize_axes(x: jax.Array, dim_nums: MuonDimensionNumbers) -> tuple[tuple[int, ...], tuple[int, ...]]:
MuonDimensionNumbers = ShardedMuonDimensionNumbers


def _normalize_axes(x: jax.Array, dim_nums: ShardedMuonDimensionNumbers) -> tuple[tuple[int, ...], tuple[int, ...]]:
"""Normalize axes in dimension numbers to non-negative int tuples."""
reduction_axes = (
(dim_nums.reduction_axis,) if isinstance(dim_nums.reduction_axis, int) else tuple(dim_nums.reduction_axis)
Expand All @@ -261,7 +264,7 @@ def _normalize_axes(x: jax.Array, dim_nums: MuonDimensionNumbers) -> tuple[tuple
return reduction_axes, output_axes


WeightDimNumOrFn = MuonDimensionNumbers | optax.Params | Callable[[optax.Params], optax.Params | None]
WeightDimNumOrFn = ShardedMuonDimensionNumbers | optax.Params | Callable[[optax.Params], optax.Params | None]


class MuonState(NamedTuple):
Expand Down Expand Up @@ -308,8 +311,8 @@ def init_fn(params: optax.Params) -> MuonState:
raise ValueError(f"ns_coeffs must have shape (3,) or (n, 3), got {ns_coeffs_.shape}")
if ns_coeffs_.ndim == 2:
# pyrefly: ignore[unsupported-operation]
if not ns_coeffs_.shape[0] <= ns_steps:
raise ValueError(f"Not enough coeffs to perform {ns_steps} steps")
if ns_coeffs_.shape[0] < ns_steps:
raise ValueError(f"Not enough coeffs to perform {ns_steps} steps, got {ns_coeffs_.shape[0]}")
# pyrefly: ignore[unsupported-operation]
ns_coeffs_ = ns_coeffs_[-ns_steps:]

Expand Down Expand Up @@ -341,13 +344,13 @@ def update_fn(

def orthogonalize_leaf(
leaf: jax.Array | optax.MaskedNode,
dim_nums: MuonDimensionNumbers | optax.MaskedNode | None,
dim_nums: ShardedMuonDimensionNumbers | optax.MaskedNode | None,
) -> jax.Array | optax.MaskedNode:
"""Orthogonalize a single leaf tensor."""
if isinstance(leaf, optax.MaskedNode) or isinstance(dim_nums, optax.MaskedNode):
return optax.MaskedNode()
if dim_nums is None:
dim_nums = MuonDimensionNumbers()
dim_nums = ShardedMuonDimensionNumbers()
return orthogonalize(
x=leaf,
ns_coeffs=state.ns_coeffs,
Expand All @@ -361,12 +364,12 @@ def orthogonalize_leaf(
if callable(weight_dimension_numbers):
resolved_dim_nums = weight_dimension_numbers(updates)
elif weight_dimension_numbers is None:
resolved_dim_nums = jax.tree.map(lambda _: MuonDimensionNumbers(), updates)
resolved_dim_nums = jax.tree.map(lambda _: ShardedMuonDimensionNumbers(), updates)
else:
resolved_dim_nums = weight_dimension_numbers

def is_leaf(x):
return x is None or isinstance(x, (MuonDimensionNumbers, optax.MaskedNode))
return x is None or isinstance(x, (ShardedMuonDimensionNumbers, optax.MaskedNode))

updates = jax.tree.map(
orthogonalize_leaf,
Expand Down Expand Up @@ -396,7 +399,7 @@ def scale_by_dual_norm(x: jax.Array, y: jax.Array) -> jax.Array:

def get_reshape_fns(
x: jax.Array,
dim_nums: MuonDimensionNumbers,
dim_nums: ShardedMuonDimensionNumbers,
use_all_to_all: bool = True,
) -> tuple[
reshape_utils.ReshapeFn,
Expand All @@ -420,7 +423,7 @@ def orthogonalize(
ns_steps: jax.typing.ArrayLike,
precond_fn: Callable[[jax.Array], jax.Array],
ns_step_fn: Callable[..., jax.Array],
dim_nums: MuonDimensionNumbers,
dim_nums: ShardedMuonDimensionNumbers,
use_all_to_all: bool = True,
) -> jax.Array:
"""Apply Newton-Schulz iterations to a single leaf tensor."""
Expand Down Expand Up @@ -469,14 +472,14 @@ def _scan_body(carry, coeffs_step):
return unreshape_fn(x_flat_orthogonalized)


def _get_shape_products(x: jax.Array, dim_nums: MuonDimensionNumbers) -> tuple[float, float]:
def _get_shape_products(x: jax.Array, dim_nums: ShardedMuonDimensionNumbers) -> tuple[float, float]:
reduction_axes, output_axes = _normalize_axes(x, dim_nums)
fan_in = math.prod(x.shape[ax] for ax in reduction_axes)
fan_out = math.prod(x.shape[ax] for ax in output_axes)
return fan_in, fan_out


def _scale_update_for_width_transfer(update: jax.Array, dim_nums: MuonDimensionNumbers):
def _scale_update_for_width_transfer(update: jax.Array, dim_nums: ShardedMuonDimensionNumbers):
"""Apply width scaling from <https://github.com/KellerJordan/Muon>."""
fan_in, fan_out = _get_shape_products(update, dim_nums)
scale = jnp.sqrt(jnp.maximum(1, fan_out / fan_in))
Expand All @@ -485,7 +488,7 @@ def _scale_update_for_width_transfer(update: jax.Array, dim_nums: MuonDimensionN

def _scale_update_for_consistent_rms(
update: jax.Array,
dim_nums: MuonDimensionNumbers,
dim_nums: ShardedMuonDimensionNumbers,
consistent_rms: jax.typing.ArrayLike,
):
"""Apply consistent RMS scaling from <https://arxiv.org/abs/2502.16982>."""
Expand All @@ -502,7 +505,7 @@ def scale_by_shape(

Args:
weight_dimension_numbers: An optional tree with the same structure as the
params of `MuonDimensionNumbers`s, specifying how to reshape the
params of `ShardedMuonDimensionNumbers`s, specifying how to reshape the
parameters before and after the orthogonalization OR a callable returning
such a tree. None implies that all parameters are 2D matrices.
consistent_rms: An optional float to activate consistent RMS scaling. If
Expand Down Expand Up @@ -531,11 +534,11 @@ def scaling_fn(update, dim_nums):
if isinstance(update, optax.MaskedNode) or isinstance(dim_nums, optax.MaskedNode):
return optax.MaskedNode()
if dim_nums is None:
dim_nums = MuonDimensionNumbers()
dim_nums = ShardedMuonDimensionNumbers()
return base_scaling_fn(update, dim_nums)

def is_leaf(x):
return x is None or isinstance(x, (MuonDimensionNumbers, optax.MaskedNode))
return x is None or isinstance(x, (ShardedMuonDimensionNumbers, optax.MaskedNode))

scaled_updates = jax.tree.map(
scaling_fn,
Expand Down Expand Up @@ -626,13 +629,14 @@ def muon(
adam_weight_decay: Weight decay factor for Adam.
adam_learning_rate: Auxiliary learning rate for the Adam optimizer. If
`None`, the learning rate for Adam defaults to the same as Muon.
muon_weight_dimension_numbers: An optional tree of `MuonDimensionNumbers`s,
specifying how to reshape the parameters for orthogonalization otherwise
muon parameters are assumed to be 2D matrices. A `None` value indicates
that the parameter is not a muon parameter and will be optimized with
Adam. A callable takes as input the params and returns a possibly masked
pytree of specs, similar to `weight_decay_mask`. If not provided, muon is
applied to all 2D parameters.
muon_weight_dimension_numbers: An optional tree of
`ShardedMuonDimensionNumbers`s, specifying how to reshape the parameters
for orthogonalization otherwise muon parameters are assumed to be 2D
matrices. A `None` value indicates that the parameter is not a muon
parameter and will be optimized with Adam. A callable takes as input the
params and returns a possibly masked pytree of specs, similar to
`weight_decay_mask`. If not provided, muon is applied to all 2D
parameters.
consistent_rms: An optional float to activate consistent RMS scaling. Scales
updates by `sqrt(max(fan_in, fan_out)) * consistent_rms` to make root mean
square (RMS) shape-independent, like AdamW. `0.2` is recommended to match
Expand Down Expand Up @@ -686,7 +690,7 @@ def muon(
def param_labels(params):
return jax.tree.map(lambda x: "muon" if x.ndim == 2 else "adam", params)

muon_weight_dimension_numbers = MuonDimensionNumbers()
muon_weight_dimension_numbers = ShardedMuonDimensionNumbers()
else:

def param_labels(params):
Expand All @@ -704,7 +708,7 @@ def populate_subtree_(dim_num, x):
populate_subtree_,
dim_nums,
params,
is_leaf=lambda x: x is None or isinstance(x, MuonDimensionNumbers),
is_leaf=lambda x: x is None or isinstance(x, ShardedMuonDimensionNumbers),
)

# We need to normalize the dimension numbers because they have to match the
Expand All @@ -720,7 +724,7 @@ def muon_weight_dim_nums_fn(params):
mask = jax.tree.map(lambda label: label == "muon", param_labels(params))

def is_leaf(x):
return x is None or isinstance(x, (MuonDimensionNumbers, optax.MaskedNode))
return x is None or isinstance(x, (ShardedMuonDimensionNumbers, optax.MaskedNode))

def populate_subtree_(dim_nums, submask):
return jax.tree.map(lambda m: dim_nums if m else optax.MaskedNode(), submask)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,21 +14,56 @@

"""Integration tests for sharded Muon optimizer."""

import os
import subprocess
import sys
import unittest

from absl.testing import parameterized
import chex
import jax
import jax.numpy as jnp

jax.config.update("jax_num_cpu_devices", 8)

from maxtext.optimizers import muon as _muon
from maxtext.optimizers import reshape_utils
import numpy as np
from optax.contrib import _muon as optax_muon
import pytest

from maxtext.optimizers import muon as _muon
from maxtext.optimizers import reshape_utils
_REQUIRED_CPU_DEVICES = 8

pytestmark = pytest.mark.cpu_only


@pytest.mark.cpu_only
def test_sharded_muon_on_cpu_mesh():
"""Runs the Muon sharding tests in a subprocess with forced 8 CPU devices."""
env = os.environ.copy()
env["XLA_FLAGS"] = env.get("XLA_FLAGS", "") + f" --xla_force_host_platform_device_count={_REQUIRED_CPU_DEVICES}"
env["JAX_PLATFORMS"] = "cpu"
result = subprocess.run(
[sys.executable, __file__],
env=env,
capture_output=True,
text=True,
check=False,
)
assert result.returncode == 0, f"stdout:\n{result.stdout}\nstderr:\n{result.stderr}"
assert "MUON_SHARDING_TESTS_PASSED" in result.stdout


class MuonShardingTest(parameterized.TestCase):
__test__ = False

def setUp(self):
super().setUp()
try:
cpu_devices = jax.devices("cpu")
except (RuntimeError, ValueError):
self.skipTest(f"MuonShardingTest requires {_REQUIRED_CPU_DEVICES} CPU devices; CPU" " backend unavailable.")
if len(cpu_devices) < _REQUIRED_CPU_DEVICES:
self.skipTest(
f"MuonShardingTest requires {_REQUIRED_CPU_DEVICES} CPU devices;" " executed via test_sharded_muon_on_cpu_mesh."
)

@parameterized.parameters(jax.sharding.AxisType.Explicit, jax.sharding.AxisType.Auto)
def test_sharded_scale_by_muon_matches_optax_contrib(self, axis_type):
Expand All @@ -49,7 +84,7 @@ def test_sharded_scale_by_muon_matches_optax_contrib(self, axis_type):
params = {"w": w}
grads = {"w": g}

local_dim_nums = {"w": _muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)}
local_dim_nums = {"w": _muon.ShardedMuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)}
optax_dim_nums = {"w": optax_muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2)}

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -90,7 +125,7 @@ def test_sharded_2d_tensor_reduction_and_output_axes(self, axis_type):
grads = {"w": g}

local_transform = _muon.scale_by_muon(
weight_dimension_numbers={"w": _muon.MuonDimensionNumbers(sharding=sharding)},
weight_dimension_numbers={"w": _muon.ShardedMuonDimensionNumbers(sharding=sharding)},
)
local_state = local_transform.init(params)
local_updates, _ = local_transform.update(grads, local_state, params)
Expand Down Expand Up @@ -121,7 +156,7 @@ def test_sharded_3d_unsharded_batch_with_sharded_matrix_axes(self, axis_type):
params = {"w": w}
grads = {"w": g}

dim_num = _muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
dim_num = _muon.ShardedMuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
optax_dim_num = optax_muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2)

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -161,7 +196,7 @@ def test_sharded_3d_unsharded_batch_with_sharded_matrix_axes_without_all_to_all(
params = {"w": w}
grads = {"w": g}

dim_num = _muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
dim_num = _muon.ShardedMuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
optax_dim_num = optax_muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2)

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -202,7 +237,7 @@ def test_sharded_3d_batch_axis(self, axis_type):
params = {"w": w}
grads = {"w": g}

dim_num = _muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
dim_num = _muon.ShardedMuonDimensionNumbers(reduction_axis=1, output_axis=2, sharding=sharding)
optax_dim_num = optax_muon.MuonDimensionNumbers(reduction_axis=1, output_axis=2)

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -244,7 +279,7 @@ def test_sharded_4d_mixed_batch_axes(self, axis_type):
params = {"w": w}
grads = {"w": g}

dim_num = _muon.MuonDimensionNumbers(reduction_axis=2, output_axis=3, sharding=sharding)
dim_num = _muon.ShardedMuonDimensionNumbers(reduction_axis=2, output_axis=3, sharding=sharding)
optax_dim_num = optax_muon.MuonDimensionNumbers(reduction_axis=2, output_axis=3)

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -284,7 +319,7 @@ def test_sharded_4d_all_sharded_batch_axes(self, axis_type):
params = {"w": w}
grads = {"w": g}

dim_num = _muon.MuonDimensionNumbers(reduction_axis=2, output_axis=3, sharding=sharding)
dim_num = _muon.ShardedMuonDimensionNumbers(reduction_axis=2, output_axis=3, sharding=sharding)
optax_dim_num = optax_muon.MuonDimensionNumbers(reduction_axis=2, output_axis=3)

local_transform = _muon.scale_by_muon(
Expand Down Expand Up @@ -318,7 +353,7 @@ def test_mixed_mesh_axis_types_raises_error(self):
params = {"w": w}

transform = _muon.scale_by_muon(
weight_dimension_numbers={"w": _muon.MuonDimensionNumbers(sharding=sharding)},
weight_dimension_numbers={"w": _muon.ShardedMuonDimensionNumbers(sharding=sharding)},
)
state = transform.init(params)
with self.assertRaisesRegex(ValueError, "Mixed mesh axis types"):
Expand Down Expand Up @@ -572,7 +607,7 @@ def test_sharded_all_to_all_with_non_zero_padding_jitted(self, axis_type):
grads = {"w": w}
params = {"w": w}

dim_nums = {"w": _muon.MuonDimensionNumbers(reduction_axis=-2, output_axis=-1, sharding=sharding)}
dim_nums = {"w": _muon.ShardedMuonDimensionNumbers(reduction_axis=-2, output_axis=-1, sharding=sharding)}
opt = _muon.muon(learning_rate=1e-3, muon_weight_dimension_numbers=dim_nums)
state = opt.init(params)

Expand All @@ -587,4 +622,9 @@ def step(p, s, g):


if __name__ == "__main__":
parameterized.absltest.main()
suite = unittest.defaultTestLoader.loadTestsFromTestCase(MuonShardingTest)
runner = unittest.TextTestRunner(verbosity=2)
res = runner.run(suite)
if not res.wasSuccessful():
sys.exit(1)
print("MUON_SHARDING_TESTS_PASSED")
Loading
Loading