diff --git a/include/xsimd/arch/common/xsimd_common_memory.hpp b/include/xsimd/arch/common/xsimd_common_memory.hpp index 046faafad..d0d58275c 100644 --- a/include/xsimd/arch/common/xsimd_common_memory.hpp +++ b/include/xsimd/arch/common/xsimd_common_memory.hpp @@ -450,6 +450,29 @@ namespace xsimd return batch::load_aligned(buffer.data()); } + template + XSIMD_INLINE batch, A> + load_complex_masked(std::complex const* mem, batch_bool mask, Mode, requires_arch) noexcept + { + // Scalar fallback: only active lanes are touched. Arches with + // hardware predicated loads should override this. + constexpr std::size_t size = batch::size; + alignas(A::alignment()) std::array buffer_real; + alignas(A::alignment()) std::array buffer_imag; + for (std::size_t i = 0; i < size; ++i) + if (mask.get(i)) + { + buffer_real[i] = mem[i].real(); + buffer_imag[i] = mem[i].imag(); + } + else + { + buffer_real[i] = T(0); + buffer_imag[i] = T(0); + } + return batch, A>::load_aligned(buffer_real.data(), buffer_imag.data()); + } + template XSIMD_INLINE void store_masked(T_out* mem, batch const& src, batch_bool_constant mask, alignment mode, requires_arch) noexcept @@ -865,6 +888,18 @@ namespace xsimd store_complex_aligned(dst, src, A {}); } + template + XSIMD_INLINE void + store_complex_masked(std::complex* mem, batch, A> const& src, batch_bool mask, Mode, requires_arch) noexcept + { + constexpr std::size_t size = batch::size; + alignas(A::alignment()) std::array, size> buffer; + src.store_aligned(buffer.data()); + for (std::size_t i = 0; i < size; ++i) + if (mask.get(i)) + mem[i] = buffer[i]; + } + // transpose template XSIMD_INLINE void transpose(batch* matrix_begin, batch* matrix_end, requires_arch) noexcept diff --git a/include/xsimd/arch/xsimd_avx.hpp b/include/xsimd/arch/xsimd_avx.hpp index 814452cea..6daa88710 100644 --- a/include/xsimd/arch/xsimd_avx.hpp +++ b/include/xsimd/arch/xsimd_avx.hpp @@ -1025,6 +1025,19 @@ namespace xsimd return _mm256_maskload_pd(mem, _mm256_castpd_si256(mask)); } + template + XSIMD_INLINE batch, A> + load_complex_masked(std::complex const* mem, batch_bool mask, Mode mode, requires_arch) noexcept + { + using mask_register_type = typename batch_bool::register_type; + mask_register_type nmask = mask.to_native(); + batch_bool lo_mask = zip_lo(batch(nmask), batch(nmask)).to_native(); + batch_bool hi_mask = zip_hi(batch(nmask), batch(nmask)).to_native(); + batch res_lo = batch::load(reinterpret_cast(mem), lo_mask, mode); + batch res_hi = batch::load(reinterpret_cast(mem) + mask.size, hi_mask, mode); + return detail::load_complex(res_lo, res_hi, A{}); + } + // 4/8-byte ints: bitcast to same-width float, reuse the vmaskmov path. template XSIMD_INLINE std::enable_if_t && (sizeof(T) == 4 || sizeof(T) == 8), batch> @@ -1201,6 +1214,20 @@ namespace xsimd } } + template + XSIMD_INLINE void + store_complex_masked(std::complex* mem, batch, A> const& src, batch_bool mask, Mode mode, requires_arch) noexcept + { + using mask_register_type = typename batch_bool::register_type; + mask_register_type nmask = mask.to_native(); + batch_bool lo_mask = zip_lo(batch(nmask), batch(nmask)).to_native(); + batch_bool hi_mask = zip_hi(batch(nmask), batch(nmask)).to_native(); + batch src_lo = zip_lo(src.real(), src.imag()); + batch src_hi = zip_hi(src.real(), src.imag()); + store_masked(reinterpret_cast(mem), src_lo, lo_mask, mode, A {}); + store_masked(reinterpret_cast(mem) + src.size, src_hi, hi_mask, mode, A {}); + } + namespace detail { // Reinterpret a constant-mask 4/8-byte load/store as same-width float diff --git a/include/xsimd/arch/xsimd_avx512f.hpp b/include/xsimd/arch/xsimd_avx512f.hpp index 658b7d448..698f9e6c2 100644 --- a/include/xsimd/arch/xsimd_avx512f.hpp +++ b/include/xsimd/arch/xsimd_avx512f.hpp @@ -372,6 +372,37 @@ namespace xsimd detail::store_masked(mem, src, mask.mask(), Mode {}); } + namespace detail + { + template + std::array, 2> zip_complex_mask(batch_bool mask) + { + using mask_register_type = typename batch_bool::register_type; + mask_register_type nmask = mask.to_native(); + + constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask(mask.size / 2); + mask_register_type lo_mask = nmask & lo_bitmask; + lo_mask |= lo_mask << (mask.size / 2); + + constexpr mask_register_type hi_bitmask = lo_bitmask << (mask.size / 2); + mask_register_type hi_mask = nmask & hi_bitmask; + hi_mask |= hi_mask >> (mask.size / 2); + + return { { lo_mask }, { hi_mask } }; + } + } + + template + XSIMD_INLINE void + store_complex_masked(std::complex* mem, batch, A> const& src, batch_bool mask, Mode mode, requires_arch) noexcept + { + auto [lo_mask, hi_mask] = detail::zip_complex_mask(mask); + batch src_lo = zip_lo(src.real(), src.imag()); + batch src_hi = zip_hi(src.real(), src.imag()); + detail::store_masked(reinterpret_cast(mem), src_lo, lo_mask, mode); + detail::store_masked(reinterpret_cast(mem) + src.size, src_hi, hi_mask, mode); + } + // abs template XSIMD_INLINE batch abs(batch const& self, requires_arch) noexcept @@ -1633,6 +1664,16 @@ namespace xsimd } } + template + XSIMD_INLINE batch, A> + load_complex_masked(std::complex const* mem, batch_bool mask, Mode mode, requires_arch) noexcept + { + auto [lo_mask, hi_mask] = detail::zip_complex_mask(mask); + batch res_lo = batch::load(reinterpret_cast(mem), lo_mask, mode); + batch res_hi = batch::load(reinterpret_cast(mem) + mask.size, hi_mask, mode); + return detail::load_complex(res_lo, res_hi, A{}); + } + // load_unaligned template >> XSIMD_INLINE batch load_unaligned(T const* mem, convert, requires_arch) noexcept diff --git a/include/xsimd/types/xsimd_batch.hpp b/include/xsimd/types/xsimd_batch.hpp index 8d4721fa9..972593460 100644 --- a/include/xsimd/types/xsimd_batch.hpp +++ b/include/xsimd/types/xsimd_batch.hpp @@ -1512,13 +1512,9 @@ namespace xsimd template template - XSIMD_INLINE void batch, A>::store(value_type* mem, batch_bool mask, Mode) const noexcept + XSIMD_INLINE void batch, A>::store(value_type* mem, batch_bool mask, Mode mode) const noexcept { - alignas(A::alignment()) std::array buffer; - store_aligned(buffer.data()); - for (std::size_t i = 0; i < size; ++i) - if (mask.get(i)) - mem[i] = buffer[i]; + kernel::store_complex_masked(mem, *this, mask, mode, A {}); } template @@ -1561,12 +1557,9 @@ namespace xsimd template template - XSIMD_INLINE batch, A> batch, A>::load(value_type const* mem, batch_bool mask, Mode) noexcept + XSIMD_INLINE batch, A> batch, A>::load(value_type const* mem, batch_bool mask, Mode mode) noexcept { - alignas(A::alignment()) std::array buffer {}; - for (std::size_t i = 0; i < size; ++i) - buffer[i] = mask.get(i) ? mem[i] : value_type(0); - return load_aligned(buffer.data()); + return kernel::load_complex_masked(mem, mask, mode, A {}); } template