| 98 | |
| 99 | template <typename T, typename... NDParams> |
| 100 | nb::ndarray<NDParams...> mlx_to_nd_array_impl( |
| 101 | mx::array a, |
| 102 | std::optional<nb::dlpack::dtype> t = {}) { |
| 103 | { |
| 104 | nb::gil_scoped_release nogil; |
| 105 | a.eval(); |
| 106 | } |
| 107 | std::vector<size_t> shape(a.shape().begin(), a.shape().end()); |
| 108 | return nb::ndarray<NDParams...>( |
| 109 | a.data<T>(), |
| 110 | a.ndim(), |
| 111 | shape.data(), |
| 112 | /* owner= */ nb::none(), |
| 113 | a.strides().data(), |
| 114 | t.value_or(nb::dtype<T>())); |
| 115 | } |
| 116 | |
| 117 | template <typename... NDParams> |
| 118 | nb::ndarray<NDParams...> mlx_to_nd_array(const mx::array& a) { |