Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions include/xsimd/arch/common/xsimd_common_memory.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -450,6 +450,29 @@ namespace xsimd
return batch<T, A>::load_aligned(buffer.data());
}

template <class A, class T, class Mode>
XSIMD_INLINE batch<std::complex<T>, A>
load_complex_masked(std::complex<T> const* mem, batch_bool<T, A> mask, Mode, requires_arch<common>) noexcept
{
// Scalar fallback: only active lanes are touched. Arches with
// hardware predicated loads should override this.
constexpr std::size_t size = batch<T, A>::size;
alignas(A::alignment()) std::array<T, size> buffer_real;
alignas(A::alignment()) std::array<T, size> 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<std::complex<T>, A>::load_aligned(buffer_real.data(), buffer_imag.data());
}

template <class A, class T_in, class T_out, bool... Values, class alignment>
XSIMD_INLINE void
store_masked(T_out* mem, batch<T_in, A> const& src, batch_bool_constant<T_in, A, Values...> mask, alignment mode, requires_arch<common>) noexcept
Expand Down Expand Up @@ -865,6 +888,17 @@ namespace xsimd
store_complex_aligned<A>(dst, src, A {});
}

template <class A, class T, class Mode>
XSIMD_INLINE void
store_complex_masked(std::complex<T>* mem, batch<std::complex<T>, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<common>) noexcept
{
alignas(A::alignment()) std::array<std::complex<T>, src.size> buffer;
src.store_aligned(buffer.data());
for (std::size_t i = 0; i < src.size; ++i)
if (mask.get(i))
mem[i] = buffer[i];
}

// transpose
template <class A, class T>
XSIMD_INLINE void transpose(batch<T, A>* matrix_begin, batch<T, A>* matrix_end, requires_arch<common>) noexcept
Expand Down
27 changes: 27 additions & 0 deletions include/xsimd/arch/xsimd_avx.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -1025,6 +1025,19 @@ namespace xsimd
return _mm256_maskload_pd(mem, _mm256_castpd_si256(mask));
}

template <class A, class T, class Mode>
XSIMD_INLINE batch<std::complex<T>, A>
load_complex_masked(std::complex<T> const* mem, batch_bool<T, A> mask, Mode mode, requires_arch<avx>) noexcept
{
using mask_register_type = typename batch_bool<T, A>::register_type;
mask_register_type nmask = mask.to_native();
batch_bool<T, A> lo_mask = zip_lo(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
batch_bool<T, A> hi_mask = zip_hi(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
batch<T, A> res_lo = batch<T, A>::load(reinterpret_cast<T const*>(mem), lo_mask, mode);
batch<T, A> res_hi = batch<T, A>::load(reinterpret_cast<T const*>(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 <class A, class T, class Mode>
XSIMD_INLINE std::enable_if_t<std::is_integral_v<T> && (sizeof(T) == 4 || sizeof(T) == 8), batch<T, A>>
Expand Down Expand Up @@ -1201,6 +1214,20 @@ namespace xsimd
}
}

template <class A, class T, class Mode>
XSIMD_INLINE void
store_complex_masked(std::complex<T>* mem, batch<std::complex<T>, A> const& src, batch_bool<T, A> mask, Mode mode, requires_arch<avx>) noexcept
{
using mask_register_type = typename batch_bool<T, A>::register_type;
mask_register_type nmask = mask.to_native();
batch_bool<T, A> lo_mask = zip_lo(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
batch_bool<T, A> hi_mask = zip_hi(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
batch<T, A> src_lo = zip_lo(src.real(), src.imag());
batch<T, A> src_hi = zip_hi(src.real(), src.imag());
store_masked(reinterpret_cast<T*>(mem), src_lo, lo_mask, mode, A {});
store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, hi_mask, mode, A {});
}

namespace detail
{
// Reinterpret a constant-mask 4/8-byte load/store as same-width float
Expand Down
43 changes: 43 additions & 0 deletions include/xsimd/arch/xsimd_avx512f.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -372,6 +372,49 @@ namespace xsimd
detail::store_masked(mem, src, mask.mask(), Mode {});
}

template <class A, class T, class Mode>
XSIMD_INLINE batch<std::complex<T>, A>
load_complex_masked(std::complex<T> const* mem, batch_bool<T, A> mask, Mode mode, requires_arch<avx512f>) noexcept
{
using mask_register_type = typename batch_bool<T, A>::register_type;
mask_register_type nmask = mask.to_native();

// manually zip mask
constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask<mask_register_type>(src.size / 2);
mask_register_type lo_mask = nmask & lo_bitmask;
lo_mask |= lo_mask << (src.size / 2);

constexpr mask_register_type hi_bitmask = lo_bitmask << (src.size / 2);
mask_register_type hi_mask = nmask & hi_bitmask;
hi_mask |= hi_mask >> (src.size / 2);

batch<T, A> res_lo = batch<T, A>::load(reinterpret_cast<T const*>(mem), lo_mask, mode);
batch<T, A> res_hi = batch<T, A>::load(reinterpret_cast<T const*>(mem) + mask.size, hi_mask, mode);
return detail::load_complex(res_lo, res_hi, A{});
}

template <class A, class T, class Mode>
XSIMD_INLINE void
store_complex_masked(std::complex<T>* mem, batch<std::complex<T>, A> const& src, batch_bool<T, A> mask, Mode mode, requires_arch<avx512f>) noexcept
{
using mask_register_type = typename batch_bool<T, A>::register_type;
mask_register_type nmask = mask.to_native();

// manually zip mask
constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask<mask_register_type>(src.size / 2);
mask_register_type lo_mask = nmask & lo_bitmask;
lo_mask |= lo_mask << (src.size / 2);

constexpr mask_register_type hi_bitmask = lo_bitmask << (src.size / 2);
mask_register_type hi_mask = nmask & hi_bitmask;
hi_mask |= hi_mask >> (src.size / 2);

batch<T, A> src_lo = zip_lo(src.real(), src.imag());
batch<T, A> src_hi = zip_hi(src.real(), src.imag());
store_masked(reinterpret_cast<T*>(mem), src_lo, batch_bool<T, A> { lo_mask }, mode);
store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, batch_bool<T, A> { hi_mask }, mode);
}

// abs
template <class A>
XSIMD_INLINE batch<float, A> abs(batch<float, A> const& self, requires_arch<avx512f>) noexcept
Expand Down
15 changes: 4 additions & 11 deletions include/xsimd/types/xsimd_batch.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -1512,13 +1512,9 @@ namespace xsimd

template <class T, class A>
template <class Mode>
XSIMD_INLINE void batch<std::complex<T>, A>::store(value_type* mem, batch_bool<T, A> mask, Mode) const noexcept
XSIMD_INLINE void batch<std::complex<T>, A>::store(value_type* mem, batch_bool<T, A> mask, Mode mode) const noexcept
{
alignas(A::alignment()) std::array<value_type, size> 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<A>(mem, *this, mask, mode, A {});
}

template <class T, class A>
Expand Down Expand Up @@ -1561,12 +1557,9 @@ namespace xsimd

template <class T, class A>
template <class Mode>
XSIMD_INLINE batch<std::complex<T>, A> batch<std::complex<T>, A>::load(value_type const* mem, batch_bool<T, A> mask, Mode) noexcept
XSIMD_INLINE batch<std::complex<T>, A> batch<std::complex<T>, A>::load(value_type const* mem, batch_bool<T, A> mask, Mode mode) noexcept
{
alignas(A::alignment()) std::array<value_type, size> 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<A>(mem, mask, mode, A {});
}

template <class T, class A>
Expand Down
Loading