From e1e1399e1b62d05160cb805dcf959c8485c1dd99 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Tue, 10 Mar 2026 11:47:25 +0100 Subject: [PATCH] Hybrid sharding (#3194) --- python/mlx/nn/utils.py | 112 +++++++++++++------------- python/tests/mlx_distributed_tests.py | 11 --- python/tests/nccl_test_distributed.py | 101 +++++++++++++++++++++++ 3 files changed, 155 insertions(+), 69 deletions(-) diff --git a/python/mlx/nn/utils.py b/python/mlx/nn/utils.py index 6d3223a7..b53e9efe 100644 --- a/python/mlx/nn/utils.py +++ b/python/mlx/nn/utils.py @@ -100,7 +100,6 @@ def average_gradients( gradients: Any, group: Optional[mx.distributed.Group] = None, all_reduce_size: int = 32 * 1024**2, - communication_type: Optional[mx.Dtype] = None, communication_stream: Optional[mx.Stream] = None, ): """Average the gradients across the distributed processes in the passed group. @@ -117,9 +116,6 @@ def average_gradients( all_reduce_size (int): Group arrays until their size in bytes exceeds this number. Perform one communication step per group of arrays. If less or equal to 0 array grouping is disabled. Default: ``32MiB``. - communication_type (Optional[mlx.core.Dtype]): If provided cast to this - type before performing the communication. Typically cast to a - smaller float to reduce the communication size. Default: ``None``. communication_stream (Optional[mlx.core.Stream]): The stream to use for the communication. If unspecified the default communication stream is used which can vary by back-end. Default: ``None``. @@ -130,13 +126,16 @@ def average_gradients( if N == 1: return gradients - def _average(x): - dt = x.dtype - x = x.astype(communication_type) if communication_type is not None else x - return mx.distributed.all_sum(x, stream=communication_stream).astype(dt) / N - if all_reduce_size <= 0: - return tree_map(_average, gradients) + return tree_map( + lambda x: mx.distributed.all_sum( + x, + group=group, + stream=communication_stream, + ) + / N, + gradients, + ) else: flat_grads = tree_flatten(gradients) @@ -148,15 +147,9 @@ def average_gradients( # We can't group them if they have mixed types if not all(dt == dtypes[0] for dt in dtypes): - return average_gradients(gradients, group, 0, communication_type) - itemsize = ( - communication_type.size - if communication_type is not None - else dtypes[0].size - ) - + return average_gradients(gradients, group, 0) # Gather the gradients in groups that are just above or equal to all_reduce_size - grad_groups = _group_by_size(keys, sizes, itemsize, all_reduce_size) + grad_groups = _group_by_size(keys, sizes, dtypes[0].size, all_reduce_size) # Concatenate-reduce-split new_flat_grads = [] @@ -165,7 +158,12 @@ def average_gradients( big_grad = mx.concatenate( [flat_grads[i][1].reshape(-1) for i in grad_group] ) - big_grad = _average(big_grad) + big_grad = ( + mx.distributed.all_sum( + big_grad, stream=communication_stream, group=group + ) + / N + ) big_grad = mx.split(big_grad, indices[1:-1]) new_flat_grads.extend( (keys[j], big_grad[i].reshape(shapes[j])) @@ -175,9 +173,9 @@ def average_gradients( return tree_unflatten(new_flat_grads) -def _clip_grads_fsdp(grads_slice, max_norm): +def _clip_grads_fsdp(grads_slice, max_norm, group=None): local_norm_sq = tree_reduce(lambda acc, g: acc + g.square().sum(), grads_slice, 0.0) - global_norm_sq = mx.distributed.all_sum(local_norm_sq) + global_norm_sq = mx.distributed.all_sum(local_norm_sq, group=group) grad_norm = mx.sqrt(global_norm_sq) normalizer = mx.minimum(max_norm / (grad_norm + 1e-6), 1.0) grads_slice = tree_map(lambda g: g * normalizer, grads_slice) @@ -189,9 +187,9 @@ def fsdp_apply_gradients( gradients, parameters, optimizer, - group=None, + fsdp_group=None, + dp_group=None, communication_size=32 * 1024**2, - communication_type=None, communication_stream=None, max_norm=None, ): @@ -208,20 +206,20 @@ def fsdp_apply_gradients( Args: gradients (Any): The Python tree containing the full gradients (it should have the same structure as ``parameters``). Each gradient's first - dimension must be divisible by the world size. + dimension must be divisible by ``fsdp_group.size()``. parameters (Any): The Python tree containing the full parameters (it should have the same structure across processes). Each parameter's first - dimension must be divisible by the world size. + dimension must be divisible by ``fsdp_group.size()``. optimizer: Optimizer with an ``apply_gradients`` method. - group (Optional[mlx.core.distributed.Group]): The group of processes for - communication. If ``None``, the global group is used. + fsdp_group (Optional[mlx.core.distributed.Group]): The group of processes + for FSDP sharding. If ``None``, the global group is used. + dp_group (Optional[mlx.core.distributed.Group]): The group of processes + for data-parallel gradient averaging. Required when ``fsdp_group`` is + smaller than the world (e.g. FSDP intra-node, DDP inter-node). Default: ``None``. communication_size (int): Group arrays until their size in bytes exceeds this number. Perform one communication step per group of arrays. If less or equal to 0 array grouping is disabled. Default: ``32MiB``. - communication_type (Optional[mlx.core.Dtype]): If provided cast to this - type before performing the communication. Typically cast to a - smaller float to reduce the communication size. Default: ``None``. communication_stream (Optional[mlx.core.Stream]): The stream to use for the communication. If unspecified the default communication stream is used which can vary by back-end. Default: ``None``. @@ -247,9 +245,8 @@ def fsdp_apply_gradients( ... ) >>> model.update(updated_params) """ - group = group or mx.distributed.init() - N = group.size() - rank = group.rank() + fsdp_group = fsdp_group or mx.distributed.init() + N = fsdp_group.size() * (dp_group.size() if dp_group is not None else 1) if N == 1: if max_norm is not None: @@ -260,45 +257,41 @@ def fsdp_apply_gradients( flat_grads = tree_flatten(gradients) flat_params = tree_flatten(parameters) - def _sum_scatter(x): - dt = x.dtype - x = x.astype(communication_type) if communication_type is not None else x - return ( - mx.distributed.sum_scatter( - x, group=group, stream=communication_stream - ).astype(dt) - / N - ) - - def _all_gather(x): - dt = x.dtype - x = x.astype(communication_type) if communication_type is not None else x - return mx.distributed.all_gather( - x, group=group, stream=communication_stream - ).astype(dt) - keys, shapes, sizes, dtypes = _extract_info(flat_grads) itemsize = dtypes[0].size groups = _group_by_size(keys, sizes, itemsize, communication_size) + S = fsdp_group.size() + fsdp_rank = fsdp_group.rank() # reduce-scatter gradients, shard parameters grad_slices = {} param_slices = {} for group_idx, arr_group in enumerate(groups): big_grad = mx.concatenate( - [flat_grads[i][1].reshape(N, -1) for i in arr_group], axis=1 + [flat_grads[i][1].reshape(S, -1) for i in arr_group], axis=1 ) - grad_slices[group_idx] = _sum_scatter(big_grad) + grad_slices[group_idx] = ( + mx.distributed.sum_scatter( + big_grad, group=fsdp_group, stream=communication_stream + ) + / N + ) + if dp_group is not None: + grad_slices[group_idx] = mx.distributed.all_sum( + grad_slices[group_idx], group=dp_group, stream=communication_stream + ) big_param = mx.concatenate( - [flat_params[i][1].reshape(N, -1) for i in arr_group], axis=1 + [flat_params[i][1].reshape(S, -1) for i in arr_group], axis=1 ) - param_slices[group_idx] = big_param[rank] + param_slices[group_idx] = big_param[fsdp_rank] # clip gradients if needed grad_norm = None if max_norm is not None: - grad_slices, grad_norm = _clip_grads_fsdp(grad_slices, max_norm) + grad_slices, grad_norm = _clip_grads_fsdp( + grad_slices, max_norm, group=fsdp_group + ) # optimizer step updated_param_slices = optimizer.apply_gradients(grad_slices, param_slices) @@ -306,9 +299,12 @@ def fsdp_apply_gradients( # all-gather and reconstruct new_flat = [] for group_idx, arr_group in enumerate(groups): - big_gathered = _all_gather(updated_param_slices[group_idx].reshape(1, -1)) - - split_sizes = [sizes[i] // N for i in arr_group] + big_gathered = mx.distributed.all_gather( + updated_param_slices[group_idx], + group=fsdp_group, + stream=communication_stream, + ) + split_sizes = [sizes[i] // S for i in arr_group] split_indices = [] acc = 0 for s in split_sizes: diff --git a/python/tests/mlx_distributed_tests.py b/python/tests/mlx_distributed_tests.py index 644da793..00d3182b 100644 --- a/python/tests/mlx_distributed_tests.py +++ b/python/tests/mlx_distributed_tests.py @@ -47,17 +47,6 @@ class MLXDistributedCommonTestCase(mlx_tests.MLXTestCase): self.assertTrue(all(mx.all(g == 1) for g in new_grads)) self.assertEqual(n_calls, 10) - n_calls = 0 - xtype = mx.float16 - new_grads = average_gradients( - grads, all_reduce_size=2 * 50, communication_type=mx.float16 - ) - mx.eval(new_grads) - self.assertEqual(len(new_grads), 10) - self.assertTrue(all(g.dtype == mx.float32 for g in new_grads)) - self.assertTrue(all(mx.all(g == 1) for g in new_grads)) - self.assertEqual(n_calls, 2) - finally: mx.distributed.all_sum = original_all_sum diff --git a/python/tests/nccl_test_distributed.py b/python/tests/nccl_test_distributed.py index f98b56cb..7a6bcce0 100644 --- a/python/tests/nccl_test_distributed.py +++ b/python/tests/nccl_test_distributed.py @@ -196,6 +196,107 @@ class TestNCCLDistributed(mlx_distributed_tests.MLXDistributedCommonTestCase): ), ) + def test_fsdp_ddp_apply_gradients(self): + world = mx.distributed.init() + N = world.size() + S = 4 + fsdp_group = world.split(world.rank() // S) + dp_group = world.split(world.rank() % S) + + self.assertEqual(fsdp_group.size(), S) + self.assertEqual(dp_group.size(), N // S) + + params = { + "w1": mx.ones((S * 10, 8)), + "w2": mx.ones((S * 20,)), + } + grads = { + "w1": mx.ones((S * 10, 8)) * 0.1, + "w2": mx.ones((S * 20,)) * 0.1, + } + + optimizer = optim.SGD(learning_rate=0.1) + updated = fsdp_apply_gradients( + grads, + params, + optimizer, + fsdp_group=fsdp_group, + dp_group=dp_group, + ) + mx.eval(updated) + + self.assertEqual(updated["w1"].shape, (S * 10, 8)) + self.assertEqual(updated["w2"].shape, (S * 20,)) + + self.assertTrue( + mx.allclose(updated["w1"], mx.ones((S * 10, 8)) * 0.99, atol=1e-6) + ) + self.assertTrue( + mx.allclose(updated["w2"], mx.ones((S * 20,)) * 0.99, atol=1e-6) + ) + + grads_big = { + "w1": mx.ones((S * 10, 8)) * 10.0, + "w2": mx.ones((S * 20,)) * 10.0, + } + + optimizer2 = optim.SGD(learning_rate=0.1) + clipped, grad_norm = fsdp_apply_gradients( + grads_big, + params, + optimizer2, + fsdp_group=fsdp_group, + dp_group=dp_group, + max_norm=1.0, + ) + mx.eval(clipped, grad_norm) + + self.assertIsNotNone(grad_norm) + expected_norm = mx.sqrt((S * 10 * 8 + S * 20) * 100.0) + self.assertTrue(mx.allclose(grad_norm, expected_norm, atol=1e-4, rtol=1e-4)) + self.assertEqual(clipped["w1"].shape, (S * 10, 8)) + self.assertEqual(clipped["w2"].shape, (S * 20,)) + + scale = 1.0 / expected_norm + expected_update = 1.0 - 0.1 * 10.0 * scale + self.assertTrue( + mx.allclose( + clipped["w1"], + mx.ones((S * 10, 8)) * expected_update, + atol=1e-4, + rtol=1e-4, + ) + ) + self.assertTrue( + mx.allclose( + clipped["w2"], + mx.ones((S * 20,)) * expected_update, + atol=1e-4, + rtol=1e-4, + ) + ) + + params_eq = {"w": mx.ones((S * 4,))} + grads_eq = {"w": mx.ones((S * 4,)) * 0.5} + + optimizer_hybrid = optim.SGD(learning_rate=0.1) + updated_hybrid = fsdp_apply_gradients( + grads_eq, + params_eq, + optimizer_hybrid, + fsdp_group=fsdp_group, + dp_group=dp_group, + ) + + optimizer_ddp = optim.SGD(learning_rate=0.1) + avg_grads = average_gradients(grads_eq) + updated_ddp = optimizer_ddp.apply_gradients(avg_grads, params_eq) + mx.eval(updated_hybrid, updated_ddp) + + self.assertTrue( + mx.allclose(updated_hybrid["w"], updated_ddp["w"], atol=1e-6, rtol=1e-4), + ) + def test_fsdp_peak_memory(self): world = mx.distributed.init() N = world.size()