21#if MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX2
29class Simd<double, 4> {
32 Simd(
double a,
double b,
double c = 0.0,
double d = 0.0)
33 :
raw(_mm256_set_pd(d, c, b, a)) {}
34 template <
class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
35 Simd(U a) :
raw(_mm256_set1_pd(a)) {}
39 :
raw(_mm256_set_m128d(high.
raw, low.
raw)) {}
43 static_assert(i >= 0 && i < 4,
"Index out of range");
49#if MOCHI_COMPILER_MSVC
50 return raw.m256d_f64[i];
58#if MOCHI_COMPILER_MSVC
60 result.raw.m256d_f64[i] = value;
63 static constexpr __m256i kMasks[] = {
64 {-1LL, 0LL, 0LL, 0LL}, {0LL, -1LL, 0LL, 0LL}, {0LL, 0LL, -1LL, 0LL}, {0LL, 0LL, 0LL, -1LL}};
65 return _mm256_blendv_pd(v.raw, _mm256_set1_pd(value), _mm256_castsi256_pd(kMasks[i]));
71 static_assert(i >= 0 && i <
kSize,
"Index out of range");
72 return Set(v, i, value);
77 static_assert(iHalf == 0 || iHalf == 1);
78 return _mm256_extractf128_pd(a.raw, iHalf);
83 auto araw = _mm256_castpd_si256(a.raw);
84 auto v = _mm256_insert_epi64(araw, 0x3FF0000000000000LL, 3);
85 return _mm256_castsi256_pd(v);
90 auto araw = _mm256_castpd_si256(a.raw);
91 auto v = _mm256_insert_epi64(araw, 0, 3);
92 return _mm256_castsi256_pd(v);
95 template <
int x,
int y,
int z,
int w>
98 x >= 0 && x < 2 && y >= 0 && y < 2 && z >= 0 && z < 2 && w >= 0 && w < 2,
99 "invalid blend index");
100 if constexpr (x == 0 && y == 0 && z == 0 && w == 0) {
102 }
else if constexpr (x == 1 && y == 1 && z == 1 && w == 1) {
105 return _mm256_blend_pd(a.raw, b.raw, x | (y << 1) | (z << 2) | (w << 3));
111 static_assert(N >= 1 && N <=
kSize,
"Unsupported N");
112 int mask = GetMSBitMask(v);
113 if constexpr (N ==
kSize) {
114 return mask == 0xFFFFFFFF;
116 int constexpr kNumBits = N *
sizeof(
Scalar);
117 auto constexpr kMustBeSet = (1UL << kNumBits) - 1;
118 return (mask & kMustBeSet) == kMustBeSet;
124 static_assert(N >= 1 && N <=
kSize,
"Unsupported N");
125 int mask = GetMSBitMask(v);
126 if constexpr (N ==
kSize) {
129 int constexpr kNumBits = N *
sizeof(
Scalar);
130 auto constexpr kMayBeSet = (1UL << kNumBits) - 1;
131 return (mask & kMayBeSet) != 0;
136 return _mm256_broadcast_sd(p);
144 template <
int N = kSize>
146 static_assert(N >= 0 && N <= 4);
147 if constexpr (N == 0) {
149 }
else if constexpr (N == 1) {
150 return Simd{*ptr, 0.0};
151 }
else if constexpr (N == 2) {
152 __m256i mask = _mm256_set_epi64x(0, 0, -1, -1);
153 return _mm256_maskload_pd(ptr, mask);
154 }
else if constexpr (N == 3) {
155 __m256i mask = _mm256_set_epi64x(0, -1, -1, -1);
156 return _mm256_maskload_pd(ptr, mask);
158 return _mm256_loadu_pd(ptr);
162#if !MOCHI_ARCH_X64_AVX512
163#if MOCHI_COMPILER_MSVC
166 static constexpr __m256i kLoadMasks[] = {
167 { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0},
168 {-1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0},
169 {-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0},
170 {-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 0, 0, 0, 0, 0, 0},
171 {-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1}};
175 static constexpr __m256i kLoadMasks[] = {
176 {0LL, 0LL, 0LL, 0LL},
177 {-1LL, 0LL, 0LL, 0LL},
178 {-1LL, -1LL, 0LL, 0LL},
179 {-1LL, -1LL, -1LL, 0LL},
180 {-1LL, -1LL, -1LL, -1LL}};
186#if MOCHI_ARCH_X64_AVX512
187 return _mm256_maskz_loadu_pd(x64_simd::kLaneMasksS8[n], ptr);
189 return _mm256_maskload_pd(ptr, kLoadMasks[n]);
194 return _mm256_i32gather_pd(ptr, indices.raw,
sizeof(
double));
198 return _mm256_i64gather_pd(ptr, indices.raw,
sizeof(
double));
201 template <
int kTupleCount = kSize>
204 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
205 constexpr int kCount0 =
Clamp(kTupleCount * 3 - 0, 0, 4);
206 constexpr int kCount1 =
Clamp(kTupleCount * 3 - 4, 0, 4);
207 constexpr int kCount2 =
Clamp(kTupleCount * 3 - 8, 0, 4);
208 auto a = Simd::Load<kCount0>(ptr).raw;
209 auto b = Simd::Load<kCount1>(kCount1 == 0 ? ptr : ptr + 4).raw;
210 auto c = Simd::Load<kCount2>(kCount2 == 0 ? ptr : ptr + 8).raw;
212 auto d = _mm256_blend_pd(a, b, 0b0100);
213 d = _mm256_blend_pd(d, c, 0b0010);
214 auto e = _mm256_permute2f128_pd(d, d, 0x01);
215 out0.raw = _mm256_blend_pd(d, e, 0b1010);
217 d = _mm256_blend_pd(a, b, 0b1001);
218 d = _mm256_blend_pd(d, c, 0b0100);
219 out1.raw = _mm256_shuffle_pd(d, d, 0b0101);
221 d = _mm256_blend_pd(a, b, 0b0010);
222 d = _mm256_blend_pd(d, c, 0b1001);
223 e = _mm256_permute2f128_pd(d, d, 0x01);
224 out2.raw = _mm256_blend_pd(d, e, 0b0101);
229 static_assert(i >= 0 && i <= 3,
"Invalid component index");
230 auto zero = _mm256_setzero_si256();
231 auto v = _mm256_insert_epi64(zero, 0x3FF0000000000000LL, i);
232 return _mm256_castsi256_pd(v);
235 template <
int N = kSize>
237 static_assert(N >= 0 && N <=
kSize);
238 if constexpr (N == 0) {
239 }
else if constexpr (N <
kSize) {
241 memcpy(ptr, &v,
sizeof(
Scalar) * N);
243 _mm256_storeu_pd(ptr, v.raw);
249#if MOCHI_ARCH_X64_AVX512
250 _mm256_mask_storeu_pd(ptr, x64_simd::kLaneMasksS8[n], v.raw);
263 auto mask = _mm256_movemask_pd(condition.raw);
265 auto const* tableRow =
266 reinterpret_cast<__m128i const*
>(x64_simd::kStoreSelectedShuffleTableD4[mask]);
267 __m256i pattern = _mm256_cvtepu8_epi32(_mm_loadl_epi64(tableRow));
268 __m256i packed = _mm256_permutevar8x32_epi32(_mm256_castpd_si256(values.raw), pattern);
269 _mm256_storeu_pd(ptr, _mm256_castsi256_pd(packed));
270 return _mm_popcnt_u32(mask);
273 template <
int kTupleCount = kSize>
275 static_assert(kTupleCount >= 1 && kTupleCount <=
kSize,
"Invalid kTupleCount");
277 auto d = _mm256_shuffle_pd(b.raw, b.raw, 0b0101);
278 auto e = _mm256_blend_pd(a.raw, c.raw, 0b0101);
279 e = _mm256_permute2f128_pd(e, e, 0x01);
280 auto f = _mm256_blend_pd(a.raw, d, 0b1010);
281 auto g = _mm256_blend_pd(d, c.raw, 0b1010);
282 constexpr int kCount0 =
Clamp(kTupleCount * 3 - 0, 0, 4);
283 constexpr int kCount1 =
Clamp(kTupleCount * 3 - 4, 0, 4);
284 constexpr int kCount2 =
Clamp(kTupleCount * 3 - 8, 0, 4);
285 Simd::Store<kCount0>(ptr, _mm256_blend_pd(f, e, 0b1100));
286 if constexpr (kCount1 > 0) {
287 Simd::Store<kCount1>(ptr + 4, _mm256_blend_pd(g, f, 0b1100));
289 if constexpr (kCount2 > 0) {
290 Simd::Store<kCount2>(ptr + 8, _mm256_blend_pd(e, g, 0b1100));
295 return _mm256_blendv_pd(b.raw, a.raw, mask.raw);
299 template <
int x = 0,
int y = 1,
int z = 2,
int w = 3>
301 static_assert(x >= 0 && x < 4,
"Invalid index");
302 static_assert(y >= 0 && y < 4,
"Invalid index");
303 static_assert(z >= 0 && z < 4,
"Invalid index");
304 static_assert(w >= 0 && w < 4,
"Invalid index");
305 if constexpr (x == 0 && y == 1 && z == 2 && w == 3) {
308 return _mm256_permute4x64_pd(v.raw, x | (y << 2) | (z << 4) | (w << 6));
313 template <
int x = 0,
int y = 1,
int z = 2,
int w = 3>
315 static_assert(x >= 0 && x < 4,
"Invalid index");
316 static_assert(y >= 0 && y < 4,
"Invalid index");
317 static_assert(z >= 0 && z < 4,
"Invalid index");
318 static_assert(w >= 0 && w < 4,
"Invalid index");
323 return _mm256_sqrt_pd(v.raw);
327 return Simd{1.0} / v;
336 return _mm256_castsi256_pd(_mm256_set1_epi64x(0x8000000000000000LL));
340 return _mm256_andnot_pd(SignBitMask().
raw, v.raw);
344 return _mm256_min_pd(a.raw, b.raw);
348 return _mm256_max_pd(a.raw, b.raw);
352 return _mm256_floor_pd(a.raw);
356 return _mm256_round_pd(v.raw, _MM_FROUND_TO_NEAREST_INT);
359#if MOCHI_ARCH_X64_SVML
361 return _mm256_cos_pd(a.raw);
365 return _mm256_sin_pd(a.raw);
369 return _mm256_tan_pd(a.raw);
373 return _mm256_acos_pd(a.raw);
377 return _mm256_asin_pd(a.raw);
381 return _mm256_atan_pd(a.raw);
385 return _mm256_exp_pd(a.raw);
389 return _mm256_log_pd(a.raw);
393 return _mm256_tanh_pd(a.raw);
398#if MOCHI_ARCH_X64_FMA
399 return {_mm256_fmadd_pd(a.raw, b.raw, c.raw)};
406#if MOCHI_ARCH_X64_FMA
407 return _mm256_fmsub_pd(a.raw, b.raw, c.raw);
414#if MOCHI_ARCH_X64_FMA
415 return _mm256_fnmadd_pd(a.raw, b.raw, c.raw);
422#if MOCHI_ARCH_X64_FMA
423 return _mm256_fnmsub_pd(a.raw, b.raw, c.raw);
431 static_assert(N >= 2 && N <= 4,
"Unsupported N");
433 if constexpr (N == 2) {
435 }
else if constexpr (N == 3) {
438 return HalfT::Get<0>(HalfT::Min(HalfT::Min(lo, HalfT::Broadcast<1>(lo)), hi));
446 static_assert(N >= 2 && N <= 4,
"Unsupported N");
448 if constexpr (N == 2) {
450 }
else if constexpr (N == 3) {
453 return HalfT::Get<0>(HalfT::Max(HalfT::Max(lo, HalfT::Broadcast<1>(lo)), hi));
461 static_assert(N >= 2 && N <= 4,
"Unsupported N");
462 if constexpr (N == 2) {
464 }
else if constexpr (N == 3) {
467 return HalfT::Get<0>(tmp) +
Get<1>(a);
468 }
else if constexpr (N == 4) {
474 return HalfT::Get<0>(tmp) + HalfT::Get<1>(tmp);
480 static_assert(N >= 2 && N <= 4,
"Unsupported N");
483 if constexpr (N == 2) {
484 return buf[0] * buf[1];
485 }
else if constexpr (N == 3) {
486 return buf[0] * buf[1] * buf[2];
487 }
else if constexpr (N == 4) {
488 return buf[0] * buf[1] * buf[2] * buf[3];
494 static_assert(N >= 2 && N <= 4,
"Unsupported N");
499 return _mm256_cmp_pd(this->
raw, rhs.raw, _CMP_LT_OQ);
503 return _mm256_cmp_pd(this->
raw, rhs.raw, _CMP_GT_OQ);
507 return _mm256_cmp_pd(this->
raw, rhs.raw, _CMP_LE_OQ);
511 return _mm256_cmp_pd(this->
raw, rhs.raw, _CMP_GE_OQ);
515 return _mm256_cmp_pd(a.raw, b.raw, _CMP_EQ_OQ);
519 return _mm256_cmp_pd(a.raw, b.raw, _CMP_NEQ_UQ);
523 return _mm256_setzero_pd();
527 auto mask = GetMSBitMask(
Equal(*
this, rhs));
528 return mask == 0xFFFFFFFF;
532 auto mask = GetMSBitMask(
NotEqual(*
this, rhs));
537 auto ones = _mm256_castsi256_pd(_mm256_set1_epi64x(-1));
538 return _mm256_xor_pd(
raw, ones);
542 return _mm256_xor_pd(
raw, SignBitMask().
raw);
546 return _mm256_add_pd(
raw, rhs.raw);
550 return _mm256_sub_pd(
raw, rhs.raw);
554 return _mm256_mul_pd(
raw, rhs.raw);
558 return _mm256_div_pd(
raw, rhs.raw);
562 return _mm256_and_pd(
raw, rhs.raw);
566 return _mm256_or_pd(
raw, rhs.raw);
570 return _mm256_xor_pd(
raw, rhs.raw);
576 return _mm256_movemask_epi8(_mm256_castpd_si256(a.raw));
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)
Simd< T, 2 > Shuffle(Simd< T, 2 > a)
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)
Simd< T, N > Blend(Simd< T, N > a, Simd< T, N > b)
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)