SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_double_4_inl.h
Go to the documentation of this file.
1/*
2 * Copyright (c) Meta Platforms, Inc. and affiliates.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 * http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17#pragma once
18
19#include "x64_simd_inl.h" // for IntelliSense
20
21#if MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX2
22
23namespace superdex {
24
25/***********************************************************************************************
26 Simd<double, 4>
27*/
28template <>
29class Simd<double, 4> {
30 public:
31 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(double, 4, __m256d);
32 Simd(double a, double b, double c = 0.0, double d = 0.0)
33 : raw(_mm256_set_pd(d, c, b, a)) {} // AVX
34 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
35 Simd(U a) : raw(_mm256_set1_pd(a)) {} // AVX
36
37 // Joint two Vec2d into a single Vec4d
38 Simd(Simd<double, 2> const& low, Simd<double, 2> const& high)
39 : raw(_mm256_set_m128d(high.raw, low.raw)) {} // AVX
40
41 template <int i>
42 [[nodiscard]] static MOCHI_FORCE_INLINE double Get(Simd v) {
43 static_assert(i >= 0 && i < 4, "Index out of range");
44 return v[i];
45 }
46
47 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
48 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
49#if MOCHI_COMPILER_MSVC
50 return raw.m256d_f64[i];
51#else
52 return raw[i];
53#endif
54 }
55
56 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, int i, Scalar value) {
57 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
58#if MOCHI_COMPILER_MSVC
59 auto result = v;
60 result.raw.m256d_f64[i] = value;
61 return result;
62#else
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])); // AVX
66#endif
67 }
68
69 template <int i>
70 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, Scalar value) {
71 static_assert(i >= 0 && i < kSize, "Index out of range");
72 return Set(v, i, value);
73 }
74
75 template <int iHalf>
76 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<double, 2> GetHalf(Simd a) {
77 static_assert(iHalf == 0 || iHalf == 1);
78 return _mm256_extractf128_pd(a.raw, iHalf); // AVX
79 }
80
81 [[nodiscard]] static MOCHI_FORCE_INLINE Simd AsPoint(Simd a) {
82 // Replace the 3rd component with an integer that has the same bits as 1.0.
83 auto araw = _mm256_castpd_si256(a.raw); // AVX
84 auto v = _mm256_insert_epi64(araw, 0x3FF0000000000000LL, 3); // AVX
85 return _mm256_castsi256_pd(v); // AVX
86 }
87
88 [[nodiscard]] static MOCHI_FORCE_INLINE Simd AsDirection(Simd a) {
89 // Replace the 3rd component with an integer that has the same bits as 0.0.
90 auto araw = _mm256_castpd_si256(a.raw); // AVX
91 auto v = _mm256_insert_epi64(araw, 0, 3); // AVX
92 return _mm256_castsi256_pd(v); // AVX
93 }
94
95 template <int x, int y, int z, int w>
96 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Blend(Simd a, Simd b) {
97 static_assert(
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) {
101 return a;
102 } else if constexpr (x == 1 && y == 1 && z == 1 && w == 1) {
103 return b;
104 } else {
105 return _mm256_blend_pd(a.raw, b.raw, x | (y << 1) | (z << 2) | (w << 3)); // SSE4.1
106 }
107 }
108
109 template <int N>
110 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
111 static_assert(N >= 1 && N <= kSize, "Unsupported N");
112 int mask = GetMSBitMask(v); // One bit for each byte in the vector
113 if constexpr (N == kSize) {
114 return mask == 0xFFFFFFFF;
115 } else {
116 int constexpr kNumBits = N * sizeof(Scalar);
117 auto constexpr kMustBeSet = (1UL << kNumBits) - 1;
118 return (mask & kMustBeSet) == kMustBeSet;
119 }
120 }
121
122 template <int N>
123 [[nodiscard]] static MOCHI_FORCE_INLINE bool AnyTrue(Simd v) {
124 static_assert(N >= 1 && N <= kSize, "Unsupported N");
125 int mask = GetMSBitMask(v); // One bit for each byte in the vector
126 if constexpr (N == kSize) {
127 return mask != 0;
128 } else {
129 int constexpr kNumBits = N * sizeof(Scalar);
130 auto constexpr kMayBeSet = (1UL << kNumBits) - 1;
131 return (mask & kMayBeSet) != 0;
132 }
133 }
134
135 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
136 return _mm256_broadcast_sd(p); // AVX
137 }
138
139 template <int i>
140 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
141 return Shuffle<i, i, i, i>(v);
142 }
143
144 template <int N = kSize>
145 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
146 static_assert(N >= 0 && N <= 4);
147 if constexpr (N == 0) {
148 return Simd::Zero();
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); // AVX
153 return _mm256_maskload_pd(ptr, mask); // AVX
154 } else if constexpr (N == 3) {
155 __m256i mask = _mm256_set_epi64x(0, -1, -1, -1); // AVX
156 return _mm256_maskload_pd(ptr, mask); // AVX
157 } else {
158 return _mm256_loadu_pd(ptr); // AVX
159 }
160 }
161
162#if !MOCHI_ARCH_X64_AVX512
163#if MOCHI_COMPILER_MSVC
164 // MSVC defines __m128i as a union. Byte arrays are required to initialize it this way.
165 // clang-format off
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}};
172 // clang-format on
173#else
174 // GCC and Clang define __m256i as 'long long' with special attributes.
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}};
181#endif
182#endif
183
184 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
185 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
186#if MOCHI_ARCH_X64_AVX512
187 return _mm256_maskz_loadu_pd(x64_simd::kLaneMasksS8[n], ptr); // AVX512VL
188#else
189 return _mm256_maskload_pd(ptr, kLoadMasks[n]); // AVX
190#endif
191 }
192
193 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int, 4> const& indices) {
194 return _mm256_i32gather_pd(ptr, indices.raw, sizeof(double)); // AVX2
195 }
196
197 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int64_t, 4> const& indices) {
198 return _mm256_i64gather_pd(ptr, indices.raw, sizeof(double)); // AVX2
199 }
200
201 template <int kTupleCount = kSize>
202 MOCHI_FORCE_INLINE static void
203 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
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; // [0,1,2,3]
209 auto b = Simd::Load<kCount1>(kCount1 == 0 ? ptr : ptr + 4).raw; // [4,5,6,7]
210 auto c = Simd::Load<kCount2>(kCount2 == 0 ? ptr : ptr + 8).raw; // [8,9,10,11]
211
212 auto d = _mm256_blend_pd(a, b, 0b0100); // [0,_,6,3]
213 d = _mm256_blend_pd(d, c, 0b0010); // [0,9,6,3]
214 auto e = _mm256_permute2f128_pd(d, d, 0x01); // [6,3,0,9]
215 out0.raw = _mm256_blend_pd(d, e, 0b1010); // [0,3,6,9]
216
217 d = _mm256_blend_pd(a, b, 0b1001); // [4,1,_,7]
218 d = _mm256_blend_pd(d, c, 0b0100); // [4,1,10,7]
219 out1.raw = _mm256_shuffle_pd(d, d, 0b0101); // [1,4,7,10]
220
221 d = _mm256_blend_pd(a, b, 0b0010); // [_,5,2,_]
222 d = _mm256_blend_pd(d, c, 0b1001); // [8,5,2,11]
223 e = _mm256_permute2f128_pd(d, d, 0x01); // [2,11,8,5]
224 out2.raw = _mm256_blend_pd(d, e, 0b0101); // [2,5,8,11]
225 }
226
227 template <int i>
228 [[nodiscard]] static MOCHI_FORCE_INLINE Simd SetBasisVector() {
229 static_assert(i >= 0 && i <= 3, "Invalid component index");
230 auto zero = _mm256_setzero_si256(); // AVX
231 auto v = _mm256_insert_epi64(zero, 0x3FF0000000000000LL, i); // AVX
232 return _mm256_castsi256_pd(v); // AVX
233 }
234
235 template <int N = kSize>
236 static MOCHI_FORCE_INLINE void Store([[maybe_unused]] Scalar* ptr, [[maybe_unused]] Simd v) {
237 static_assert(N >= 0 && N <= kSize);
238 if constexpr (N == 0) {
239 } else if constexpr (N < kSize) {
240 // About 3X faster than a masked store on older AMD CPUs. About the same on others.
241 memcpy(ptr, &v, sizeof(Scalar) * N);
242 } else {
243 _mm256_storeu_pd(ptr, v.raw); // AVX
244 }
245 }
246
247 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
248 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
249#if MOCHI_ARCH_X64_AVX512
250 _mm256_mask_storeu_pd(ptr, x64_simd::kLaneMasksS8[n], v.raw); // AVX512VL
251#else
252 switch (n) { // clang-format off
253 case 1: Store<1>(ptr, v); break;
254 case 2: Store<2>(ptr, v); break;
255 case 3: Store<3>(ptr, v); break;
256 case 4: Store<4>(ptr, v); break;
257 MOCHI_UNLIKELY default: break;
258 } // clang-format on
259#endif
260 }
261
262 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
263 auto mask = _mm256_movemask_pd(condition.raw);
264 // Load 8 bytes from the table, then zero-exend to get the shuffle pattern.
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);
271 }
272
273 template <int kTupleCount = kSize>
274 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
275 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
276 // a = [0,3,6,9], b = [1,4,7,10], c = [2,5,8,11]
277 auto d = _mm256_shuffle_pd(b.raw, b.raw, 0b0101); // [4,1,10,7]
278 auto e = _mm256_blend_pd(a.raw, c.raw, 0b0101); // [2,3,8,9]
279 e = _mm256_permute2f128_pd(e, e, 0x01); // [8,9,2,3]
280 auto f = _mm256_blend_pd(a.raw, d, 0b1010); // [0,1,6,7]
281 auto g = _mm256_blend_pd(d, c.raw, 0b1010); // [4,5,10,11]
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)); // [0,1,2,3]
286 if constexpr (kCount1 > 0) {
287 Simd::Store<kCount1>(ptr + 4, _mm256_blend_pd(g, f, 0b1100)); // [4,5,6,7]
288 }
289 if constexpr (kCount2 > 0) {
290 Simd::Store<kCount2>(ptr + 8, _mm256_blend_pd(e, g, 0b1100)); // [8,9,10,11]
291 }
292 }
293
294 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
295 return _mm256_blendv_pd(b.raw, a.raw, mask.raw); // AVX
296 }
297
298 // return Simd{v[x], v[y], v[z], v[w]}
299 template <int x = 0, int y = 1, int z = 2, int w = 3>
300 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Shuffle(Simd v) {
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) {
306 return v;
307 } else {
308 return _mm256_permute4x64_pd(v.raw, x | (y << 2) | (z << 4) | (w << 6)); // AVX2
309 }
310 }
311
312 // return Simd{a[x], a[y], b[z], b[w]}
313 template <int x = 0, int y = 1, int z = 2, int w = 3>
314 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Shuffle(Simd a, Simd b) {
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");
319 return Simd{Get<x>(a), Get<y>(a), Get<z>(b), Get<w>(b)};
320 }
321
322 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sqrt(Simd v) {
323 return _mm256_sqrt_pd(v.raw); // AVX
324 }
325
326 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpApprox(Simd v) {
327 return Simd{1.0} / v;
328 }
329
330 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpSqrtApprox(Simd v) {
331 return Simd{1.0} / Sqrt(v);
332 }
333
334 // Broadcast the value -0.0. Use this in bitwise operations to affect just the sign bit.
335 [[nodiscard]] static MOCHI_FORCE_INLINE Simd SignBitMask() {
336 return _mm256_castsi256_pd(_mm256_set1_epi64x(0x8000000000000000LL)); // AVX, AVX
337 }
338
339 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Abs(Simd v) {
340 return _mm256_andnot_pd(SignBitMask().raw, v.raw); // AVX
341 }
342
343 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
344 return _mm256_min_pd(a.raw, b.raw); // AVX
345 }
346
347 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
348 return _mm256_max_pd(a.raw, b.raw); // AVX
349 }
350
351 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Floor(Simd a) {
352 return _mm256_floor_pd(a.raw); // AVX
353 }
354
355 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FastRound(Simd v) {
356 return _mm256_round_pd(v.raw, _MM_FROUND_TO_NEAREST_INT); // AVX
357 }
358
359#if MOCHI_ARCH_X64_SVML
360 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Cos(Simd a) {
361 return _mm256_cos_pd(a.raw); // AVX
362 }
363
364 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sin(Simd a) {
365 return _mm256_sin_pd(a.raw); // AVX
366 }
367
368 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tan(Simd a) {
369 return _mm256_tan_pd(a.raw); // AVX
370 }
371
372 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ACos(Simd a) {
373 return _mm256_acos_pd(a.raw); // AVX
374 }
375
376 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ASin(Simd a) {
377 return _mm256_asin_pd(a.raw); // AVX
378 }
379
380 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ATan(Simd a) {
381 return _mm256_atan_pd(a.raw); // AVX
382 }
383
384 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Exp(Simd a) {
385 return _mm256_exp_pd(a.raw); // AVX
386 }
387
388 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Ln(Simd a) {
389 return _mm256_log_pd(a.raw); // AVX
390 }
391
392 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tanh(Simd a) {
393 return _mm256_tanh_pd(a.raw); // AVX
394 }
395#endif // MOCHI_ARCH_X64_SVML
396
397 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulAdd(Simd a, Simd b, Simd c) {
398#if MOCHI_ARCH_X64_FMA
399 return {_mm256_fmadd_pd(a.raw, b.raw, c.raw)}; // FMA
400#else
401 return (a * b) + c;
402#endif
403 }
404
405 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulSub(Simd a, Simd b, Simd c) {
406#if MOCHI_ARCH_X64_FMA
407 return _mm256_fmsub_pd(a.raw, b.raw, c.raw); // FMA
408#else
409 return (a * b) - c;
410#endif
411 }
412
413 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulAdd(Simd a, Simd b, Simd c) {
414#if MOCHI_ARCH_X64_FMA
415 return _mm256_fnmadd_pd(a.raw, b.raw, c.raw); // FMA
416#else
417 return -(a * b) + c;
418#endif
419 }
420
421 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulSub(Simd a, Simd b, Simd c) {
422#if MOCHI_ARCH_X64_FMA
423 return _mm256_fnmsub_pd(a.raw, b.raw, c.raw); // FMA
424#else
425 return -(a * b) - c;
426#endif
427 }
428
429 template <int N = 4>
430 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
431 static_assert(N >= 2 && N <= 4, "Unsupported N");
432 using HalfT = Simd<Scalar, 2>;
433 if constexpr (N == 2) {
434 return Get<0>(Min(a, Broadcast<1>(a)));
435 } else if constexpr (N == 3) {
436 auto lo = GetHalf<0>(a);
437 auto hi = GetHalf<1>(a);
438 return HalfT::Get<0>(HalfT::Min(HalfT::Min(lo, HalfT::Broadcast<1>(lo)), hi));
439 } else {
440 return HalfT::HMin(HalfT::Min(GetHalf<0>(a), GetHalf<1>(a)));
441 }
442 }
443
444 template <int N = 4>
445 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
446 static_assert(N >= 2 && N <= 4, "Unsupported N");
447 using HalfT = Simd<Scalar, 2>;
448 if constexpr (N == 2) {
449 return Get<0>(Max(a, Broadcast<1>(a)));
450 } else if constexpr (N == 3) {
451 auto lo = GetHalf<0>(a);
452 auto hi = GetHalf<1>(a);
453 return HalfT::Get<0>(HalfT::Max(HalfT::Max(lo, HalfT::Broadcast<1>(lo)), hi));
454 } else {
455 return HalfT::HMax(HalfT::Max(GetHalf<0>(a), GetHalf<1>(a)));
456 }
457 }
458
459 template <int N>
460 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
461 static_assert(N >= 2 && N <= 4, "Unsupported N");
462 if constexpr (N == 2) {
463 return Get<0>(a) + Get<1>(a);
464 } else if constexpr (N == 3) {
465 using HalfT = Simd<double, 2>;
466 HalfT tmp = GetHalf<0>(a) + GetHalf<1>(a);
467 return HalfT::Get<0>(tmp) + Get<1>(a); // (a[0] + a[2]) + a[1]
468 } else if constexpr (N == 4) {
469 // PERF NOTE: An alternate implementation could be: Get<0>(a) + Get<1>(a) + Get<2>(a) +
470 // Get<3>(a). In comparison, this implementation saves two instructions, while maintaining
471 // identical throughput and latency, when benchmarked on an AMD CPU.
472 using HalfT = Simd<double, 2>;
473 HalfT tmp = GetHalf<0>(a) + GetHalf<1>(a);
474 return HalfT::Get<0>(tmp) + HalfT::Get<1>(tmp); // (a[0] + a[2]) + (a[1] + a[3]);
475 }
476 }
477
478 template <int N>
479 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HProd(Simd a) {
480 static_assert(N >= 2 && N <= 4, "Unsupported N");
481 alignas(alignof(Simd)) Scalar buf[4];
482 Store(buf, a);
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];
489 }
490 }
491
492 template <int N>
493 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Dot(Simd a, Simd b) {
494 static_assert(N >= 2 && N <= 4, "Unsupported N");
495 return Simd{HSum<N>(a * b)};
496 }
497
498 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
499 return _mm256_cmp_pd(this->raw, rhs.raw, _CMP_LT_OQ); // AVX
500 }
501
502 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
503 return _mm256_cmp_pd(this->raw, rhs.raw, _CMP_GT_OQ); // AVX
504 }
505
506 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
507 return _mm256_cmp_pd(this->raw, rhs.raw, _CMP_LE_OQ); // AVX
508 }
509
510 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
511 return _mm256_cmp_pd(this->raw, rhs.raw, _CMP_GE_OQ); // AVX
512 }
513
514 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
515 return _mm256_cmp_pd(a.raw, b.raw, _CMP_EQ_OQ); // AVX
516 }
517
518 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
519 return _mm256_cmp_pd(a.raw, b.raw, _CMP_NEQ_UQ); // AVX
520 }
521
522 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
523 return _mm256_setzero_pd(); // AVX
524 }
525
526 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
527 auto mask = GetMSBitMask(Equal(*this, rhs));
528 return mask == 0xFFFFFFFF; // All values equal
529 }
530
531 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
532 auto mask = GetMSBitMask(NotEqual(*this, rhs));
533 return mask != 0; // Any values not equal
534 }
535
536 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
537 auto ones = _mm256_castsi256_pd(_mm256_set1_epi64x(-1)); // AVX, AVX
538 return _mm256_xor_pd(raw, ones); // AVX
539 }
540
541 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
542 return _mm256_xor_pd(raw, SignBitMask().raw); // AVX, AVX
543 }
544
545 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
546 return _mm256_add_pd(raw, rhs.raw); // AVX
547 }
548
549 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
550 return _mm256_sub_pd(raw, rhs.raw); // AVX
551 }
552
553 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
554 return _mm256_mul_pd(raw, rhs.raw); // AVX
555 }
556
557 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
558 return _mm256_div_pd(raw, rhs.raw); // AVX
559 }
560
561 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
562 return _mm256_and_pd(raw, rhs.raw); // AVX
563 }
564
565 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
566 return _mm256_or_pd(raw, rhs.raw); // AVX
567 }
568
569 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
570 return _mm256_xor_pd(raw, rhs.raw); // AVX
571 }
572
573 private:
574 // Integer mask with the most significant bit of each byte in the vector
575 [[nodiscard]] static MOCHI_FORCE_INLINE int GetMSBitMask(Simd a) {
576 return _mm256_movemask_epi8(_mm256_castpd_si256(a.raw)); // AVX2, AVX
577 }
578};
579
580} // namespace superdex
581
582#endif // MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX2
Simd operator&(Simd rhs) const
bool operator==(Simd rhs) const
NativeType raw
Definition simd.h:174
Simd operator>(Simd rhs) const
Simd operator*(Simd rhs) const
Simd operator^(Simd rhs) const
Simd operator-() 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
Definition simd.h:96
Simd operator+(Simd rhs) const
Simd operator~() const
Simd operator/(Simd rhs) const
Simd operator<=(Simd rhs) const
Scalar operator[](int i) const
#define MOCHI_ASSERT_VERBOSE(condition_without_side_effects,...)
Definition debug.h:102
#define MOCHI_UNLIKELY
#define MOCHI_FORCE_INLINE
T Dot(Simd< T, N > a, Simd< T, N > b)
Definition simd.h:673
constexpr T ACos(T a)
Simd< T, 2 > Shuffle(Simd< T, 2 > a)
Definition simd_inl.h:273
V LoadIndexed(typename V::Scalar const *ptr, Simd< I, V::kSize > indices)
Definition simd_inl.h:200
constexpr T const & Min(T const &a, T const &b)
constexpr auto Equal(T const &a, T const &b)
constexpr T Sin(T a)
T HSum(Simd< T, N > a)
Definition simd_inl.h:377
T HMin(Simd< T, N > a)
Definition simd_inl.h:389
bool AllTrue(T const &a)
Definition basic_utils.h:60
constexpr auto MulAdd(A a, B b, C c)
Simd< T, N > Tanh(Simd< T, N > a)
Definition simd_inl.h:709
Simd< T, N > Set(Simd< T, N > a, T value)
Definition simd_inl.h:313
constexpr auto NotEqual(T const &a, T const &b)
constexpr T Exp(T a)
constexpr T Cos(T a)
T HMax(Simd< T, N > a)
Definition simd_inl.h:395
V Broadcast(typename V::Scalar a)
Definition simd_inl.h:115
constexpr T Abs(T a)
Definition basic_utils.h:50
constexpr auto MulSub(A a, B b, C c)
Simd< T, N > Blend(Simd< T, N > a, Simd< T, N > b)
Definition simd_inl.h:288
constexpr T Tan(T a)
bool AnyTrue(T const &a)
Definition basic_utils.h:66
constexpr T Select(bool condition, T a, T b)
constexpr T Sqrt(T a)
constexpr T ATan(T a)
Simd< T, N > Ln(Simd< T, N > a)
Definition simd_inl.h:699
constexpr auto NegMulAdd(A a, B b, C c)
T Get(Simd< T, N > v)
Definition simd_inl.h:303
constexpr T Floor(T a)
constexpr T ASin(T a)
constexpr ValT Clamp(ValT value, MinT min, MaxT max)
Simd< T, N/2 > GetHalf(Simd< T, N > a)
Definition simd_inl.h:308
T HProd(Simd< T, N > a)
Definition simd_inl.h:383
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)
Definition simd_inl.h:207
Simd< T, N > FastRound(Simd< T, N > a)
Definition simd_inl.h:416
void StoreTransposed(T *ptr, Simd< T, N > a, Simd< T, N > b, Simd< T, N > c)
Definition simd_inl.h:245
void Store(T *ptr, Simd< T, N > a)
Definition simd_inl.h:213
Simd< T, N > RcpSqrtApprox(Simd< T, N > a)
Definition simd_inl.h:357
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)
Definition simd_inl.h:225
V Load(typename V::Scalar const *ptr)
Definition simd_inl.h:184
#define MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(T, N, NativeT)