diff --git a/python/src/ops.cpp b/python/src/ops.cpp index d8892a45d2..bdcf3dde4b 100644 --- a/python/src/ops.cpp +++ b/python/src/ops.cpp @@ -3386,14 +3386,14 @@ void init_ops(nb::module_& m) { mx::StreamOrDevice s) { std::vector arrays = nb::cast>(arrays_); - return mx::meshgrid(arrays, sparse, indexing, s); + return nb::tuple(nb::cast(mx::meshgrid(arrays, sparse, indexing, s))); }, "arrays"_a, "sparse"_a = false, "indexing"_a = "xy", "stream"_a = nb::none(), nb::sig( - "def meshgrid(*arrays: array, sparse: bool | None = False, indexing: str | None = 'xy', stream: StreamOrDevice = None) -> array"), + "def meshgrid(*arrays: array, sparse: bool | None = False, indexing: str | None = 'xy', stream: StreamOrDevice = None) -> tuple[array, ...]"), R"pbdoc( Generate multidimensional coordinate grids from 1-D coordinate arrays @@ -3406,7 +3406,7 @@ void init_ops(nb::module_& m) { Defaults to ``'xy'``. Returns: - list(array): The output arrays. + tuple(array): The output arrays. )pbdoc"); m.def( "repeat", diff --git a/python/tests/test_ops.py b/python/tests/test_ops.py index 1b46237ce5..83f95ea5fc 100644 --- a/python/tests/test_ops.py +++ b/python/tests/test_ops.py @@ -2132,6 +2132,11 @@ def test_meshgrid(self): x = mx.array([1, 2, 3], dtype=mx.int32) y = np.array([1, 2, 3], dtype=np.int32) + # Test return type is a tuple + self.assertIsInstance(mx.meshgrid(x), tuple) + self.assertIsInstance(mx.meshgrid(x, x), tuple) + self.assertIsInstance(mx.meshgrid(x, x, x, sparse=True), tuple) + # Test single input a_mlx = mx.meshgrid(x) a_np = np.meshgrid(y)