Skip to content

Commit 2969f40

Browse files
committed
Simplify common_memory load_masked
1 parent c4c12c5 commit 2969f40

3 files changed

Lines changed: 73 additions & 202 deletions

File tree

include/xsimd/arch/common/xsimd_common_memory.hpp

Lines changed: 63 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
#define XSIMD_COMMON_MEMORY_HPP
1414

1515
#include "../../types/xsimd_batch_constant.hpp"
16+
#include "../../utils/xsimd_type_traits.hpp"
1617
#include "./xsimd_common_details.hpp"
1718

1819
#include <algorithm>
@@ -360,88 +361,81 @@ namespace xsimd
360361
return load_unaligned<A>(mem, convert<T> {}, A {});
361362
}
362363

363-
template <class A, class T_in, class T_out, bool... Values, class alignment>
364-
XSIMD_INLINE batch<T_out, A>
365-
load_masked(T_in const* mem, batch_bool_constant<T_out, A, Values...>, convert<T_out>, alignment, requires_arch<common>) noexcept
366-
{
367-
constexpr std::size_t size = batch<T_out, A>::size;
368-
alignas(A::alignment()) std::array<T_out, size> buffer {};
369-
constexpr bool mask[size] = { Values... };
370-
371-
for (std::size_t i = 0; i < size; ++i)
372-
buffer[i] = mask[i] ? static_cast<T_out>(mem[i]) : T_out(0);
373-
374-
return batch<T_out, A>::load(buffer.data(), aligned_mode {});
375-
}
376-
377-
template <class A, class T_in, class T_out, bool... Values, class alignment>
378-
XSIMD_INLINE void
379-
store_masked(T_out* mem, batch<T_in, A> const& src, batch_bool_constant<T_in, A, Values...>, alignment, requires_arch<common>) noexcept
364+
namespace detail
380365
{
381-
constexpr std::size_t size = batch<T_in, A>::size;
382-
constexpr bool mask[size] = { Values... };
366+
// Compile-time dispatch tag for the common `load_masked`/ `store_masked`
367+
// implementations: true iff we can use the int->float bitcast path (matching-size
368+
// integer T_in/T_out with a SIMD register available for the matching
369+
// floating-point type), false otherwise (use the scalar buffer fallback).
370+
template <class A, class T_in, class T_out>
371+
using common_masked_via_fp = std::integral_constant<bool,
372+
std::is_same<T_in, T_out>::value
373+
&& std::is_integral<T_out>::value
374+
&& !std::is_void<sized_fp_t<sizeof(T_out)>>::value
375+
&& types::has_simd_register<sized_fp_t<sizeof(T_out)>, A>::value>;
383376

384-
for (std::size_t i = 0; i < size; ++i)
385-
if (mask[i])
386-
{
387-
mem[i] = static_cast<T_out>(src.get(i));
388-
}
389-
}
377+
// Scalar-buffer fallback: works for any T_in/T_out.
378+
template <class A, class T_in, class T_out, bool... Values, class alignment>
379+
XSIMD_INLINE batch<T_out, A>
380+
load_masked_common(T_in const* mem, batch_bool_constant<T_out, A, Values...>, convert<T_out>, alignment, std::false_type /* via_fp */) noexcept
381+
{
382+
constexpr std::size_t size = batch<T_out, A>::size;
383+
alignas(A::alignment()) std::array<T_out, size> buffer {};
384+
constexpr bool mask[size] = { Values... };
390385

391-
template <class A, bool... Values, class Mode>
392-
XSIMD_INLINE batch<int32_t, A> load_masked(int32_t const* mem, batch_bool_constant<int32_t, A, Values...>, convert<int32_t>, Mode, requires_arch<A>) noexcept
393-
{
394-
const auto f = load_masked<A>(reinterpret_cast<const float*>(mem), batch_bool_constant<float, A, Values...> {}, convert<float> {}, Mode {}, A {});
395-
return bitwise_cast<int32_t>(f);
396-
}
386+
for (std::size_t i = 0; i < size; ++i)
387+
buffer[i] = mask[i] ? static_cast<T_out>(mem[i]) : T_out(0);
397388

398-
template <class A, bool... Values, class Mode>
399-
XSIMD_INLINE batch<uint32_t, A> load_masked(uint32_t const* mem, batch_bool_constant<uint32_t, A, Values...>, convert<uint32_t>, Mode, requires_arch<A>) noexcept
400-
{
401-
const auto f = load_masked<A>(reinterpret_cast<const float*>(mem), batch_bool_constant<float, A, Values...> {}, convert<float> {}, Mode {}, A {});
402-
return bitwise_cast<uint32_t>(f);
403-
}
389+
return batch<T_out, A>::load(buffer.data(), aligned_mode {});
390+
}
404391

405-
template <class A, bool... Values, class Mode>
406-
XSIMD_INLINE std::enable_if_t<types::has_simd_register<double, A>::value, batch<int64_t, A>>
407-
load_masked(int64_t const* mem, batch_bool_constant<int64_t, A, Values...>, convert<int64_t>, Mode, requires_arch<A>) noexcept
408-
{
409-
const auto d = load_masked<A>(reinterpret_cast<const double*>(mem), batch_bool_constant<double, A, Values...> {}, convert<double> {}, Mode {}, A {});
410-
return bitwise_cast<int64_t>(d);
411-
}
392+
// Integer-via-float bitcast: T_in == T_out == integral T with a matching
393+
// `sized_fp_t<sizeof(T)>` for which the arch has a SIMD register.
394+
// Dispatches to the floating `load_masked` (which is arch-specialized) and bitcasts back.
395+
template <class A, class T, bool... Values, class Mode>
396+
XSIMD_INLINE batch<T, A>
397+
load_masked_common(T const* mem, batch_bool_constant<T, A, Values...>, convert<T>, Mode, std::true_type /* via_fp */) noexcept
398+
{
399+
using fp_t = sized_fp_t<sizeof(T)>;
400+
const auto f = ::xsimd::kernel::load_masked<A>(reinterpret_cast<const fp_t*>(mem), batch_bool_constant<fp_t, A, Values...> {}, convert<fp_t> {}, Mode {}, A {});
401+
return bitwise_cast<T>(f);
402+
}
412403

413-
template <class A, bool... Values, class Mode>
414-
XSIMD_INLINE std::enable_if_t<types::has_simd_register<double, A>::value, batch<uint64_t, A>>
415-
load_masked(uint64_t const* mem, batch_bool_constant<uint64_t, A, Values...>, convert<uint64_t>, Mode, requires_arch<A>) noexcept
416-
{
417-
const auto d = load_masked<A>(reinterpret_cast<const double*>(mem), batch_bool_constant<double, A, Values...> {}, convert<double> {}, Mode {}, A {});
418-
return bitwise_cast<uint64_t>(d);
419-
}
404+
template <class A, class T_in, class T_out, bool... Values, class alignment>
405+
XSIMD_INLINE void
406+
store_masked_common(T_out* mem, batch<T_in, A> const& src, batch_bool_constant<T_in, A, Values...>, alignment, std::false_type /* via_fp */) noexcept
407+
{
408+
constexpr std::size_t size = batch<T_in, A>::size;
409+
constexpr bool mask[size] = { Values... };
420410

421-
template <class A, bool... Values, class Mode>
422-
XSIMD_INLINE void store_masked(int32_t* mem, batch<int32_t, A> const& src, batch_bool_constant<int32_t, A, Values...>, Mode, requires_arch<A>) noexcept
423-
{
424-
store_masked<A>(reinterpret_cast<float*>(mem), bitwise_cast<float>(src), batch_bool_constant<float, A, Values...> {}, Mode {}, A {});
425-
}
411+
for (std::size_t i = 0; i < size; ++i)
412+
if (mask[i])
413+
{
414+
mem[i] = static_cast<T_out>(src.get(i));
415+
}
416+
}
426417

427-
template <class A, bool... Values, class Mode>
428-
XSIMD_INLINE void store_masked(uint32_t* mem, batch<uint32_t, A> const& src, batch_bool_constant<uint32_t, A, Values...>, Mode, requires_arch<A>) noexcept
429-
{
430-
store_masked<A>(reinterpret_cast<float*>(mem), bitwise_cast<float>(src), batch_bool_constant<float, A, Values...> {}, Mode {}, A {});
431-
}
418+
template <class A, class T, bool... Values, class Mode>
419+
XSIMD_INLINE void
420+
store_masked_common(T* mem, batch<T, A> const& src, batch_bool_constant<T, A, Values...>, Mode, std::true_type /* via_fp */) noexcept
421+
{
422+
using fp_t = sized_fp_t<sizeof(T)>;
423+
::xsimd::kernel::store_masked<A>(reinterpret_cast<fp_t*>(mem), bitwise_cast<fp_t>(src), batch_bool_constant<fp_t, A, Values...> {}, Mode {}, A {});
424+
}
425+
} // namespace detail
432426

433-
template <class A, bool... Values, class Mode>
434-
XSIMD_INLINE std::enable_if_t<types::has_simd_register<double, A>::value>
435-
store_masked(int64_t* mem, batch<int64_t, A> const& src, batch_bool_constant<int64_t, A, Values...>, Mode, requires_arch<A>) noexcept
427+
template <class A, class T_in, class T_out, bool... Values, class alignment>
428+
XSIMD_INLINE batch<T_out, A>
429+
load_masked(T_in const* mem, batch_bool_constant<T_out, A, Values...> mask, convert<T_out> cvt, alignment mode, requires_arch<common>) noexcept
436430
{
437-
store_masked<A>(reinterpret_cast<double*>(mem), bitwise_cast<double>(src), batch_bool_constant<double, A, Values...> {}, Mode {}, A {});
431+
return detail::load_masked_common(mem, mask, cvt, mode, detail::common_masked_via_fp<A, T_in, T_out> {});
438432
}
439433

440-
template <class A, bool... Values, class Mode>
441-
XSIMD_INLINE std::enable_if_t<types::has_simd_register<double, A>::value>
442-
store_masked(uint64_t* mem, batch<uint64_t, A> const& src, batch_bool_constant<uint64_t, A, Values...>, Mode, requires_arch<A>) noexcept
434+
template <class A, class T_in, class T_out, bool... Values, class alignment>
435+
XSIMD_INLINE void
436+
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
443437
{
444-
store_masked<A>(reinterpret_cast<double*>(mem), bitwise_cast<double>(src), batch_bool_constant<double, A, Values...> {}, Mode {}, A {});
438+
detail::store_masked_common(mem, src, mask, mode, detail::common_masked_via_fp<A, T_in, T_out> {});
445439
}
446440

447441
template <class A, class T_in, class T_out>

include/xsimd/arch/xsimd_avx512f.hpp

Lines changed: 10 additions & 117 deletions
Original file line numberDiff line numberDiff line change
@@ -298,25 +298,11 @@ namespace xsimd
298298

299299
} // namespace detail
300300

301-
// The AVX512F masked-load logic lives in this plain `*_avx512f` helper
302-
// (no `requires_arch` tag) and is exposed through the concrete
303-
// element-type overloads below.
304-
//
305-
// Why not a single generic `load_masked(T const*, ..., requires_arch<avx512f>)`?
306-
// It is ambiguous against the concrete-type / generic-arch overloads in
307-
// xsimd_common_memory.hpp (e.g. `load_masked(int32_t const*, ...,
308-
// requires_arch<A>)`): the avx512f overload is more specialized on the
309-
// architecture while the common one is more specialized on the pointer
310-
// type, so partial ordering cannot pick a winner. When AVX512DQ/BW is
311-
// available a fully concrete `requires_arch<avx512dq>` overload is the
312-
// unique best match and hides this, but a pure-AVX512F target (the
313-
// `avx512f` preset) has no such tie-breaker and the call fails to
314-
// compile. Concrete element-type `requires_arch<avx512f>` overloads make
315-
// the avx512f candidate the unique best match for every integer type.
316-
template <class A, class T, bool... Values, class Mode>
317-
XSIMD_INLINE batch<T, A> load_masked_avx512f(T const* mem,
318-
batch_bool_constant<T, A, Values...> mask,
319-
Mode) noexcept
301+
template <class A, class T, bool... Values, class Mode,
302+
typename = std::enable_if_t<(sizeof(T) >= 4)>>
303+
XSIMD_INLINE batch<T, A> load_masked(T const* mem,
304+
batch_bool_constant<T, A, Values...> mask,
305+
convert<T>, Mode, requires_arch<avx512f>) noexcept
320306
{
321307
constexpr auto half = batch<T, A>::size / 2;
322308
XSIMD_IF_CONSTEXPR(mask.countl_zero() >= half) // lower-half AVX2 forwarding
@@ -338,61 +324,12 @@ namespace xsimd
338324
}
339325
}
340326

341-
template <class A, bool... Values, class Mode>
342-
XSIMD_INLINE batch<int32_t, A> load_masked(int32_t const* mem,
343-
batch_bool_constant<int32_t, A, Values...> mask,
344-
convert<int32_t>, Mode, requires_arch<avx512f>) noexcept
345-
{
346-
return load_masked_avx512f(mem, mask, Mode {});
347-
}
348-
349-
template <class A, bool... Values, class Mode>
350-
XSIMD_INLINE batch<uint32_t, A> load_masked(uint32_t const* mem,
351-
batch_bool_constant<uint32_t, A, Values...> mask,
352-
convert<uint32_t>, Mode, requires_arch<avx512f>) noexcept
353-
{
354-
return load_masked_avx512f(mem, mask, Mode {});
355-
}
356-
357-
template <class A, bool... Values, class Mode>
358-
XSIMD_INLINE batch<int64_t, A> load_masked(int64_t const* mem,
359-
batch_bool_constant<int64_t, A, Values...> mask,
360-
convert<int64_t>, Mode, requires_arch<avx512f>) noexcept
361-
{
362-
return load_masked_avx512f(mem, mask, Mode {});
363-
}
364-
365-
template <class A, bool... Values, class Mode>
366-
XSIMD_INLINE batch<uint64_t, A> load_masked(uint64_t const* mem,
367-
batch_bool_constant<uint64_t, A, Values...> mask,
368-
convert<uint64_t>, Mode, requires_arch<avx512f>) noexcept
369-
{
370-
return load_masked_avx512f(mem, mask, Mode {});
371-
}
372-
373-
// Non-integer element types only (float, double, ...): integer types
374-
// are handled by the concrete overloads above. gcc 10's partial
375-
// ordering cannot break the tie between a concrete-type avx512f
376-
// overload and a generic-T avx512f overload, so this catch-all must
377-
// exclude the integer types we already specialized.
378327
template <class A, class T, bool... Values, class Mode,
379-
typename = std::enable_if_t<(sizeof(T) >= 4) && !std::is_integral<T>::value>>
380-
XSIMD_INLINE batch<T, A> load_masked(T const* mem,
381-
batch_bool_constant<T, A, Values...> mask,
382-
convert<T>, Mode, requires_arch<avx512f>) noexcept
383-
{
384-
return load_masked_avx512f(mem, mask, Mode {});
385-
}
386-
387-
// Same ambiguity as load_masked above (see comment there): factor the
388-
// AVX512F store logic into a plain helper and expose it via concrete
389-
// element-type `requires_arch<avx512f>` overloads so a pure-AVX512F
390-
// target has a unique best match.
391-
template <class A, class T, bool... Values, class Mode>
392-
XSIMD_INLINE void store_masked_avx512f(T* mem,
393-
batch<T, A> const& src,
394-
batch_bool_constant<T, A, Values...> mask,
395-
Mode) noexcept
328+
typename = std::enable_if_t<(sizeof(T) >= 4)>>
329+
XSIMD_INLINE void store_masked(T* mem,
330+
batch<T, A> const& src,
331+
batch_bool_constant<T, A, Values...> mask,
332+
Mode, requires_arch<avx512f>) noexcept
396333
{
397334
constexpr auto half = batch<T, A>::size / 2;
398335
XSIMD_IF_CONSTEXPR(mask.countl_zero() >= half) // lower-half AVX2 forwarding
@@ -414,50 +351,6 @@ namespace xsimd
414351
}
415352
}
416353

417-
template <class A, bool... Values, class Mode>
418-
XSIMD_INLINE void store_masked(int32_t* mem, batch<int32_t, A> const& src,
419-
batch_bool_constant<int32_t, A, Values...> mask,
420-
Mode, requires_arch<avx512f>) noexcept
421-
{
422-
store_masked_avx512f(mem, src, mask, Mode {});
423-
}
424-
425-
template <class A, bool... Values, class Mode>
426-
XSIMD_INLINE void store_masked(uint32_t* mem, batch<uint32_t, A> const& src,
427-
batch_bool_constant<uint32_t, A, Values...> mask,
428-
Mode, requires_arch<avx512f>) noexcept
429-
{
430-
store_masked_avx512f(mem, src, mask, Mode {});
431-
}
432-
433-
template <class A, bool... Values, class Mode>
434-
XSIMD_INLINE void store_masked(int64_t* mem, batch<int64_t, A> const& src,
435-
batch_bool_constant<int64_t, A, Values...> mask,
436-
Mode, requires_arch<avx512f>) noexcept
437-
{
438-
store_masked_avx512f(mem, src, mask, Mode {});
439-
}
440-
441-
template <class A, bool... Values, class Mode>
442-
XSIMD_INLINE void store_masked(uint64_t* mem, batch<uint64_t, A> const& src,
443-
batch_bool_constant<uint64_t, A, Values...> mask,
444-
Mode, requires_arch<avx512f>) noexcept
445-
{
446-
store_masked_avx512f(mem, src, mask, Mode {});
447-
}
448-
449-
// Non-integer element types only: see load_masked above for the gcc 10
450-
// partial ordering rationale.
451-
template <class A, class T, bool... Values, class Mode,
452-
typename = std::enable_if_t<(sizeof(T) >= 4) && !std::is_integral<T>::value>>
453-
XSIMD_INLINE void store_masked(T* mem,
454-
batch<T, A> const& src,
455-
batch_bool_constant<T, A, Values...> mask,
456-
Mode, requires_arch<avx512f>) noexcept
457-
{
458-
store_masked_avx512f(mem, src, mask, Mode {});
459-
}
460-
461354
// abs
462355
template <class A>
463356
XSIMD_INLINE batch<float, A> abs(batch<float, A> const& self, requires_arch<avx512f>) noexcept

0 commit comments

Comments
 (0)