| 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.""" |