| 2126 | self.assertEqual(arr_pass.tolist(), [4, 5, 6]) |
| 2127 | |
| 2128 | def test_asarray(self): |
| 2129 | # List inputs |
| 2130 | self.assertEqual(mx.asarray([1, 2, 3]).tolist(), [1, 2, 3]) |
| 2131 | self.assertEqual(mx.asarray([[1, 2], [3, 4]]).tolist(), [[1, 2], [3, 4]]) |
| 2132 | |
| 2133 | # Tuple inputs |
| 2134 | self.assertEqual(mx.asarray((1, 2, 3)).tolist(), [1, 2, 3]) |
| 2135 | self.assertEqual(mx.asarray(((1, 2), (3, 4))).tolist(), [[1, 2], [3, 4]]) |
| 2136 | |
| 2137 | # Mixed nesting |
| 2138 | self.assertEqual(mx.asarray([(1, 2), (3, 4)]).tolist(), [[1, 2], [3, 4]]) |
| 2139 | self.assertEqual(mx.asarray(([1, 2], [3, 4])).tolist(), [[1, 2], [3, 4]]) |
| 2140 | |
| 2141 | # Scalar inputs |
| 2142 | self.assertEqual(mx.asarray(42).item(), 42) |
| 2143 | self.assertEqual(mx.asarray(3.14).item(), 3.140000104904175) |
| 2144 | self.assertEqual(mx.asarray(True).item(), True) |
| 2145 | self.assertEqual(mx.asarray(1 + 2j).item(), (1 + 2j)) |
| 2146 | |
| 2147 | # MLX array inputs |
| 2148 | arr = mx.array([1, 2, 3]) |
| 2149 | self.assertEqual(mx.asarray(arr).tolist(), [1, 2, 3]) |
| 2150 | |
| 2151 | arr_int = mx.array([1, 2, 3], dtype=mx.int32) |
| 2152 | arr_float = mx.asarray(arr_int, dtype=mx.float32) |
| 2153 | self.assertEqual(arr_float.dtype, mx.float32) |
| 2154 | self.assertEqual(arr_float.tolist(), [1.0, 2.0, 3.0]) |
| 2155 | |
| 2156 | # NumPy array inputs |
| 2157 | np_arr = np.array([1.0, 2.0, 3.0], dtype=np.float32) |
| 2158 | mx_arr = mx.asarray(np_arr) |
| 2159 | self.assertEqual(mx_arr.tolist(), [1.0, 2.0, 3.0]) |
| 2160 | self.assertEqual(mx_arr.dtype, mx.float32) |
| 2161 | |
| 2162 | # dtype parameter |
| 2163 | self.assertEqual(mx.asarray([1, 2, 3], dtype=mx.float32).dtype, mx.float32) |
| 2164 | self.assertEqual(mx.asarray(42, dtype=mx.float16).dtype, mx.float16) |
| 2165 | |
| 2166 | def test_to_scalar(self): |
| 2167 | a = mx.array(1) |