21#if MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX512
29class Simd<int64_t, 8> {
31 static_assert(
sizeof(int64_t) ==
sizeof(
long long));
44 :
raw(_mm512_set_epi64(h, g, f, e, d, c, b, a)) {}
46 template <
class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
47 Simd(U a) :
raw(_mm512_set1_epi64(static_cast<long long>(a))) {}
50 :
raw(_mm512_inserti64x4(_mm512_castsi256_si512(low.
raw), high.
raw, 1)) {}
54 static_assert(i >= 0 && i <
kSize,
"Index out of range");
60#if MOCHI_COMPILER_MSVC
61 return raw.m512i_i64[i];
69 static_assert(iHalf == 0 || iHalf == 1);
70 if constexpr (iHalf == 0) {
71 return _mm512_castsi512_si256(a.raw);
73 return _mm512_extracti64x4_epi64(a.raw, 1);
79 static_assert(N >= 1 && N <=
kSize,
"Unsupported N");
80 auto const mask = ToMask(v);
81 if constexpr (N ==
kSize) {
82 return _kortestc_mask8_u8(mask, mask) != 0;
84 constexpr auto kLanes = LaneMask<N>();
85 return (mask & kLanes) == kLanes;
95 static_assert(i >= 0 && i <
kSize,
"Index out of range");
96 if constexpr (i == 0) {
97 return _mm512_broadcastq_epi64(_mm512_castsi512_si128(v.raw));
99 constexpr int kLane = i % 2;
100 constexpr int kGroup = i / 2;
102 _mm512_shuffle_i64x2(v.raw, v.raw, _MM_SHUFFLE(kGroup, kGroup, kGroup, kGroup));
103 return _mm512_shuffle_epi32(
105 static_cast<_MM_PERM_ENUM
>(
106 kLane == 0 ? _MM_SHUFFLE(1, 0, 1, 0) : _MM_SHUFFLE(3, 2, 3, 2)));
110 template <
int N = kSize>
112 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
115 if constexpr (N <= 4) {
116 return HalfT::template
HMin<N>(lo);
119 if constexpr (N == 5) {
121 }
else if constexpr (N ==
kSize) {
122 return HalfT::template
HMin<4>(HalfT::Min(lo, hi));
124 return _mm512_mask_reduce_min_epi64(LaneMask<N>(), a.raw);
129 template <
int N = kSize>
131 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
134 if constexpr (N <= 4) {
135 return HalfT::template
HMax<N>(lo);
138 if constexpr (N == 5) {
140 }
else if constexpr (N ==
kSize) {
141 return HalfT::template
HMax<4>(HalfT::Max(lo, hi));
143 return _mm512_mask_reduce_max_epi64(LaneMask<N>(), a.raw);
150 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
153 if constexpr (N <= 4) {
154 return HalfT::template
HSum<N>(lo);
157 if constexpr (N == 5) {
158 return HalfT::template
HSum<4>(lo) + HalfT::template
Get<0>(hi);
159 }
else if constexpr (N ==
kSize) {
160 return HalfT::template
HSum<4>(lo + hi);
162 return _mm512_mask_reduce_add_epi64(LaneMask<N>(), a.raw);
167 template <
int N = kSize>
169 static_assert(N >= 0 && N <=
kSize);
170 if constexpr (N == 0) {
172 }
else if constexpr (N == 1) {
173 return _mm512_zextsi128_si512(_mm_loadl_epi64(
reinterpret_cast<__m128i const*
>(ptr)));
174 }
else if constexpr (N == 2) {
175 return _mm512_zextsi128_si512(_mm_loadu_si128(
reinterpret_cast<__m128i const*
>(ptr)));
176 }
else if constexpr (N < 4) {
177 return _mm512_zextsi256_si512(
178 _mm256_maskz_loadu_epi64(
static_cast<__mmask8
>((uint32_t{1} << N) - 1), ptr));
179 }
else if constexpr (N == 4) {
180 return _mm512_zextsi256_si512(_mm256_loadu_si256(
reinterpret_cast<__m256i const*
>(ptr)));
181 }
else if constexpr (N <
kSize) {
182 return _mm512_maskz_loadu_epi64(LaneMask<N>(), ptr);
184 return _mm512_loadu_si512(ptr);
190 auto const mask =
static_cast<__mmask8
>((uint32_t{1} << n) - 1);
191 return _mm512_maskz_loadu_epi64(mask, ptr);
194 template <
int kTupleCount = kSize>
197 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
198 if constexpr (kTupleCount == 1) {
204 constexpr int kTotalCount = kTupleCount * 3;
206 if constexpr (kTupleCount <= 5) {
207 auto const index0 = LoadTransposeIndices<kTupleCount, 0>();
208 auto const index1 = LoadTransposeIndices<kTupleCount, 1>();
209 auto const index2 = LoadTransposeIndices<kTupleCount, 2>();
210 if constexpr (kTupleCount <= 2) {
211 out0.raw = _mm512_permutexvar_epi64(index0, x0);
212 out1.raw = _mm512_permutexvar_epi64(index1, x0);
213 out2.raw = _mm512_permutexvar_epi64(index2, x0);
216 out0.raw = _mm512_permutex2var_epi64(x0, index0, x1);
217 out1.raw = _mm512_permutex2var_epi64(x0, index1, x1);
218 out2.raw = _mm512_permutex2var_epi64(x0, index2, x1);
222 constexpr int kCount2 = kTotalCount - 2 *
kSize;
224 auto const index0 = _mm512_setr_epi64(0, 3, 6, 9, 12, 15, 0, 0);
225 auto const index1 = _mm512_setr_epi64(1, 4, 7, 10, 13, 0, 0, 0);
226 auto const index2 = _mm512_setr_epi64(2, 5, 8, 11, 14, 0, 0, 0);
227 if constexpr (kTupleCount == 6) {
228 out0.raw = _mm512_maskz_permutex2var_epi64(LaneMask<kTupleCount>(), x0, index0, x1);
229 constexpr int kZeroIndex =
kSize + kCount2;
230 auto const partial1 = _mm512_permutex2var_epi64(x0, index1, x1);
231 auto const finalIndex1 = _mm512_setr_epi64(0, 1, 2, 3, 4, 8, kZeroIndex, kZeroIndex);
232 out1.raw = _mm512_permutex2var_epi64(partial1, finalIndex1, x2);
233 auto const partial2 = _mm512_permutex2var_epi64(x0, index2, x1);
234 auto const finalIndex2 = _mm512_setr_epi64(0, 1, 2, 3, 4, 9, kZeroIndex, kZeroIndex);
235 out2.raw = _mm512_permutex2var_epi64(partial2, finalIndex2, x2);
237 constexpr int kZeroIndex =
kSize + kCount2;
238 auto const partial0 = _mm512_permutex2var_epi64(x0, index0, x1);
239 auto const finalIndex0 = _mm512_setr_epi64(
240 0, 1, 2, 3, 4, 5, kTupleCount > 6 ? 10 : kZeroIndex, kTupleCount > 7 ? 13 : kZeroIndex);
241 out0.raw = _mm512_permutex2var_epi64(partial0, finalIndex0, x2);
242 auto const partial1 = _mm512_permutex2var_epi64(x0, index1, x1);
243 auto const finalIndex1 = _mm512_setr_epi64(
244 0, 1, 2, 3, 4, 8, kTupleCount > 6 ? 11 : kZeroIndex, kTupleCount > 7 ? 14 : kZeroIndex);
245 out1.raw = _mm512_permutex2var_epi64(partial1, finalIndex1, x2);
246 auto const partial2 = _mm512_permutex2var_epi64(x0, index2, x1);
247 auto const finalIndex2 = _mm512_setr_epi64(
248 0, 1, 2, 3, 4, 9, kTupleCount > 6 ? 12 : kZeroIndex, kTupleCount > 7 ? 15 : kZeroIndex);
249 out2.raw = _mm512_permutex2var_epi64(partial2, finalIndex2, x2);
255 return _mm512_min_epi64(a.raw, b.raw);
259 return _mm512_max_epi64(a.raw, b.raw);
263 return _mm512_mask_blend_epi64(ToMask(mask), b.raw, a.raw);
266 template <
int x0,
int x1,
int x2,
int x3,
int x4,
int x5,
int x6,
int x7>
269 x0 >= 0 && x0 < kSize && x1 >= 0 && x1 < kSize && x2 >= 0 && x2 < kSize && x3 >= 0 &&
270 x3 < kSize && x4 >= 0 && x4 < kSize && x5 >= 0 && x5 < kSize && x6 >= 0 && x6 <
kSize &&
271 x7 >= 0 && x7 <
kSize,
274 x0 == 0 && x1 == 1 && x2 == 2 && x3 == 3 && x4 == 4 && x5 == 5 && x6 == 6 && x7 == 7) {
277 auto const indices = _mm512_setr_epi64(x0, x1, x2, x3, x4, x5, x6, x7);
278 return _mm512_permutexvar_epi64(indices, a.raw);
282 template <
int x0,
int x1,
int x2,
int x3,
int x4,
int x5,
int x6,
int x7>
285 x0 >= 0 && x0 < kSize && x1 >= 0 && x1 < kSize && x2 >= 0 && x2 < kSize && x3 >= 0 &&
286 x3 < kSize && x4 >= 0 && x4 < kSize && x5 >= 0 && x5 < kSize && x6 >= 0 && x6 <
kSize &&
287 x7 >= 0 && x7 <
kSize,
290 x0 == 0 && x1 == 1 && x2 == 2 && x3 == 3 && x4 == 0 && x5 == 1 && x6 == 2 && x7 == 3) {
295 return _mm512_permutex2var_epi64(a.raw, indices, b.raw);
299 template <
int N = kSize>
301 static_assert(N >= 0 && N <=
kSize);
302 if constexpr (N == 0) {
303 }
else if constexpr (N == 1) {
304 _mm_storel_epi64(
reinterpret_cast<__m128i*
>(ptr), _mm512_castsi512_si128(v.raw));
305 }
else if constexpr (N == 2) {
306 _mm_storeu_si128(
reinterpret_cast<__m128i*
>(ptr), _mm512_castsi512_si128(v.raw));
307 }
else if constexpr (N < 4) {
308 _mm256_mask_storeu_epi64(
309 ptr,
static_cast<__mmask8
>((uint32_t{1} << N) - 1), _mm512_castsi512_si256(v.raw));
310 }
else if constexpr (N == 4) {
311 _mm256_storeu_si256(
reinterpret_cast<__m256i*
>(ptr), _mm512_castsi512_si256(v.raw));
312 }
else if constexpr (N <
kSize) {
313 _mm512_mask_storeu_epi64(ptr, LaneMask<N>(), v.raw);
315 _mm512_storeu_si512(ptr, v.raw);
321 auto const mask =
static_cast<__mmask8
>((uint32_t{1} << n) - 1);
322 _mm512_mask_storeu_epi64(ptr, mask, v.raw);
326 auto const mask = ToMask(condition);
327 _mm512_mask_compressstoreu_epi64(ptr, mask, values.raw);
328 return _mm_popcnt_u32(mask);
331 template <
int kTupleCount = kSize>
333 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
334 if constexpr (kTupleCount == 1) {
340 constexpr int kTotalCount = kTupleCount * 3;
342 _mm512_permutex2var_epi64(a.raw, _mm512_setr_epi64(0, 8, 0, 1, 9, 0, 2, 10), b.raw);
344 _mm512_permutex2var_epi64(ab0, _mm512_setr_epi64(0, 1, 8, 3, 4, 9, 6, 7), c.raw);
346 if constexpr (kTotalCount >
kSize) {
348 _mm512_permutex2var_epi64(a.raw, _mm512_setr_epi64(0, 3, 11, 0, 4, 12, 0, 5), b.raw);
350 _mm512_permutex2var_epi64(ab1, _mm512_setr_epi64(10, 1, 2, 11, 4, 5, 12, 7), c.raw);
353 if constexpr (kTotalCount > 2 *
kSize) {
355 _mm512_permutex2var_epi64(a.raw, _mm512_setr_epi64(13, 0, 6, 14, 0, 7, 15, 0), b.raw);
357 _mm512_permutex2var_epi64(ab2, _mm512_setr_epi64(0, 13, 2, 3, 14, 5, 6, 15), c.raw);
363 return _mm512_setzero_si512();
367 return FromMask(_mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_LT));
371 return FromMask(_mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_GT));
375 return FromMask(_mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_LE));
379 return FromMask(_mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_GE));
383 return FromMask(_mm512_cmp_epi64_mask(a.raw, b.raw, _MM_CMPINT_EQ));
387 return FromMask(_mm512_cmp_epi64_mask(a.raw, b.raw, _MM_CMPINT_NE));
391 return _mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_EQ) == __mmask8{0xFF};
395 return _mm512_cmp_epi64_mask(
raw, rhs.raw, _MM_CMPINT_NE) != 0;
399 return _mm512_xor_si512(
raw, _mm512_set1_epi64(-1));
403 return _mm512_sub_epi64(_mm512_setzero_si512(),
raw);
407 return _mm512_add_epi64(
raw, rhs.raw);
411 return _mm512_sub_epi64(
raw, rhs.raw);
415 return _mm512_mullo_epi64(
raw, rhs.raw);
419#if MOCHI_ARCH_X64_SVML
420 return _mm512_div_epi64(
raw, rhs.raw);
435 return _mm512_and_si512(
raw, rhs.raw);
439 return _mm512_or_si512(
raw, rhs.raw);
443 return _mm512_xor_si512(
raw, rhs.raw);
447 return _mm512_sll_epi64(
raw, _mm_cvtsi32_si128(rhs));
450 template <
int kShift>
452 static_assert(kShift >= 0 && kShift < 64,
"Shift amount out-of-range");
453 if constexpr (kShift == 0) {
456 return _mm512_srai_epi64(a.raw, kShift);
461 template <
int kTupleCount,
int kComponent>
463 constexpr int kZeroIndex = kTupleCount * 3;
464 return _mm512_setr_epi64(
466 kTupleCount > 1 ? 3 + kComponent : kZeroIndex,
467 kTupleCount > 2 ? 6 + kComponent : kZeroIndex,
468 kTupleCount > 3 ? 9 + kComponent : kZeroIndex,
469 kTupleCount > 4 ? 12 + kComponent : kZeroIndex,
470 kTupleCount > 5 ? 15 + kComponent : kZeroIndex,
471 kTupleCount > 6 ? 18 + kComponent : kZeroIndex,
472 kTupleCount > 7 ? 21 + kComponent : kZeroIndex);
477 [[nodiscard]]
static constexpr __mmask8 LaneMask() {
478 static_assert(N >= 0 && N <=
kSize);
479 if constexpr (N ==
kSize) {
480 return __mmask8{0xFF};
482 return static_cast<__mmask8
>((uint32_t{1} << N) - 1);
488 auto const mask = _mm512_movepi64_mask(a.raw);
490 _mm512_cmpeq_epi64_mask(a.raw, _mm512_movm_epi64(mask)) == LaneMask<kSize>(),
491 "Expected a canonical logical mask");
497 return _mm512_movm_epi64(mask);
Simd operator&(Simd rhs) const
bool operator==(Simd rhs) const
Simd operator>(Simd rhs) const
Simd operator<<(int shift) const
Simd operator*(Simd rhs) const
Simd operator^(Simd rhs) const
Simd operator>=(Simd rhs) const
Simd operator<(Simd rhs) const
bool operator!=(Simd rhs) const
Simd operator|(Simd rhs) const
static constexpr int kSize
Simd operator+(Simd rhs) const
Simd operator/(Simd rhs) const
Simd operator<=(Simd rhs) const
Scalar operator[](int i) const
#define MOCHI_ASSERT_VERBOSE(condition_without_side_effects,...)
Simd< T, 2 > Shuffle(Simd< T, 2 > a)
constexpr T const & Min(T const &a, T const &b)
constexpr auto Equal(T const &a, T const &b)
constexpr auto NotEqual(T const &a, T const &b)
V Broadcast(typename V::Scalar a)
constexpr T Select(bool condition, T a, T b)
constexpr ValT Clamp(ValT value, MinT min, MaxT max)
Simd< T, N/2 > GetHalf(Simd< T, N > a)
constexpr T const & Max(T const &a, T const &b)
void LoadTransposed(T const *ptr, Simd< T, N > &out0, Simd< T, N > &out1, Simd< T, N > &out2)
void StoreTransposed(T *ptr, Simd< T, N > a, Simd< T, N > b, Simd< T, N > c)
void Store(T *ptr, Simd< T, N > a)
int StoreSelected(T *ptr, Simd< MaskT, N > condition, Simd< T, N > values)
V Load(typename V::Scalar const *ptr)
Simd< T, N > ShiftRight(Simd< T, N > a)
#define MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(T, N, NativeT)