[CUDA] Fsdp (easy) (#3130)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user