21#if MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX512
29class Simd<double, 8> {
41 :
raw(_mm512_set_pd(h, g, f, e, d, c, b, a)) {}
43 template <
class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
44 Simd(U a) :
raw(_mm512_set1_pd(a)) {}
47 :
raw(_mm512_insertf64x4(_mm512_castpd256_pd512(low.
raw), high.
raw, 1)) {}
51 static_assert(i >= 0 && i <
kSize,
"Index out of range");
57#if MOCHI_COMPILER_MSVC
58 return raw.m512d_f64[i];
66 static_assert(iHalf == 0 || iHalf == 1);
67 if constexpr (iHalf == 0) {
68 return _mm512_castpd512_pd256(a.raw);
70 return _mm512_extractf64x4_pd(a.raw, 1);
76 auto const mask =
static_cast<__mmask8
>(1u << i);
77 return _mm512_mask_broadcastsd_pd(v.raw, mask, _mm_set_sd(value));
82 static_assert(i >= 0 && i <
kSize,
"Index out of range");
83 constexpr auto kMask =
static_cast<__mmask8
>(1u << i);
84 return _mm512_mask_broadcastsd_pd(v.raw, kMask, _mm_set_sd(value));
89 static_assert(N >= 1 && N <=
kSize,
"Unsupported N");
90 auto const mask = ToMask(v);
91 if constexpr (N ==
kSize) {
92 return _kortestc_mask8_u8(mask, mask) != 0;
94 constexpr auto kLanes = LaneMask<N>();
95 return (mask & kLanes) == kLanes;
101 static_assert(N >= 1 && N <=
kSize,
"Unsupported N");
102 auto const mask = ToMask(v);
103 if constexpr (N ==
kSize) {
104 return _kortestz_mask8_u8(mask, mask) == 0;
106 return (mask & LaneMask<N>()) != 0;
111 return _mm512_set1_pd(*p);
116 static_assert(i >= 0 && i <
kSize,
"Index out of range");
117 if constexpr (i == 0) {
118 return _mm512_broadcastsd_pd(_mm512_castpd512_pd128(v.raw));
120 constexpr int kLane = i % 2;
121 constexpr int kGroup = i / 2;
123 _mm512_shuffle_f64x2(v.raw, v.raw, _MM_SHUFFLE(kGroup, kGroup, kGroup, kGroup));
124 return _mm512_permute_pd(group, kLane == 0 ? 0x00 : 0xFF);
128 template <
int N = kSize>
130 static_assert(N >= 0 && N <=
kSize);
131 if constexpr (N == 0) {
133 }
else if constexpr (N == 1) {
134 return _mm512_zextpd128_pd512(_mm_load_sd(ptr));
135 }
else if constexpr (N == 2) {
136 return _mm512_zextpd128_pd512(_mm_loadu_pd(ptr));
137 }
else if constexpr (N == 3) {
138 return _mm512_zextpd256_pd512(_mm256_maskz_loadu_pd(LaneMask<N>(), ptr));
139 }
else if constexpr (N == 4) {
140 return _mm512_zextpd256_pd512(_mm256_loadu_pd(ptr));
141 }
else if constexpr (N <
kSize) {
142 return _mm512_maskz_loadu_pd(LaneMask<N>(), ptr);
144 return _mm512_loadu_pd(ptr);
150 return _mm512_maskz_loadu_pd(LaneMask(n), ptr);
154 return _mm512_i32gather_pd(indices.raw, ptr,
sizeof(
double));
158 return _mm512_i64gather_pd(indices.raw, ptr,
sizeof(
double));
161 template <
int kTupleCount = kSize>
164 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
165 if constexpr (kTupleCount == 1) {
171 constexpr int kTotalCount = kTupleCount * 3;
173 if constexpr (kTupleCount <= 5) {
174 auto const index0 = LoadTransposeIndices<kTupleCount, 0>();
175 auto const index1 = LoadTransposeIndices<kTupleCount, 1>();
176 auto const index2 = LoadTransposeIndices<kTupleCount, 2>();
177 if constexpr (kTupleCount <= 2) {
178 out0.raw = _mm512_permutexvar_pd(index0, x0);
179 out1.raw = _mm512_permutexvar_pd(index1, x0);
180 out2.raw = _mm512_permutexvar_pd(index2, x0);
183 out0.raw = _mm512_permutex2var_pd(x0, index0, x1);
184 out1.raw = _mm512_permutex2var_pd(x0, index1, x1);
185 out2.raw = _mm512_permutex2var_pd(x0, index2, x1);
189 constexpr int kCount2 = kTotalCount - 2 *
kSize;
191 auto const index0 = _mm512_setr_epi64(0, 3, 6, 9, 12, 15, 0, 0);
192 auto const index1 = _mm512_setr_epi64(1, 4, 7, 10, 13, 0, 0, 0);
193 auto const index2 = _mm512_setr_epi64(2, 5, 8, 11, 14, 0, 0, 0);
194 if constexpr (kTupleCount == 6) {
195 out0.raw = _mm512_maskz_permutex2var_pd(LaneMask<kTupleCount>(), x0, index0, x1);
196 constexpr int kZeroIndex =
kSize + kCount2;
197 auto const partial1 = _mm512_permutex2var_pd(x0, index1, x1);
198 auto const finalIndex1 = _mm512_setr_epi64(0, 1, 2, 3, 4, 8, kZeroIndex, kZeroIndex);
199 out1.raw = _mm512_permutex2var_pd(partial1, finalIndex1, x2);
200 auto const partial2 = _mm512_permutex2var_pd(x0, index2, x1);
201 auto const finalIndex2 = _mm512_setr_epi64(0, 1, 2, 3, 4, 9, kZeroIndex, kZeroIndex);
202 out2.raw = _mm512_permutex2var_pd(partial2, finalIndex2, x2);
204 constexpr int kZeroIndex =
kSize + kCount2;
205 auto const partial0 = _mm512_permutex2var_pd(x0, index0, x1);
206 auto const finalIndex0 = _mm512_setr_epi64(
207 0, 1, 2, 3, 4, 5, kTupleCount > 6 ? 10 : kZeroIndex, kTupleCount > 7 ? 13 : kZeroIndex);
208 out0.raw = _mm512_permutex2var_pd(partial0, finalIndex0, x2);
209 auto const partial1 = _mm512_permutex2var_pd(x0, index1, x1);
210 auto const finalIndex1 = _mm512_setr_epi64(
211 0, 1, 2, 3, 4, 8, kTupleCount > 6 ? 11 : kZeroIndex, kTupleCount > 7 ? 14 : kZeroIndex);
212 out1.raw = _mm512_permutex2var_pd(partial1, finalIndex1, x2);
213 auto const partial2 = _mm512_permutex2var_pd(x0, index2, x1);
214 auto const finalIndex2 = _mm512_setr_epi64(
215 0, 1, 2, 3, 4, 9, kTupleCount > 6 ? 12 : kZeroIndex, kTupleCount > 7 ? 15 : kZeroIndex);
216 out2.raw = _mm512_permutex2var_pd(partial2, finalIndex2, x2);
221 template <
int N = kSize>
223 static_assert(N >= 0 && N <=
kSize);
224 if constexpr (N == 0) {
225 }
else if constexpr (N == 1) {
226 _mm_store_sd(ptr, _mm512_castpd512_pd128(v.raw));
227 }
else if constexpr (N == 2) {
228 _mm_storeu_pd(ptr, _mm512_castpd512_pd128(v.raw));
229 }
else if constexpr (N == 3) {
230 _mm256_mask_storeu_pd(
231 ptr,
static_cast<__mmask8
>((uint32_t{1} << N) - 1), _mm512_castpd512_pd256(v.raw));
232 }
else if constexpr (N == 4) {
233 _mm256_storeu_pd(ptr, _mm512_castpd512_pd256(v.raw));
234 }
else if constexpr (N <
kSize) {
235 _mm512_mask_storeu_pd(ptr, LaneMask(N), v.raw);
237 _mm512_storeu_pd(ptr, v.raw);
243 _mm512_mask_storeu_pd(ptr, LaneMask(n), v.raw);
247 __mmask8
const mask = ToMask(condition);
248 _mm512_mask_compressstoreu_pd(ptr, mask, values.raw);
249 return _mm_popcnt_u32(mask);
252 template <
int kTupleCount = kSize>
254 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
255 if constexpr (kTupleCount == 1) {
261 constexpr int kTotalCount = kTupleCount * 3;
263 _mm512_permutex2var_pd(a.raw, _mm512_setr_epi64(0, 8, 0, 1, 9, 0, 2, 10), b.raw);
264 auto const x0 = _mm512_permutex2var_pd(ab0, _mm512_setr_epi64(0, 1, 8, 3, 4, 9, 6, 7), c.raw);
266 if constexpr (kTotalCount >
kSize) {
268 _mm512_permutex2var_pd(a.raw, _mm512_setr_epi64(0, 3, 11, 0, 4, 12, 0, 5), b.raw);
270 _mm512_permutex2var_pd(ab1, _mm512_setr_epi64(10, 1, 2, 11, 4, 5, 12, 7), c.raw);
273 if constexpr (kTotalCount > 2 *
kSize) {
275 _mm512_permutex2var_pd(a.raw, _mm512_setr_epi64(13, 0, 6, 14, 0, 7, 15, 0), b.raw);
277 _mm512_permutex2var_pd(ab2, _mm512_setr_epi64(0, 13, 2, 3, 14, 5, 6, 15), c.raw);
283 return _mm512_mask_blend_pd(ToMask(mask), b.raw, a.raw);
287 return _mm512_sqrt_pd(v.raw);
291 return _mm512_rcp14_pd(v.raw);
295 return _mm512_rsqrt14_pd(v.raw);
299 return _mm512_castsi512_pd(
300 _mm512_set1_epi64(
static_cast<long long>(0x8000000000000000ULL)));
304 return _mm512_abs_pd(v.raw);
308 return _mm512_min_pd(a.raw, b.raw);
312 return _mm512_max_pd(a.raw, b.raw);
316 return _mm512_floor_pd(a.raw);
320 return _mm512_roundscale_pd(v.raw, _MM_FROUND_TO_NEAREST_INT);
323#if MOCHI_ARCH_X64_SVML
325 return _mm512_cos_pd(a.raw);
329 return _mm512_sin_pd(a.raw);
333 return _mm512_tan_pd(a.raw);
337 return _mm512_acos_pd(a.raw);
341 return _mm512_asin_pd(a.raw);
345 return _mm512_atan_pd(a.raw);
349 return _mm512_exp_pd(a.raw);
353 return _mm512_log_pd(a.raw);
357 return _mm512_tanh_pd(a.raw);
362 return _mm512_fmadd_pd(a.raw, b.raw, c.raw);
366 return _mm512_fmsub_pd(a.raw, b.raw, c.raw);
370 return _mm512_fnmadd_pd(a.raw, b.raw, c.raw);
374 return _mm512_fnmsub_pd(a.raw, b.raw, c.raw);
377 template <
int N = kSize>
379 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
382 if constexpr (N <= 4) {
383 return HalfT::template
HMin<N>(lo);
386 if constexpr (N == 5) {
388 }
else if constexpr (N == 8) {
389 return HalfT::template
HMin<4>(HalfT::Min(lo, hi));
391 return _mm512_mask_reduce_min_pd(LaneMask<N>(), a.raw);
396 template <
int N = kSize>
398 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
401 if constexpr (N <= 4) {
402 return HalfT::template
HMax<N>(lo);
405 if constexpr (N == 5) {
407 }
else if constexpr (N == 8) {
408 return HalfT::template
HMax<4>(HalfT::Max(lo, hi));
410 return _mm512_mask_reduce_max_pd(LaneMask<N>(), a.raw);
417 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
420 if constexpr (N <= 4) {
421 return HalfT::template
HSum<N>(lo);
424 if constexpr (N == 5) {
425 return HalfT::template
HSum<4>(lo) + HalfT::template
Get<0>(hi);
426 }
else if constexpr (N ==
kSize) {
427 return HalfT::template
HSum<4>(lo + hi);
429 return _mm512_mask_reduce_add_pd(LaneMask<N>(), a.raw);
436 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
439 if constexpr (N <= 4) {
440 return HalfT::template
HProd<N>(lo);
443 if constexpr (N == 5) {
445 }
else if constexpr (N ==
kSize) {
446 return HalfT::template
HProd<4>(lo * hi);
448 return _mm512_mask_reduce_mul_pd(LaneMask<N>(), a.raw);
455 static_assert(N >= 2 && N <=
kSize,
"Unsupported N");
460 return FromMask(_mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_LT_OQ));
464 return FromMask(_mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_GT_OQ));
468 return FromMask(_mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_LE_OQ));
472 return FromMask(_mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_GE_OQ));
476 return FromMask(_mm512_cmp_pd_mask(a.raw, b.raw, _CMP_EQ_OQ));
480 return FromMask(_mm512_cmp_pd_mask(a.raw, b.raw, _CMP_NEQ_UQ));
484 return _mm512_setzero_pd();
488 return _mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_EQ_OQ) ==
static_cast<__mmask8
>(0xFFu);
492 return _mm512_cmp_pd_mask(
raw, rhs.raw, _CMP_NEQ_UQ) != 0;
496 return _mm512_castsi512_pd(
497 _mm512_xor_si512(_mm512_castpd_si512(
raw), _mm512_set1_epi64(-1)));
501 return _mm512_xor_pd(
raw, SignBitMask().
raw);
505 return _mm512_add_pd(
raw, rhs.raw);
509 return _mm512_sub_pd(
raw, rhs.raw);
513 return _mm512_mul_pd(
raw, rhs.raw);
517 return _mm512_div_pd(
raw, rhs.raw);
521 return _mm512_and_pd(
raw, rhs.raw);
525 return _mm512_or_pd(
raw, rhs.raw);
529 return _mm512_xor_pd(
raw, rhs.raw);
533 template <
int kTupleCount,
int kComponent>
535 constexpr int kZeroIndex = kTupleCount * 3;
536 return _mm512_setr_epi64(
538 kTupleCount > 1 ? 3 + kComponent : kZeroIndex,
539 kTupleCount > 2 ? 6 + kComponent : kZeroIndex,
540 kTupleCount > 3 ? 9 + kComponent : kZeroIndex,
541 kTupleCount > 4 ? 12 + kComponent : kZeroIndex,
542 kTupleCount > 5 ? 15 + kComponent : kZeroIndex,
543 kTupleCount > 6 ? 18 + kComponent : kZeroIndex,
544 kTupleCount > 7 ? 21 + kComponent : kZeroIndex);
549 [[nodiscard]]
static constexpr __mmask8 LaneMask() {
550 static_assert(N >= 0 && N <=
kSize);
555 [[nodiscard]]
static constexpr __mmask8 LaneMask(
int n) {
557 return static_cast<__mmask8
>((uint32_t{1} << n) - 1);
562 auto const bits = _mm512_castpd_si512(a.raw);
563 auto const mask = _mm512_movepi64_mask(bits);
565 _mm512_cmpeq_epi64_mask(bits, _mm512_movm_epi64(mask)) == LaneMask<kSize>(),
566 "Expected a canonical logical mask");
572 return _mm512_castsi512_pd(_mm512_movm_epi64(mask));
Simd operator&(Simd rhs) const
bool operator==(Simd rhs) const
Simd operator>(Simd rhs) 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,...)
T Dot(Simd< T, N > a, Simd< T, N > b)
V LoadIndexed(typename V::Scalar const *ptr, Simd< I, V::kSize > indices)
constexpr T const & Min(T const &a, T const &b)
constexpr auto Equal(T const &a, T const &b)
constexpr auto MulAdd(A a, B b, C c)
Simd< T, N > Tanh(Simd< T, N > a)
Simd< T, N > Set(Simd< T, N > a, T value)
constexpr auto NotEqual(T const &a, T const &b)
V Broadcast(typename V::Scalar a)
constexpr auto MulSub(A a, B b, C c)
constexpr T Select(bool condition, T a, T b)
Simd< T, N > Ln(Simd< T, N > a)
constexpr auto NegMulAdd(A a, B b, C c)
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)
Simd< T, N > FastRound(Simd< T, N > a)
void StoreTransposed(T *ptr, Simd< T, N > a, Simd< T, N > b, Simd< T, N > c)
void Store(T *ptr, Simd< T, N > a)
Simd< T, N > RcpSqrtApprox(Simd< T, N > a)
constexpr T RcpApprox(T a)
constexpr auto NegMulSub(A a, B b, C c)
int StoreSelected(T *ptr, Simd< MaskT, N > condition, Simd< T, N > values)
V Load(typename V::Scalar const *ptr)
#define MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(T, N, NativeT)