Skip to content

Commit d30e970

Browse files
WIP
1 parent fb39929 commit d30e970

1 file changed

Lines changed: 24 additions & 26 deletions

File tree

include/xsimd/arch/xsimd_avx512f.hpp

Lines changed: 24 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -372,26 +372,35 @@ 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
375+
namespace detail
378376
{
379-
using mask_register_type = typename batch_bool<T, A>::register_type;
380-
mask_register_type nmask = mask.to_native();
377+
template <class A, class T>
378+
std::array<batch_bool<T, A>, 2> zip_complex_mask(batch_bool<T, A> mask)
379+
{
380+
using mask_register_type = typename batch_bool<T, A>::register_type;
381+
mask_register_type nmask = mask.to_native();
382+
383+
constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask<mask_register_type>(mask.size / 2);
384+
mask_register_type lo_mask = nmask & lo_bitmask;
385+
lo_mask |= lo_mask << (mask.size / 2);
381386

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);
387+
constexpr mask_register_type hi_bitmask = lo_bitmask << (mask.size / 2);
388+
mask_register_type hi_mask = nmask & hi_bitmask;
389+
hi_mask |= hi_mask >> (mask.size / 2);
386390

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);
391+
return { { lo_mask }, { hi_mask } };
392+
}
393+
}
390394

395+
template <class A, class T, class Mode>
396+
XSIMD_INLINE void
397+
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
398+
{
399+
auto [lo_mask, hi_mask] = detail::zip_complex_mask(mask);
391400
batch<T, A> src_lo = zip_lo(src.real(), src.imag());
392401
batch<T, A> src_hi = zip_hi(src.real(), src.imag());
393-
detail::store_masked(reinterpret_cast<T*>(mem), src_lo, batch_bool<T, A> { lo_mask }, mode);
394-
detail::store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, batch_bool<T, A> { hi_mask }, mode);
402+
detail::store_masked(reinterpret_cast<T*>(mem), src_lo, lo_mask, mode);
403+
detail::store_masked(reinterpret_cast<T*>(mem) + src.size, src_hi, hi_mask, mode);
395404
}
396405

397406
// abs
@@ -1659,18 +1668,7 @@ namespace xsimd
16591668
XSIMD_INLINE batch<std::complex<T>, A>
16601669
load_complex_masked(std::complex<T> const* mem, batch_bool<T, A> mask, Mode mode, requires_arch<avx512f>) noexcept
16611670
{
1662-
using mask_register_type = typename batch_bool<T, A>::register_type;
1663-
mask_register_type nmask = mask.to_native();
1664-
1665-
// manually zip mask
1666-
constexpr mask_register_type lo_bitmask = xsimd::utils::make_low_mask<mask_register_type>(mask.size / 2);
1667-
mask_register_type lo_mask = nmask & lo_bitmask;
1668-
lo_mask |= lo_mask << (mask.size / 2);
1669-
1670-
constexpr mask_register_type hi_bitmask = lo_bitmask << (mask.size / 2);
1671-
mask_register_type hi_mask = nmask & hi_bitmask;
1672-
hi_mask |= hi_mask >> (mask.size / 2);
1673-
1671+
auto [lo_mask, hi_mask] = detail::zip_complex_mask(mask);
16741672
batch<T, A> res_lo = batch<T, A>::load(reinterpret_cast<T const*>(mem), lo_mask, mode);
16751673
batch<T, A> res_hi = batch<T, A>::load(reinterpret_cast<T const*>(mem) + mask.size, hi_mask, mode);
16761674
return detail::load_complex(res_lo, res_hi, A{});

0 commit comments

Comments
 (0)