diff --git a/.github/workflows/integration_test_8gpu_features.yaml b/.github/workflows/integration_test_8gpu_features.yaml index 2a0abfcc0a..38700df521 100644 --- a/.github/workflows/integration_test_8gpu_features.yaml +++ b/.github/workflows/integration_test_8gpu_features.yaml @@ -89,6 +89,10 @@ jobs: end=$(date +%s) echo "pip install torchao took $((end - start)) seconds" + # Exercise distributed optimizer tests that are skipped in CPU unit-test CI. + python -m pytest tests/unit_tests/test_distributed_muon.py \ + --durations=20 -vv + sudo mkdir -p "$RUNNER_TEMP/artifacts-to-be-uploaded" sudo chown -R $(id -u):$(id -g) "$RUNNER_TEMP/artifacts-to-be-uploaded" diff --git a/tests/unit_tests/test_distributed_muon.py b/tests/unit_tests/test_distributed_muon.py new file mode 100644 index 0000000000..d94e265d45 --- /dev/null +++ b/tests/unit_tests/test_distributed_muon.py @@ -0,0 +1,1462 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import unittest +from unittest.mock import patch + +import torch +import torch.distributed as dist +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh +from torch.distributed.tensor import distribute_tensor, DTensor, Replicate, Shard +from torch.distributed.tensor.placement_types import _StridedShard +from torch.testing._internal.distributed._tensor.common_dtensor import ( + DTensorTestBase, + with_comms, +) +from torchtitan.components.checkpoint_utils import ( + get_flat_optim_state_dict, + init_optim_state, + load_flat_optim_state_dict, +) +from torchtitan.components.distributed_optimizers.flex_optimizer_reshard import ( + BucketConfig, + BucketSpec, +) +from torchtitan.components.distributed_optimizers.muon import ( + _adjust_learning_rate, + _has_dim0_sharded_storage, + _has_replicated_storage, + DistributedMuon, + Owned, +) +from torchtitan.components.distributed_optimizers.muon_parameter_prep import ( + BatchedMatrixComputeView, + build_distributed_muon, + MuonComputeSharding, +) + + +# Allow a few BF16 quantization steps across different GEMM schedules. +_BATCHED_BF16_DIRECTION_ATOL = 2e-2 + + +def _assert_exact(actual: torch.Tensor, expected: torch.Tensor) -> None: + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def _assert_batched_muon_update_close( + actual_before: torch.Tensor, + actual_after: torch.Tensor, + expected_before: torch.Tensor, + expected_after: torch.Tensor, + *, + lr: float, + weight_decay: float, + compute_matrix_shape: torch.Size, +) -> None: + decay = 1 - lr * weight_decay + adjusted_lr = _adjust_learning_rate(lr, None, compute_matrix_shape) + actual_update = (actual_before * decay - actual_after) / adjusted_lr + expected_update = (expected_before * decay - expected_after) / adjusted_lr + torch.testing.assert_close( + actual_update, + expected_update, + rtol=0, + atol=_BATCHED_BF16_DIRECTION_ATOL, + ) + + +class TestDistributedMuonStoragePolicy(unittest.TestCase): + def test_rejects_placement_subclasses(self): + class UnsupportedShard(Shard): + pass + + class UnsupportedReplicate(Replicate): + pass + + class FakeParameter: + ndim = 3 + + parameter = FakeParameter() + parameter.placements = (UnsupportedShard(0),) + self.assertFalse(_has_dim0_sharded_storage(parameter)) + parameter.placements = (UnsupportedReplicate(),) + self.assertFalse(_has_replicated_storage(parameter)) + + def test_rejects_nan_hyperparameters(self): + valid_group = { + "lr": 0.03, + "weight_decay": 0.2, + "momentum": 0.8, + "nesterov": True, + "ns_coefficients": (3.4445, -4.7750, 2.0315), + "eps": 1e-7, + "ns_steps": 2, + "adjust_lr_fn": None, + "fused": False, + "foreach": False, + } + optimizer = object.__new__(DistributedMuon) + for name in ("lr", "weight_decay", "momentum", "eps"): + with self.subTest(name=name): + optimizer.param_groups = [{**valid_group, name: float("nan")}] + with self.assertRaisesRegex( + ValueError, "unsupported DistributedMuon group 0" + ): + optimizer._validate_groups() + + +class _DistributedMuonTestBase(DTensorTestBase): + @property + def world_size(self): + return 2 + + @property + def device_type(self): + return "cuda" + + @property + def mesh(self): + if not hasattr(self, "_mesh"): + self._mesh = init_device_mesh( + self.device_type, + (self.world_size,), + mesh_dim_names=("dp_shard",), + ) + return self._mesh + + @property + def device(self): + return torch.device("cuda", self.rank) + + def _parameter(self, value: torch.Tensor) -> torch.nn.Parameter: + return torch.nn.Parameter( + distribute_tensor(value.clone(), self.mesh, (Shard(0),)) + ) + + def _optimizer( + self, + redistributed: torch.nn.Parameter, + local_blocks: torch.nn.Parameter, + *, + local_num_matrices: int = 2, + owner_rank: int = 1, + ) -> DistributedMuon: + return build_distributed_muon( + [ + { + "params": [redistributed], + "param_names": ["layers.0.redistributed"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + }, + { + "params": [local_blocks], + "param_names": ["layers.0.local_blocks"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=local_num_matrices, + matrices_flattened_into_dim=0, + ), + placement=Shard(0), + ), + }, + ], + bucket_configs=[ + BucketConfig( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.redistributed": owner_rank}, + mesh_axes=("dp_shard",), + name="layers.0", + ) + ], + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + + def _set_grads( + self, + redistributed: torch.nn.Parameter, + local_blocks: torch.nn.Parameter, + redistributed_grad: torch.Tensor, + local_blocks_grad: torch.Tensor, + ) -> None: + redistributed.grad = distribute_tensor( + redistributed_grad.clone(), self.mesh, (Shard(0),) + ) + local_blocks.grad = distribute_tensor( + local_blocks_grad.clone(), self.mesh, (Shard(0),) + ) + + def _assert_matches_reference( + self, + optimizer: DistributedMuon, + redistributed: torch.nn.Parameter, + local_blocks: torch.nn.Parameter, + reference_optimizer: torch.optim.Muon, + reference_redistributed: torch.nn.Parameter, + reference_local_blocks: tuple[torch.nn.Parameter, torch.nn.Parameter], + *, + local_blocks_before: torch.Tensor, + reference_local_blocks_before: tuple[torch.Tensor, torch.Tensor], + ) -> None: + rank = self.mesh.get_local_rank() + expected_redistributed = reference_redistributed.detach().chunk( + self.world_size, dim=0 + )[rank] + expected_local_blocks = reference_local_blocks[rank].detach() + _assert_exact(redistributed.to_local(), expected_redistributed) + _assert_batched_muon_update_close( + local_blocks_before, + local_blocks.to_local(), + reference_local_blocks_before[rank], + expected_local_blocks, + lr=0.03, + weight_decay=0.2, + compute_matrix_shape=expected_local_blocks.shape, + ) + + for param in (redistributed, local_blocks): + self.assertIsInstance(param, DTensor) + self.assertEqual(param.placements, (Shard(0),)) + + redistributed_momentum = optimizer.state[redistributed]["momentum_buffer"] + self.assertIsInstance(redistributed_momentum, DTensor) + self.assertEqual(redistributed_momentum.placements, (Shard(0),)) + expected_redistributed_momentum = ( + reference_optimizer.state[reference_redistributed]["momentum_buffer"] + .detach() + .chunk(self.world_size, dim=0)[rank] + ) + _assert_exact( + redistributed_momentum.to_local(), expected_redistributed_momentum + ) + + local_blocks_momentum = optimizer.state[local_blocks]["momentum_buffer"] + self.assertIsInstance(local_blocks_momentum, DTensor) + self.assertEqual(local_blocks_momentum.placements, (Shard(0),)) + expected_local_blocks_momentum = reference_optimizer.state[ + reference_local_blocks[rank] + ]["momentum_buffer"].detach() + _assert_exact(local_blocks_momentum.to_local(), expected_local_blocks_momentum) + + +@unittest.skipUnless(torch.cuda.device_count() >= 1, "requires one CUDA device") +class TestDistributedMuonSingleRank(_DistributedMuonTestBase): + @property + def world_size(self): + return 1 + + @with_comms + def test_owned_compute_accepts_static_owner_for_replicated_storage(self): + value = torch.arange(12, device=self.device).reshape(4, 3).float() + parameter = torch.nn.Parameter( + distribute_tensor(value.clone(), self.mesh, (Replicate(),)) + ) + optimizer = build_distributed_muon( + [ + { + "params": [parameter], + "param_names": ["layers.0.weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.weight": 0}, + mesh=self.mesh, + ) + ], + ns_steps=1, + ) + parameter.grad = distribute_tensor( + torch.ones_like(value), self.mesh, (Replicate(),) + ) + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + + collective.assert_not_called() + + +@unittest.skipUnless(torch.cuda.device_count() >= 2, "requires two CUDA devices") +class TestDistributedMuon(_DistributedMuonTestBase): + @with_comms + def test_constructor_strictly_validates_strided_storage_shards(self): + mesh = init_device_mesh(self.device_type, (self.world_size, 1)) + + def make_parameter(value, dim): + placements = ( + _StridedShard(dim, split_factor=self.world_size), + Shard(dim), + ) + parameter = torch.nn.Parameter(distribute_tensor(value, mesh, placements)) + self.assertEqual(parameter.placements, placements) + return parameter + + def build(parameter, name, compute_placement, owner_rank=None): + fqn = f"layers.0.{name}" + owners = {} if owner_rank is None else {fqn: owner_rank} + return build_distributed_muon( + [ + { + "params": [parameter], + "param_names": [fqn], + "compute_sharding": MuonComputeSharding( + placement=compute_placement + ), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn=owners, + mesh=self.mesh, + ) + ], + ) + + local_blocks = make_parameter( + torch.arange(24, device=self.device).reshape(4, 2, 3).float(), 0 + ) + optimizer = build(local_blocks, "local_blocks", Shard(0)) + self.assertIs(optimizer._plans[0].local_items[0].param, local_blocks) + local_blocks.grad = distribute_tensor( + torch.ones(4, 2, 3, device=self.device), + mesh, + local_blocks.placements, + ) + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single" + ) as collective: + optimizer.step() + collective.assert_not_called() + self.assertEqual( + optimizer.state[local_blocks]["momentum_buffer"].placements, + local_blocks.placements, + ) + local_blocks._local_tensor = local_blocks.to_local().clone() + with self.assertRaisesRegex(RuntimeError, "local storage changed"): + optimizer.step() + + backing = torch.arange(24, device=self.device).float() + first_region = backing[:12].view(2, 2, 3) + second_region = backing[12:].view(2, 2, 3) + shared_storage = torch.nn.Parameter( + DTensor.from_local( + first_region, + self.mesh, + (Shard(0),), + run_check=False, + ) + ) + shared_optimizer = build(shared_storage, "shared_storage", Shard(0)) + shared_storage.grad = DTensor.from_local( + torch.ones_like(first_region), + self.mesh, + (Shard(0),), + run_check=False, + ) + self.assertEqual( + shared_storage.to_local().untyped_storage().data_ptr(), + second_region.untyped_storage().data_ptr(), + ) + self.assertNotEqual( + shared_storage.to_local().data_ptr(), second_region.data_ptr() + ) + shared_storage._local_tensor = second_region + with self.assertRaisesRegex(RuntimeError, "local storage changed"): + shared_optimizer.step() + + dim1_sharded = make_parameter( + torch.arange(24, device=self.device).reshape(2, 4, 3).float(), 1 + ) + with self.assertRaisesRegex(ValueError, "storage-to-compute layout"): + build(dim1_sharded, "dim1_sharded", Shard(0)) + + owned = make_parameter( + torch.arange(12, device=self.device).reshape(4, 3).float(), 0 + ) + with self.assertRaisesRegex(ValueError, "storage-to-compute layout"): + build(owned, "owned", Owned(), owner_rank=0) + + @with_comms + def test_constructor_requires_exact_bucket_coverage_without_creating_state(self): + redistributed = self._parameter( + torch.arange(12, device=self.device).reshape(4, 3).float() + ) + local_blocks = self._parameter( + torch.arange(12, 24, device=self.device).reshape(4, 3).float() + ) + + optimizer = self._optimizer(redistributed, local_blocks) + self.assertEqual(len(optimizer.state), 0) + redistribution = optimizer._plans[0].redistribution_plans[0] + self.assertTrue( + all( + route.destination_participants == (1,) + for route in redistribution.storage_to_compute_routes + ) + ) + redistributed_before = redistributed.to_local().clone() + redistributed.grad = distribute_tensor( + torch.ones(4, 3, device=self.device), self.mesh, (Shard(0),) + ) + with patch("torch.distributed.all_reduce") as validation_collective: + with self.assertRaisesRegex(RuntimeError, "layers.0.local_blocks"): + optimizer.step() + validation_collective.assert_not_called() + self.assertEqual(len(optimizer.state), 0) + _assert_exact(redistributed.to_local(), redistributed_before) + redistributed.grad = None + + with self.assertRaisesRegex(ValueError, "must match one bucket"): + build_distributed_muon( + [ + { + "params": [redistributed, local_blocks], + "param_names": [ + "layers.0.redistributed", + "layers.0.local_blocks", + ], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=2, matrices_flattened_into_dim=0 + ), + placement=Shard(0), + ), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("*.redistributed",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ) + ], + ) + + with self.assertRaisesRegex(ValueError, "must match one bucket"): + build_distributed_muon( + [ + { + "params": [redistributed, local_blocks], + "param_names": [ + "layers.0.redistributed", + "layers.0.local_blocks", + ], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=2, matrices_flattened_into_dim=0 + ), + placement=Shard(0), + ), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ), + BucketSpec( + patterns=("*.local_blocks",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ), + ], + ) + + local_blocks_before = local_blocks.to_local().clone() + init_optim_state(optimizer) + self.assertEqual(len(optimizer.state), 2) + _assert_exact(redistributed.to_local(), redistributed_before) + _assert_exact(local_blocks.to_local(), local_blocks_before) + flat_state = get_flat_optim_state_dict(optimizer) + self.assertIn("state.layers.0.redistributed.momentum_buffer", flat_state) + self.assertIn("state.layers.0.local_blocks.momentum_buffer", flat_state) + + @with_comms + def test_flat_state_dict_loads_after_group_membership_changes(self): + values = { + "layers.0.a": torch.arange(12, device=self.device).reshape(4, 3).float(), + "layers.0.b": torch.arange(12, 24, device=self.device) + .reshape(4, 3) + .float(), + } + + def build(names, *, storage_is_compute_ready=False): + parameters = [ + torch.nn.Parameter( + distribute_tensor(values[name].clone(), self.mesh, (Shard(0),)) + ) + for name in names + ] + compute_sharding = ( + MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView(num_matrices=2), + placement=Shard(0), + ) + if storage_is_compute_ready + else MuonComputeSharding(placement=Owned()) + ) + optimizer = build_distributed_muon( + [ + { + "params": parameters, + "param_names": names, + "compute_sharding": compute_sharding, + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn=( + {} if storage_is_compute_ready else dict.fromkeys(names, 0) + ), + mesh=self.mesh, + ) + ], + ns_steps=2, + ) + return parameters, optimizer + + source_parameters, source_optimizer = build(("layers.0.a", "layers.0.b")) + for name, parameter in zip( + ("layers.0.a", "layers.0.b"), source_parameters, strict=True + ): + parameter.grad = distribute_tensor( + torch.ones_like(values[name]), + self.mesh, + (Shard(0),), + ) + source_optimizer.step() + flat_state_dict = get_flat_optim_state_dict(source_optimizer) + self.assertIn( + "state.layers.0.a._distributed_muon_layout_fingerprint", + flat_state_dict, + ) + + target_parameters, target_optimizer = build(("layers.0.a",)) + init_optim_state(target_optimizer) + load_flat_optim_state_dict(target_optimizer, flat_state_dict) + _assert_exact( + target_optimizer.state[target_parameters[0]]["momentum_buffer"].to_local(), + source_optimizer.state[source_parameters[0]]["momentum_buffer"].to_local(), + ) + + _, changed_layout_optimizer = build( + ("layers.0.a",), storage_is_compute_ready=True + ) + init_optim_state(changed_layout_optimizer) + with self.assertRaisesRegex(ValueError, "compute layout"): + load_flat_optim_state_dict( + changed_layout_optimizer, + flat_state_dict, + ) + + @with_comms + def test_constructor_rejects_storage_shards_that_split_matrices(self): + parameter = self._parameter( + torch.arange(36, device=self.device).reshape(12, 3).float() + ) + with self.assertRaisesRegex(ValueError, "not aligned"): + build_distributed_muon( + [ + { + "params": [parameter], + "param_names": ["layers.0.wq.weight"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=3, matrices_flattened_into_dim=0 + ), + placement=Shard(0), + ), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ) + ], + ) + + @with_comms + def test_constructor_requires_valid_owner_assignments(self): + first = self._parameter( + torch.arange(12, device=self.device).reshape(4, 3).float() + ) + second = self._parameter( + torch.arange(12, 24, device=self.device).reshape(4, 3).float() + ) + params = [ + { + "params": [first, second], + "param_names": ["layers.0.first", "layers.0.second"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ] + + with self.assertRaisesRegex(ValueError, "storage-to-compute layout"): + build_distributed_muon( + [ + { + "params": [first], + "param_names": ["layers.0.first"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=2 + ), + placement=Owned(), + ), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.first": 0}, + mesh=self.mesh, + ) + ], + ) + + with self.assertRaisesRegex(ValueError, "exactly cover"): + build_distributed_muon( + params, + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.first": 0}, + mesh=self.mesh, + ) + ], + ) + self.assertIn("compute_sharding", params[0]) + + with self.assertRaisesRegex(ValueError, "outside its process group"): + build_distributed_muon( + params, + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={ + "layers.0.first": 0, + "layers.0.second": self.world_size, + }, + mesh=self.mesh, + ) + ], + ) + + @with_comms + def test_constructor_rejects_cross_rank_plan_mismatch(self): + redistributed = self._parameter( + torch.arange(12, device=self.device).reshape(4, 3).float() + ) + with self.assertRaisesRegex(RuntimeError, "plans differ across ranks"): + build_distributed_muon( + [ + { + "params": [redistributed], + "param_names": ["layers.0.redistributed"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.redistributed": self.rank}, + mesh=self.mesh, + ) + ], + ) + + with self.assertRaisesRegex(RuntimeError, "plans differ across ranks"): + build_distributed_muon( + [ + { + "params": [redistributed], + "param_names": ["layers.0.redistributed"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.redistributed": 0}, + mesh=self.mesh, + ) + ], + lr=0.01 if self.rank == 0 else 0.02, + ) + + @with_comms + def test_constructor_accepts_uneven_storage_shards(self): + redistributed = self._parameter( + torch.arange(15, device=self.device).reshape(5, 3).float() + ) + optimizer = build_distributed_muon( + [ + { + "params": [redistributed], + "param_names": ["layers.0.redistributed"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.redistributed": 0}, + mesh=self.mesh, + ) + ], + ) + + schedule = optimizer._plans[0].storage_to_compute_schedule + self.assertEqual(schedule.input_buffer_numel, 9 if self.rank == 0 else 6) + + @with_comms + def test_shard1_owned_matches_plain_muon(self): + value = torch.arange(15, device=self.device).reshape(3, 5).float().div_(10) + placement = (Shard(1),) + parameter = torch.nn.Parameter( + distribute_tensor(value.clone(), self.mesh, placement) + ) + optimizer = build_distributed_muon( + [ + { + "params": [parameter], + "param_names": ["layers.0.weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.weight": 1}, + mesh=self.mesh, + ) + ], + lr=0.03, + momentum=0.8, + ns_steps=2, + ) + grad = value.flip((0, 1)).contiguous() + parameter.grad = distribute_tensor(grad, self.mesh, placement) + + reference = torch.nn.Parameter(value.clone()) + reference.grad = grad.clone() + reference_optimizer = torch.optim.Muon( + [reference], + lr=0.03, + momentum=0.8, + ns_steps=2, + ) + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard." + "dist.all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + reference_optimizer.step() + + self.assertEqual(collective.call_count, 2) + expected_parameter = distribute_tensor(reference.detach(), self.mesh, placement) + expected_momentum = distribute_tensor( + reference_optimizer.state[reference]["momentum_buffer"], + self.mesh, + placement, + ) + momentum = optimizer.state[parameter]["momentum_buffer"] + self.assertEqual(parameter.placements, placement) + self.assertEqual(momentum.placements, placement) + _assert_exact(parameter.to_local(), expected_parameter.to_local()) + _assert_exact( + momentum.to_local(), + expected_momentum.to_local(), + ) + + @with_comms + def test_replicated_storage_matches_plain_muon_without_redistribution(self): + values = [ + torch.arange(offset, offset + 12, device=self.device) + .reshape(4, 3) + .float() + .div_(10) + for offset in (1, 13) + ] + owned, batched = ( + torch.nn.Parameter( + distribute_tensor(value.clone(), self.mesh, (Replicate(),)) + ) + for value in values + ) + optimizer = build_distributed_muon( + [ + { + "params": [owned], + "param_names": ["layers.0.owned"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + }, + { + "params": [batched], + "param_names": ["layers.0.batched"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView(num_matrices=2), + placement=Shard(0), + ), + }, + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ) + ], + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + + grads = [value.flip(0).contiguous() for value in values] + for param, grad in zip((owned, batched), grads, strict=True): + param.grad = distribute_tensor(grad, self.mesh, (Replicate(),)) + + owned_reference = torch.nn.Parameter(values[0].clone()) + batched_references = [ + torch.nn.Parameter(matrix.clone()) + for matrix in values[1].view(2, 2, 3).unbind() + ] + references = [owned_reference, *batched_references] + reference_optimizer = torch.optim.Muon( + references, + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + batched_before = batched.to_local().clone() + reference_batched_before = torch.stack( + [reference.detach().clone() for reference in batched_references] + ).view(batched.shape) + owned_reference.grad = grads[0] + for reference, grad in zip( + batched_references, + grads[1].view(2, 2, 3).unbind(), + strict=True, + ): + reference.grad = grad + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + reference_optimizer.step() + + collective.assert_not_called() + reference_values = ( + owned_reference, + torch.stack(batched_references).view(batched.shape), + ) + reference_momenta = ( + reference_optimizer.state[owned_reference]["momentum_buffer"], + torch.stack( + [ + reference_optimizer.state[reference]["momentum_buffer"] + for reference in batched_references + ] + ).view(batched.shape), + ) + _assert_exact(owned.to_local(), reference_values[0]) + _assert_batched_muon_update_close( + batched_before, + batched.to_local(), + reference_batched_before, + reference_values[1], + lr=0.03, + weight_decay=0.2, + compute_matrix_shape=batched_references[0].shape, + ) + for param, reference_momentum in zip( + (owned, batched), + reference_momenta, + strict=True, + ): + self.assertEqual(param.placements, (Replicate(),)) + self.assertEqual(param.grad.placements, (Replicate(),)) + momentum = optimizer.state[param]["momentum_buffer"] + self.assertEqual(momentum.placements, (Replicate(),)) + _assert_exact( + momentum.to_local(), + reference_momentum, + ) + + @with_comms + def test_step_rejects_gradient_with_reordered_mesh(self): + value = torch.arange(12, device=self.device).reshape(4, 3).float() + parameter = self._parameter(value) + optimizer = build_distributed_muon( + [ + { + "params": [parameter], + "param_names": ["layers.0.weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.weight": 0}, + mesh=self.mesh, + ) + ], + ns_steps=1, + ) + reversed_mesh = DeviceMesh( + self.device_type, + torch.arange(self.world_size - 1, -1, -1), + mesh_dim_names=("dp_shard",), + ) + self.assertEqual( + tuple(dist.get_process_group_ranks(reversed_mesh.get_group())), + tuple(dist.get_process_group_ranks(self.mesh.get_group())), + ) + self.assertNotEqual( + tuple(reversed_mesh.mesh.flatten().tolist()), + tuple(self.mesh.mesh.flatten().tolist()), + ) + parameter.grad = distribute_tensor( + value.flip(0).contiguous(), reversed_mesh, (Shard(0),) + ) + parameter_before = parameter.to_local().clone() + + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard." + "dist.all_to_all_single" + ) as collective: + with self.assertRaisesRegex( + RuntimeError, "gradient storage layout changed" + ): + optimizer.step() + + collective.assert_not_called() + self.assertEqual(len(optimizer.state), 0) + _assert_exact(parameter.to_local(), parameter_before) + + @with_comms + def test_step_matches_plain_muon_and_continues_from_state_dict(self): + redistributed_value = ( + torch.arange(12, device=self.device).reshape(4, 3).float().div_(10).add_(1) + ) + local_blocks_value = ( + torch.arange(12, 24, device=self.device).reshape(4, 3).float().div_(10) + ) + redistributed = self._parameter(redistributed_value) + local_blocks = self._parameter(local_blocks_value) + optimizer = self._optimizer(redistributed, local_blocks) + self.assertEqual(len(optimizer.state), 0) + + reference_redistributed = torch.nn.Parameter(redistributed_value.clone()) + reference_local_blocks = tuple( + torch.nn.Parameter(block.clone()) + for block in local_blocks_value.chunk(self.world_size, dim=0) + ) + reference_optimizer = torch.optim.Muon( + [reference_redistributed, *reference_local_blocks], + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + local_blocks_before = local_blocks.to_local().clone() + reference_local_blocks_before = tuple( + reference.detach().clone() for reference in reference_local_blocks + ) + + first_redistributed_grad = ( + torch.arange(1, 13, device=self.device).reshape(4, 3).float().div_(17) + ) + first_local_blocks_grad = ( + torch.arange(13, 25, device=self.device).reshape(4, 3).float().div_(19) + ) + self._set_grads( + redistributed, + local_blocks, + first_redistributed_grad, + first_local_blocks_grad, + ) + reference_redistributed.grad = first_redistributed_grad.clone() + for parameter, grad in zip( + reference_local_blocks, + first_local_blocks_grad.chunk(self.world_size, dim=0), + ): + parameter.grad = grad.clone() + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + self.assertEqual(collective.call_count, 2) + reference_optimizer.step() + self._assert_matches_reference( + optimizer, + redistributed, + local_blocks, + reference_optimizer, + reference_redistributed, + reference_local_blocks, + local_blocks_before=local_blocks_before, + reference_local_blocks_before=reference_local_blocks_before, + ) + + local_blocks.grad = None + with self.assertRaisesRegex(RuntimeError, "layers.0.local_blocks"): + optimizer.step() + + state_dict = optimizer.state_dict() + flat_state_dict = get_flat_optim_state_dict(optimizer) + self.assertTrue( + all( + "compute_sharding" not in group and "_compute_placement" not in group + for group in state_dict["param_groups"] + ) + ) + resumed_redistributed_value = redistributed.full_tensor().detach() + resumed_local_blocks_value = local_blocks.full_tensor().detach() + resumed_redistributed = self._parameter(resumed_redistributed_value) + resumed_local_blocks = self._parameter(resumed_local_blocks_value) + changed_view_optimizer = self._optimizer( + self._parameter(resumed_redistributed_value), + self._parameter(resumed_local_blocks_value), + local_num_matrices=4, + ) + with self.assertRaisesRegex(ValueError, "compute layout"): + changed_view_optimizer.load_state_dict(state_dict) + + resumed_optimizer = self._optimizer( + resumed_redistributed, + resumed_local_blocks, + owner_rank=0, + ) + init_optim_state(resumed_optimizer) + load_flat_optim_state_dict(resumed_optimizer, flat_state_dict) + _assert_exact(resumed_redistributed.to_local(), redistributed.to_local()) + _assert_exact(resumed_local_blocks.to_local(), local_blocks.to_local()) + + second_redistributed_grad = first_redistributed_grad.flip(0).contiguous() + second_local_blocks_grad = first_local_blocks_grad.flip(0).contiguous() + local_blocks_before = resumed_local_blocks.to_local().clone() + reference_local_blocks_before = tuple( + reference.detach().clone() for reference in reference_local_blocks + ) + self._set_grads( + resumed_redistributed, + resumed_local_blocks, + second_redistributed_grad, + second_local_blocks_grad, + ) + reference_redistributed.grad = second_redistributed_grad.clone() + for parameter, grad in zip( + reference_local_blocks, + second_local_blocks_grad.chunk(self.world_size, dim=0), + ): + parameter.grad = grad.clone() + + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + resumed_optimizer.step() + self.assertEqual(collective.call_count, 2) + reference_optimizer.step() + self._assert_matches_reference( + resumed_optimizer, + resumed_redistributed, + resumed_local_blocks, + reference_optimizer, + reference_redistributed, + reference_local_blocks, + local_blocks_before=local_blocks_before, + reference_local_blocks_before=reference_local_blocks_before, + ) + + +@unittest.skipUnless(torch.cuda.device_count() >= 4, "requires four CUDA devices") +class TestDistributedMuonUnevenShards(_DistributedMuonTestBase): + @property + def world_size(self): + return 4 + + @with_comms + def test_all_ranks_reject_shards_that_split_matrices(self): + parameter = self._parameter( + torch.arange(45, device=self.device).reshape(15, 3).float() + ) + with patch.object(DistributedMuon, "__init__", return_value=None) as init: + with self.assertRaisesRegex(ValueError, "not aligned"): + build_distributed_muon( + [ + { + "params": [parameter], + "param_names": ["layers.0.wq.weight"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=5, + matrices_flattened_into_dim=0, + ), + placement=Shard(0), + ), + } + ], + bucket_spec=(), + ) + init.assert_not_called() + + +@unittest.skipUnless(torch.cuda.device_count() >= 4, "requires four CUDA devices") +class TestDistributedMuonBucketMeshes(_DistributedMuonTestBase): + @property + def world_size(self): + return 4 + + @property + def mesh(self): + if not hasattr(self, "_mesh"): + self._mesh = init_device_mesh( + self.device_type, + (2, 2), + mesh_dim_names=("fsdp", "tp"), + ) + return self._mesh + + @with_comms + def test_bucket_config_mixes_dense_and_expert_storage_meshes(self): + dense_mesh = init_device_mesh( + self.device_type, + (self.world_size,), + mesh_dim_names=("dp_shard",), + ) + expert_mesh = init_device_mesh( + self.device_type, + (2, 2), + mesh_dim_names=("efsdp", "ep"), + ) + dense_value = ( + torch.arange(24, device=self.device).reshape(8, 3).float().div_(10) + ) + expert_value = ( + torch.arange(48, device=self.device).reshape(8, 2, 3).float().div_(10) + ) + dense = torch.nn.Parameter( + distribute_tensor(dense_value.clone(), dense_mesh, (Shard(0),)) + ) + expert_placements = ( + _StridedShard(0, split_factor=expert_mesh["ep"].size()), + Shard(0), + ) + expert = torch.nn.Parameter( + distribute_tensor( + expert_value.clone(), + expert_mesh, + expert_placements, + ) + ) + optimizer = build_distributed_muon( + [ + { + "params": [dense], + "param_names": ["layers.0.dense.weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + }, + { + "params": [expert], + "param_names": ["layers.0.experts.weight"], + "compute_sharding": MuonComputeSharding(placement=Shard(0)), + }, + ], + bucket_configs=[ + BucketConfig( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.dense.weight": 1}, + mesh_axes=("dp_shard",), + name="layers.0", + ) + ], + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + + dense_grad = ( + torch.arange(1, 25, device=self.device).reshape(8, 3).float().div_(17) + ) + expert_grad = expert_value.flip((0, 1, 2)).contiguous().div_(19) + dense.grad = distribute_tensor(dense_grad.clone(), dense_mesh, (Shard(0),)) + expert.grad = distribute_tensor( + expert_grad.clone(), + expert_mesh, + expert_placements, + ) + + reference_dense = torch.nn.Parameter(dense_value.clone()) + reference_experts = tuple( + torch.nn.Parameter(matrix.clone()) for matrix in expert.to_local() + ) + reference_dense.grad = dense_grad.clone() + for parameter, grad in zip( + reference_experts, + expert.grad.to_local(), + strict=True, + ): + parameter.grad = grad.clone() + reference_optimizer = torch.optim.Muon( + [reference_dense, *reference_experts], + lr=0.03, + weight_decay=0.2, + momentum=0.8, + nesterov=True, + ns_steps=2, + ) + expert_before = expert.to_local().clone() + reference_experts_before = torch.stack( + [parameter.detach().clone() for parameter in reference_experts] + ) + + optimizer.step() + reference_optimizer.step() + + rank = dense_mesh.get_local_rank() + _assert_exact( + dense.to_local(), + reference_dense.detach().chunk(self.world_size, dim=0)[rank], + ) + _assert_batched_muon_update_close( + expert_before, + expert.to_local(), + reference_experts_before, + torch.stack([parameter.detach() for parameter in reference_experts]), + lr=0.03, + weight_decay=0.2, + compute_matrix_shape=reference_experts[0].shape, + ) + dense_momentum = optimizer.state[dense]["momentum_buffer"] + reference_dense_momentum = reference_optimizer.state[reference_dense][ + "momentum_buffer" + ] + _assert_exact( + dense_momentum.to_local(), + reference_dense_momentum.chunk(self.world_size, dim=0)[rank], + ) + expert_momentum = optimizer.state[expert]["momentum_buffer"] + _assert_exact( + expert_momentum.to_local(), + torch.stack( + [ + reference_optimizer.state[parameter]["momentum_buffer"] + for parameter in reference_experts + ] + ), + ) + + @with_comms + def test_distinct_bucket_meshes_use_mesh_local_owners(self): + fsdp_mesh = self.mesh["fsdp"] + tp_mesh = self.mesh["tp"] + meshes = (fsdp_mesh, tp_mesh) + values = ( + torch.arange(15, device=self.device).reshape(5, 3).float().div_(10), + torch.arange(20, device=self.device).reshape(4, 5).float().div_(10), + ) + params = [ + torch.nn.Parameter(distribute_tensor(value.clone(), mesh, (Shard(0),))) + for value, mesh in zip(values, meshes, strict=True) + ] + names = ("layers.0.fsdp", "layers.1.tp") + optimizer = build_distributed_muon( + [ + { + "params": [param], + "param_names": [name], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + for param, name in zip(params, names, strict=True) + ], + bucket_spec=[ + BucketSpec( + patterns=(name,), + owner_rank_by_fqn={name: 1}, + mesh=mesh, + ) + for name, mesh in zip(names, meshes, strict=True) + ], + ns_steps=1, + ) + + for param, value, mesh in zip(params, values, meshes, strict=True): + param.grad = distribute_tensor(torch.ones_like(value), mesh, (Shard(0),)) + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + + for plan, mesh in zip(optimizer._plans, meshes, strict=True): + participants = tuple(dist.get_process_group_ranks(mesh.get_group())) + route = plan.redistribution_plans[0].storage_to_compute_routes[0] + self.assertEqual(route.destination_participants, (participants[1],)) + self.assertEqual( + sum( + call.kwargs["group"] is mesh.get_group() + for call in collective.call_args_list + ), + 2, + ) + + +@unittest.skipUnless(torch.cuda.device_count() >= 2, "requires two CUDA devices") +class TestDistributedMuonPipeline(_DistributedMuonTestBase): + @with_comms + def test_local_only_bucket_does_not_reuse_inflight_slot(self): + values = [ + torch.arange(offset, offset + 12, device=self.device) + .reshape(4, 3) + .float() + .div_(10) + for offset in (0, 12, 24) + ] + distributed_0, local_blocks, distributed_2 = map(self._parameter, values) + optimizer = build_distributed_muon( + [ + { + "params": [distributed_0, distributed_2], + "param_names": [ + "layers.0.redistributed", + "layers.2.redistributed", + ], + "compute_sharding": MuonComputeSharding(placement=Owned()), + }, + { + "params": [local_blocks], + "param_names": ["layers.1.local_blocks"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=2, matrices_flattened_into_dim=0 + ), + placement=Shard(0), + ), + }, + ], + bucket_spec=[ + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={"layers.0.redistributed": 0}, + mesh=self.mesh, + ), + BucketSpec( + patterns=("layers.1.*",), + owner_rank_by_fqn={}, + mesh=self.mesh, + ), + BucketSpec( + patterns=("layers.2.*",), + owner_rank_by_fqn={"layers.2.redistributed": 0}, + mesh=self.mesh, + ), + ], + lr=0.03, + momentum=0.8, + ns_steps=1, + ) + grads = [ + torch.full_like(value, index + 1) for index, value in enumerate(values) + ] + for param, grad in zip( + (distributed_0, local_blocks, distributed_2), grads, strict=True + ): + param.grad = distribute_tensor(grad, self.mesh, (Shard(0),)) + + rank = self.mesh.get_local_rank() + references = [ + torch.nn.Parameter(values[0].clone()), + torch.nn.Parameter(values[1].chunk(self.world_size, dim=0)[rank].clone()), + torch.nn.Parameter(values[2].clone()), + ] + reference = torch.optim.Muon(references, lr=0.03, momentum=0.8, ns_steps=1) + local_blocks_before = local_blocks.to_local().clone() + reference_local_blocks_before = references[1].detach().clone() + references[0].grad = grads[0].clone() + references[1].grad = grads[1].chunk(self.world_size, dim=0)[rank].clone() + references[2].grad = grads[2].clone() + + all_to_all_single = dist.all_to_all_single + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "all_to_all_single", + wraps=all_to_all_single, + ) as collective: + optimizer.step() + reference.step() + + self.assertEqual(collective.call_count, 4) + splits = [ + ( + tuple(call.kwargs["input_split_sizes"]), + tuple(call.kwargs["output_split_sizes"]), + ) + for call in collective.call_args_list + ] + self.assertEqual(splits[0], splits[1]) + self.assertEqual(splits[2], splits[3]) + self.assertNotEqual(splits[0], splits[2]) + _assert_exact( + distributed_0.to_local(), references[0].chunk(self.world_size, dim=0)[rank] + ) + _assert_batched_muon_update_close( + local_blocks_before, + local_blocks.to_local(), + reference_local_blocks_before, + references[1], + lr=0.03, + weight_decay=0.1, + compute_matrix_shape=references[1].shape, + ) + _assert_exact( + distributed_2.to_local(), references[2].chunk(self.world_size, dim=0)[rank] + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit_tests/test_distributed_muon_math.py b/tests/unit_tests/test_distributed_muon_math.py new file mode 100644 index 0000000000..a65e0892ba --- /dev/null +++ b/tests/unit_tests/test_distributed_muon_math.py @@ -0,0 +1,111 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import itertools +import unittest + +import torch +from torchtitan.components.distributed_optimizers.muon import ( + _adjust_learning_rate, + _compute_muon_update, +) + + +class TestDistributedMuonMath(unittest.TestCase): + def test_two_steps_match_torch_muon(self): + optimizer_kwargs = { + "lr": 0.03, + "weight_decay": 0.2, + "momentum": 0.8, + "nesterov": True, + "ns_coefficients": (3.4445, -4.7750, 2.0315), + "eps": 1e-7, + "ns_steps": 3, + "adjust_lr_fn": "original", + } + + for dtype, shape in itertools.product( + (torch.bfloat16, torch.float32), ((3, 5), (5, 3)) + ): + with self.subTest(dtype=dtype, shape=shape): + generator = torch.Generator().manual_seed(4) + initial = torch.randn(shape, generator=generator, dtype=dtype) + gradients = [ + torch.randn(shape, generator=generator, dtype=dtype) + for _ in range(2) + ] + + reference_param = torch.nn.Parameter(initial.clone()) + reference = torch.optim.Muon([reference_param], **optimizer_kwargs) + actual_param = initial.clone() + actual_momentum = torch.zeros_like(actual_param) + + for gradient in gradients: + reference_param.grad = gradient.clone() + reference.step() + + actual_momentum.lerp_(gradient, 1 - optimizer_kwargs["momentum"]) + prepared = torch.lerp( + gradient, + actual_momentum, + optimizer_kwargs["momentum"], + ) + update = _compute_muon_update( + prepared, + out=torch.empty_like(prepared), + ns_coefficients=optimizer_kwargs["ns_coefficients"], + ns_steps=optimizer_kwargs["ns_steps"], + eps=optimizer_kwargs["eps"], + ) + adjusted_lr = _adjust_learning_rate( + optimizer_kwargs["lr"], + optimizer_kwargs["adjust_lr_fn"], + prepared.shape, + ) + actual_param.mul_( + 1 - optimizer_kwargs["lr"] * optimizer_kwargs["weight_decay"] + ) + actual_param.add_(update, alpha=-adjusted_lr) + + self.assertTrue(torch.equal(actual_param, reference_param)) + self.assertTrue( + torch.equal( + actual_momentum, + reference.state[reference_param]["momentum_buffer"], + ) + ) + + def test_batched_update_matches_independent_matrices(self): + kwargs = { + "ns_coefficients": (3.4445, -4.7750, 2.0315), + "ns_steps": 3, + "eps": 1e-7, + } + + for shape in ((8, 3, 4), (8, 4, 3)): + with self.subTest(shape=shape): + generator = torch.Generator().manual_seed(5) + prepared = torch.randn(shape, generator=generator) + batched = _compute_muon_update( + prepared, out=torch.empty_like(prepared), **kwargs + ) + independent = torch.stack( + [ + _compute_muon_update( + matrix, out=torch.empty_like(matrix), **kwargs + ) + for matrix in prepared + ] + ) + + # Batched and independent matrix multiplications may use + # different BF16 reduction orders. + torch.testing.assert_close( + batched, + independent, + rtol=0, + atol=2e-2, + ) diff --git a/tests/unit_tests/test_flex_optimizer_reshard.py b/tests/unit_tests/test_flex_optimizer_reshard.py new file mode 100644 index 0000000000..b180353f25 --- /dev/null +++ b/tests/unit_tests/test_flex_optimizer_reshard.py @@ -0,0 +1,735 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import unittest +from contextlib import nullcontext +from dataclasses import dataclass +from unittest.mock import MagicMock, Mock, patch + +import torch +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.tensor import DTensor, Shard +from torchtitan.components.distributed_optimizers.flex_optimizer_reshard import ( + _bind_bucket_configs, + _BucketedRedistributionRuntime, + _BucketPlan, + _BucketWork, + _BufferSlot, + _build_bucket_plans, + _build_owned_bucket_plans, + _compute_redistributed, + _finalize_redistributed, + _lower_packed_all_to_all, + _PackedAllToAllSchedule, + _ParticipantPartition, + _prepare_redistributed, + _RedistributionGroup, + _RedistributionPlan, + _tensor_region_view, + _TensorRegion, + _TensorRegionRoute, + assign_balanced_owners, + BucketConfig, + BucketSpec, +) + + +def _fragmented_partition_plan( + participants: tuple[int, ...], + *, + num_compute_units: int = 5, + compute_unit_rows: int = 3, + num_columns: int = 2, +) -> _RedistributionPlan: + num_participants = len(participants) + compute_ranges = tuple( + Shard.local_shard_size_and_offset( + num_compute_units, + num_participants, + rank, + ) + for rank in range(num_participants) + ) + compute_partitions = tuple( + _ParticipantPartition( + participant=participant, + tensor_shape=(num_local_units, compute_unit_rows, num_columns), + logical_regions=( + _TensorRegion( + offsets=(unit_offset, 0, 0), + shape=(num_local_units, compute_unit_rows, num_columns), + ), + ) + if num_local_units + else (), + ) + for participant, (num_local_units, unit_offset) in zip( + participants, + compute_ranges, + strict=True, + ) + ) + + storage_partitions = [] + forward_routes = [] + reverse_routes = [] + for storage_rank, participant in enumerate(participants): + num_local_rows, storage_offset = Shard.local_shard_size_and_offset( + num_compute_units * compute_unit_rows, + num_participants, + storage_rank, + ) + logical_regions = [] + fragment_start = storage_offset + storage_end = storage_offset + num_local_rows + while fragment_start < storage_end: + compute_unit = fragment_start // compute_unit_rows + row = fragment_start % compute_unit_rows + num_rows = min(compute_unit_rows - row, storage_end - fragment_start) + logical_region = _TensorRegion( + offsets=(compute_unit, row, 0), + shape=(1, num_rows, num_columns), + ) + storage_region = _TensorRegion( + offsets=(fragment_start - storage_offset, 0), + shape=(num_rows, num_columns), + ) + compute_rank = next( + rank + for rank, (num_local_units, unit_offset) in enumerate(compute_ranges) + if unit_offset <= compute_unit < unit_offset + num_local_units + ) + compute_unit_offset = compute_ranges[compute_rank][1] + compute_region = _TensorRegion( + offsets=(compute_unit - compute_unit_offset, row, 0), + shape=(1, num_rows, num_columns), + ) + compute_participant = participants[compute_rank] + logical_regions.append(logical_region) + forward_routes.append( + _TensorRegionRoute( + logical_region=logical_region, + source_region=storage_region, + destination_region=compute_region, + source_participants=(participant,), + destination_participants=(compute_participant,), + ) + ) + reverse_routes.append( + _TensorRegionRoute( + logical_region=logical_region, + source_region=compute_region, + destination_region=storage_region, + source_participants=(compute_participant,), + destination_participants=(participant,), + ) + ) + fragment_start += num_rows + storage_partitions.append( + _ParticipantPartition( + participant=participant, + tensor_shape=(num_local_rows, num_columns), + logical_regions=tuple(logical_regions), + ) + ) + + return _RedistributionPlan( + participants=participants, + logical_shape=(num_compute_units, compute_unit_rows, num_columns), + storage_partitions=tuple(storage_partitions), + compute_partitions=compute_partitions, + storage_to_compute_routes=tuple(forward_routes), + compute_to_storage_routes=tuple(reverse_routes), + ) + + +class TestFlexOptimizerReshard(unittest.TestCase): + def test_runtime_enqueues_return_before_next_compute(self): + runtime = _BucketedRedistributionRuntime(torch.device("cuda")) + caller_stream = Mock() + context = Mock() + context.device_handle.current_stream.return_value = caller_stream + context.device_handle.stream.side_effect = lambda _stream: nullcontext() + context.device_handle.Event.side_effect = Mock + context.slots = (Mock(), Mock()) + for slot in context.slots: + slot.communication_buffers.return_value = (object(), object()) + runtime._context = context + + plans = tuple( + Mock(redistributed_items=(object(),), local_items=()) for _ in range(3) + ) + plan_names = {plan: f"bucket_{index}" for index, plan in enumerate(plans)} + events = [] + for plan in plans: + plan.storage_to_compute_schedule.execute.side_effect = ( + lambda *, _plan=plan, **_kwargs: events.append( + ("gather", plan_names[_plan]) + ) + ) + plan.compute_to_storage_schedule.execute.side_effect = ( + lambda *, _plan=plan, **_kwargs: events.append( + ("return", plan_names[_plan]) + ) + ) + + def compute_redistributed(work, *_args, **_kwargs): + events.append(("compute", plan_names[work.plan])) + + original_release = runtime._release + + def release(work, caller): + events.append(("release", plan_names[work.plan])) + original_release(work, caller) + + with patch( + "torchtitan.components.distributed_optimizers." + "flex_optimizer_reshard._prepare_redistributed" + ), patch( + "torchtitan.components.distributed_optimizers." + "flex_optimizer_reshard._compute_redistributed", + side_effect=compute_redistributed, + ), patch( + "torchtitan.components.distributed_optimizers." + "flex_optimizer_reshard._finalize_redistributed" + ), patch.object( + runtime, + "_release", + side_effect=release, + ): + runtime.run( + plans, + local_tensor_spec=Mock(), + prepare=Mock(), + compute=Mock(), + finalize=Mock(), + ) + + self.assertEqual( + events, + [ + ("gather", "bucket_0"), + ("compute", "bucket_0"), + ("gather", "bucket_1"), + ("return", "bucket_0"), + ("compute", "bucket_1"), + ("release", "bucket_0"), + ("gather", "bucket_2"), + ("return", "bucket_1"), + ("compute", "bucket_2"), + ("release", "bucket_1"), + ("return", "bucket_2"), + ("release", "bucket_2"), + ], + ) + self.assertEqual(context.slots[0].communication_buffers.call_count, 2) + self.assertEqual(context.slots[1].communication_buffers.call_count, 1) + + def test_balanced_owner_assignment(self): + self.assertEqual( + assign_balanced_owners( + [("a", "b"), ("c",)], + {"a": 8, "b": 4, "c": 4}, + num_ranks=2, + initial_memory_by_rank=(0, 4), + ), + ({"a": 0, "b": 1}, {"c": 0}), + ) + + owners = {"a": 0} + config = BucketConfig( + patterns=("a",), + owner_rank_by_fqn=owners, + mesh_axes=("optimizer",), + ) + owners["a"] = 1 + self.assertEqual(config.owner_rank_by_fqn, {"a": 0}) + + def test_bucket_config_requires_exactly_one_mesh_axis(self): + for mesh_axes in ((), ("optimizer", "replicate")): + with self.subTest(mesh_axes=mesh_axes), self.assertRaisesRegex( + ValueError, "exactly one mesh axis" + ): + BucketConfig( + patterns=("*",), + owner_rank_by_fqn={}, + mesh_axes=mesh_axes, + ) + + def test_bucket_config_uses_redistributed_parameters_to_resolve_mesh(self): + redistributed_mesh = Mock(spec=DeviceMesh) + redistributed_mesh.ndim = 1 + redistributed_storage_mesh = MagicMock(spec=DeviceMesh) + redistributed_storage_mesh.__getitem__.return_value = redistributed_mesh + compute_ready_storage_mesh = MagicMock(spec=DeviceMesh) + redistributed = Mock(device_mesh=redistributed_storage_mesh) + compute_ready = Mock(device_mesh=compute_ready_storage_mesh) + + specs = _bind_bucket_configs( + ( + BucketConfig( + patterns=("layers.*",), + owner_rank_by_fqn={"layers.redistributed": 0}, + mesh_axes=("optimizer",), + ), + ), + { + "layers.redistributed": redistributed, + "layers.compute_ready": compute_ready, + }, + ) + + self.assertEqual(len(specs), 1) + self.assertIs(specs[0].mesh, redistributed_mesh) + redistributed_storage_mesh.__getitem__.assert_called_once_with(("optimizer",)) + compute_ready_storage_mesh.__getitem__.assert_not_called() + + def test_bucket_planner_preserves_empty_local_storage_region(self): + @dataclass(frozen=True) + class Item: + fqn: str + tensor: DTensor + + tensor = Mock(spec=DTensor) + tensor.shape = torch.Size((2, 3)) + tensor.to_local.return_value = torch.empty(0, 3) + item = Item("layers.0.weight", tensor) + regions = ( + ((3,), _TensorRegion(offsets=(0, 0), shape=(2, 3))), + ((7,), _TensorRegion(offsets=(2, 0), shape=(0, 3))), + ) + group = _RedistributionGroup( + process_group=object(), + participants=(3, 7), + local_participant=7, + ) + mesh = Mock(spec=DeviceMesh) + mesh.ndim = 1 + + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "get_process_group_ranks", + return_value=[3, 7], + ), patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard." + "_redistribution_group", + return_value=group, + ), patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard." + "_dtensor_storage_regions", + return_value=regions, + ): + result = _build_owned_bucket_plans( + (item,), + ( + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={item.fqn: 0}, + mesh=mesh, + ), + ), + get_fqn=lambda value: value.fqn, + storage_is_compute_ready=lambda _value: False, + get_storage_dtensor=lambda value: value.tensor, + ) + + plan = result.plans[0] + self.assertEqual(result.ordered_items, (item,)) + self.assertEqual( + plan.storage_to_compute_schedule.input_spans_by_parameter[0][0].numel, + 0, + ) + self.assertEqual( + plan.compute_to_storage_schedule.output_spans_by_parameter[0][0].numel, + 0, + ) + + def test_bucket_planner_accepts_owner_free_redistribution(self): + @dataclass(frozen=True) + class Item: + fqn: str + tensor: DTensor + + participants = (3, 7) + first = _TensorRegion(offsets=(0, 0), shape=(1, 2)) + second = _TensorRegion(offsets=(1, 0), shape=(1, 2)) + local = _TensorRegion(offsets=(0, 0), shape=(1, 2)) + redistribution_plan = _RedistributionPlan( + participants=participants, + logical_shape=(2, 2), + storage_partitions=( + _ParticipantPartition(3, (1, 2), (first,)), + _ParticipantPartition(7, (1, 2), (second,)), + ), + compute_partitions=( + _ParticipantPartition(3, (1, 2), (second,)), + _ParticipantPartition(7, (1, 2), (first,)), + ), + storage_to_compute_routes=( + _TensorRegionRoute(first, local, local, (3,), (7,)), + _TensorRegionRoute(second, local, local, (7,), (3,)), + ), + compute_to_storage_routes=( + _TensorRegionRoute(first, local, local, (7,), (3,)), + _TensorRegionRoute(second, local, local, (3,), (7,)), + ), + ) + tensor = Mock(spec=DTensor) + tensor.to_local.return_value = torch.empty(1, 2) + item = Item("layers.0.weight", tensor) + group = _RedistributionGroup( + process_group=object(), + participants=participants, + local_participant=3, + ) + mesh = Mock(spec=DeviceMesh) + mesh.ndim = 1 + received_owner_ranks = [] + + def build_plan(_item, received_group, owner_rank): + self.assertIs(_item, item) + self.assertIs(received_group, group) + received_owner_ranks.append(owner_rank) + return redistribution_plan + + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard." + "_redistribution_group", + return_value=group, + ), patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "get_process_group_ranks", + return_value=list(participants), + ): + result = _build_bucket_plans( + (item,), + ( + BucketSpec( + patterns=("layers.0.*",), + owner_rank_by_fqn={}, + mesh=mesh, + ), + ), + get_fqn=lambda value: value.fqn, + requires_owner=lambda _value: False, + get_storage_dtensor=lambda value: value.tensor, + build_redistribution_plan=build_plan, + ) + + self.assertEqual(received_owner_ranks, [None]) + self.assertEqual(result.ordered_items, (item,)) + plan = result.plans[0] + self.assertEqual(plan.local_items, ()) + self.assertEqual(plan.redistributed_items, (item,)) + self.assertIs(plan.redistribution_plans[0], redistribution_plan) + self.assertEqual( + plan.storage_to_compute_schedule.input_split_sizes, + (0, 2), + ) + self.assertEqual( + plan.storage_to_compute_schedule.output_split_sizes, + (0, 2), + ) + + def test_transport_neutral_routes_lower_to_packed_all_to_all(self): + first = _TensorRegion(offsets=(0, 0), shape=(2, 3)) + second = _TensorRegion(offsets=(2, 0), shape=(2, 3)) + local = _TensorRegion(offsets=(0, 0), shape=(2, 3)) + full = _TensorRegion(offsets=(0, 0), shape=(4, 3)) + plan = _RedistributionPlan( + participants=(3, 7), + logical_shape=(4, 3), + storage_partitions=( + _ParticipantPartition(3, (2, 3), (first,)), + _ParticipantPartition(7, (2, 3), (second,)), + ), + compute_partitions=( + _ParticipantPartition(3, (0,), ()), + _ParticipantPartition(7, (4, 3), (full,)), + ), + storage_to_compute_routes=( + _TensorRegionRoute(first, local, first, (3,), (7,)), + _TensorRegionRoute(second, local, second, (7,), (7,)), + ), + compute_to_storage_routes=( + _TensorRegionRoute(first, first, local, (7,), (3,)), + _TensorRegionRoute(second, second, local, (7,), (7,)), + ), + ) + + with patch( + "torchtitan.components.distributed_optimizers.flex_optimizer_reshard.dist." + "get_process_group_ranks", + return_value=[3, 7], + ): + forward = _lower_packed_all_to_all( + (plan,), + storage_to_compute=True, + process_group=object(), + local_participant=7, + ) + reverse = _lower_packed_all_to_all( + (plan,), + storage_to_compute=False, + process_group=object(), + local_participant=7, + ) + + self.assertIsInstance(forward, _PackedAllToAllSchedule) + self.assertEqual(forward.input_split_sizes, (0, 6)) + self.assertEqual(forward.output_split_sizes, (6, 6)) + self.assertEqual( + tuple(span.region for span in forward.output_spans_by_parameter[0]), + (first, second), + ) + self.assertEqual(reverse.input_split_sizes, (6, 6)) + self.assertEqual(reverse.output_split_sizes, (0, 6)) + self.assertEqual( + tuple(span.region for span in reverse.output_spans_by_parameter[0]), + (local,), + ) + + def test_partitioned_compute_supports_fragmented_storage_and_empty_participant( + self, + ): + participants = (3, 7, 11, 13) + plan = _fragmented_partition_plan(participants) + + self.assertEqual(plan.logical_shape, (5, 3, 2)) + self.assertEqual(plan.storage_partition(13).tensor_shape, (3, 2)) + self.assertEqual(plan.compute_partition(13).tensor_shape, (0, 3, 2)) + self.assertEqual(plan.compute_partition(13).logical_regions, ()) + self.assertTrue( + any( + len(route.source_region.shape) == 2 + and len(route.destination_region.shape) == 3 + for route in plan.storage_to_compute_routes + ) + ) + + forward = _lower_packed_all_to_all( + (plan,), + storage_to_compute=True, + process_group=object(), + local_participant=13, + ) + reverse = _lower_packed_all_to_all( + (plan,), + storage_to_compute=False, + process_group=object(), + local_participant=13, + ) + + self.assertEqual(forward.input_split_sizes, (0, 0, 6, 0)) + self.assertEqual(forward.output_split_sizes, (0, 0, 0, 0)) + self.assertEqual(reverse.input_split_sizes, (0, 0, 0, 0)) + self.assertEqual(reverse.output_split_sizes, (0, 0, 6, 0)) + + item = object() + group = _RedistributionGroup( + process_group=object(), + participants=participants, + local_participant=13, + ) + bucket = _BucketPlan( + local_items=(), + redistributed_items=(item,), + redistribution_plans=(plan,), + group=group, + storage_to_compute_schedule=forward, + compute_to_storage_schedule=reverse, + dtype=torch.float32, + device=torch.device("cpu"), + ) + slot = _BufferSlot() + compute = Mock() + _compute_redistributed( + _BucketWork( + plan=bucket, + slot=slot, + storage_buffer=torch.empty(reverse.output_buffer_numel), + compute_fragment_buffer=torch.empty(0), + ), + slot, + compute=compute, + ) + compute.assert_not_called() + + def test_prepare_and_finalize_assemble_multiple_endpoint_spans(self): + participants = (3, 7, 11, 13) + redistribution_plan = _fragmented_partition_plan(participants) + local_participant = 7 + forward = _lower_packed_all_to_all( + (redistribution_plan,), + storage_to_compute=True, + process_group=object(), + local_participant=local_participant, + ) + reverse = _lower_packed_all_to_all( + (redistribution_plan,), + storage_to_compute=False, + process_group=object(), + local_participant=local_participant, + ) + item = object() + group = _RedistributionGroup( + process_group=object(), + participants=participants, + local_participant=local_participant, + ) + bucket = _BucketPlan( + local_items=(), + redistributed_items=(item,), + redistribution_plans=(redistribution_plan,), + group=group, + storage_to_compute_schedule=forward, + compute_to_storage_schedule=reverse, + dtype=torch.float32, + device=torch.device("cpu"), + ) + slot = _BufferSlot() + prepared = torch.arange(8, dtype=torch.float32).view(4, 2) + packed_forward = torch.empty(forward.input_buffer_numel) + + def prepare(_item, out): + self.assertIs(_item, item) + out.copy_(prepared) + + _prepare_redistributed( + bucket, + slot, + packed_forward, + prepare=prepare, + ) + self.assertGreater( + len(forward.input_spans_by_parameter[0]), + 1, + ) + torch.testing.assert_close(packed_forward, prepared.flatten()) + + expected_update = prepared.add(100) + packed_reverse = torch.empty(reverse.output_buffer_numel) + for span in reverse.output_spans_by_parameter[0]: + packed_reverse[span.buffer_offset : span.buffer_offset + span.numel].copy_( + _tensor_region_view(expected_update, span.region).reshape(-1) + ) + work = _BucketWork( + plan=bucket, + slot=slot, + storage_buffer=packed_reverse, + compute_fragment_buffer=torch.empty(reverse.input_buffer_numel), + ) + finalized = [] + _finalize_redistributed( + work, + slot, + finalize=lambda _item, update: finalized.append((_item, update.clone())), + ) + self.assertEqual(len(reverse.output_spans_by_parameter[0]), 2) + self.assertEqual(len(finalized), 1) + self.assertIs(finalized[0][0], item) + torch.testing.assert_close(finalized[0][1], expected_update) + + def test_routes_require_an_exact_nonoverlapping_partition(self): + def plan(regions): + routes = tuple( + _TensorRegionRoute(region, region, region, (3,), (3,)) + for region in regions + ) + return _RedistributionPlan( + participants=(3, 7), + logical_shape=(2, 3), + storage_partitions=( + _ParticipantPartition(3, (2, 3), regions), + _ParticipantPartition(7, (0,), ()), + ), + compute_partitions=( + _ParticipantPartition(3, (2, 3), regions), + _ParticipantPartition(7, (0,), ()), + ), + storage_to_compute_routes=routes, + compute_to_storage_routes=routes, + ) + + invalid_partitions = ( + ( + (_TensorRegion((0, 0), (3, 3)),), + ValueError, + "in bounds", + ), + ( + ( + _TensorRegion((0, 0), (2, 3)), + _TensorRegion((0, 0), (2, 3)), + ), + NotImplementedError, + "overlapping", + ), + ( + (_TensorRegion((0, 0), (1, 3)),), + ValueError, + "do not cover", + ), + ) + for regions, error, message in invalid_partitions: + with self.subTest(message=message), self.assertRaisesRegex(error, message): + plan(regions) + + first = _TensorRegion((0, 0), (1, 3)) + second = _TensorRegion((1, 0), (1, 3)) + local = _TensorRegion((0, 0), (1, 3)) + split_routes = ( + _TensorRegionRoute(first, local, local, (3,), (3,)), + _TensorRegionRoute(second, local, local, (7,), (7,)), + ) + split_plan = _RedistributionPlan( + participants=(3, 7), + logical_shape=(2, 3), + storage_partitions=( + _ParticipantPartition(3, (1, 3), (first,)), + _ParticipantPartition(7, (1, 3), (second,)), + ), + compute_partitions=( + _ParticipantPartition(3, (1, 3), (first,)), + _ParticipantPartition(7, (1, 3), (second,)), + ), + storage_to_compute_routes=split_routes, + compute_to_storage_routes=split_routes, + ) + self.assertEqual( + tuple( + partition.tensor_shape for partition in split_plan.compute_partitions + ), + ((1, 3), (1, 3)), + ) + + full = _TensorRegion((0, 0), (2, 3)) + forward_routes = ( + _TensorRegionRoute(first, first, first, (3,), (3,)), + _TensorRegionRoute(second, second, second, (3,), (3,)), + ) + swapped_reverse_routes = ( + _TensorRegionRoute(first, second, first, (3,), (3,)), + _TensorRegionRoute(second, first, second, (3,), (3,)), + ) + with self.assertRaisesRegex(ValueError, "must exactly invert"): + _RedistributionPlan( + participants=(3, 7), + logical_shape=(2, 3), + storage_partitions=( + _ParticipantPartition(3, (2, 3), (full,)), + _ParticipantPartition(7, (0,), ()), + ), + compute_partitions=( + _ParticipantPartition(3, (2, 3), (full,)), + _ParticipantPartition(7, (0,), ()), + ), + storage_to_compute_routes=forward_routes, + compute_to_storage_routes=swapped_reverse_routes, + ) + + with self.assertRaisesRegex(ValueError, "participants must be unique"): + _TensorRegionRoute(first, first, first, (3, 3), (3,)) diff --git a/tests/unit_tests/test_kimi_k2_7_muon_config.py b/tests/unit_tests/test_kimi_k2_7_muon_config.py new file mode 100644 index 0000000000..37a1bbf0d4 --- /dev/null +++ b/tests/unit_tests/test_kimi_k2_7_muon_config.py @@ -0,0 +1,275 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import json +import unittest + +import torch +from torch.distributed.tensor import Shard +from torchtitan.components.distributed_optimizers.flex_optimizer_reshard import ( + assign_balanced_owners, +) +from torchtitan.components.distributed_optimizers.muon import Owned +from torchtitan.components.distributed_optimizers.muon_parameter_prep import ( + BatchedMatrixComputeView, + MuonComputeSharding, +) +from torchtitan.components.optimizer import OptimizersContainer +from torchtitan.distributed.activation_checkpoint import FullAC +from torchtitan.models.kimi_k2_7.config_registry import ( + kimi_k2_5_muon, + moonlight_16b_a3b_muon, +) + + +class _KimiMuonConfigTests: + config_factory = None + num_layers = 0 + num_heads = 0 + num_owner_ranks = 0 + expert_parallel_degree = 0 + attention_projections: tuple[str, ...] = () + owned_attention_projections: frozenset[str] = frozenset() + + @classmethod + def setUpClass(cls): + assert cls.config_factory is not None + cls.config = cls.config_factory() + assert cls.config.model_spec is not None + with torch.device("meta"): + cls.model = cls.config.model_spec.model.build() + + @classmethod + def tearDownClass(cls): + del cls.model + + def test_parameter_routing(self): + optimizer_config = self.config.optimizer + impl_kwargs = OptimizersContainer._build_impl_kwargs(optimizer_config) + groups_by_optimizer, _ = OptimizersContainer._build_param_groups( + self.model, + optimizer_config.param_groups, + impl_kwargs, + ) + model_names = set(dict(self.model.named_parameters())) + expected_muon_names = set() + suffix_counts = ( + *( + (f".attention.{projection}.weight", self.num_layers) + for projection in self.attention_projections + ), + (".feed_forward.w1.weight", 1), + (".feed_forward.w2.weight", 1), + (".feed_forward.w3.weight", 1), + ( + ".moe.routed_experts.inner_experts.w1_EFD", + self.num_layers - 1, + ), + ( + ".moe.routed_experts.inner_experts.w2_EDF", + self.num_layers - 1, + ), + ( + ".moe.routed_experts.inner_experts.w3_EFD", + self.num_layers - 1, + ), + (".moe.router.gate.weight", self.num_layers - 1), + (".moe.shared_experts.w1.weight", self.num_layers - 1), + (".moe.shared_experts.w2.weight", self.num_layers - 1), + (".moe.shared_experts.w3.weight", self.num_layers - 1), + ) + for suffix, count in suffix_counts: + names = {name for name in model_names if name.endswith(suffix)} + self.assertEqual(len(names), count, suffix) + expected_muon_names.update(names) + + muon_groups = groups_by_optimizer["DistributedMuon"] + muon_names = {name for group in muon_groups for name in group["param_names"]} + adamw_names = { + name + for group in groups_by_optimizer["AdamW"] + for name in group["param_names"] + } + self.assertEqual(muon_names, expected_muon_names) + self.assertEqual(adamw_names, model_names - expected_muon_names) + self.assertEqual( + {group["adjust_lr_fn"] for group in muon_groups}, + {"match_rms_adamw"}, + ) + + representative_suffixes = ( + *( + f".attention.{projection}.weight" + for projection in self.attention_projections + ), + ".feed_forward.w1.weight", + ".moe.router.gate.weight", + ".moe.shared_experts.w1.weight", + ".moe.routed_experts.inner_experts.w1_EFD", + ) + group_by_suffix = { + suffix: next( + group + for group in muon_groups + if group["param_names"][0].endswith(suffix) + ) + for suffix in representative_suffixes + } + per_head = MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=self.num_heads, + matrices_flattened_into_dim=0, + ), + placement=Shard(0), + ) + expected_sharding = { + f".attention.{projection}.weight": ( + MuonComputeSharding(placement=Owned()) + if projection in self.owned_attention_projections + else per_head + ) + for projection in self.attention_projections + } + expected_sharding.update( + { + ".feed_forward.w1.weight": MuonComputeSharding(placement=Owned()), + ".moe.router.gate.weight": MuonComputeSharding(placement=Owned()), + ".moe.shared_experts.w1.weight": MuonComputeSharding(placement=Owned()), + ".moe.routed_experts.inner_experts.w1_EFD": MuonComputeSharding( + placement=Shard(0) + ), + } + ) + for suffix, group in group_by_suffix.items(): + self.assertEqual(group["compute_sharding"], expected_sharding[suffix]) + + def test_bucket_and_parallelism_config(self): + optimizer_config = self.config.optimizer + bucket_configs = optimizer_config.optimizer_init_kwargs["DistributedMuon"][ + "bucket_configs" + ] + bucket_layer_ids = ((0,),) + tuple( + tuple( + range( + first_layer_id, + min(first_layer_id + 2, self.num_layers), + ) + ) + for first_layer_id in range(1, self.num_layers, 2) + ) + self.assertEqual(len(bucket_configs), len(bucket_layer_ids)) + expected_bucket_patterns = [] + expected_owned_fqns = set() + for layer_ids, bucket in zip(bucket_layer_ids, bucket_configs, strict=True): + expected = () + expected_owners = set() + for layer_id in layer_ids: + prefix = f"layers.{layer_id}" + attention_fqns = tuple( + f"{prefix}.attention.{projection}.weight" + for projection in self.attention_projections + ) + expected += attention_fqns + expected_owners.update( + f"{prefix}.attention.{projection}.weight" + for projection in self.owned_attention_projections + ) + if not layer_id: + dense_fqns = tuple( + f"{prefix}.feed_forward.{projection}.weight" + for projection in ("w1", "w2", "w3") + ) + expected += dense_fqns + expected_owners.update(dense_fqns) + else: + expert_fqns = tuple( + f"{prefix}.moe.routed_experts.inner_experts.{projection}" + for projection in ("w1_EFD", "w2_EDF", "w3_EFD") + ) + router_fqn = f"{prefix}.moe.router.gate.weight" + shared_fqns = tuple( + f"{prefix}.moe.shared_experts.{projection}.weight" + for projection in ("w1", "w2", "w3") + ) + expected += expert_fqns + (router_fqn,) + shared_fqns + expected_owners.update((router_fqn, *shared_fqns)) + self.assertEqual( + bucket.name, + "layers." + "-".join(map(str, layer_ids)), + ) + self.assertEqual(bucket.patterns, expected) + self.assertEqual(bucket.mesh_axes, ("dp_shard",)) + self.assertEqual( + set(bucket.owner_rank_by_fqn), + expected_owners, + ) + expected_bucket_patterns.append(expected) + expected_owned_fqns.update(expected_owners) + self.assertTrue( + all( + rank in range(self.num_owner_ranks) + for rank in bucket.owner_rank_by_fqn.values() + ) + ) + + parameter_numel_by_fqn = { + fqn: parameter.numel() + for fqn, parameter in self.model.named_parameters() + if fqn in expected_owned_fqns + } + self.assertEqual(set(parameter_numel_by_fqn), expected_owned_fqns) + self.assertEqual( + tuple(dict(bucket.owner_rank_by_fqn) for bucket in bucket_configs), + assign_balanced_owners( + expected_bucket_patterns, + parameter_numel_by_fqn, + num_ranks=self.num_owner_ranks, + ), + ) + + parallelism = self.config.parallelism + self.assertEqual(parallelism.data_parallel_replicate_degree, 1) + self.assertEqual( + parallelism.data_parallel_shard_degree, + self.num_owner_ranks, + ) + self.assertEqual( + parallelism.expert_parallel_degree, + self.expert_parallel_degree, + ) + self.assertEqual(parallelism.tensor_parallel_degree, 1) + self.assertEqual(parallelism.context_parallel_degree, 1) + self.assertEqual(parallelism.pipeline_parallel_degree, 1) + self.assertFalse(parallelism.enable_sequence_parallel) + self.assertEqual(parallelism.spmd_backend, "spmd_types") + self.assertIsInstance(self.config.activation_checkpoint, FullAC.Config) + + def test_config_is_json_serializable(self): + json.dumps(self.config.to_dict()) + + +class TestKimiK25MuonConfig(_KimiMuonConfigTests, unittest.TestCase): + config_factory = staticmethod(kimi_k2_5_muon) + num_layers = 61 + num_heads = 64 + num_owner_ranks = 64 + expert_parallel_degree = 8 + attention_projections = ("wq_a", "wq_b", "wkv_a", "wkv_b", "wo") + owned_attention_projections = frozenset(("wq_a", "wkv_a", "wo")) + + +class TestMoonlightMuonConfig(_KimiMuonConfigTests, unittest.TestCase): + config_factory = staticmethod(moonlight_16b_a3b_muon) + num_layers = 27 + num_heads = 16 + num_owner_ranks = 8 + expert_parallel_degree = 4 + attention_projections = ("wq", "wkv_a", "wkv_b", "wo") + owned_attention_projections = frozenset(("wkv_a", "wo")) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit_tests/test_muon_parameter_prep.py b/tests/unit_tests/test_muon_parameter_prep.py new file mode 100644 index 0000000000..8641025d0a --- /dev/null +++ b/tests/unit_tests/test_muon_parameter_prep.py @@ -0,0 +1,205 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import unittest +from unittest import mock + +import torch +from torch.distributed.tensor import DTensor, Shard +from torch.distributed.tensor.placement_types import _StridedShard +from torchtitan.components.distributed_optimizers.muon import DistributedMuon, Owned +from torchtitan.components.distributed_optimizers.muon_parameter_prep import ( + BatchedMatrixComputeView, + build_distributed_muon, + MuonComputeSharding, +) + + +class TestMuonParameterPrep(unittest.TestCase): + def test_batched_matrix_view_validation(self): + with self.assertRaisesRegex(ValueError, "positive integer"): + BatchedMatrixComputeView(0) + with self.assertRaisesRegex(ValueError, "only matrices_flattened_into_dim=0"): + BatchedMatrixComputeView(3, 1) + + def test_builder_compiles_layout_without_mutating_caller_group(self): + view = BatchedMatrixComputeView(num_matrices=3, matrices_flattened_into_dim=0) + compute_sharding = MuonComputeSharding( + view_before_placement=view, + placement=Shard(0), + ) + storage = torch.arange(24).reshape(6, 4) + other_storage = torch.empty(9, 5) + group = { + "params": [storage, other_storage], + "param_names": [ + "layers.0.wq.weight", + "layers.0.wkv_b.weight", + ], + "compute_sharding": compute_sharding, + } + identity_storage = torch.empty(4, 3) + identity_group = { + "params": [identity_storage], + "param_names": ["layers.0.wkv_a.weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + bucket_spec = () + + with mock.patch.object(DistributedMuon, "__init__", return_value=None) as init: + optimizer = build_distributed_muon( + [group, identity_group], + bucket_spec=bucket_spec, + lr=0.1, + ) + + self.assertIsInstance(optimizer, DistributedMuon) + core_groups = init.call_args.args[0] + prepared = init.call_args.kwargs["_prepared_compute_views"] + init.assert_called_once_with( + core_groups, + bucket_spec=bucket_spec, + _prepared_compute_views=prepared, + lr=0.1, + ) + self.assertIsNot(core_groups[0], group) + self.assertIsNot(core_groups[1], identity_group) + self.assertIs(group["compute_sharding"], compute_sharding) + self.assertNotIn("compute_sharding", core_groups[0]) + self.assertEqual(core_groups[0]["_compute_placement"], Shard(0)) + self.assertEqual(core_groups[1]["_compute_placement"], Owned()) + self.assertFalse(any(value is view for value in core_groups[0].values())) + self.assertEqual( + prepared["layers.0.wq.weight"].global_compute_shape, + torch.Size((3, 2, 4)), + ) + self.assertEqual( + prepared["layers.0.wq.weight"].local_storage_view.shape, + torch.Size((3, 2, 4)), + ) + self.assertEqual( + prepared["layers.0.wkv_b.weight"].global_compute_shape, + torch.Size((3, 3, 5)), + ) + self.assertEqual( + prepared["layers.0.wkv_b.weight"].local_storage_view.shape, + torch.Size((3, 3, 5)), + ) + self.assertEqual( + prepared["layers.0.wq.weight"].local_storage_view.data_ptr(), + storage.data_ptr(), + ) + self.assertEqual( + prepared["layers.0.wkv_a.weight"].global_compute_shape, + identity_storage.shape, + ) + self.assertEqual( + prepared["layers.0.wkv_a.weight"].local_storage_view.shape, + identity_storage.shape, + ) + self.assertIs( + prepared["layers.0.wkv_a.weight"].local_storage_view, + identity_storage, + ) + + def test_builder_validates_global_shape_and_aligned_names(self): + for shape in ((2, 3, 4), (5, 4)): + with self.subTest(shape=shape): + with self.assertRaisesRegex(ValueError, "cannot be viewed"): + build_distributed_muon( + [ + { + "params": [torch.empty(shape)], + "param_names": ["layers.0.wq.weight"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + 2, 0 + ), + placement=Shard(0), + ), + } + ], + bucket_spec=(), + ) + + with self.assertRaisesRegex(ValueError, "must be aligned"): + build_distributed_muon( + [ + { + "params": [torch.empty(6, 4)], + "param_names": [], + "compute_sharding": MuonComputeSharding(placement=Shard(0)), + } + ], + bucket_spec=(), + ) + + def test_builder_requires_dtensor_storage(self): + with self.assertRaisesRegex(TypeError, "DTensor parameters"): + build_distributed_muon( + [ + { + "params": [torch.empty(2, 2)], + "param_names": ["weight"], + "compute_sharding": MuonComputeSharding(placement=Owned()), + } + ], + bucket_spec=(), + ) + + def test_builder_validates_all_storage_shards_before_construction(self): + param = mock.Mock(spec=DTensor) + param.shape = torch.Size((15, 3)) + param.ndim = 2 + param.placements = (Shard(0),) + param.device_mesh = mock.Mock(shape=(4,)) + param.to_local.return_value = torch.empty(3, 3) + + with mock.patch.object(DistributedMuon, "__init__", return_value=None) as init: + with self.assertRaisesRegex(ValueError, "not aligned"): + build_distributed_muon( + [ + { + "params": [param], + "param_names": ["layers.0.wq.weight"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView(5), + placement=Shard(0), + ), + } + ], + bucket_spec=(), + ) + + param.to_local.assert_not_called() + init.assert_not_called() + + def test_builder_rejects_strided_storage_shard_for_batched_matrices(self): + param = mock.Mock(spec=DTensor) + param.shape = torch.Size((6, 4)) + param.placements = ( + _StridedShard(0, split_factor=2), + Shard(0), + ) + + with self.assertRaisesRegex(ValueError, "Shard or Replicate"): + build_distributed_muon( + [ + { + "params": [param], + "param_names": ["layers.0.wq.weight"], + "compute_sharding": MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView(3), + placement=Shard(0), + ), + } + ], + bucket_spec=(), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/torchtitan/components/distributed_optimizers/__init__.py b/torchtitan/components/distributed_optimizers/__init__.py new file mode 100644 index 0000000000..1e174886ff --- /dev/null +++ b/torchtitan/components/distributed_optimizers/__init__.py @@ -0,0 +1,7 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Distributed optimizer implementations and redistribution runtimes.""" diff --git a/torchtitan/components/distributed_optimizers/flex_optimizer_reshard.py b/torchtitan/components/distributed_optimizers/flex_optimizer_reshard.py new file mode 100644 index 0000000000..7d8511d152 --- /dev/null +++ b/torchtitan/components/distributed_optimizers/flex_optimizer_reshard.py @@ -0,0 +1,1515 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Public bucket contracts and private resharding for distributed optimizers.""" + +from __future__ import annotations + +import fnmatch +import hashlib +import heapq +import math +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field +from types import ModuleType +from typing import Any, Generic, TypeVar + +import torch +import torch.distributed as dist +from torch import Tensor +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.tensor import DTensor, Shard + + +__all__ = ["BucketConfig", "BucketSpec", "assign_balanced_owners"] + + +@dataclass(frozen=True, slots=True) +class BucketConfig: + """Static bucket configuration resolved after runtime meshes exist. + + ``mesh_axes`` contains exactly one storage mesh axis name. When + ``owner_rank_by_fqn`` is nonempty, its redistributed parameters determine + the resolved mesh; otherwise all parameters matching ``patterns`` do. For + DistributedMuon, owner-assigned parameters use ``Owned``. An owner-free + parameter may compute locally or redistribute on the resolved bucket mesh + without a designated owner. + """ + + patterns: tuple[str, ...] + owner_rank_by_fqn: Mapping[str, int] + mesh_axes: tuple[str, ...] + name: str = "" + + def __post_init__(self) -> None: + mesh_axes = tuple(self.mesh_axes) + if len(mesh_axes) != 1: + raise ValueError( + "BucketConfig mesh_axes must contain exactly one mesh axis" + ) + object.__setattr__(self, "patterns", tuple(self.patterns)) + object.__setattr__(self, "owner_rank_by_fqn", dict(self.owner_rank_by_fqn)) + object.__setattr__(self, "mesh_axes", mesh_axes) + + def bind(self, mesh: DeviceMesh) -> BucketSpec: + return BucketSpec( + patterns=self.patterns, + owner_rank_by_fqn=self.owner_rank_by_fqn, + mesh=mesh, + name=self.name, + ) + + +@dataclass(frozen=True, slots=True) +class BucketSpec: + """One ordered optimizer-work bucket selected by canonical FQN. + + Patterns use case-sensitive ``fnmatch`` syntax. Every optimizer FQN must + match exactly one bucket, and sequence order controls execution order. + ``mesh`` is the bucket's exact one-dimensional communication mesh. + ``owner_rank_by_fqn`` must exactly cover parameters whose compute transition + requires a designated owner and uses mesh-local ranks. Owner-free + transitions have no entry whether they compute locally or redistribute. A + redundant rank-0 entry is accepted for a local transition on a one-rank + mesh, where sharded storage may normalize to replication. + ``name`` is diagnostic metadata only. + """ + + patterns: tuple[str, ...] + owner_rank_by_fqn: Mapping[str, int] + mesh: DeviceMesh + name: str = "" + + def __post_init__(self) -> None: + if self.mesh.ndim != 1: + raise ValueError("BucketSpec mesh must be one-dimensional") + object.__setattr__(self, "patterns", tuple(self.patterns)) + object.__setattr__(self, "owner_rank_by_fqn", dict(self.owner_rank_by_fqn)) + + +def assign_balanced_owners( + bucket_fqns: Sequence[Sequence[str]], + memory_estimate_by_fqn: Mapping[str, int], + *, + num_ranks: int, + initial_memory_by_rank: Sequence[int] | None = None, +) -> tuple[dict[str, int], ...]: + """Greedily balance selected FQNs across group-local rank indices. + + Only FQNs present in ``memory_estimate_by_fqn`` receive owners. One running + load vector balances cumulatively across buckets; FQN and rank ordering + make equal-load assignments deterministic. + """ + initial_memory_by_rank = initial_memory_by_rank or (0,) * num_ranks + rank_loads = list(zip(initial_memory_by_rank, range(num_ranks), strict=True)) + heapq.heapify(rank_loads) + owners_by_bucket = [] + for bucket in bucket_fqns: + bucket_owners = {} + candidates = (fqn for fqn in bucket if fqn in memory_estimate_by_fqn) + for fqn in sorted( + candidates, key=lambda name: (-memory_estimate_by_fqn[name], name) + ): + load, rank = heapq.heappop(rank_loads) + bucket_owners[fqn] = rank + heapq.heappush(rank_loads, (load + memory_estimate_by_fqn[fqn], rank)) + owners_by_bucket.append(bucket_owners) + return tuple(owners_by_bucket) + + +_ItemT = TypeVar("_ItemT") + + +def _bind_bucket_configs( + configs: Sequence[BucketConfig], + storage_by_fqn: Mapping[str, DTensor], +) -> tuple[BucketSpec, ...]: + """Bind static configs to storage meshes after model parallelization.""" + specs = [] + for config in configs: + candidates = tuple(config.owner_rank_by_fqn) or tuple( + fqn + for fqn in storage_by_fqn + if any(fnmatch.fnmatchcase(fqn, pattern) for pattern in config.patterns) + ) + if not candidates: + continue + + meshes = [] + for fqn in candidates: + storage_mesh = storage_by_fqn[fqn].device_mesh + meshes.append(storage_mesh[config.mesh_axes]) + + mesh = meshes[0] + if any(not torch.equal(candidate.mesh, mesh.mesh) for candidate in meshes[1:]): + raise ValueError( + f"bucket {config.name!r} resolves to inconsistent communication meshes" + ) + specs.append(config.bind(mesh)) + return tuple(specs) + + +def _resolve_buckets( + items: Sequence[_ItemT], + specs: Sequence[BucketSpec], + *, + get_fqn: Callable[[_ItemT], str], +) -> tuple[tuple[_ItemT, ...], ...]: + resolved: list[list[_ItemT]] = [[] for _ in specs] + for item in items: + name = get_fqn(item) + matches = [ + index + for index, spec in enumerate(specs) + if any(fnmatch.fnmatchcase(name, pattern) for pattern in spec.patterns) + ] + if len(matches) != 1: + raise ValueError(f"optimizer parameter {name!r} must match one bucket") + resolved[matches[0]].append(item) + return tuple(tuple(bucket) for bucket in resolved) + + +@dataclass(frozen=True, slots=True) +class _TensorRegion: + """One rectangular tensor region.""" + + offsets: tuple[int, ...] + shape: tuple[int, ...] + + @property + def numel(self) -> int: + return math.prod(self.shape) + + +@dataclass(frozen=True, slots=True) +class _ParticipantPartition: + """One participant's tensor shape and its global logical regions.""" + + participant: int + tensor_shape: tuple[int, ...] + logical_regions: tuple[_TensorRegion, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "tensor_shape", tuple(self.tensor_shape)) + object.__setattr__(self, "logical_regions", tuple(self.logical_regions)) + if any(size < 0 for size in self.tensor_shape): + raise ValueError("participant tensor shape must be nonnegative") + + +@dataclass(frozen=True, slots=True) +class _TensorRegionRoute: + """Map one logical region between differently shaped endpoint tensors.""" + + logical_region: _TensorRegion + source_region: _TensorRegion + destination_region: _TensorRegion + source_participants: tuple[int, ...] + destination_participants: tuple[int, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "source_participants", tuple(self.source_participants)) + object.__setattr__( + self, "destination_participants", tuple(self.destination_participants) + ) + if not ( + self.logical_region.numel + == self.source_region.numel + == self.destination_region.numel + ): + raise ValueError("route endpoint regions must have equal numel") + if len(set(self.source_participants)) != len(self.source_participants) or len( + set(self.destination_participants) + ) != len(self.destination_participants): + raise ValueError("route endpoint participants must be unique") + + +@dataclass(frozen=True, slots=True) +class _RedistributionPlan: + """Transport-neutral exact region partitions in both directions.""" + + participants: tuple[int, ...] + logical_shape: tuple[int, ...] + storage_partitions: tuple[_ParticipantPartition, ...] + compute_partitions: tuple[_ParticipantPartition, ...] + storage_to_compute_routes: tuple[_TensorRegionRoute, ...] + compute_to_storage_routes: tuple[_TensorRegionRoute, ...] + + def __post_init__(self) -> None: + object.__setattr__(self, "participants", tuple(self.participants)) + object.__setattr__(self, "logical_shape", tuple(self.logical_shape)) + object.__setattr__(self, "storage_partitions", tuple(self.storage_partitions)) + object.__setattr__(self, "compute_partitions", tuple(self.compute_partitions)) + object.__setattr__( + self, "storage_to_compute_routes", tuple(self.storage_to_compute_routes) + ) + object.__setattr__( + self, "compute_to_storage_routes", tuple(self.compute_to_storage_routes) + ) + if len(set(self.participants)) != len(self.participants): + raise ValueError("redistribution participants must be unique") + for name, partitions in ( + ("storage", self.storage_partitions), + ("compute", self.compute_partitions), + ): + if tuple(partition.participant for partition in partitions) != tuple( + self.participants + ): + raise ValueError( + f"{name} partitions must follow the redistribution participants" + ) + _validate_tensor_region_partition( + tuple( + region + for partition in partitions + for region in partition.logical_regions + ), + self.logical_shape, + direction=f"{name} logical partition", + ) + + all_routes = self.storage_to_compute_routes + self.compute_to_storage_routes + if any( + not route.source_participants or not route.destination_participants + for route in all_routes + ): + raise ValueError("redistribution routes require sources and destinations") + participant_set = set(self.participants) + if any( + participant not in participant_set + for route in all_routes + for participant in ( + route.source_participants + route.destination_participants + ) + ): + raise ValueError("redistribution route references an unknown participant") + + def route_key(route: _TensorRegionRoute) -> tuple[Any, ...]: + return ( + route.logical_region.offsets, + route.logical_region.shape, + route.source_region.offsets, + route.source_region.shape, + route.destination_region.offsets, + route.destination_region.shape, + route.source_participants, + route.destination_participants, + ) + + mirrored_forward_keys = sorted( + ( + route.logical_region.offsets, + route.logical_region.shape, + route.destination_region.offsets, + route.destination_region.shape, + route.source_region.offsets, + route.source_region.shape, + route.destination_participants, + route.source_participants, + ) + for route in self.storage_to_compute_routes + ) + reverse_keys = sorted(map(route_key, self.compute_to_storage_routes)) + if mirrored_forward_keys != reverse_keys: + raise ValueError( + "compute-to-storage routes must exactly invert storage-to-compute" + ) + for direction, routes in ( + ("storage-to-compute", self.storage_to_compute_routes), + ("compute-to-storage", self.compute_to_storage_routes), + ): + _validate_tensor_region_partition( + tuple(route.logical_region for route in routes), + self.logical_shape, + direction=direction, + ) + + storage_by_participant = { + partition.participant: partition for partition in self.storage_partitions + } + compute_by_participant = { + partition.participant: partition for partition in self.compute_partitions + } + for participant in self.participants: + storage = storage_by_participant[participant] + compute = compute_by_participant[participant] + endpoint_specs = ( + ( + "storage-to-compute source", + self.storage_to_compute_routes, + "source_participants", + "source_region", + storage, + ), + ( + "storage-to-compute destination", + self.storage_to_compute_routes, + "destination_participants", + "destination_region", + compute, + ), + ( + "compute-to-storage source", + self.compute_to_storage_routes, + "source_participants", + "source_region", + compute, + ), + ( + "compute-to-storage destination", + self.compute_to_storage_routes, + "destination_participants", + "destination_region", + storage, + ), + ) + for ( + direction, + routes, + participants_attr, + region_attr, + partition, + ) in endpoint_specs: + participant_routes = tuple( + route + for route in routes + if participant in getattr(route, participants_attr) + ) + _validate_regions_cover_partition( + tuple(route.logical_region for route in participant_routes), + partition.logical_regions, + self.logical_shape, + direction=f"{direction} {participant}", + ) + _validate_tensor_region_partition( + tuple(getattr(route, region_attr) for route in participant_routes), + partition.tensor_shape, + direction=f"{direction} tensor {participant}", + ) + + def storage_partition(self, participant: int) -> _ParticipantPartition: + return next( + partition + for partition in self.storage_partitions + if partition.participant == participant + ) + + def compute_partition(self, participant: int) -> _ParticipantPartition: + return next( + partition + for partition in self.compute_partitions + if partition.participant == participant + ) + + +def _validate_tensor_region_partition( + regions: tuple[_TensorRegion, ...], + logical_shape: tuple[int, ...], + *, + direction: str, +) -> None: + if any(size < 0 for size in logical_shape) or any( + len(region.offsets) != len(logical_shape) + or len(region.shape) != len(logical_shape) + or any( + offset < 0 or size < 0 or offset + size > logical_size + for offset, size, logical_size in zip( + region.offsets, region.shape, logical_shape, strict=True + ) + ) + for region in regions + ): + raise ValueError(f"{direction} regions must be in bounds") + + positive_regions = tuple(region for region in regions if region.numel) + for index, first in enumerate(positive_regions): + for second in positive_regions[index + 1 :]: + if all( + max(first_offset, second_offset) + < min( + first_offset + first_size, + second_offset + second_size, + ) + for first_offset, first_size, second_offset, second_size in zip( + first.offsets, + first.shape, + second.offsets, + second.shape, + strict=True, + ) + ): + raise NotImplementedError( + "overlapping logical tensor regions are not supported" + ) + + if sum(region.numel for region in regions) != math.prod(logical_shape): + raise ValueError(f"{direction} regions do not cover the logical tensor") + + +def _tensor_region_intersection_numel( + first: _TensorRegion, second: _TensorRegion +) -> int: + if len(first.shape) != len(second.shape): + return 0 + intersection_shape = tuple( + max( + 0, + min(first_offset + first_size, second_offset + second_size) + - max(first_offset, second_offset), + ) + for first_offset, first_size, second_offset, second_size in zip( + first.offsets, + first.shape, + second.offsets, + second.shape, + strict=True, + ) + ) + return math.prod(intersection_shape) + + +def _validate_regions_cover_partition( + regions: tuple[_TensorRegion, ...], + expected: tuple[_TensorRegion, ...], + logical_shape: tuple[int, ...], + *, + direction: str, +) -> None: + """Validate that nonoverlapping regions exactly cover an expected partition.""" + if any( + len(region.offsets) != len(logical_shape) + or len(region.shape) != len(logical_shape) + or any( + offset < 0 or size < 0 or offset + size > logical_size + for offset, size, logical_size in zip( + region.offsets, region.shape, logical_shape, strict=True + ) + ) + for region in regions + ): + raise ValueError(f"{direction} regions must be in bounds") + + positive_regions = tuple(region for region in regions if region.numel) + for index, first in enumerate(positive_regions): + for second in positive_regions[index + 1 :]: + if _tensor_region_intersection_numel(first, second): + raise NotImplementedError( + "overlapping logical tensor regions are not supported" + ) + + if sum(region.numel for region in regions) != sum( + region.numel for region in expected + ): + raise ValueError(f"{direction} regions do not cover the participant partition") + for region in positive_regions: + if ( + sum( + _tensor_region_intersection_numel(region, expected_region) + for expected_region in expected + ) + != region.numel + ): + raise ValueError(f"{direction} regions leave the participant partition") + + +def _build_owned_redistribution_plan( + storage_regions: Sequence[tuple[tuple[int, ...], _TensorRegion]], + *, + participants: tuple[int, ...], + owner: int, + logical_shape: tuple[int, ...], +) -> _RedistributionPlan: + """Build mirrored routes from one canonical region-to-holders mapping.""" + storage_partitions = [] + storage_mapping_by_participant = {} + for holders, logical_region in storage_regions: + if len(holders) != 1: + raise NotImplementedError( + "redistributed optimizer storage requires one holder per region" + ) + participant = holders[0] + tensor_region = _TensorRegion( + offsets=(0,) * len(logical_region.shape), + shape=logical_region.shape, + ) + storage_mapping_by_participant[participant] = (logical_region, tensor_region) + for participant in participants: + logical_region, tensor_region = storage_mapping_by_participant[participant] + storage_partitions.append( + _ParticipantPartition( + participant=participant, + tensor_shape=tensor_region.shape, + logical_regions=(logical_region,), + ) + ) + + full_region = _TensorRegion( + offsets=(0,) * len(logical_shape), + shape=logical_shape, + ) + compute_partitions = tuple( + _ParticipantPartition( + participant=participant, + tensor_shape=logical_shape if participant == owner else (0,), + logical_regions=(full_region,) if participant == owner else (), + ) + for participant in participants + ) + return _RedistributionPlan( + participants=participants, + logical_shape=logical_shape, + storage_partitions=tuple(storage_partitions), + compute_partitions=compute_partitions, + storage_to_compute_routes=tuple( + _TensorRegionRoute( + logical_region=logical_region, + source_region=storage_mapping_by_participant[holders[0]][1], + destination_region=logical_region, + source_participants=holders, + destination_participants=(owner,), + ) + for holders, logical_region in storage_regions + ), + compute_to_storage_routes=tuple( + _TensorRegionRoute( + logical_region=logical_region, + source_region=logical_region, + destination_region=storage_mapping_by_participant[holders[0]][1], + source_participants=(owner,), + destination_participants=holders, + ) + for holders, logical_region in storage_regions + ), + ) + + +@dataclass(frozen=True, slots=True) +class _PackedSpan: + """Physical packed-buffer location for an endpoint tensor region.""" + + region: _TensorRegion + buffer_offset: int + + @property + def numel(self) -> int: + return self.region.numel + + +@dataclass(frozen=True, slots=True) +class _PackedAllToAllSchedule: + process_group: dist.ProcessGroup + participants: tuple[int, ...] + local_participant: int + input_split_sizes: tuple[int, ...] + output_split_sizes: tuple[int, ...] + input_spans_by_parameter: tuple[tuple[_PackedSpan, ...], ...] + output_spans_by_parameter: tuple[tuple[_PackedSpan, ...], ...] + + @property + def input_buffer_numel(self) -> int: + return sum(self.input_split_sizes) + + @property + def output_buffer_numel(self) -> int: + return sum(self.output_split_sizes) + + def execute(self, output: Tensor, input: Tensor) -> None: + dist.all_to_all_single( + output[: self.output_buffer_numel], + input[: self.input_buffer_numel], + output_split_sizes=list(self.output_split_sizes), + input_split_sizes=list(self.input_split_sizes), + group=self.process_group, + ) + + +@dataclass(frozen=True, slots=True) +class _LocalSchedule: + participants: tuple[int, ...] = () + local_participant: int = -1 + input_spans_by_parameter: tuple[tuple[_PackedSpan, ...], ...] = () + output_spans_by_parameter: tuple[tuple[_PackedSpan, ...], ...] = () + input_buffer_numel: int = 0 + output_buffer_numel: int = 0 + + def execute(self, output: Tensor, input: Tensor) -> None: + output[: self.output_buffer_numel].copy_(input[: self.input_buffer_numel]) + + +_CommunicationSchedule = _PackedAllToAllSchedule | _LocalSchedule + + +@dataclass(frozen=True, slots=True) +class _RedistributionGroup: + process_group: dist.ProcessGroup + participants: tuple[int, ...] + local_participant: int + + +@dataclass(slots=True) +class _BucketPlan(Generic[_ItemT]): + local_items: tuple[_ItemT, ...] + redistributed_items: tuple[_ItemT, ...] + redistribution_plans: tuple[_RedistributionPlan, ...] + group: _RedistributionGroup + storage_to_compute_schedule: _CommunicationSchedule + compute_to_storage_schedule: _CommunicationSchedule + dtype: torch.dtype + device: torch.device + + +@dataclass(frozen=True, slots=True) +class _BucketPlanningResult(Generic[_ItemT]): + plans: tuple[_BucketPlan[_ItemT], ...] + ordered_items: tuple[_ItemT, ...] + + +@dataclass(slots=True) +class _BucketWork(Generic[_ItemT]): + plan: _BucketPlan[_ItemT] + slot: _BufferSlot + storage_buffer: Tensor + compute_fragment_buffer: Tensor + forward_ready: torch.Event | None = None + compute_done: torch.Event | None = None + done: torch.Event | None = None + + +@dataclass(slots=True) +class _BufferSlot: + storage_exchange_storage: dict[tuple[torch.device, torch.dtype], Tensor] = field( + default_factory=dict + ) + compute_exchange_storage: dict[tuple[torch.device, torch.dtype], Tensor] = field( + default_factory=dict + ) + compute_storage: dict[tuple[torch.device, torch.dtype], Tensor] = field( + default_factory=dict + ) + storage_partition_storage: dict[tuple[torch.device, torch.dtype], Tensor] = field( + default_factory=dict + ) + + @staticmethod + def _ensure_capacity( + storage: dict[tuple[torch.device, torch.dtype], Tensor], + *, + numel: int, + dtype: torch.dtype, + device: torch.device, + ) -> Tensor: + key = (device, dtype) + buffer = storage.get(key) + if buffer is None or buffer.numel() < numel: + buffer = torch.empty(numel, dtype=dtype, device=device) + storage[key] = buffer + return buffer[:numel] + + def communication_buffers(self, plan: _BucketPlan[Any]) -> tuple[Tensor, Tensor]: + to_compute = plan.storage_to_compute_schedule + to_storage = plan.compute_to_storage_schedule + return ( + self._ensure_capacity( + self.storage_exchange_storage, + numel=max( + to_compute.input_buffer_numel, + to_storage.output_buffer_numel, + ), + dtype=plan.dtype, + device=plan.device, + ), + self._ensure_capacity( + self.compute_exchange_storage, + numel=max( + to_compute.output_buffer_numel, + to_storage.input_buffer_numel, + ), + dtype=plan.dtype, + device=plan.device, + ), + ) + + def compute_buffer( + self, + shape: torch.Size | tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, + ) -> Tensor: + return self._ensure_capacity( + self.compute_storage, + numel=math.prod(shape), + dtype=dtype, + device=device, + ).view(shape) + + def storage_partition_buffer( + self, + shape: torch.Size | tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, + ) -> Tensor: + return self._ensure_capacity( + self.storage_partition_storage, + numel=math.prod(shape), + dtype=dtype, + device=device, + ).view(shape) + + +@dataclass(slots=True) +class _CommunicationContext: + device_handle: ModuleType + transfer_stream: torch.Stream + slots: tuple[_BufferSlot, _BufferSlot] + + @classmethod + def create(cls, device: torch.device) -> _CommunicationContext: + device_handle = torch.get_device_module(device) + transfer_stream = device_handle.Stream(device=device, priority=0) + return cls( + device_handle=device_handle, + transfer_stream=transfer_stream, + slots=(_BufferSlot(), _BufferSlot()), + ) + + +class _BucketedRedistributionRuntime(Generic[_ItemT]): + """Execute bucket plans with one-bucket-ahead communication prefetch. + + ``prepare`` writes a Muon input into runtime-owned scratch, ``compute`` + updates its runtime-owned input in place, and ``finalize`` consumes a + runtime-owned result before reuse. Callbacks run under the stream selected + by the runtime and must not retain tensors, synchronize, or call + ``Tensor.record_stream()``. + """ + + def __init__(self, device: torch.device) -> None: + # pyrefly: ignore [read-only] + self._device = device + self._context: _CommunicationContext | None = None + + def run( + self, + plans: Sequence[_BucketPlan[_ItemT]], + *, + local_tensor_spec: Callable[ + [_ItemT], tuple[torch.Size, torch.dtype, torch.device] + ], + prepare: Callable[[_ItemT, Tensor], None], + compute: Callable[[_ItemT, Tensor], None], + finalize: Callable[[_ItemT, Tensor], None], + ) -> None: + if self._context is None: + self._context = _CommunicationContext.create(self._device) + context = self._context + handle = context.device_handle + caller = handle.current_stream(self._device) + context.transfer_stream.wait_stream(caller) + + previous: _BucketWork[_ItemT] | None = None + redistributed_index = 0 + try: + for plan in plans: + slot = context.slots[redistributed_index % 2] + if not plan.redistributed_items: + with handle.stream(caller): + self._compute_local( + plan, + slot, + local_tensor_spec=local_tensor_spec, + prepare=prepare, + compute=compute, + finalize=finalize, + ) + continue + + work = self._enqueue_storage_to_compute( + plan, + slot, + context, + prepare=prepare, + ) + redistributed_index += 1 + # Keep collective launches ahead of redistributed work: + # gather(current) -> return(previous) -> compute(current). + if previous is not None: + self._enqueue_compute_to_storage( + previous, context, finalize=finalize + ) + + self._compute_bucket( + work, + slot, + caller, + context, + local_tensor_spec=local_tensor_spec, + prepare=prepare, + compute=compute, + finalize=finalize, + ) + if previous is not None: + self._release(previous, caller) + previous = work + + if previous is not None: + self._enqueue_compute_to_storage(previous, context, finalize=finalize) + self._release(previous, caller) + except Exception: + # Preserve allocator lifetime ordering for work already enqueued on + # either stream. This is an error-path drain, not synchronization. + context.transfer_stream.wait_stream(caller) + caller.wait_stream(context.transfer_stream) + raise + + @staticmethod + def _enqueue_storage_to_compute( + plan: _BucketPlan[_ItemT], + slot: _BufferSlot, + context: _CommunicationContext, + *, + prepare: Callable[[_ItemT, Tensor], None], + ) -> _BucketWork[_ItemT]: + handle = context.device_handle + transfer = context.transfer_stream + with handle.stream(transfer): + storage_buffer, compute_fragment_buffer = slot.communication_buffers(plan) + work = _BucketWork(plan, slot, storage_buffer, compute_fragment_buffer) + _prepare_redistributed( + plan, + slot, + storage_buffer, + prepare=prepare, + ) + plan.storage_to_compute_schedule.execute( + output=compute_fragment_buffer, + input=storage_buffer, + ) + work.forward_ready = handle.Event() + work.forward_ready.record(transfer) + return work + + @staticmethod + def _compute_bucket( + work: _BucketWork[_ItemT], + slot: _BufferSlot, + caller_stream: torch.Stream, + context: _CommunicationContext, + *, + local_tensor_spec: Callable[ + [_ItemT], tuple[torch.Size, torch.dtype, torch.device] + ], + prepare: Callable[[_ItemT, Tensor], None], + compute: Callable[[_ItemT, Tensor], None], + finalize: Callable[[_ItemT, Tensor], None], + ) -> None: + assert work.forward_ready is not None + handle = context.device_handle + with handle.stream(caller_stream): + _BucketedRedistributionRuntime._compute_local( + work.plan, + slot, + local_tensor_spec=local_tensor_spec, + prepare=prepare, + compute=compute, + finalize=finalize, + ) + caller_stream.wait_event(work.forward_ready) + _compute_redistributed( + work, + slot, + compute=compute, + ) + work.compute_done = handle.Event() + work.compute_done.record(caller_stream) + + @staticmethod + def _enqueue_compute_to_storage( + work: _BucketWork[_ItemT], + context: _CommunicationContext, + *, + finalize: Callable[[_ItemT, Tensor], None], + ) -> None: + assert work.compute_done is not None + handle = context.device_handle + transfer = context.transfer_stream + with handle.stream(transfer): + transfer.wait_event(work.compute_done) + work.plan.compute_to_storage_schedule.execute( + output=work.storage_buffer, + input=work.compute_fragment_buffer, + ) + _finalize_redistributed(work, work.slot, finalize=finalize) + work.done = handle.Event() + work.done.record(transfer) + + @staticmethod + def _release(work: _BucketWork[_ItemT], caller_stream: torch.Stream) -> None: + assert work.done is not None + caller_stream.wait_event(work.done) + + @staticmethod + def _compute_local( + plan: _BucketPlan[_ItemT], + slot: _BufferSlot, + *, + local_tensor_spec: Callable[ + [_ItemT], tuple[torch.Size, torch.dtype, torch.device] + ], + prepare: Callable[[_ItemT, Tensor], None], + compute: Callable[[_ItemT, Tensor], None], + finalize: Callable[[_ItemT, Tensor], None], + ) -> None: + for item in plan.local_items: + shape, dtype, device = local_tensor_spec(item) + prepared = slot.compute_buffer(shape, dtype=dtype, device=device) + prepare(item, prepared) + compute(item, prepared) + finalize(item, prepared) + + +def _prepare_redistributed( + plan: _BucketPlan[_ItemT], + slot: _BufferSlot, + storage_buffer: Tensor, + *, + prepare: Callable[[_ItemT, Tensor], None], +) -> None: + schedule = plan.storage_to_compute_schedule + participant = plan.group.local_participant + for index, (item, redistribution_plan) in enumerate( + zip( + plan.redistributed_items, + plan.redistribution_plans, + strict=True, + ) + ): + partition = redistribution_plan.storage_partition(participant) + prepared = slot.storage_partition_buffer( + partition.tensor_shape, + dtype=plan.dtype, + device=plan.device, + ) + prepare(item, prepared) + spans = schedule.input_spans_by_parameter[index] + for span in spans: + packed = storage_buffer[ + span.buffer_offset : span.buffer_offset + span.numel + ] + packed.copy_(_tensor_region_view(prepared, span.region).reshape(-1)) + + +def _compute_redistributed( + work: _BucketWork[_ItemT], + slot: _BufferSlot, + *, + compute: Callable[[_ItemT, Tensor], None], +) -> None: + plan = work.plan + participant = plan.group.local_participant + to_compute = plan.storage_to_compute_schedule + to_storage = plan.compute_to_storage_schedule + for index, (item, redistribution_plan) in enumerate( + zip( + plan.redistributed_items, + plan.redistribution_plans, + strict=True, + ) + ): + partition = redistribution_plan.compute_partition(participant) + if not math.prod(partition.tensor_shape): + continue + received_spans = to_compute.output_spans_by_parameter[index] + compute_tensor = slot.compute_buffer( + partition.tensor_shape, + dtype=plan.dtype, + device=plan.device, + ) + for span in received_spans: + received = work.compute_fragment_buffer[ + span.buffer_offset : span.buffer_offset + span.numel + ] + _tensor_region_view(compute_tensor, span.region).copy_( + received.view(span.region.shape) + ) + + compute(item, compute_tensor) + + for span in to_storage.input_spans_by_parameter[index]: + packed = work.compute_fragment_buffer[ + span.buffer_offset : span.buffer_offset + span.numel + ] + packed.copy_(_tensor_region_view(compute_tensor, span.region).reshape(-1)) + + +def _finalize_redistributed( + work: _BucketWork[_ItemT], + slot: _BufferSlot, + *, + finalize: Callable[[_ItemT, Tensor], None], +) -> None: + plan = work.plan + participant = plan.group.local_participant + schedule = work.plan.compute_to_storage_schedule + for index, (item, redistribution_plan) in enumerate( + zip( + plan.redistributed_items, + plan.redistribution_plans, + strict=True, + ) + ): + partition = redistribution_plan.storage_partition(participant) + update = slot.storage_partition_buffer( + partition.tensor_shape, + dtype=plan.dtype, + device=plan.device, + ) + spans = schedule.output_spans_by_parameter[index] + for span in spans: + packed = work.storage_buffer[ + span.buffer_offset : span.buffer_offset + span.numel + ] + _tensor_region_view(update, span.region).copy_( + packed.view(span.region.shape) + ) + finalize(item, update) + + +def _resolve_routes_to_transfers( + routes: tuple[_TensorRegionRoute, ...], participants: tuple[int, ...] +) -> tuple[tuple[int, int, _TensorRegion, _TensorRegion], ...]: + participant_order = { + participant: index for index, participant in enumerate(participants) + } + transfers = [] + for route in routes: + sources = tuple( + sorted(route.source_participants, key=participant_order.__getitem__) + ) + for destination in route.destination_participants: + source = destination if destination in sources else sources[0] + transfers.append( + (source, destination, route.source_region, route.destination_region) + ) + return tuple(transfers) + + +def _packed_spans_by_parameter( + indexed_spans: list[tuple[int, _PackedSpan]], parameter_count: int +) -> tuple[tuple[_PackedSpan, ...], ...]: + return tuple( + tuple( + span + for span_parameter_index, span in indexed_spans + if span_parameter_index == parameter_index + ) + for parameter_index in range(parameter_count) + ) + + +def _lower_packed_all_to_all( + redistribution_plans: tuple[_RedistributionPlan, ...], + *, + storage_to_compute: bool, + process_group: dist.ProcessGroup, + local_participant: int, +) -> _PackedAllToAllSchedule: + """Lower nonempty plans with one shared participant order to packed A2A.""" + participants = redistribution_plans[0].participants + if storage_to_compute: + routes_by_parameter = tuple( + plan.storage_to_compute_routes for plan in redistribution_plans + ) + else: + routes_by_parameter = tuple( + plan.compute_to_storage_routes for plan in redistribution_plans + ) + transfers_by_parameter = tuple( + _resolve_routes_to_transfers(routes, participants) + for routes in routes_by_parameter + ) + + input_split_sizes = [] + input_spans = [] + input_cursor = 0 + for destination in participants: + split_start = input_cursor + for parameter_index, transfers in enumerate(transfers_by_parameter): + for ( + source, + transfer_destination, + source_region, + _destination_region, + ) in transfers: + if source != local_participant or transfer_destination != destination: + continue + input_spans.append( + (parameter_index, _PackedSpan(source_region, input_cursor)) + ) + input_cursor += source_region.numel + input_split_sizes.append(input_cursor - split_start) + + output_split_sizes = [] + output_spans = [] + output_cursor = 0 + for source in participants: + split_start = output_cursor + for parameter_index, transfers in enumerate(transfers_by_parameter): + for ( + transfer_source, + destination, + _source_region, + destination_region, + ) in transfers: + if transfer_source != source or destination != local_participant: + continue + output_spans.append( + (parameter_index, _PackedSpan(destination_region, output_cursor)) + ) + output_cursor += destination_region.numel + output_split_sizes.append(output_cursor - split_start) + + return _PackedAllToAllSchedule( + process_group=process_group, + participants=participants, + local_participant=local_participant, + input_split_sizes=tuple(input_split_sizes), + output_split_sizes=tuple(output_split_sizes), + input_spans_by_parameter=_packed_spans_by_parameter( + input_spans, len(redistribution_plans) + ), + output_spans_by_parameter=_packed_spans_by_parameter( + output_spans, len(redistribution_plans) + ), + ) + + +def _device_mesh_ranks(mesh: DeviceMesh) -> tuple[int, ...]: + # Process groups can canonicalize rank order, but DeviceMesh order defines + # which logical shard each global rank holds. + return tuple(mesh.mesh.flatten().tolist()) + + +def _redistribution_group(mesh: DeviceMesh) -> _RedistributionGroup: + process_group = mesh.get_group() + participants = tuple(dist.get_process_group_ranks(process_group)) + return _RedistributionGroup( + process_group=process_group, + participants=participants, + local_participant=participants[dist.get_rank(process_group)], + ) + + +def _dtensor_storage_region_for_participant( + tensor: DTensor, + participant: int, +) -> _TensorRegion: + mesh_shape = tuple(tensor.device_mesh.shape) + flat_mesh_index = _device_mesh_ranks(tensor.device_mesh).index(participant) + mesh_coordinate = [0] * len(mesh_shape) + for mesh_axis in range(len(mesh_shape) - 1, -1, -1): + flat_mesh_index, mesh_coordinate[mesh_axis] = divmod( + flat_mesh_index, mesh_shape[mesh_axis] + ) + + local_shape = list(tensor.shape) + global_offsets = [0] * tensor.ndim + for mesh_axis, placement in enumerate(tensor.placements): + if type(placement) is not Shard: + raise ValueError( + "redistributed optimizer storage requires exact Shard placements" + ) + tensor_dim = placement.dim % tensor.ndim + local_size, global_offset = Shard.local_shard_size_and_offset( + tensor.shape[tensor_dim], + mesh_shape[mesh_axis], + mesh_coordinate[mesh_axis], + ) + local_shape[tensor_dim] = local_size + global_offsets[tensor_dim] = global_offset + return _TensorRegion( + offsets=tuple(global_offsets), + shape=tuple(local_shape), + ) + + +def _dtensor_storage_regions( + tensor: DTensor, + participants: tuple[int, ...], +) -> tuple[tuple[tuple[int, ...], _TensorRegion], ...]: + storage_participants = _device_mesh_ranks(tensor.device_mesh) + if storage_participants != participants: + raise ValueError( + "bucket mesh participants must match redistributed DTensor storage" + ) + return tuple( + ( + (participant,), + _dtensor_storage_region_for_participant(tensor, participant), + ) + for participant in participants + ) + + +def _build_bucket_plans( + items: Sequence[_ItemT], + specs: Sequence[BucketSpec], + *, + get_fqn: Callable[[_ItemT], str], + requires_owner: Callable[[_ItemT], bool], + get_storage_dtensor: Callable[[_ItemT], DTensor], + build_redistribution_plan: Callable[ + [_ItemT, _RedistributionGroup, int | None], + _RedistributionPlan | None, + ], +) -> _BucketPlanningResult[_ItemT]: + """Build ordered local and redistributed optimizer bucket plans. + + ``build_redistribution_plan`` receives a mesh-local owner rank exactly when + ``requires_owner`` is true. It returns ``None`` when storage is already + compute-ready or a transport-neutral plan for redistribution. This keeps + bucket ordering, owner validation, dtype validation, and packed + communication independent of a particular optimizer compute placement. + """ + resolved = _resolve_buckets(items, specs, get_fqn=get_fqn) + plans = [] + ordered_items = [] + for spec, bucket in zip(specs, resolved, strict=True): + if not bucket: + continue + group = _redistribution_group(spec.mesh) + sorted_bucket = tuple(sorted(bucket, key=get_fqn)) + expected_owners = { + get_fqn(item) for item in sorted_bucket if requires_owner(item) + } + provided_owners = set(spec.owner_rank_by_fqn) + missing_owners = expected_owners - provided_owners + if missing_owners: + raise ValueError( + f"bucket {spec.name!r} owner assignment must exactly cover " + "owner-requiring parameters; " + f"missing={sorted(missing_owners)}, " + f"extra={sorted(provided_owners - expected_owners)}" + ) + + local_items_list = [] + redistributed_items_list = [] + redistribution_plans = [] + for item in sorted_bucket: + needs_owner = requires_owner(item) + owner_rank = spec.owner_rank_by_fqn[get_fqn(item)] if needs_owner else None + if owner_rank is not None and owner_rank not in range( + len(group.participants) + ): + raise ValueError( + f"bucket {spec.name!r} has owner outside its process group" + ) + item_plan = build_redistribution_plan(item, group, owner_rank) + if item_plan is None: + local_items_list.append(item) + continue + if item_plan.participants != group.participants: + raise ValueError( + f"bucket {spec.name!r} redistribution participants do not " + "match its process group" + ) + local_tensor = get_storage_dtensor(item).to_local() + storage_partition = item_plan.storage_partition(group.local_participant) + if tuple(local_tensor.shape) != storage_partition.tensor_shape: + raise ValueError( + f"bucket {spec.name!r} storage partition does not match its mesh" + ) + redistributed_items_list.append(item) + redistribution_plans.append(item_plan) + + local_items = tuple(local_items_list) + redistributed_items = tuple(redistributed_items_list) + # Size-one sharded storage may normalize to Replicate. In that case a + # static rank-0 owner entry is equivalent to the resolved local compute. + redundant_owners = { + get_fqn(item) + for item in local_items + if len(group.participants) == 1 + and spec.owner_rank_by_fqn.get(get_fqn(item)) == 0 + } + effective_provided_owners = provided_owners - redundant_owners + if effective_provided_owners != expected_owners: + raise ValueError( + f"bucket {spec.name!r} owner assignment must exactly cover " + "owner-requiring parameters; " + f"missing={sorted(expected_owners - effective_provided_owners)}, " + f"extra={sorted(effective_provided_owners - expected_owners)}" + ) + ordered_items.extend(local_items) + ordered_items.extend(redistributed_items) + + if not redistributed_items: + tensor = get_storage_dtensor(local_items[0]).to_local() + plans.append( + _BucketPlan( + local_items=local_items, + redistributed_items=(), + redistribution_plans=(), + group=group, + storage_to_compute_schedule=_LocalSchedule(), + compute_to_storage_schedule=_LocalSchedule(), + dtype=tensor.dtype, + device=tensor.device, + ) + ) + continue + + storage_dtensors = [get_storage_dtensor(item) for item in redistributed_items] + local_tensors = [tensor.to_local() for tensor in storage_dtensors] + dtype = local_tensors[0].dtype + device = local_tensors[0].device + if any( + tensor.dtype != dtype or tensor.device != device for tensor in local_tensors + ): + raise ValueError(f"bucket {spec.name!r} mixes dtype or device") + + redistribution_plans_tuple = tuple(redistribution_plans) + plans.append( + _BucketPlan( + local_items=local_items, + redistributed_items=redistributed_items, + redistribution_plans=redistribution_plans_tuple, + group=group, + storage_to_compute_schedule=_lower_packed_all_to_all( + redistribution_plans_tuple, + storage_to_compute=True, + process_group=group.process_group, + local_participant=group.local_participant, + ), + compute_to_storage_schedule=_lower_packed_all_to_all( + redistribution_plans_tuple, + storage_to_compute=False, + process_group=group.process_group, + local_participant=group.local_participant, + ), + dtype=dtype, + device=device, + ) + ) + + return _BucketPlanningResult( + plans=tuple(plans), + ordered_items=tuple(ordered_items), + ) + + +def _build_owned_bucket_plans( + items: Sequence[_ItemT], + specs: Sequence[BucketSpec], + *, + get_fqn: Callable[[_ItemT], str], + storage_is_compute_ready: Callable[[_ItemT], bool], + get_storage_dtensor: Callable[[_ItemT], DTensor], +) -> _BucketPlanningResult[_ItemT]: + """Compatibility wrapper for local and whole-tensor-owned compute.""" + + def build_owned_redistribution_plan( + item: _ItemT, + group: _RedistributionGroup, + owner_rank: int | None, + ) -> _RedistributionPlan | None: + if storage_is_compute_ready(item): + return None + assert owner_rank is not None + tensor = get_storage_dtensor(item) + return _build_owned_redistribution_plan( + _dtensor_storage_regions(tensor, group.participants), + participants=group.participants, + owner=group.participants[owner_rank], + logical_shape=tuple(tensor.shape), + ) + + return _build_bucket_plans( + items, + specs, + get_fqn=get_fqn, + requires_owner=lambda item: not storage_is_compute_ready(item), + get_storage_dtensor=get_storage_dtensor, + build_redistribution_plan=build_owned_redistribution_plan, + ) + + +def _validate_bucket_plans_across_ranks( + plans: Sequence[_BucketPlan[_ItemT]], + *, + item_signature: Callable[[_ItemT], tuple[Any, ...]], +) -> None: + """Collectively verify rank-stable plans before runtime communication. + + Every rank must provide the same plan count and process-group order so all + workers enter these validation collectives in the same sequence. + """ + for plan in plans: + description = ( + str(plan.dtype), + plan.device.type, + tuple( + _redistribution_plan_key(redistribution_plan) + for redistribution_plan in plan.redistribution_plans + ), + [ + item_signature(item) + for item in plan.local_items + plan.redistributed_items + ], + ) + digest = hashlib.sha256(repr(description).encode()).digest() + plan_hash = int.from_bytes(digest[:7], "little") + local_hash = torch.tensor(plan_hash, dtype=torch.int64, device=plan.device) + process_group = plan.group.process_group + gathered = [ + torch.empty_like(local_hash) + for _ in range(dist.get_world_size(process_group)) + ] + dist.all_gather(gathered, local_hash, group=process_group) + if any(value.item() != plan_hash for value in gathered): + raise RuntimeError("optimizer bucket plans differ across ranks") + + +def _tensor_region_view(tensor: Tensor, region: _TensorRegion) -> Tensor: + view = tensor[ + tuple( + slice(offset, offset + size) + for offset, size in zip(region.offsets, region.shape, strict=True) + ) + ] + return view + + +def _redistribution_plan_key(plan: _RedistributionPlan) -> tuple[Any, ...]: + def partition_key(partition: _ParticipantPartition) -> tuple[Any, ...]: + return ( + partition.participant, + partition.tensor_shape, + tuple( + (region.offsets, region.shape) for region in partition.logical_regions + ), + ) + + def route_key(route: _TensorRegionRoute) -> tuple[Any, ...]: + return ( + route.logical_region.offsets, + route.logical_region.shape, + route.source_region.offsets, + route.source_region.shape, + route.destination_region.offsets, + route.destination_region.shape, + route.source_participants, + route.destination_participants, + ) + + return ( + plan.participants, + plan.logical_shape, + tuple(map(partition_key, plan.storage_partitions)), + tuple(map(partition_key, plan.compute_partitions)), + tuple(map(route_key, plan.storage_to_compute_routes)), + tuple(map(route_key, plan.compute_to_storage_routes)), + ) diff --git a/torchtitan/components/distributed_optimizers/muon.py b/torchtitan/components/distributed_optimizers/muon.py new file mode 100644 index 0000000000..eb7ab1b295 --- /dev/null +++ b/torchtitan/components/distributed_optimizers/muon.py @@ -0,0 +1,741 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Muon compute placement and the internal DistributedMuon runtime.""" + +from __future__ import annotations + +import hashlib +import math +from collections.abc import Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass +from typing import Any, cast, overload + +import torch +from torch import Tensor +from torch.distributed.tensor import DTensor, Replicate, Shard +from torch.distributed.tensor.placement_types import _StridedShard +from torch.optim import Optimizer + +from .flex_optimizer_reshard import ( + _BucketedRedistributionRuntime, + _build_bucket_plans, + _build_owned_redistribution_plan, + _device_mesh_ranks, + _dtensor_storage_regions, + _RedistributionGroup, + _RedistributionPlan, + _validate_bucket_plans_across_ranks, + assign_balanced_owners, + BucketSpec, +) + + +__all__ = ["BucketSpec", "assign_balanced_owners", "Owned"] + + +@dataclass(frozen=True, slots=True) +class Owned: + """Require complete 2D matrix compute. + + This is a Muon compute placement, not a DTensor storage placement. + Replicated storage computes locally; sharded storage uses the parameter's + mesh-local owner from ``BucketSpec.owner_rank_by_fqn``. + """ + + +class DistributedMuon(Optimizer): + """Internal runtime constructed through ``build_distributed_muon``. + + Parameter groups, FQNs, storage layouts, compute layouts, and bucket plans + are frozen after construction. Every configured parameter must have a + layout-compatible DTensor gradient before each rank enters ``step()``. + + Batched matrix compute views use batched BF16 kernels. They implement the + same mathematical update as ``torch.optim.Muon`` running one matrix at a + time, but bitwise equality across the two kernel schedules is not part of + the contract. + """ + + def __init__( + self, + params: Iterable[dict[str, Any]], + *, + bucket_spec: Sequence[BucketSpec], + _prepared_compute_views: Mapping[str, _PreparedParameterComputeView], + lr: float = 1e-3, + weight_decay: float = 0.1, + momentum: float = 0.95, + nesterov: bool = True, + ns_coefficients: tuple[float, float, float] = (3.4445, -4.7750, 2.0315), + eps: float = 1e-7, + ns_steps: int = 5, + adjust_lr_fn: str | None = None, + ) -> None: + defaults = { + "lr": lr, + "weight_decay": weight_decay, + "momentum": momentum, + "nesterov": nesterov, + "ns_coefficients": ns_coefficients, + "eps": eps, + "ns_steps": ns_steps, + "adjust_lr_fn": adjust_lr_fn, + } + self._first_step_validated = False + self._prepared_compute_views = dict(_prepared_compute_views) + super().__init__(params, defaults) + tensor_device = self._validate_parameter_storage() + group_compute_placements = [] + for group in self.param_groups: + compute_placement = group.pop("_compute_placement", None) + group_compute_placements.append(compute_placement) + self._group_compute_placements = tuple(group_compute_placements) + + self._specs = tuple(bucket_spec) + self._validate_groups() + self._initialize_plan() + self._validate_plan_across_ranks() + self._redistribution_runtime = _BucketedRedistributionRuntime[ + _ParameterComputeLayout + ](tensor_device) + self._set_checkpoint_layout_fingerprints() + + @overload + def step(self, closure: None = None) -> None: + ... + + @overload + def step(self, closure: Callable[[], float]) -> float: + ... + + @torch.no_grad() + def step(self, closure: Callable[[], float] | None = None) -> float | None: + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + + self._preflight_step() + self._redistribution_runtime.run( + self._plans, + local_tensor_spec=self._local_tensor_spec, + prepare=self._prepare_local, + compute=self._compute_update, + finalize=self._apply_update, + ) + return loss + + def add_param_group(self, param_group: dict[str, Any]) -> None: + if hasattr(self, "_plans"): + raise RuntimeError("DistributedMuon parameter groups are frozen") + super().add_param_group(param_group) + + def state_dict(self) -> dict[str, Any]: + state_dict = super().state_dict() + for saved_group, current_group in zip( + state_dict["param_groups"], self.param_groups, strict=True + ): + for param_id, fqn in zip( + saved_group["params"], current_group["param_names"], strict=True + ): + state_dict["state"].setdefault(param_id, {})[ + _LAYOUT_FINGERPRINT_KEY + ] = self._layout_fingerprints_by_fqn[fqn] + return state_dict + + def load_state_dict(self, state_dict: dict[str, Any]) -> None: + saved_groups = state_dict.get("param_groups", ()) + saved_state = state_dict.get("state", {}) + for saved_group, current_group in zip( + saved_groups, self.param_groups, strict=True + ): + for param_id, fqn in zip( + saved_group["params"], current_group["param_names"], strict=True + ): + fingerprint = saved_state.get(param_id, {}).get(_LAYOUT_FINGERPRINT_KEY) + if fingerprint != self._layout_fingerprints_by_fqn[fqn]: + raise ValueError( + "checkpoint changed DistributedMuon's compute layout" + ) + super().load_state_dict(state_dict) + self._validate_plan_across_ranks() + self._first_step_validated = False + + def _validate_groups(self) -> None: + for group_index, group in enumerate(self.param_groups): + ns_steps = group["ns_steps"] + coefficients = group["ns_coefficients"] + if ( + group.get("fused") + or group.get("foreach") + or any( + not 0 <= group[name] + for name in ("lr", "weight_decay", "momentum", "eps") + ) + or not isinstance(ns_steps, int) + or not 0 <= ns_steps < 100 + or len(coefficients) != 3 + or not all(isinstance(value, (int, float)) for value in coefficients) + or group["adjust_lr_fn"] + not in (None, "original", "match_rms_adamw", "spectral_unclamped") + ): + raise ValueError(f"unsupported DistributedMuon group {group_index}") + + def _validate_parameter_storage(self) -> torch.device: + local_devices = set() + for group in self.param_groups: + for param in group["params"]: + if not isinstance(param, DTensor): + raise TypeError("DistributedMuon requires DTensor parameters") + local_device = param.to_local().device + local_devices.add(local_device) + if len(local_devices) != 1 or next(iter(local_devices)).type != "cuda": + raise ValueError("DistributedMuon requires one CUDA device per process") + return local_devices.pop() + + def _build_parameter_compute_layouts( + self, + ) -> tuple[_ParameterComputeLayout, ...]: + parameters = [] + seen_names = set() + seen_params = set() + for group_index, group in enumerate(self.param_groups): + params = group["params"] + names = group["param_names"] + for fqn, param in zip(names, params, strict=True): + if fqn in seen_names or id(param) in seen_params: + raise ValueError(f"duplicate Muon parameter {fqn!r}") + seen_names.add(fqn) + seen_params.add(id(param)) + parameters.append((group_index, fqn, param)) + + compute_layouts = [] + for group_index, fqn, param in parameters: + compute_placement = self._group_compute_placements[group_index] + prepared = self._prepared_compute_views[fqn] + global_compute_shape = torch.Size(prepared.global_compute_shape) + local_storage_view = prepared.local_storage_view + resolved_transition = _resolve_storage_to_compute_transition( + fqn, + param, + global_compute_shape, + local_storage_view, + compute_placement, + ) + compute_layouts.append( + _ParameterComputeLayout( + fqn=fqn, + param=param, + group_index=group_index, + compute_view_key=prepared.compute_view_key, + global_compute_shape=global_compute_shape, + local_storage_view=local_storage_view, + local_storage_signature=_local_storage_signature(param.to_local()), + compute_placement_key=resolved_transition.fingerprint_key, + storage_to_compute_transition=resolved_transition.storage_to_compute_transition, + ) + ) + return tuple(compute_layouts) + + def _initialize_plan(self) -> None: + compute_layouts = self._build_parameter_compute_layouts() + result = _build_bucket_plans( + compute_layouts, + self._specs, + get_fqn=lambda item: item.fqn, + requires_owner=lambda item: item.requires_owner, + get_storage_dtensor=lambda item: item.param, + build_redistribution_plan=_build_parameter_redistribution_plan, + ) + self._plans = result.plans + self._parameter_compute_layouts = result.ordered_items + + def _set_checkpoint_layout_fingerprints(self) -> None: + self._layout_fingerprints_by_fqn = {} + for layout in self._parameter_compute_layouts: + descriptor = ( + layout.fqn, + tuple(layout.param.shape), + layout.compute_view_key, + tuple(layout.global_compute_shape), + layout.compute_placement_key, + ) + self._layout_fingerprints_by_fqn[layout.fqn] = ( + _LAYOUT_FINGERPRINT_VERSION, + # Optimizer.load_state_dict rebuilds iterable state values via + # type(value)(generator), which round-trips bytes but not strings. + hashlib.sha256(repr(descriptor).encode()).digest(), + ) + + def _validate_plan_across_ranks(self) -> None: + _validate_bucket_plans_across_ranks( + self._plans, + item_signature=self._plan_item_signature, + ) + + def _plan_item_signature( + self, compute_layout: _ParameterComputeLayout + ) -> tuple[Any, ...]: + return ( + compute_layout.fqn, + compute_layout.group_index, + tuple(compute_layout.param.shape), + tuple(compute_layout.param.stride()), + str(compute_layout.param.dtype), + compute_layout.param.to_local().device.type, + tuple(compute_layout.global_compute_shape), + type(compute_layout.storage_to_compute_transition).__name__, + compute_layout.compute_view_key, + compute_layout.compute_placement_key, + _device_mesh_ranks(compute_layout.param.device_mesh), + tuple(map(str, compute_layout.param.placements)), + self._group_signature(compute_layout), + ) + + def _group(self, compute_layout: _ParameterComputeLayout) -> dict[str, Any]: + return self.param_groups[compute_layout.group_index] + + def _group_signature( + self, compute_layout: _ParameterComputeLayout + ) -> tuple[Any, ...]: + group = self._group(compute_layout) + return tuple( + group[key] + for key in ( + "lr", + "weight_decay", + "momentum", + "nesterov", + "ns_coefficients", + "eps", + "ns_steps", + "adjust_lr_fn", + ) + ) + + def _preflight_step(self) -> None: + """Fail the local worker before bucket communication on invalid input. + + TorchTitan's elastic launcher terminates peer workers after this error + escapes. Do not add a validation collective to the optimizer hot path. + """ + initialize_state = not self._first_step_validated + missing_gradients = [ + compute_layout.fqn + for compute_layout in self._parameter_compute_layouts + if compute_layout.param.grad is None + ] + if missing_gradients: + raise RuntimeError( + "DistributedMuon requires every configured gradient before " + f"step(); missing gradients: {missing_gradients}" + ) + + for compute_layout in self._parameter_compute_layouts: + if ( + compute_layout.storage_is_compute_ready + and _local_storage_signature(compute_layout.param.to_local()) + != compute_layout.local_storage_signature + ): + raise RuntimeError( + f"parameter local storage changed for {compute_layout.fqn!r}; " + "rebuild DistributedMuon" + ) + gradients = [] + for compute_layout in self._parameter_compute_layouts: + grad = self._gradient(compute_layout) + gradients.append((compute_layout, grad)) + if initialize_state: + self._validate_momentum(compute_layout) + + # State creation happens only after every gradient and existing state + # tensor has passed validation, so a deterministic input error cannot + # partially update an earlier bucket. + if initialize_state: + for compute_layout, grad in gradients: + self._momentum(compute_layout, grad) + self._first_step_validated = True + + @staticmethod + def _has_storage_layout( + tensor: DTensor, compute_layout: _ParameterComputeLayout + ) -> bool: + local = tensor.to_local() + param_local = compute_layout.param.to_local() + return ( + tensor.shape == compute_layout.param.shape + and tensor.stride() == compute_layout.param.stride() + and _device_mesh_ranks(tensor.device_mesh) + == _device_mesh_ranks(compute_layout.param.device_mesh) + and tensor.placements == compute_layout.param.placements + and local.shape == param_local.shape + and local.stride() == param_local.stride() + and local.dtype == param_local.dtype + and local.device == param_local.device + and local.is_contiguous() + ) + + def _gradient(self, compute_layout: _ParameterComputeLayout) -> DTensor: + grad = compute_layout.param.grad + if not isinstance(grad, DTensor) or not self._has_storage_layout( + grad, compute_layout + ): + raise RuntimeError( + f"gradient storage layout changed for {compute_layout.fqn!r}" + ) + return grad + + def _validate_momentum(self, compute_layout: _ParameterComputeLayout) -> None: + state = self.state.get(compute_layout.param, {}) + fingerprint = state.get(_LAYOUT_FINGERPRINT_KEY) + if ( + fingerprint is not None + and fingerprint != self._layout_fingerprints_by_fqn[compute_layout.fqn] + ): + raise RuntimeError( + f"optimizer state layout changed for {compute_layout.fqn!r}" + ) + momentum = state.get("momentum_buffer") + if momentum is None: + return + if not isinstance(momentum, DTensor) or not self._has_storage_layout( + momentum, compute_layout + ): + raise RuntimeError( + f"momentum storage layout changed for {compute_layout.fqn!r}" + ) + + def _momentum( + self, compute_layout: _ParameterComputeLayout, grad: DTensor + ) -> DTensor: + state = self.state[compute_layout.param] + state[_LAYOUT_FINGERPRINT_KEY] = self._layout_fingerprints_by_fqn[ + compute_layout.fqn + ] + if "momentum_buffer" not in state: + state["momentum_buffer"] = torch.zeros_like( + grad, memory_format=torch.preserve_format + ) + return state["momentum_buffer"] + + def _update_local_momentum( + self, compute_layout: _ParameterComputeLayout + ) -> tuple[Tensor, Tensor, dict[str, Any]]: + grad = cast(DTensor, compute_layout.param.grad) + momentum = cast(DTensor, self.state[compute_layout.param]["momentum_buffer"]) + local_grad = grad.to_local().view_as(compute_layout.local_storage_view) + local_momentum = momentum.to_local().view_as(compute_layout.local_storage_view) + group = self._group(compute_layout) + local_momentum.lerp_(local_grad, 1 - group["momentum"]) + torch.autograd.graph.increment_version(momentum) + return local_grad, local_momentum, group + + @staticmethod + def _write_prepared( + group: dict[str, Any], grad: Tensor, momentum: Tensor, out: Tensor + ) -> None: + if group["nesterov"]: + torch.lerp( + grad, + momentum, + group["momentum"], + out=out, + ) + else: + out.copy_(momentum) + + def _prepare_local( + self, compute_layout: _ParameterComputeLayout, out: Tensor + ) -> None: + grad, momentum, group = self._update_local_momentum(compute_layout) + self._write_prepared(group, grad, momentum, out) + + def _compute_update( + self, compute_layout: _ParameterComputeLayout, compute: Tensor + ) -> None: + group = self._group(compute_layout) + _compute_muon_update( + compute, + ns_coefficients=group["ns_coefficients"], + ns_steps=group["ns_steps"], + eps=group["eps"], + out=compute, + ) + + def _apply_update( + self, compute_layout: _ParameterComputeLayout, direction: Tensor + ) -> None: + group = self._group(compute_layout) + local_param = ( + compute_layout.local_storage_view + if compute_layout.storage_is_compute_ready + else compute_layout.param.to_local() + ) + adjusted_lr = _adjust_learning_rate( + group["lr"], + group["adjust_lr_fn"], + compute_layout.global_compute_shape, + ) + local_param.mul_(1 - group["lr"] * group["weight_decay"]) + local_param.add_(direction, alpha=-adjusted_lr) + torch.autograd.graph.increment_version(compute_layout.param) + + @staticmethod + def _local_tensor_spec( + compute_layout: _ParameterComputeLayout, + ) -> tuple[torch.Size, torch.dtype, torch.device]: + tensor = compute_layout.local_storage_view + return tensor.shape, tensor.dtype, tensor.device + + +_LAYOUT_FINGERPRINT_KEY = "_distributed_muon_layout_fingerprint" +_LAYOUT_FINGERPRINT_VERSION = 1 + + +@dataclass(frozen=True, slots=True) +class _PreparedParameterComputeView: + compute_view_key: tuple[Any, ...] + global_compute_shape: torch.Size + local_storage_view: Tensor + + +@dataclass(frozen=True, slots=True) +class _ParameterComputeLayout: + fqn: str + param: DTensor + group_index: int + compute_view_key: tuple[Any, ...] + global_compute_shape: torch.Size + local_storage_view: Tensor + local_storage_signature: tuple[Any, ...] + compute_placement_key: tuple[Any, ...] + storage_to_compute_transition: _StorageToComputeTransition + + @property + def storage_is_compute_ready(self) -> bool: + return isinstance( + self.storage_to_compute_transition, _NoRedistributionTransition + ) + + @property + def requires_owner(self) -> bool: + return isinstance( + self.storage_to_compute_transition, _OwnedRedistributionTransition + ) + + +@dataclass(frozen=True, slots=True) +class _NoRedistributionTransition: + pass + + +@dataclass(frozen=True, slots=True) +class _OwnedRedistributionTransition: + pass + + +_StorageToComputeTransition = ( + _NoRedistributionTransition | _OwnedRedistributionTransition +) + + +@dataclass(frozen=True, slots=True) +class _ResolvedStorageToComputeTransition: + fingerprint_key: tuple[Any, ...] + storage_to_compute_transition: _StorageToComputeTransition + + +def _build_parameter_redistribution_plan( + compute_layout: _ParameterComputeLayout, + group: _RedistributionGroup, + owner_rank: int | None, +) -> _RedistributionPlan | None: + transition = compute_layout.storage_to_compute_transition + if isinstance(transition, _NoRedistributionTransition): + return None + + assert isinstance(transition, _OwnedRedistributionTransition) + assert owner_rank is not None + return _build_owned_redistribution_plan( + _dtensor_storage_regions(compute_layout.param, group.participants), + participants=group.participants, + owner=group.participants[owner_rank], + logical_shape=tuple(compute_layout.param.shape), + ) + + +def _resolve_storage_to_compute_transition( + fqn: str, + param: DTensor, + global_compute_shape: torch.Size, + local_storage_view: Tensor, + compute_placement: object, +) -> _ResolvedStorageToComputeTransition: + """Validate and canonicalize one storage-to-compute transition. + + ``Owned`` accepts a 2D matrix stored as exact ``Shard`` on a 1D mesh and + redistributes it to its configured owner. ``Shard(0)`` accepts matrix + batches whose storage shards keep each matrix whole. Fully replicated + storage computes locally under either compatible declaration, including + when a size-one mesh axis has normalized sharded storage to replication. + """ + local = param.to_local() + if torch.is_complex(param) or param.ndim < 2 or not local.is_contiguous(): + raise ValueError(f"Muon parameter {fqn!r} has unsupported shape or storage") + + replicated_storage = _has_replicated_storage(param) + if isinstance(compute_placement, Shard): + if ( + len(global_compute_shape) >= 3 + and _normalize_dim(compute_placement.dim, len(global_compute_shape)) == 0 + and ( + ( + replicated_storage + and local_storage_view.shape == global_compute_shape + ) + or ( + local_storage_view.shape[1:] == global_compute_shape[1:] + and _has_dim0_sharded_storage(param) + ) + ) + ): + return _ResolvedStorageToComputeTransition( + fingerprint_key=("shard", 0), + storage_to_compute_transition=_NoRedistributionTransition(), + ) + elif ( + isinstance(compute_placement, Owned) + and len(global_compute_shape) == 2 + and param.ndim == 2 + ): + if replicated_storage: + return _ResolvedStorageToComputeTransition( + fingerprint_key=("owned",), + storage_to_compute_transition=_NoRedistributionTransition(), + ) + if _has_owned_sharded_storage(param): + return _ResolvedStorageToComputeTransition( + fingerprint_key=("owned",), + storage_to_compute_transition=_OwnedRedistributionTransition(), + ) + raise ValueError(f"unsupported storage-to-compute layout for {fqn!r}") + + +def _has_replicated_storage(param: DTensor) -> bool: + return all(type(placement) is Replicate for placement in param.placements) + + +def _has_dim0_sharded_storage(param: DTensor) -> bool: + has_shard = False + for placement in param.placements: + # FSDP2 emits _StridedShard when a later TP/EP axis already shards + # this dimension. Keep the allowlist exact so new placements fail closed. + if type(placement) in (Shard, _StridedShard): + shard = cast(Shard | _StridedShard, placement) + if shard.dim % param.ndim != 0: + return False + has_shard = True + elif type(placement) is not Replicate: + return False + return has_shard + + +def _has_owned_sharded_storage(param: DTensor) -> bool: + return ( + param.device_mesh.ndim == 1 + and len(param.placements) == 1 + and type(param.placements[0]) is Shard + ) + + +def _normalize_dim(dim: int, ndim: int) -> int: + normalized = dim if dim >= 0 else dim + ndim + if normalized < 0 or normalized >= ndim: + raise ValueError(f"dimension {dim} is invalid for a rank-{ndim} tensor") + return normalized + + +def _local_storage_signature(tensor: Tensor) -> tuple[Any, ...]: + return ( + tensor.data_ptr(), + tensor.storage_offset(), + tuple(tensor.shape), + tuple(tensor.stride()), + tensor.dtype, + tensor.device, + ) + + +# Keep the functional math aligned with torch.optim.Muon while owning the +# implementation here so the distributed runtime has no Muon dependency. +def _adjust_learning_rate( + lr: float, + adjust_lr_fn: str | None, + compute_matrix_shape: torch.Size, +) -> float: + rows, columns = compute_matrix_shape[-2:] + if adjust_lr_fn is None or adjust_lr_fn == "original": + ratio = math.sqrt(max(1.0, rows / columns)) + elif adjust_lr_fn == "match_rms_adamw": + ratio = 0.2 * math.sqrt(max(rows, columns)) + elif adjust_lr_fn == "spectral_unclamped": + ratio = math.sqrt(rows / columns) + else: + raise ValueError(f"unsupported adjust_lr_fn {adjust_lr_fn!r}") + return lr * ratio + + +def _compute_muon_update( + prepared: Tensor, + *, + ns_coefficients: tuple[float, float, float], + ns_steps: int, + eps: float, + out: Tensor, +) -> Tensor: + direction = _zeropower_via_newtonschulz( + prepared, + ns_coefficients=ns_coefficients, + ns_steps=ns_steps, + eps=eps, + ) + out.copy_(direction) + return out + + +def _zeropower_via_newtonschulz( + update: Tensor, + *, + ns_coefficients: tuple[float, float, float], + ns_steps: int, + eps: float, +) -> Tensor: + """Compute Muon's approximate polar factor without using torch.optim.Muon.""" + a, b, c = ns_coefficients + result = update.to(dtype=torch.bfloat16, copy=True) + transposed = result.shape[-2] > result.shape[-1] + if transposed: + result = result.transpose(-2, -1) + result.div_(result.norm(dim=(-2, -1), keepdim=True).clamp_min(eps)) + + if result.ndim == 2: + for _ in range(ns_steps): + gram = result @ result.T + gram_update = torch.addmm(gram, gram, gram, beta=b, alpha=c) + result = torch.addmm(result, gram_update, result, beta=a) + else: + original_shape = result.shape + matrices = result.reshape(-1, *original_shape[-2:]) + # Batching reduces launch overhead, but bmm/baddbmm and independent + # mm/addmm calls can use different BF16 reduction orders. + for _ in range(ns_steps): + gram = matrices @ matrices.transpose(-2, -1) + gram_update = torch.baddbmm(gram, gram, gram, beta=b, alpha=c) + matrices = torch.baddbmm(matrices, gram_update, matrices, beta=a) + result = matrices.reshape(original_shape) + + return result.transpose(-2, -1) if transposed else result diff --git a/torchtitan/components/distributed_optimizers/muon_parameter_prep.py b/torchtitan/components/distributed_optimizers/muon_parameter_prep.py new file mode 100644 index 0000000000..f50830caf1 --- /dev/null +++ b/torchtitan/components/distributed_optimizers/muon_parameter_prep.py @@ -0,0 +1,264 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Public configuration and construction for DistributedMuon.""" + +from __future__ import annotations + +from collections.abc import Iterable, Sequence +from dataclasses import dataclass +from itertools import product +from typing import Any + +import torch +from torch import Tensor +from torch.distributed.tensor import DTensor, Replicate, Shard +from torch.distributed.tensor._utils import _compute_local_shape_and_global_offset + +from .flex_optimizer_reshard import _bind_bucket_configs, BucketConfig, BucketSpec +from .muon import _PreparedParameterComputeView, DistributedMuon, Owned + + +__all__ = [ + "BatchedMatrixComputeView", + "MuonComputeSharding", + "build_distributed_muon", +] + + +@dataclass(frozen=True, slots=True) +class BatchedMatrixComputeView: + """View 2D storage as matrices with batch and rows flattened into dim 0.""" + + num_matrices: int + matrices_flattened_into_dim: int = 0 + + def __post_init__(self) -> None: + if ( + isinstance(self.num_matrices, bool) + or not isinstance(self.num_matrices, int) + or self.num_matrices <= 0 + ): + raise ValueError("num_matrices must be a positive integer") + if self.matrices_flattened_into_dim != 0: + raise ValueError("only matrices_flattened_into_dim=0 is supported") + + def _resolve(self, storage_shape: torch.Size) -> _ResolvedBatchedMatrixView: + if ( + len(storage_shape) != 2 + or storage_shape[0] == 0 + or storage_shape[0] % self.num_matrices + ): + raise ValueError( + f"storage shape {tuple(storage_shape)} cannot be viewed as " + f"{self.num_matrices} matrices" + ) + return _ResolvedBatchedMatrixView( + matrix_rows=storage_shape[0] // self.num_matrices, + matrix_columns=storage_shape[1], + ) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class MuonComputeSharding: + """Define the logical Muon tensor and its compute placement. + + ``Owned`` requires one rank to compute a complete 2D matrix when storage + is sharded. ``Shard(0)`` computes rank-3-or-higher matrix batches locally + when storage shards preserve complete matrices. + """ + + # Applied before compute placement, so placement dimensions refer to the + # viewed tensor. A future view_after_placement mode can apply a local view + # after redistribution; that ordering is not supported yet. + view_before_placement: BatchedMatrixComputeView | None = None + placement: Owned | Shard + + def __post_init__(self) -> None: + if not isinstance(self.placement, (Owned, Shard)) or ( + self.view_before_placement is not None + and not isinstance(self.view_before_placement, BatchedMatrixComputeView) + ): + raise TypeError( + "MuonComputeSharding requires a supported view and placement" + ) + + def to_dict(self) -> dict: + return {"repr": repr(self)} + + +def build_distributed_muon( + params: Iterable[dict[str, Any]], + *, + bucket_spec: Sequence[BucketSpec] | None = None, + bucket_configs: Sequence[BucketConfig] | None = None, + **kwargs: Any, +) -> DistributedMuon: + """Prepare named DTensor parameter groups and construct DistributedMuon. + + Every group must provide aligned ``params`` and ``param_names`` plus one + ``compute_sharding`` contract. Exactly one of ``bucket_spec`` or + ``bucket_configs`` is required. Parameter groups and layouts are frozen + after construction because optimizer state and collectives depend on them. + """ + if (bucket_spec is None) == (bucket_configs is None): + raise ValueError("provide exactly one of bucket_spec or bucket_configs") + + prepared_params = [] + parameters_to_prepare = [] + for param_group in params: + group = dict(param_group) + compute_sharding = group.pop("compute_sharding") + compute_view = compute_sharding.view_before_placement + group["_compute_placement"] = compute_sharding.placement + raw_params = group.get("params", ()) + group_params = ( + (raw_params,) if isinstance(raw_params, Tensor) else tuple(raw_params) + ) + raw_param_names = group.get("param_names") + param_names = () if raw_param_names is None else tuple(raw_param_names) + if raw_param_names is None or len(group_params) != len(param_names): + raise ValueError("params and param_names must be aligned") + group["params"] = group_params + group["param_names"] = param_names + + for param, fqn in zip(group_params, param_names, strict=True): + parameters_to_prepare.append((param, fqn, compute_view)) + prepared_params.append(group) + + if bucket_configs is not None: + storage_by_fqn = { + fqn: param + for param, fqn, _compute_view in parameters_to_prepare + if isinstance(param, DTensor) + } + if len(storage_by_fqn) != len(parameters_to_prepare): + raise TypeError("bucket_configs require named DTensor parameters") + bucket_spec = _bind_bucket_configs(bucket_configs, storage_by_fqn) + else: + assert bucket_spec is not None + bucket_spec = tuple(bucket_spec) + + prepared_compute_views = {} + for param, fqn, compute_view in parameters_to_prepare: + global_storage_shape = torch.Size(param.shape) + resolved_view = None + if compute_view is not None and any( + type(placement) not in (Shard, Replicate) + for placement in getattr(param, "placements", ()) + ): + raise ValueError( + f"batched-matrix Muon parameter {fqn!r} requires exact " + "Shard or Replicate storage placements" + ) + if compute_view is not None: + resolved_view = compute_view._resolve(global_storage_shape) + if isinstance(param, DTensor): + _validate_batched_matrix_storage_alignment( + fqn, + param, + resolved_view, + ) + local_storage = param.to_local() if isinstance(param, DTensor) else param + compute_storage = ( + local_storage.detach() if isinstance(param, DTensor) else local_storage + ) + local_storage_shape = torch.Size(local_storage.shape) + if compute_view is None: + compute_view_key = ("identity",) + global_compute_shape = global_storage_shape + local_storage_view = compute_storage + else: + compute_view_key = ( + "batched_matrix", + compute_view.num_matrices, + compute_view.matrices_flattened_into_dim, + ) + assert resolved_view is not None + global_compute_shape = torch.Size( + ( + compute_view.num_matrices, + resolved_view.matrix_rows, + resolved_view.matrix_columns, + ) + ) + local_storage_view = compute_storage.view( + resolved_view.compute_shape(local_storage_shape) + ) + prepared_compute_views[fqn] = _PreparedParameterComputeView( + compute_view_key=compute_view_key, + global_compute_shape=global_compute_shape, + local_storage_view=local_storage_view, + ) + + return DistributedMuon( + prepared_params, + bucket_spec=bucket_spec, + _prepared_compute_views=prepared_compute_views, + **kwargs, + ) + + +@dataclass(frozen=True, slots=True) +class _ResolvedBatchedMatrixView: + matrix_rows: int + matrix_columns: int + + def compute_shape(self, storage_shape: torch.Size) -> torch.Size: + if ( + len(storage_shape) != 2 + or storage_shape[0] % self.matrix_rows + or storage_shape[1] != self.matrix_columns + ): + raise ValueError( + f"storage shape {tuple(storage_shape)} is not aligned to " + f"matrix shape {(self.matrix_rows, self.matrix_columns)}" + ) + return torch.Size( + ( + storage_shape[0] // self.matrix_rows, + self.matrix_rows, + self.matrix_columns, + ) + ) + + +def _validate_batched_matrix_storage_alignment( + fqn: str, + param: DTensor, + resolved_view: _ResolvedBatchedMatrixView, +) -> None: + """Validate every storage shard from globally identical DTensor metadata.""" + for placement in param.placements: + if type(placement) is Replicate: + continue + assert type(placement) is Shard + if placement.dim % param.ndim != 0: + raise ValueError( + f"batched-matrix Muon parameter {fqn!r} requires storage " + "shards along tensor dimension 0" + ) + + matrix_rows = resolved_view.matrix_rows + # Every rank must validate all coordinates before DistributedMuon performs + # collectives; checking only the local shard could strand its peers. + coordinates = product( + *(range(mesh_axis_size) for mesh_axis_size in param.device_mesh.shape) + ) + for coordinate in coordinates: + local_shape, global_offset = _compute_local_shape_and_global_offset( + param.shape, + param.device_mesh.shape, + list(coordinate), + param.placements, + ) + if local_shape[0] and ( + local_shape[0] % matrix_rows or global_offset[0] % matrix_rows + ): + raise ValueError( + f"batched-matrix Muon parameter {fqn!r} storage shards are not " + f"aligned to matrix rows of size {matrix_rows}" + ) diff --git a/torchtitan/components/optimizer.py b/torchtitan/components/optimizer.py index 1a7afaf9a6..ed410de0e6 100644 --- a/torchtitan/components/optimizer.py +++ b/torchtitan/components/optimizer.py @@ -23,6 +23,9 @@ init_optim_state, load_flat_optim_state_dict, ) +from torchtitan.components.distributed_optimizers.muon_parameter_prep import ( + build_distributed_muon, +) from torchtitan.config import Configurable from torchtitan.distributed import ParallelDims from torchtitan.tools.logging import logger @@ -105,6 +108,13 @@ class Config(Configurable.Config): regex pattern and a self-contained optimizer setup. Patterns are checked in order; first match wins.""" + optimizer_init_kwargs: dict[str, dict[str, Any]] = field(default_factory=dict) + """Programmatic optimizer-wide constructor arguments keyed by name. + + Use this for instance-wide objects such as communication bucket specs; + parameter-group hyperparameters belong in ``ParamGroupConfig``. + """ + implementation: Literal[ "for-loop", "foreach", "fused", "fused_opt_states_bf16" ] = "fused" @@ -127,14 +137,15 @@ class Config(Configurable.Config): model_parts: list[nn.Module] @staticmethod - def _resolve_optimizer_cls(name: str) -> type: - optimizer_classes = { + def _resolve_optimizer_factory(name: str) -> Callable[..., Optimizer]: + optimizer_factories: dict[str, Callable[..., Optimizer]] = { "Adam": torch.optim.Adam, "AdamW": torch.optim.AdamW, + "DistributedMuon": build_distributed_muon, } - if name not in optimizer_classes: + if name not in optimizer_factories: raise NotImplementedError(f"Optimizer {name} not added.") - return optimizer_classes[name] + return optimizer_factories[name] @staticmethod def _build_impl_kwargs(config: Config) -> dict[str, Any]: @@ -214,8 +225,11 @@ def __init__(self, config: Config, *, model_parts: list[nn.Module]) -> None: model, param_group_configs, impl_kwargs ) for opt_name, opt_param_groups in groups_by_opt_name.items(): - optimizer = self._resolve_optimizer_cls(opt_name)(opt_param_groups) - self.optimizers.append(optimizer) + optimizer = self._resolve_optimizer_factory(opt_name)( + opt_param_groups, + **config.optimizer_init_kwargs.get(opt_name, {}), + ) + self.optimizers.append(cast(T, optimizer)) self._log_optimizer(optimizer, part_idx, patterns_by_opt_name[opt_name]) for group in opt_param_groups: all_params.extend(group["params"]) diff --git a/torchtitan/config/configurable.py b/torchtitan/config/configurable.py index e79d797016..83ade47577 100644 --- a/torchtitan/config/configurable.py +++ b/torchtitan/config/configurable.py @@ -50,7 +50,10 @@ def _convert(val): if hasattr(val, "to_dict"): return val.to_dict() elif dataclasses.is_dataclass(val): - return dataclasses.asdict(val) + return { + f.name: _convert(getattr(val, f.name)) + for f in dataclasses.fields(val) + } elif isinstance(val, (list, tuple)): return type(val)(_convert(v) for v in val) elif isinstance(val, dict): diff --git a/torchtitan/models/kimi_k2_7/config_registry.py b/torchtitan/models/kimi_k2_7/config_registry.py index 5dae9b381b..2951262448 100644 --- a/torchtitan/models/kimi_k2_7/config_registry.py +++ b/torchtitan/models/kimi_k2_7/config_registry.py @@ -4,11 +4,28 @@ # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. +from collections.abc import Mapping + +from torch.distributed.tensor import Shard + from torchtitan.components.checkpoint import CheckpointManager +from torchtitan.components.distributed_optimizers.flex_optimizer_reshard import ( + assign_balanced_owners, + BucketConfig, +) +from torchtitan.components.distributed_optimizers.muon import Owned +from torchtitan.components.distributed_optimizers.muon_parameter_prep import ( + BatchedMatrixComputeView, + MuonComputeSharding, +) from torchtitan.components.loss import ChunkedLossWrapper, CrossEntropyLoss from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor -from torchtitan.components.optimizer import default_adamw +from torchtitan.components.optimizer import ( + default_adamw, + OptimizersContainer, + ParamGroupConfig, +) from torchtitan.components.tokenizer import MultiModalTokenizer from torchtitan.config import CompileConfig, ParallelismConfig, TrainingConfig from torchtitan.distributed.activation_checkpoint import FullAC, SelectiveAC @@ -40,6 +57,215 @@ def _mm_dataloader(dataset: str, **kwargs) -> MMDataLoader.Config: ) +def _per_head_muon_sharding(num_heads: int) -> MuonComputeSharding: + return MuonComputeSharding( + view_before_placement=BatchedMatrixComputeView( + num_matrices=num_heads, + matrices_flattened_into_dim=0, + ), + placement=Shard(0), + ) + + +def _kimi_text_distributed_muon_optimizer( + *, + num_layers: int, + num_owner_ranks: int, + lr: float, + attention_shardings: Mapping[str, MuonComputeSharding], + owned_parameter_numel_by_suffix: Mapping[str, int], +) -> OptimizersContainer.Config: + muon_kwargs = { + "lr": lr, + "weight_decay": 0.1, + "foreach": False, + # Kimi K2's MuonClip recipe uses 0.2 * sqrt(max(rows, columns)) + # for shape-consistent AdamW-scale updates instead of Muon's original + # aspect-ratio scaling. + "adjust_lr_fn": "match_rms_adamw", + } + adamw_kwargs = { + "lr": lr, + "betas": (0.9, 0.95), + "eps": 1e-8, + "weight_decay": 0.1, + } + expert_projections = ("w1_EFD", "w2_EDF", "w3_EFD") + param_groups = [ + ParamGroupConfig( + pattern=rf"attention\.{projection}\.weight$", + optimizer_name="DistributedMuon", + optimizer_kwargs={ + **muon_kwargs, + "compute_sharding": compute_sharding, + }, + ) + for projection, compute_sharding in attention_shardings.items() + ] + for projection in expert_projections: + param_groups.append( + ParamGroupConfig( + pattern=rf"routed_experts\.inner_experts\.{projection}$", + optimizer_name="DistributedMuon", + optimizer_kwargs={ + **muon_kwargs, + "compute_sharding": MuonComputeSharding(placement=Shard(0)), + }, + ) + ) + for pattern in ( + r"feed_forward\.w[123]\.weight$", + # Keep the 2D router gate on Muon: this follows the Kimi team's + # matrix-parameter rule, and Moonlight reports a larger SVD-entropy + # gain over AdamW for MoE router weights. + r"moe\.router\.gate\.weight$", + r"moe\.shared_experts\.w[123]\.weight$", + ): + param_groups.append( + ParamGroupConfig( + pattern=pattern, + optimizer_name="DistributedMuon", + optimizer_kwargs={ + **muon_kwargs, + "compute_sharding": MuonComputeSharding(placement=Owned()), + }, + ) + ) + param_groups.append( + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs=adamw_kwargs, + ) + ) + + def layer_fqns(layer_id: int) -> tuple[str, ...]: + prefix = f"layers.{layer_id}" + fqns = tuple( + f"{prefix}.attention.{projection}.weight" + for projection in attention_shardings + ) + if not layer_id: + fqns += tuple( + f"{prefix}.feed_forward.{projection}.weight" + for projection in ("w1", "w2", "w3") + ) + else: + fqns += tuple( + f"{prefix}.moe.routed_experts.inner_experts.{projection}" + for projection in expert_projections + ) + fqns += (f"{prefix}.moe.router.gate.weight",) + fqns += tuple( + f"{prefix}.moe.shared_experts.{projection}.weight" + for projection in ("w1", "w2", "w3") + ) + return fqns + + layer_bucket_fqns = tuple(layer_fqns(layer_id) for layer_id in range(num_layers)) + # Layer 0 has a much larger dense MLP, so keep it separate while amortizing + # collective launch overhead across pairs of MoE layers. + bucket_layer_ids = ((0,),) + tuple( + tuple(range(first_layer_id, min(first_layer_id + 2, num_layers))) + for first_layer_id in range(1, num_layers, 2) + ) + bucket_fqns = tuple( + tuple(fqn for layer_id in layer_ids for fqn in layer_bucket_fqns[layer_id]) + for layer_ids in bucket_layer_ids + ) + owner_rank_by_bucket = assign_balanced_owners( + bucket_fqns, + { + f"layers.{layer_id}.{suffix}": numel + for layer_id, fqns in enumerate(layer_bucket_fqns) + for suffix, numel in owned_parameter_numel_by_suffix.items() + if f"layers.{layer_id}.{suffix}" in fqns + }, + num_ranks=num_owner_ranks, + ) + bucket_configs = tuple( + BucketConfig( + name="layers." + "-".join(map(str, layer_ids)), + patterns=fqns, + owner_rank_by_fqn=owners, + mesh_axes=("dp_shard",), + ) + for layer_ids, fqns, owners in zip( + bucket_layer_ids, + bucket_fqns, + owner_rank_by_bucket, + strict=True, + ) + ) + return OptimizersContainer.Config( + implementation="foreach", + param_groups=param_groups, + optimizer_init_kwargs={ + "DistributedMuon": { + "bucket_configs": bucket_configs, + } + }, + ) + + +def _moonlight_distributed_muon_optimizer( + *, num_owner_ranks: int +) -> OptimizersContainer.Config: + per_head = _per_head_muon_sharding(num_heads=16) + return _kimi_text_distributed_muon_optimizer( + num_layers=27, + num_owner_ranks=num_owner_ranks, + lr=3e-4, + attention_shardings={ + "wq": per_head, + "wkv_a": MuonComputeSharding(placement=Owned()), + "wkv_b": per_head, + "wo": MuonComputeSharding(placement=Owned()), + }, + owned_parameter_numel_by_suffix={ + "attention.wkv_a.weight": 576 * 2048, + "attention.wo.weight": 2048 * 2048, + "feed_forward.w1.weight": 11264 * 2048, + "feed_forward.w2.weight": 2048 * 11264, + "feed_forward.w3.weight": 11264 * 2048, + "moe.router.gate.weight": 64 * 2048, + "moe.shared_experts.w1.weight": 2816 * 2048, + "moe.shared_experts.w2.weight": 2048 * 2816, + "moe.shared_experts.w3.weight": 2816 * 2048, + }, + ) + + +def _kimi_k2_5_distributed_muon_optimizer( + *, num_owner_ranks: int +) -> OptimizersContainer.Config: + per_head = _per_head_muon_sharding(num_heads=64) + return _kimi_text_distributed_muon_optimizer( + num_layers=61, + num_owner_ranks=num_owner_ranks, + lr=2.2e-4, + attention_shardings={ + "wq_a": MuonComputeSharding(placement=Owned()), + "wq_b": per_head, + "wkv_a": MuonComputeSharding(placement=Owned()), + "wkv_b": per_head, + "wo": MuonComputeSharding(placement=Owned()), + }, + owned_parameter_numel_by_suffix={ + "attention.wq_a.weight": 1536 * 7168, + "attention.wkv_a.weight": 576 * 7168, + "attention.wo.weight": 7168 * 8192, + "feed_forward.w1.weight": 18432 * 7168, + "feed_forward.w2.weight": 7168 * 18432, + "feed_forward.w3.weight": 18432 * 7168, + "moe.router.gate.weight": 384 * 7168, + "moe.shared_experts.w1.weight": 2048 * 7168, + "moe.shared_experts.w2.weight": 7168 * 2048, + "moe.shared_experts.w3.weight": 2048 * 7168, + }, + ) + + def kimi_k2_5_debugmodel() -> Trainer.Config: model_spec = model_registry("debugmodel") return Trainer.Config( @@ -108,6 +334,26 @@ def moonlight_16b_a3b() -> Trainer.Config: ) +def moonlight_16b_a3b_muon() -> Trainer.Config: + """Moonlight 16B-A3B with DistributedMuon for matrix parameters.""" + config = moonlight_16b_a3b() + num_owner_ranks = 8 + config.optimizer = _moonlight_distributed_muon_optimizer( + num_owner_ranks=num_owner_ranks + ) + config.parallelism = ParallelismConfig( + data_parallel_replicate_degree=1, + data_parallel_shard_degree=num_owner_ranks, + tensor_parallel_degree=1, + context_parallel_degree=1, + pipeline_parallel_degree=1, + expert_parallel_degree=4, + enable_sequence_parallel=False, + spmd_backend="spmd_types", + ) + return config + + def kimi_vl_a3b() -> Trainer.Config: """Kimi-VL A3B: Moonlight text tower + 2D MoonViT vision (image-text).""" model_spec = model_registry("Kimi-VL-A3B", attn_backend="flex") @@ -181,3 +427,23 @@ def kimi_k2_5() -> Trainer.Config: activation_checkpoint=FullAC.Config(), compile=compile_config, ) + + +def kimi_k2_5_muon() -> Trainer.Config: + """Full Kimi K2.5 with DistributedMuon for text-tower matrices.""" + config = kimi_k2_5() + num_owner_ranks = 64 + config.optimizer = _kimi_k2_5_distributed_muon_optimizer( + num_owner_ranks=num_owner_ranks, + ) + config.parallelism = ParallelismConfig( + data_parallel_replicate_degree=1, + data_parallel_shard_degree=num_owner_ranks, + tensor_parallel_degree=1, + context_parallel_degree=1, + pipeline_parallel_degree=1, + expert_parallel_degree=8, + enable_sequence_parallel=False, + spmd_backend="spmd_types", + ) + return config