Fix np bfloat16 misinterpreted as complex (#3146)

Co-authored-by: Cheng <[email protected]>
This commit is contained in:
Kellen Sun
2026-03-31 08:04:55 +09:00
committed by GitHub
co-authored by Cheng
parent 0ff1115a46
commit 8a6d28713c
6 changed files with 118 additions and 56 deletions
+32
View File
@@ -15,6 +15,13 @@ try:
except ImportError as e:
has_torch = False
try:
import ml_dtypes
has_ml_dtypes = True
except ImportError:
has_ml_dtypes = False
class TestBF16(mlx_tests.MLXTestCase):
def __test_ops(
@@ -191,6 +198,31 @@ class TestBF16(mlx_tests.MLXTestCase):
self.assertEqual(a_mx.dtype, mx.bfloat16)
self.assertTrue(mx.array_equal(a_mx, expected))
@unittest.skipIf(not has_ml_dtypes, "requires ml_dtypes")
def test_conversion_ml_dtypes(self):
x_scalar = np.array(1.5, dtype=ml_dtypes.bfloat16)
a_scalar = mx.array(x_scalar)
self.assertEqual(a_scalar.dtype, mx.bfloat16)
self.assertEqual(a_scalar.shape, ())
self.assertEqual(a_scalar.item(), 1.5)
data = [1.5, 2.5, 3.5]
x_vector = np.array(data, dtype=ml_dtypes.bfloat16)
a_vector = mx.array(x_vector)
expected = mx.array(data, dtype=mx.bfloat16)
self.assertEqual(a_vector.dtype, mx.bfloat16)
self.assertEqual(a_vector.shape, (3,))
self.assertTrue(mx.array_equal(a_vector, expected))
a_cast = mx.array(x_scalar, dtype=mx.float32)
self.assertEqual(a_cast.dtype, mx.float32)
self.assertEqual(a_cast.item(), 1.5)
a_asarray = mx.asarray(x_vector)
self.assertEqual(a_asarray.dtype, mx.bfloat16)
self.assertEqual(a_asarray.shape, (3,))
self.assertTrue(mx.array_equal(a_asarray, expected))
if __name__ == "__main__":
mlx_tests.MLXTestRunner()