From 333d0603a1e171d6af7e31a21c8f85a9d718ccdf Mon Sep 17 00:00:00 2001 From: John Elliott <54456354+PhysicistJohn@users.noreply.github.com> Date: Tue, 11 Aug 2026 21:27:38 -0700 Subject: [PATCH] Preserve lanes in complex-like conversions --- mlx/types/complex.h | 41 ++++++++++++++++++++++---- tests/array_tests.cpp | 67 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 102 insertions(+), 6 deletions(-) diff --git a/mlx/types/complex.h b/mlx/types/complex.h index 51101cc97b..8c8cd66fb5 100644 --- a/mlx/types/complex.h +++ b/mlx/types/complex.h @@ -2,6 +2,9 @@ #pragma once #include +#include +#include + #include "mlx/types/half_types.h" namespace mlx::core { @@ -9,6 +12,22 @@ namespace mlx::core { struct complex64_t; struct complex128_t; +namespace detail { + +template +concept complex_like = requires(T& value) { + value.real(); + value.imag(); +}; + +template +concept convertible_complex_like = complex_like && requires(const T& value) { + { value.real() } -> std::convertible_to; + { value.imag() } -> std::convertible_to; +}; + +} // namespace detail + template inline constexpr bool can_convert_to_complex128 = !std::is_same_v && std::is_convertible_v; @@ -18,9 +37,14 @@ struct complex128_t : public std::complex { complex128_t(double v, double u) : std::complex(v, u) {}; complex128_t(std::complex v) : std::complex(v) {}; - template < - typename T, - typename = typename std::enable_if>::type> + template + requires( + !std::same_as, complex128_t> && + detail::convertible_complex_like) + complex128_t(const T& x) : std::complex(x.real(), x.imag()){}; + + template + requires(can_convert_to_complex128 && !detail::complex_like) complex128_t(T x) : std::complex(x){}; operator float() const { @@ -37,9 +61,14 @@ struct complex64_t : public std::complex { complex64_t(float v, float u) : std::complex(v, u) {}; complex64_t(std::complex v) : std::complex(v) {}; - template < - typename T, - typename = typename std::enable_if>::type> + template + requires( + !std::same_as, complex64_t> && + detail::convertible_complex_like) + complex64_t(const T& x) : std::complex(x.real(), x.imag()){}; + + template + requires(can_convert_to_complex64 && !detail::complex_like) complex64_t(T x) : std::complex(x){}; operator float() const { diff --git a/tests/array_tests.cpp b/tests/array_tests.cpp index 68a4bed3a8..9b870b5c00 100644 --- a/tests/array_tests.cpp +++ b/tests/array_tests.cpp @@ -10,6 +10,30 @@ using namespace mlx::core; +namespace { + +struct ComplexLike { + double real() const { + return 1.25; + } + double imag() const { + return -2.5; + } + operator float() const { + return 99.0f; + } +}; + +struct NonConvertibleLane {}; + +struct UnsupportedComplexLike { + NonConvertibleLane real() const; + NonConvertibleLane imag() const; + operator float() const; +}; + +} // namespace + TEST_CASE("test array basics") { // Scalar array x(1.0); @@ -117,6 +141,49 @@ TEST_CASE("test array basics") { } } +TEST_CASE("test complex-like conversions") { + static_assert(can_convert_to_complex64); + static_assert(can_convert_to_complex128); + static_assert(can_convert_to_complex64); + static_assert(can_convert_to_complex128); + static_assert(can_convert_to_complex64); + static_assert(can_convert_to_complex128); + static_assert(!can_convert_to_complex64); + static_assert(!can_convert_to_complex128); + static_assert(!std::constructible_from); + static_assert(!std::constructible_from); + static_assert(sizeof(complex64_t) == sizeof(std::complex)); + static_assert(alignof(complex64_t) == alignof(std::complex)); + static_assert(sizeof(complex128_t) == sizeof(std::complex)); + static_assert(alignof(complex128_t) == alignof(std::complex)); + + const complex64_t custom64 = ComplexLike{}; + CHECK_EQ(custom64.real(), 1.25f); + CHECK_EQ(custom64.imag(), -2.5f); + const complex128_t custom128 = ComplexLike{}; + CHECK_EQ(custom128.real(), 1.25); + CHECK_EQ(custom128.imag(), -2.5); + + const array custom_array(ComplexLike{}, complex64); + CHECK_EQ(custom_array.dtype(), complex64); + CHECK_EQ(custom_array.item(), complex64_t{1.25f, -2.5f}); + + const complex64_t from128 = complex128_t{3.5, -4.25}; + const complex128_t from64 = complex64_t{5.5f, -6.75f}; + CHECK_EQ(from128, complex64_t{3.5f, -4.25f}); + CHECK_EQ(from64, complex128_t{5.5, -6.75}); + + const complex64_t cross64 = std::complex{7.0, -8.0}; + const complex128_t cross128 = std::complex{9.0f, -10.0f}; + CHECK_EQ(cross64, complex64_t{7.0f, -8.0f}); + CHECK_EQ(cross128, complex128_t{9.0, -10.0}); + + const complex64_t scalar64 = 11; + const complex128_t scalar128 = 12; + CHECK_EQ(scalar64, complex64_t{11.0f, 0.0f}); + CHECK_EQ(scalar128, complex128_t{12.0, 0.0}); +} + TEST_CASE("test array types") { #define basic_dtype_test(T, mlx_type) \ T val = 42; \