| 666 | self.assertEqual(x.tolist(), vals) |
| 667 | |
| 668 | def test_array_np_conversion(self): |
| 669 | # Shape test |
| 670 | a = np.array([]) |
| 671 | x = mx.array(a) |
| 672 | self.assertEqual(x.size, 0) |
| 673 | self.assertEqual(x.shape, (0,)) |
| 674 | self.assertEqual(x.dtype, mx.float32) |
| 675 | |
| 676 | a = np.array([[], [], []]) |
| 677 | x = mx.array(a) |
| 678 | self.assertEqual(x.size, 0) |
| 679 | self.assertEqual(x.shape, (3, 0)) |
| 680 | self.assertEqual(x.dtype, mx.float32) |
| 681 | |
| 682 | a = np.array([[[], []], [[], []], [[], []]]) |
| 683 | x = mx.array(a) |
| 684 | self.assertEqual(x.size, 0) |
| 685 | self.assertEqual(x.shape, (3, 2, 0)) |
| 686 | self.assertEqual(x.dtype, mx.float32) |
| 687 | |
| 688 | # Content test |
| 689 | a = 2.0 * np.ones((3, 5, 4)) |
| 690 | x = mx.array(a) |
| 691 | self.assertEqual(x.dtype, mx.float32) |
| 692 | self.assertEqual(x.ndim, 3) |
| 693 | self.assertEqual(x.shape, (3, 5, 4)) |
| 694 | |
| 695 | y = np.asarray(x) |
| 696 | self.assertTrue(np.allclose(a, y)) |
| 697 | |
| 698 | a = np.array(3, dtype=np.int32) |
| 699 | x = mx.array(a) |
| 700 | self.assertEqual(x.dtype, mx.int32) |
| 701 | self.assertEqual(x.ndim, 0) |
| 702 | self.assertEqual(x.shape, ()) |
| 703 | self.assertEqual(x.item(), 3) |
| 704 | |
| 705 | # mlx to numpy test |
| 706 | x = mx.array([True, False, True]) |
| 707 | y = np.asarray(x) |
| 708 | self.assertEqual(y.dtype, np.bool_) |
| 709 | self.assertEqual(y.ndim, 1) |
| 710 | self.assertEqual(y.shape, (3,)) |
| 711 | self.assertEqual(y[0], True) |
| 712 | self.assertEqual(y[1], False) |
| 713 | self.assertEqual(y[2], True) |
| 714 | |
| 715 | # complex64 mx <-> np |
| 716 | cvals = [0j, 1, 1 + 1j] |
| 717 | x = np.array(cvals) |
| 718 | y = mx.array(x) |
| 719 | self.assertEqual(y.dtype, mx.complex64) |
| 720 | self.assertEqual(y.shape, (3,)) |
| 721 | self.assertEqual(y.tolist(), cvals) |
| 722 | |
| 723 | y = mx.array([0j, 1, 1 + 1j]) |
| 724 | x = np.asarray(y) |
| 725 | self.assertEqual(x.dtype, np.complex64) |