[CUDA] Fsdp (easy) (#3130)

This commit is contained in:
Anastasiia Filippova
2026-03-01 23:29:09 +01:00
committed by GitHub
parent 6482d13dd3
commit 72e04f7fb7
4 changed files with 311 additions and 20 deletions
+126
View File
@@ -1,8 +1,10 @@
# Copyright © 2024 Apple Inc.
import mlx.core as mx
import mlx.optimizers as optim
import mlx_distributed_tests
import mlx_tests
from mlx.nn.utils import average_gradients, fsdp_apply_gradients
class TestNCCLDistributed(mlx_distributed_tests.MLXDistributedCommonTestCase):
@@ -114,6 +116,130 @@ class TestNCCLDistributed(mlx_distributed_tests.MLXDistributedCommonTestCase):
self.assertEqual(y.shape, (sub.size() * 2, 2, 4))
self.assertTrue(mx.all(y == 1))
def test_fsdp_apply_gradients(self):
world = mx.distributed.init()
N = world.size()
params = {
"w1": mx.ones((N * 10, 8)),
"w2": mx.ones((N * 20,)),
}
grads = {
"w1": mx.ones((N * 10, 8)) * 0.1,
"w2": mx.ones((N * 20,)) * 0.1,
}
optimizer = optim.SGD(learning_rate=0.1)
updated_params_fsdp = fsdp_apply_gradients(grads, params, optimizer)
mx.eval(updated_params_fsdp)
self.assertEqual(updated_params_fsdp["w1"].shape, (N * 10, 8))
self.assertEqual(updated_params_fsdp["w2"].shape, (N * 20,))
self.assertTrue(
mx.allclose(
updated_params_fsdp["w1"], mx.ones((N * 10, 8)) * 0.99, atol=1e-6
)
)
self.assertTrue(
mx.allclose(updated_params_fsdp["w2"], mx.ones((N * 20,)) * 0.99, atol=1e-6)
)
grads = {
"w1": mx.ones((N * 10, 8)) * 10.0,
"w2": mx.ones((N * 20,)) * 10.0,
}
new_params_clipped, grad_norm = fsdp_apply_gradients(
grads, params, optimizer, max_norm=1.0
)
mx.eval(new_params_clipped, grad_norm)
self.assertIsNotNone(grad_norm)
expected_norm = mx.sqrt((N * 10 * 8 + N * 20) * 100.0)
self.assertTrue(mx.allclose(grad_norm, expected_norm, atol=1e-4, rtol=1e-4))
self.assertEqual(new_params_clipped["w1"].shape, (N * 10, 8))
self.assertEqual(new_params_clipped["w2"].shape, (N * 20,))
scale = 1.0 / expected_norm
expected_update = 1.0 - 0.1 * 10.0 * scale
self.assertTrue(
mx.allclose(
new_params_clipped["w1"],
mx.ones((N * 10, 8)) * expected_update,
atol=1e-4,
rtol=1e-4,
)
)
self.assertTrue(
mx.allclose(
new_params_clipped["w2"],
mx.ones((N * 20,)) * expected_update,
atol=1e-4,
rtol=1e-4,
)
)
params = {"w": mx.ones((N * 4,))}
grads = {"w": mx.ones((N * 4,)) * 0.5}
optimizer_fsdp = optim.SGD(learning_rate=0.1)
updated_params_fsdp = fsdp_apply_gradients(grads, params, optimizer_fsdp)
optimizer_ddp = optim.SGD(learning_rate=0.1)
avg_grads = average_gradients(grads)
updated_params_ddp = optimizer_ddp.apply_gradients(avg_grads, params)
mx.eval(updated_params_ddp, updated_params_fsdp)
self.assertTrue(
mx.allclose(
updated_params_fsdp["w"], updated_params_ddp["w"], atol=1e-6, rtol=1e-4
),
)
def test_fsdp_peak_memory(self):
world = mx.distributed.init()
N = world.size()
mx.random.seed(42)
params = {
"w1": mx.random.normal((N * 1024, 1024)),
"w2": mx.random.normal((N * 2048, 512)),
}
grads = {
"w1": mx.random.normal((N * 1024, 1024)),
"w2": mx.random.normal((N * 2048, 512)),
}
mx.eval(params, grads)
optimizer_ddp = optim.Adam(learning_rate=0.01)
optimizer_fsdp = optim.Adam(learning_rate=0.01)
def pseudo_step_ddp(grads, params, optimizer):
grads = average_gradients(grads)
grads, grad_norm = optim.clip_grad_norm(grads, max_norm=1.0)
params = optimizer.apply_gradients(grads, params)
return grad_norm, params
def pseudo_step_fsdp(grads, params, optimizer):
params, grad_norm = fsdp_apply_gradients(
grads, params, optimizer, max_norm=1.0
)
return grad_norm, params
mx.reset_peak_memory()
for i in range(10):
grad_norm, params = pseudo_step_ddp(grads, params, optimizer_ddp)
mx.eval(grad_norm, params)
ddp_peak_memory = mx.get_peak_memory()
mx.reset_peak_memory()
for i in range(10):
grad_norm, params = pseudo_step_fsdp(grads, params, optimizer_fsdp)
mx.eval(grad_norm, params)
fsdp_peak_memory = mx.get_peak_memory()
self.assertTrue(fsdp_peak_memory < ddp_peak_memory)
if __name__ == "__main__":
mlx_tests.MLXTestRunner()