MCPcopy Create free account
hub / github.com/ml-explore/mlx / test_groups

Method test_groups

python/tests/mpi_test_distributed.py:15–29  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

13 cls.rtol = 1e-4
14
15 def test_groups(self):
16 world = mx.distributed.init()
17 self.assertEqual(world.size(), 8)
18 self.assertTrue(0 <= world.rank() < 8)
19
20 world2 = mx.distributed.init()
21 self.assertEqual(world.size(), world2.size())
22 self.assertEqual(world.rank(), world2.rank())
23
24 sub = world.split(world.rank() % 2)
25 self.assertEqual(sub.size(), 4)
26 self.assertEqual(sub.rank(), world.rank() // 2)
27
28 sub = world.split(world.rank() // 2)
29 self.assertEqual(sub.size(), 2)
30
31 def test_all_reduce_extra(self):
32 world = mx.distributed.init()

Callers

nothing calls this directly

Calls 4

initMethod · 0.45
sizeMethod · 0.45
rankMethod · 0.45
splitMethod · 0.45

Tested by

no test coverage detected