diff --git a/include/xsimd/arch/common/xsimd_common_memory.hpp b/include/xsimd/arch/common/xsimd_common_memory.hpp index 046faafad..39e24a2b1 100644 --- a/include/xsimd/arch/common/xsimd_common_memory.hpp +++ b/include/xsimd/arch/common/xsimd_common_memory.hpp @@ -865,6 +865,39 @@ 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 mode, requires_arch) noexcept + { + // Generic fallback: mask and real /imag part are zipped before + // calling the generic masked store routine. + using mask_register_type = typename batch_bool::register_type; + mask_register_type nmask = mask.to_native(); + batch_bool lo_mask; + batch_bool hi_mask; + + // Generic zip_lo/hi of batch_bool depending on native register type { + if constexpr (A::has_scalar_mask()) + { + constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask(src.size / 2); + constexpr mask_register_type hi_bitmask = lo_bitmask << (src.size / 2); + lo_mask = nmask & lo_bitmask; + lo_mask |= lo_mask << (src.size / 2); + hi_mask = nmask & hi_bitmask; + hi_mask |= hi_mask >> (src.size / 2); + } + else + { + lo_mask = zip_lo(batch(nmask), batch(nmask)).to_native(); + 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()); + src_lo.store(reinterpret_cast(mem), lo_mask, mode); + src_hi.store(reinterpret_cast(mem) + src.size, hi_mask, mode); + } + // transpose template XSIMD_INLINE void transpose(batch* matrix_begin, batch* matrix_end, requires_arch) noexcept diff --git a/include/xsimd/types/xsimd_avx512f_register.hpp b/include/xsimd/types/xsimd_avx512f_register.hpp index c54161209..230ad9429 100644 --- a/include/xsimd/types/xsimd_avx512f_register.hpp +++ b/include/xsimd/types/xsimd_avx512f_register.hpp @@ -29,6 +29,7 @@ namespace xsimd static constexpr bool available() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 64; } static constexpr bool requires_alignment() noexcept { return true; } + static constexpr bool has_scalar_mask() noexcept { return true; } static constexpr char const* name() noexcept { return "avx512f"; } }; diff --git a/include/xsimd/types/xsimd_avx_register.hpp b/include/xsimd/types/xsimd_avx_register.hpp index 515b60901..dae0200f7 100644 --- a/include/xsimd/types/xsimd_avx_register.hpp +++ b/include/xsimd/types/xsimd_avx_register.hpp @@ -30,6 +30,7 @@ namespace xsimd static constexpr std::size_t alignment() noexcept { return 32; } static constexpr bool requires_alignment() noexcept { return true; } static constexpr char const* name() noexcept { return "avx"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; /** diff --git a/include/xsimd/types/xsimd_batch.hpp b/include/xsimd/types/xsimd_batch.hpp index 8d4721fa9..046897116 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 diff --git a/include/xsimd/types/xsimd_emulated_register.hpp b/include/xsimd/types/xsimd_emulated_register.hpp index 6bfc04e94..cb928f1bf 100644 --- a/include/xsimd/types/xsimd_emulated_register.hpp +++ b/include/xsimd/types/xsimd_emulated_register.hpp @@ -38,6 +38,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return false; } static constexpr std::size_t alignment() noexcept { return 8; } static constexpr char const* name() noexcept { return "emulated"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; namespace types diff --git a/include/xsimd/types/xsimd_neon_register.hpp b/include/xsimd/types/xsimd_neon_register.hpp index 0732ea7cc..bbdf7b030 100644 --- a/include/xsimd/types/xsimd_neon_register.hpp +++ b/include/xsimd/types/xsimd_neon_register.hpp @@ -39,6 +39,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "arm32+neon"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; #if XSIMD_WITH_NEON diff --git a/include/xsimd/types/xsimd_rvv_register.hpp b/include/xsimd/types/xsimd_rvv_register.hpp index 76b9b2a32..4f5491a8e 100644 --- a/include/xsimd/types/xsimd_rvv_register.hpp +++ b/include/xsimd/types/xsimd_rvv_register.hpp @@ -39,6 +39,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "riscv+rvv"; } + static constexpr bool has_scalar_mask() noexcept { return true; } }; } diff --git a/include/xsimd/types/xsimd_sse2_register.hpp b/include/xsimd/types/xsimd_sse2_register.hpp index 48c1fd53f..69a76c0c5 100644 --- a/include/xsimd/types/xsimd_sse2_register.hpp +++ b/include/xsimd/types/xsimd_sse2_register.hpp @@ -34,6 +34,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "sse2"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; #if XSIMD_WITH_SSE2 diff --git a/include/xsimd/types/xsimd_sve_register.hpp b/include/xsimd/types/xsimd_sve_register.hpp index e4920d13b..3c51b6677 100644 --- a/include/xsimd/types/xsimd_sve_register.hpp +++ b/include/xsimd/types/xsimd_sve_register.hpp @@ -37,6 +37,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "arm64+sve"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; } diff --git a/include/xsimd/types/xsimd_vsx_register.hpp b/include/xsimd/types/xsimd_vsx_register.hpp index 36b933902..0809be796 100644 --- a/include/xsimd/types/xsimd_vsx_register.hpp +++ b/include/xsimd/types/xsimd_vsx_register.hpp @@ -33,6 +33,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "vmx+vsx"; } + static constexpr bool has_scalar_mask() noexcept { return true; } }; #if XSIMD_WITH_VSX diff --git a/include/xsimd/types/xsimd_vxe_register.hpp b/include/xsimd/types/xsimd_vxe_register.hpp index ba051baa5..31c78d7fe 100644 --- a/include/xsimd/types/xsimd_vxe_register.hpp +++ b/include/xsimd/types/xsimd_vxe_register.hpp @@ -35,6 +35,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "vxe"; } + static constexpr bool has_scalar_mask() noexcept { return true; } }; #if XSIMD_WITH_VXE diff --git a/include/xsimd/types/xsimd_wasm_register.hpp b/include/xsimd/types/xsimd_wasm_register.hpp index 5091b7636..ebbc3af00 100644 --- a/include/xsimd/types/xsimd_wasm_register.hpp +++ b/include/xsimd/types/xsimd_wasm_register.hpp @@ -34,6 +34,7 @@ namespace xsimd static constexpr bool requires_alignment() noexcept { return true; } static constexpr std::size_t alignment() noexcept { return 16; } static constexpr char const* name() noexcept { return "wasm"; } + static constexpr bool has_scalar_mask() noexcept { return false; } }; #if XSIMD_WITH_WASM