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

Method test_split

python/tests/test_ops.py:1310–1329  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

1308 self.assertEqual(b.shape, a.shape)
1309
1310 def test_split(self):
1311 a = mx.array([1, 2, 3])
1312 splits = mx.split(a, 3)
1313 for e, x in enumerate(splits):
1314 self.assertEqual(x.item(), e + 1)
1315
1316 a = mx.array([[1, 2], [3, 4], [5, 6]])
1317 x, y, z = mx.split(a, 3, axis=0)
1318 self.assertEqual(x.tolist(), [[1, 2]])
1319 self.assertEqual(y.tolist(), [[3, 4]])
1320 self.assertEqual(z.tolist(), [[5, 6]])
1321
1322 with self.assertRaises(ValueError):
1323 mx.split(a, 3, axis=2)
1324
1325 a = mx.arange(8)
1326 x, y, z = mx.split(a, [1, 5])
1327 self.assertEqual(x.tolist(), [0])
1328 self.assertEqual(y.tolist(), [1, 2, 3, 4])
1329 self.assertEqual(z.tolist(), [5, 6, 7])
1330
1331 def test_split_invalid_num_splits(self):
1332 """Regression: split with num_splits <= 0 should raise, not crash."""

Callers

nothing calls this directly

Calls 3

itemMethod · 0.80
arrayMethod · 0.60
splitMethod · 0.45

Tested by

no test coverage detected