Skip to content

Commit e1b148e

Browse files
Improve generic version of complex masked store
Use the usual kernel mechanism which allows for specialization. Implement specialization for avx and avx512. Follow-up to #1391
1 parent 67e96b0 commit e1b148e

3 files changed

Lines changed: 47 additions & 0 deletions

File tree

include/xsimd/arch/common/xsimd_common_memory.hpp

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -865,6 +865,17 @@ namespace xsimd
865865
store_complex_aligned<A>(dst, src, A {});
866866
}
867867

868+
template <class A, class T, class Mode>
869+
XSIMD_INLINE void
870+
store_complex_masked(std::complex<T>* mem, batch<std::complex<T>, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<common>) noexcept
871+
{
872+
alignas(A::alignment()) std::array<std::complex<T>, src.size> buffer;
873+
src.store_aligned(buffer.data());
874+
for (std::size_t i = 0; i < src.size; ++i)
875+
if (mask.get(i))
876+
mem[i] = buffer[i];
877+
}
878+
868879
// transpose
869880
template <class A, class T>
870881
XSIMD_INLINE void transpose(batch<T, A>* matrix_begin, batch<T, A>* matrix_end, requires_arch<common>) noexcept

include/xsimd/arch/xsimd_avx.hpp

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1201,6 +1201,20 @@ namespace xsimd
12011201
}
12021202
}
12031203

1204+
template <class A, class T, class Mode>
1205+
XSIMD_INLINE void
1206+
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
1207+
{
1208+
using mask_register_type = typename batch_bool<T, A>::register_type;
1209+
mask_register_type nmask = mask.to_native();
1210+
batch_bool<T, A> lo_mask = zip_lo(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
1211+
batch_bool<T, A> hi_mask = zip_hi(batch<T, A>(nmask), batch<T, A>(nmask)).to_native();
1212+
batch<T, A> src_lo = zip_lo(src.real(), src.imag());
1213+
batch<T, A> src_hi = zip_hi(src.real(), src.imag());
1214+
store_masked(reinterpret_cast<T*>(mem), src_lo, lo_mask, mode, A {});
1215+
store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, hi_mask, mode, A {});
1216+
}
1217+
12041218
namespace detail
12051219
{
12061220
// Reinterpret a constant-mask 4/8-byte load/store as same-width float

include/xsimd/arch/xsimd_avx512f.hpp

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -372,6 +372,28 @@ namespace xsimd
372372
detail::store_masked(mem, src, mask.mask(), Mode {});
373373
}
374374

375+
template <class A, class T, class Mode>
376+
XSIMD_INLINE void
377+
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
378+
{
379+
using mask_register_type = typename batch_bool<T, A>::register_type;
380+
mask_register_type nmask = mask.to_native();
381+
382+
// manually zip mask
383+
constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask<mask_register_type>(src.size / 2);
384+
mask_register_type lo_mask = nmask & lo_bitmask;
385+
lo_mask |= lo_mask << (src.size / 2);
386+
387+
constexpr mask_register_type hi_bitmask = lo_bitmask << (src.size / 2);
388+
mask_register_type hi_mask = nmask & hi_bitmask;
389+
hi_mask |= hi_mask >> (src.size / 2);
390+
391+
batch<T, A> src_lo = zip_lo(src.real(), src.imag());
392+
batch<T, A> src_hi = zip_hi(src.real(), src.imag());
393+
store_masked(reinterpret_cast<T*>(mem), src_lo, batch_bool<T, A> { lo_mask }, mode);
394+
store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, batch_bool<T, A> { hi_mask }, mode);
395+
}
396+
375397
// abs
376398
template <class A>
377399
XSIMD_INLINE batch<float, A> abs(batch<float, A> const& self, requires_arch<avx512f>) noexcept

0 commit comments

Comments
 (0)