Skip to content

Commit 801631c

Browse files
committed
feat: masked load/store capability traits + AVX integer masked memory
- has_mask_load/store capability traits (xsimd::detail), size-keyed per arch with inheritance-following; report a real predicated masked load/store vs the emulated scalar common fallback (pure compile-time, build-independent) - AVX / avx_128: integer 4/8-byte runtime masked load/store via float bitcast, reusing the vmaskmov path (lifts pure-AVX / fma3<avx> off the scalar fallback) - AVX512BW: native vmovdqu8/16 masked load/store for 8/16-bit integers; avx512vl_128/256 native masked overloads - dedup: constant-mask overloads with no compile-time advantage forward to the runtime path instead of duplicating the intrinsic call - sve: drop the dead pmask helper (constant masks forward via as_batch_bool)
1 parent f5161a4 commit 801631c

14 files changed

Lines changed: 358 additions & 52 deletions

docs/source/api/data_transfer.rst

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -90,5 +90,5 @@ The following empty types are used for tag dispatching:
9090
.. [#m] Masked ``load`` / ``store`` come in two flavours. The
9191
:cpp:class:`batch_bool_constant` overload encodes the mask in the type and
9292
is resolved at compile time. The runtime :cpp:class:`batch_bool` overload
93-
accepts a mask computed at runtime. Prefer the compile-time mask whenever
94-
the selection is known at compile time.
93+
accepts a mask computed at runtime. For performance reasons, prefer the
94+
compile-time mask whenever possible.

include/xsimd/arch/common/xsimd_common_memory.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -442,7 +442,7 @@ namespace xsimd
442442
load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<common>) noexcept
443443
{
444444
// Scalar fallback: only active lanes are touched. Arches with
445-
// hardware predicated loads override this.
445+
// hardware predicated loads should override this.
446446
constexpr std::size_t size = batch<T, A>::size;
447447
alignas(A::alignment()) std::array<T, size> buffer;
448448
for (std::size_t i = 0; i < size; ++i)
@@ -462,7 +462,7 @@ namespace xsimd
462462
store_masked(T* mem, batch<T, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<common>) noexcept
463463
{
464464
// Scalar fallback: only active lanes are touched. Arches with
465-
// hardware predicated stores override this.
465+
// hardware predicated stores should override this.
466466
constexpr std::size_t size = batch<T, A>::size;
467467
alignas(A::alignment()) std::array<T, size> src_buf;
468468
src.store_aligned(src_buf.data());

include/xsimd/arch/xsimd_avx.hpp

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1002,6 +1002,21 @@ namespace xsimd
10021002
return _mm256_maskload_pd(mem, _mm256_castpd_si256(mask));
10031003
}
10041004

1005+
// 4/8-byte ints: bitcast to same-width float, reuse the vmaskmov path.
1006+
template <class A, class T, class Mode>
1007+
XSIMD_INLINE std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 4 || sizeof(T) == 8), batch<T, A>>
1008+
load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<avx>) noexcept
1009+
{
1010+
XSIMD_IF_CONSTEXPR(sizeof(T) == 4)
1011+
{
1012+
return bitwise_cast<T>(batch<float, A>(_mm256_maskload_ps(reinterpret_cast<float const*>(mem), __m256i(mask))));
1013+
}
1014+
else
1015+
{
1016+
return bitwise_cast<T>(batch<double, A>(_mm256_maskload_pd(reinterpret_cast<double const*>(mem), __m256i(mask))));
1017+
}
1018+
}
1019+
10051020
// load_masked (single overload for float/double)
10061021
template <class A, class T, bool... Values, class Mode, class = std::enable_if_t<std::is_floating_point<T>::value>>
10071022
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool_constant<T, A, Values...> mask, convert<T>, Mode, requires_arch<avx>) noexcept
@@ -1096,6 +1111,21 @@ namespace xsimd
10961111
detail::maskstore(mem, mask, src);
10971112
}
10981113

1114+
// 4/8-byte ints: bitcast to same-width float, reuse the vmaskmov path.
1115+
template <class A, class T, class Mode>
1116+
XSIMD_INLINE std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 4 || sizeof(T) == 8), void>
1117+
store_masked(T* mem, batch<T, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<avx>) noexcept
1118+
{
1119+
XSIMD_IF_CONSTEXPR(sizeof(T) == 4)
1120+
{
1121+
_mm256_maskstore_ps(reinterpret_cast<float*>(mem), __m256i(mask), bitwise_cast<float>(src));
1122+
}
1123+
else
1124+
{
1125+
_mm256_maskstore_pd(reinterpret_cast<double*>(mem), __m256i(mask), bitwise_cast<double>(src));
1126+
}
1127+
}
1128+
10991129
// lt
11001130
template <class A>
11011131
XSIMD_INLINE batch_bool<float, A> lt(batch<float, A> const& self, batch<float, A> const& other, requires_arch<avx>) noexcept

include/xsimd/arch/xsimd_avx2.hpp

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -126,10 +126,12 @@ namespace xsimd
126126
{
127127
XSIMD_IF_CONSTEXPR(sizeof(T) == 4)
128128
{
129+
static_assert(sizeof(int) == 4, "_mm256_maskload_epi32 requires a 4-byte int");
129130
return _mm256_maskload_epi32(reinterpret_cast<int const*>(mem), mask);
130131
}
131132
else
132133
{
134+
static_assert(sizeof(long long) == 8, "_mm256_maskload_epi64 requires an 8-byte long long");
133135
return _mm256_maskload_epi64(reinterpret_cast<long long const*>(mem), mask);
134136
}
135137
}
@@ -139,10 +141,12 @@ namespace xsimd
139141
{
140142
XSIMD_IF_CONSTEXPR(sizeof(T) == 4)
141143
{
144+
static_assert(sizeof(int) == 4, "_mm256_maskstore_epi32 requires a 4-byte int");
142145
_mm256_maskstore_epi32(reinterpret_cast<int*>(mem), mask, src);
143146
}
144147
else
145148
{
149+
static_assert(sizeof(long long) == 8, "_mm256_maskstore_epi64 requires an 8-byte long long");
146150
_mm256_maskstore_epi64(reinterpret_cast<long long*>(mem), mask, src);
147151
}
148152
}
@@ -153,11 +157,12 @@ namespace xsimd
153157
}
154158
}
155159

160+
// no half-split shortcut for load; forward to runtime
156161
template <class A, class T, bool... Values, class Mode>
157162
XSIMD_INLINE std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 4 || sizeof(T) == 8), batch<T, A>>
158163
load_masked(T const* mem, batch_bool_constant<T, A, Values...> mask, convert<T>, Mode, requires_arch<avx2>) noexcept
159164
{
160-
return detail::maskload(mem, mask.as_batch());
165+
return load_masked(mem, mask.as_batch_bool(), convert<T> {}, Mode {}, avx2 {});
161166
}
162167

163168
template <class A, class T, class Mode>

include/xsimd/arch/xsimd_avx2_128.hpp

Lines changed: 22 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,25 @@ namespace xsimd
2424
{
2525
using namespace types;
2626

27+
// Defined in xsimd_avx512vl_128.hpp (included after this header). The
28+
// masked load/store half-fold below forwards to the 128-bit sized-batch
29+
// arch, which is avx512vl_128 when the 256-bit one stays in the AVX2
30+
// lineage (e.g. avxvnni). That unqualified dependent call resolves by
31+
// ordinary lookup here, so the avx512vl_128 overloads must be visible at
32+
// this point; declarations only, never instantiated unless AVX512VL is on.
33+
template <class A, class T, bool... V, class Mode,
34+
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
35+
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool_constant<T, A, V...> mask, convert<T>, Mode, requires_arch<avx512vl_128>) noexcept;
36+
template <class A, class T, class Mode,
37+
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
38+
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<avx512vl_128>) noexcept;
39+
template <class A, class T, bool... V, class Mode,
40+
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
41+
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool_constant<T, A, V...> mask, Mode, requires_arch<avx512vl_128>) noexcept;
42+
template <class A, class T, class Mode,
43+
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
44+
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<avx512vl_128>) noexcept;
45+
2746
// select
2847
template <class A, class T, bool... Values, class = std::enable_if_t<std::is_integral<T>::value>>
2948
XSIMD_INLINE batch<T, A> select(batch_bool_constant<T, A, Values...> const&, batch<T, A> const& true_br, batch<T, A> const& false_br, requires_arch<avx2_128>) noexcept
@@ -122,18 +141,19 @@ namespace xsimd
122141
}
123142
}
124143

144+
// constant masks gain nothing on a single register; forward to runtime
125145
template <class A, class T, bool... Values, class Mode,
126146
typename = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
127147
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool_constant<T, A, Values...> mask, convert<T>, Mode, requires_arch<avx2_128>) noexcept
128148
{
129-
return detail::maskload_avx2_128(mem, mask.as_batch());
149+
return load_masked(mem, mask.as_batch_bool(), convert<T> {}, Mode {}, avx2_128 {});
130150
}
131151

132152
template <class A, class T, bool... Values, class Mode,
133153
typename = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
134154
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool_constant<T, A, Values...> mask, Mode, requires_arch<avx2_128>) noexcept
135155
{
136-
detail::maskstore_avx2_128(mem, mask.as_batch(), __m128i(src));
156+
store_masked(mem, src, mask.as_batch_bool(), Mode {}, avx2_128 {});
137157
}
138158

139159
template <class A, class T, class Mode>

include/xsimd/arch/xsimd_avx512bw.hpp

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -378,6 +378,53 @@ namespace xsimd
378378
}
379379
}
380380

381+
// load_masked / store_masked: native vmovdqu8 / vmovdqu16 predication for
382+
// 8/16-bit, replacing the common scalar fallback. No aligned masked 8/16
383+
// intrinsic exists and masked moves never fault, so loadu fits both modes.
384+
template <class A, class T, class Mode,
385+
class = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 1 || sizeof(T) == 2)>>
386+
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<avx512bw>) noexcept
387+
{
388+
XSIMD_IF_CONSTEXPR(sizeof(T) == 1)
389+
{
390+
return _mm512_maskz_loadu_epi8((__mmask64)mask.mask(), mem);
391+
}
392+
else
393+
{
394+
return _mm512_maskz_loadu_epi16((__mmask32)mask.mask(), mem);
395+
}
396+
}
397+
398+
template <class A, class T, class Mode,
399+
class = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 1 || sizeof(T) == 2)>>
400+
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<avx512bw>) noexcept
401+
{
402+
XSIMD_IF_CONSTEXPR(sizeof(T) == 1)
403+
{
404+
_mm512_mask_storeu_epi8((void*)mem, (__mmask64)mask.mask(), src);
405+
}
406+
else
407+
{
408+
_mm512_mask_storeu_epi16((void*)mem, (__mmask32)mask.mask(), src);
409+
}
410+
}
411+
412+
// Constant masks reuse the runtime overloads; as_batch_bool() also avoids
413+
// batch_bool_constant::mask() truncating a 64-lane int8 mask to int.
414+
template <class A, class T, bool... Values, class Mode,
415+
class = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 1 || sizeof(T) == 2)>>
416+
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool_constant<T, A, Values...> mask, convert<T>, Mode, requires_arch<avx512bw>) noexcept
417+
{
418+
return load_masked(mem, mask.as_batch_bool(), convert<T> {}, Mode {}, avx512bw {});
419+
}
420+
421+
template <class A, class T, bool... Values, class Mode,
422+
class = std::enable_if_t<std::is_integral<T>::value && (sizeof(T) == 1 || sizeof(T) == 2)>>
423+
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool_constant<T, A, Values...> mask, Mode, requires_arch<avx512bw>) noexcept
424+
{
425+
store_masked(mem, src, mask.as_batch_bool(), Mode {}, avx512bw {});
426+
}
427+
381428
// max
382429
template <class A, class T, class = std::enable_if_t<std::is_integral<T>::value>>
383430
XSIMD_INLINE batch<T, A> max(batch<T, A> const& self, batch<T, A> const& other, requires_arch<avx512bw>) noexcept

include/xsimd/arch/xsimd_avx512f.hpp

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -356,7 +356,8 @@ namespace xsimd
356356

357357
// Runtime-mask load/store: same native k-register path as the constant
358358
// overloads above, minus the compile-time half-forwarding. 8/16-bit
359-
// elements fall back to the common scalar path.
359+
// elements are handled natively by avx512bw (vmovdqu8 / vmovdqu16);
360+
// without AVX512BW they fall back to the common scalar path.
360361
template <class A, class T, class Mode,
361362
typename = std::enable_if_t<(sizeof(T) >= 4)>>
362363
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<avx512f>) noexcept

include/xsimd/arch/xsimd_avx512vl_128.hpp

Lines changed: 36 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -212,92 +212,124 @@ namespace xsimd
212212
XSIMD_INLINE __m128i maskload128(T const* mem, uint64_t m, Mode) noexcept
213213
{
214214
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
215+
{
215216
return _mm_maskz_load_epi32((__mmask8)m, mem);
217+
}
216218
else
219+
{
217220
return _mm_maskz_loadu_epi32((__mmask8)m, mem);
221+
}
218222
}
219223
template <class T, class Mode, enable_sized_integral_t<T, 8> = 0>
220224
XSIMD_INLINE __m128i maskload128(T const* mem, uint64_t m, Mode) noexcept
221225
{
222226
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
227+
{
223228
return _mm_maskz_load_epi64((__mmask8)m, mem);
229+
}
224230
else
231+
{
225232
return _mm_maskz_loadu_epi64((__mmask8)m, mem);
233+
}
226234
}
227235
template <class Mode>
228236
XSIMD_INLINE __m128 maskload128(float const* mem, uint64_t m, Mode) noexcept
229237
{
230238
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
239+
{
231240
return _mm_maskz_load_ps((__mmask8)m, mem);
241+
}
232242
else
243+
{
233244
return _mm_maskz_loadu_ps((__mmask8)m, mem);
245+
}
234246
}
235247
template <class Mode>
236248
XSIMD_INLINE __m128d maskload128(double const* mem, uint64_t m, Mode) noexcept
237249
{
238250
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
251+
{
239252
return _mm_maskz_load_pd((__mmask8)m, mem);
253+
}
240254
else
255+
{
241256
return _mm_maskz_loadu_pd((__mmask8)m, mem);
257+
}
242258
}
243259

244260
template <class T, class Mode, enable_sized_integral_t<T, 4> = 0>
245261
XSIMD_INLINE void maskstore128(T* mem, __m128i src, uint64_t m, Mode) noexcept
246262
{
247263
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
264+
{
248265
_mm_mask_store_epi32(mem, (__mmask8)m, src);
266+
}
249267
else
268+
{
250269
_mm_mask_storeu_epi32(mem, (__mmask8)m, src);
270+
}
251271
}
252272
template <class T, class Mode, enable_sized_integral_t<T, 8> = 0>
253273
XSIMD_INLINE void maskstore128(T* mem, __m128i src, uint64_t m, Mode) noexcept
254274
{
255275
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
276+
{
256277
_mm_mask_store_epi64(mem, (__mmask8)m, src);
278+
}
257279
else
280+
{
258281
_mm_mask_storeu_epi64(mem, (__mmask8)m, src);
282+
}
259283
}
260284
template <class Mode>
261285
XSIMD_INLINE void maskstore128(float* mem, __m128 src, uint64_t m, Mode) noexcept
262286
{
263287
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
288+
{
264289
_mm_mask_store_ps(mem, (__mmask8)m, src);
290+
}
265291
else
292+
{
266293
_mm_mask_storeu_ps(mem, (__mmask8)m, src);
294+
}
267295
}
268296
template <class Mode>
269297
XSIMD_INLINE void maskstore128(double* mem, __m128d src, uint64_t m, Mode) noexcept
270298
{
271299
XSIMD_IF_CONSTEXPR(std::is_same<Mode, aligned_mode>::value)
300+
{
272301
_mm_mask_store_pd(mem, (__mmask8)m, src);
302+
}
273303
else
304+
{
274305
_mm_mask_storeu_pd(mem, (__mmask8)m, src);
306+
}
275307
}
276308
}
277309

278310
template <class A, class T, bool... V, class Mode,
279-
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
311+
typename>
280312
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool_constant<T, A, V...> mask, convert<T>, Mode, requires_arch<avx512vl_128>) noexcept
281313
{
282314
return detail::maskload128(mem, mask.mask(), Mode {});
283315
}
284316

285317
template <class A, class T, class Mode,
286-
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
318+
typename>
287319
XSIMD_INLINE batch<T, A> load_masked(T const* mem, batch_bool<T, A> mask, convert<T>, Mode, requires_arch<avx512vl_128>) noexcept
288320
{
289321
return detail::maskload128(mem, mask.mask(), Mode {});
290322
}
291323

292324
template <class A, class T, bool... V, class Mode,
293-
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
325+
typename>
294326
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool_constant<T, A, V...> mask, Mode, requires_arch<avx512vl_128>) noexcept
295327
{
296328
detail::maskstore128(mem, src, mask.mask(), Mode {});
297329
}
298330

299331
template <class A, class T, class Mode,
300-
typename = std::enable_if_t<std::is_arithmetic<T>::value && (sizeof(T) == 4 || sizeof(T) == 8)>>
332+
typename>
301333
XSIMD_INLINE void store_masked(T* mem, batch<T, A> const& src, batch_bool<T, A> mask, Mode, requires_arch<avx512vl_128>) noexcept
302334
{
303335
detail::maskstore128(mem, src, mask.mask(), Mode {});

0 commit comments

Comments
 (0)