SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_double_8_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_AVX512
22
23namespace superdex {
24
25/***********************************************************************************************
26 Simd<double, 8>
27*/
28template <>
29class Simd<double, 8> {
30 public:
31 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(double, 8, __m512d);
32 Simd(
33 double a,
34 double b,
35 double c = 0.0,
36 double d = 0.0,
37 double e = 0.0,
38 double f = 0.0,
39 double g = 0.0,
40 double h = 0.0)
41 : raw(_mm512_set_pd(h, g, f, e, d, c, b, a)) {} // AVX512F
42
43 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
44 Simd(U a) : raw(_mm512_set1_pd(a)) {} // AVX512F
45
46 Simd(Simd<double, 4> const& low, Simd<double, 4> const& high)
47 : raw(_mm512_insertf64x4(_mm512_castpd256_pd512(low.raw), high.raw, 1)) {} // AVX512F
48
49 template <int i>
50 [[nodiscard]] static MOCHI_FORCE_INLINE double Get(Simd v) {
51 static_assert(i >= 0 && i < kSize, "Index out of range");
52 return v[i];
53 }
54
55 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
56 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
57#if MOCHI_COMPILER_MSVC
58 return raw.m512d_f64[i];
59#else
60 return raw[i];
61#endif
62 }
63
64 template <int iHalf>
65 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<double, 4> GetHalf(Simd a) {
66 static_assert(iHalf == 0 || iHalf == 1);
67 if constexpr (iHalf == 0) {
68 return _mm512_castpd512_pd256(a.raw); // AVX512F
69 } else {
70 return _mm512_extractf64x4_pd(a.raw, 1); // AVX512F
71 }
72 }
73
74 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, int i, Scalar value) {
75 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
76 auto const mask = static_cast<__mmask8>(1u << i);
77 return _mm512_mask_broadcastsd_pd(v.raw, mask, _mm_set_sd(value));
78 }
79
80 template <int i>
81 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, Scalar 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));
85 }
86
87 template <int N>
88 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
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;
93 } else {
94 constexpr auto kLanes = LaneMask<N>();
95 return (mask & kLanes) == kLanes;
96 }
97 }
98
99 template <int N>
100 [[nodiscard]] static MOCHI_FORCE_INLINE bool AnyTrue(Simd v) {
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;
105 } else {
106 return (mask & LaneMask<N>()) != 0;
107 }
108 }
109
110 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
111 return _mm512_set1_pd(*p); // AVX512F
112 }
113
114 template <int i>
115 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
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));
119 } else {
120 constexpr int kLane = i % 2;
121 constexpr int kGroup = i / 2;
122 auto const group =
123 _mm512_shuffle_f64x2(v.raw, v.raw, _MM_SHUFFLE(kGroup, kGroup, kGroup, kGroup));
124 return _mm512_permute_pd(group, kLane == 0 ? 0x00 : 0xFF);
125 }
126 }
127
128 template <int N = kSize>
129 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
130 static_assert(N >= 0 && N <= kSize);
131 if constexpr (N == 0) {
132 return Zero();
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); // AVX512F
143 } else {
144 return _mm512_loadu_pd(ptr); // AVX512F
145 }
146 }
147
148 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
149 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
150 return _mm512_maskz_loadu_pd(LaneMask(n), ptr); // AVX512F
151 }
152
153 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int, 8> const& indices) {
154 return _mm512_i32gather_pd(indices.raw, ptr, sizeof(double)); // AVX512F
155 }
156
157 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int64_t, 8> const& indices) {
158 return _mm512_i64gather_pd(indices.raw, ptr, sizeof(double)); // AVX512F
159 }
160
161 template <int kTupleCount = kSize>
162 MOCHI_FORCE_INLINE static void
163 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
164 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
165 if constexpr (kTupleCount == 1) {
166 out0 = Load<1>(ptr);
167 out1 = Load<1>(ptr + 1);
168 out2 = Load<1>(ptr + 2);
169 return;
170 }
171 constexpr int kTotalCount = kTupleCount * 3;
172 auto const x0 = Load<Clamp(kTotalCount, 0, kSize)>(ptr).raw;
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);
181 } else {
182 auto const x1 = Load<kTotalCount - kSize>(ptr + kSize).raw;
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);
186 }
187 } else {
188 auto const x1 = Load<kSize>(ptr + kSize).raw;
189 constexpr int kCount2 = kTotalCount - 2 * kSize;
190 auto const x2 = Load<kCount2>(ptr + 2 * kSize).raw;
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);
203 } else {
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);
217 }
218 }
219 }
220
221 template <int N = kSize>
222 static MOCHI_FORCE_INLINE void Store([[maybe_unused]] Scalar* ptr, [[maybe_unused]] Simd v) {
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); // AVX512F
236 } else {
237 _mm512_storeu_pd(ptr, v.raw); // AVX512F
238 }
239 }
240
241 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
242 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
243 _mm512_mask_storeu_pd(ptr, LaneMask(n), v.raw); // AVX512F
244 }
245
246 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
247 __mmask8 const mask = ToMask(condition);
248 _mm512_mask_compressstoreu_pd(ptr, mask, values.raw); // AVX512F
249 return _mm_popcnt_u32(mask);
250 }
251
252 template <int kTupleCount = kSize>
253 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
254 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
255 if constexpr (kTupleCount == 1) {
256 Store<1>(ptr, a);
257 Store<1>(ptr + 1, b);
258 Store<1>(ptr + 2, c);
259 return;
260 }
261 constexpr int kTotalCount = kTupleCount * 3;
262 auto const ab0 =
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) {
267 auto const ab1 =
268 _mm512_permutex2var_pd(a.raw, _mm512_setr_epi64(0, 3, 11, 0, 4, 12, 0, 5), b.raw);
269 auto const x1 =
270 _mm512_permutex2var_pd(ab1, _mm512_setr_epi64(10, 1, 2, 11, 4, 5, 12, 7), c.raw);
271 Store<Clamp(kTotalCount - kSize, 0, kSize)>(ptr + kSize, x1);
272 }
273 if constexpr (kTotalCount > 2 * kSize) {
274 auto const ab2 =
275 _mm512_permutex2var_pd(a.raw, _mm512_setr_epi64(13, 0, 6, 14, 0, 7, 15, 0), b.raw);
276 auto const x2 =
277 _mm512_permutex2var_pd(ab2, _mm512_setr_epi64(0, 13, 2, 3, 14, 5, 6, 15), c.raw);
278 Store<kTotalCount - 2 * kSize>(ptr + 2 * kSize, x2);
279 }
280 }
281
282 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
283 return _mm512_mask_blend_pd(ToMask(mask), b.raw, a.raw); // AVX512F
284 }
285
286 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sqrt(Simd v) {
287 return _mm512_sqrt_pd(v.raw); // AVX512F
288 }
289
290 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpApprox(Simd v) {
291 return _mm512_rcp14_pd(v.raw); // AVX512F
292 }
293
294 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpSqrtApprox(Simd v) {
295 return _mm512_rsqrt14_pd(v.raw); // AVX512F
296 }
297
298 [[nodiscard]] static MOCHI_FORCE_INLINE Simd SignBitMask() {
299 return _mm512_castsi512_pd(
300 _mm512_set1_epi64(static_cast<long long>(0x8000000000000000ULL))); // AVX512F
301 }
302
303 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Abs(Simd v) {
304 return _mm512_abs_pd(v.raw); // AVX512F
305 }
306
307 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
308 return _mm512_min_pd(a.raw, b.raw); // AVX512F
309 }
310
311 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
312 return _mm512_max_pd(a.raw, b.raw); // AVX512F
313 }
314
315 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Floor(Simd a) {
316 return _mm512_floor_pd(a.raw); // AVX512F
317 }
318
319 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FastRound(Simd v) {
320 return _mm512_roundscale_pd(v.raw, _MM_FROUND_TO_NEAREST_INT); // AVX512F
321 }
322
323#if MOCHI_ARCH_X64_SVML
324 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Cos(Simd a) {
325 return _mm512_cos_pd(a.raw);
326 }
327
328 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sin(Simd a) {
329 return _mm512_sin_pd(a.raw);
330 }
331
332 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tan(Simd a) {
333 return _mm512_tan_pd(a.raw);
334 }
335
336 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ACos(Simd a) {
337 return _mm512_acos_pd(a.raw);
338 }
339
340 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ASin(Simd a) {
341 return _mm512_asin_pd(a.raw);
342 }
343
344 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ATan(Simd a) {
345 return _mm512_atan_pd(a.raw);
346 }
347
348 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Exp(Simd a) {
349 return _mm512_exp_pd(a.raw);
350 }
351
352 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Ln(Simd a) {
353 return _mm512_log_pd(a.raw);
354 }
355
356 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tanh(Simd a) {
357 return _mm512_tanh_pd(a.raw);
358 }
359#endif // MOCHI_ARCH_X64_SVML
360
361 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulAdd(Simd a, Simd b, Simd c) {
362 return _mm512_fmadd_pd(a.raw, b.raw, c.raw); // AVX512F
363 }
364
365 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulSub(Simd a, Simd b, Simd c) {
366 return _mm512_fmsub_pd(a.raw, b.raw, c.raw); // AVX512F
367 }
368
369 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulAdd(Simd a, Simd b, Simd c) {
370 return _mm512_fnmadd_pd(a.raw, b.raw, c.raw); // AVX512F
371 }
372
373 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulSub(Simd a, Simd b, Simd c) {
374 return _mm512_fnmsub_pd(a.raw, b.raw, c.raw); // AVX512F
375 }
376
377 template <int N = kSize>
378 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
379 static_assert(N >= 2 && N <= kSize, "Unsupported N");
380 using HalfT = Simd<Scalar, 4>;
381 auto const lo = GetHalf<0>(a);
382 if constexpr (N <= 4) {
383 return HalfT::template HMin<N>(lo);
384 } else {
385 auto const hi = GetHalf<1>(a);
386 if constexpr (N == 5) {
387 return superdex::Min(HalfT::template HMin<4>(lo), HalfT::template Get<0>(hi));
388 } else if constexpr (N == 8) {
389 return HalfT::template HMin<4>(HalfT::Min(lo, hi));
390 } else {
391 return _mm512_mask_reduce_min_pd(LaneMask<N>(), a.raw);
392 }
393 }
394 }
395
396 template <int N = kSize>
397 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
398 static_assert(N >= 2 && N <= kSize, "Unsupported N");
399 using HalfT = Simd<Scalar, 4>;
400 auto const lo = GetHalf<0>(a);
401 if constexpr (N <= 4) {
402 return HalfT::template HMax<N>(lo);
403 } else {
404 auto const hi = GetHalf<1>(a);
405 if constexpr (N == 5) {
406 return superdex::Max(HalfT::template HMax<4>(lo), HalfT::template Get<0>(hi));
407 } else if constexpr (N == 8) {
408 return HalfT::template HMax<4>(HalfT::Max(lo, hi));
409 } else {
410 return _mm512_mask_reduce_max_pd(LaneMask<N>(), a.raw);
411 }
412 }
413 }
414
415 template <int N>
416 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
417 static_assert(N >= 2 && N <= kSize, "Unsupported N");
418 using HalfT = Simd<Scalar, 4>;
419 auto const lo = GetHalf<0>(a);
420 if constexpr (N <= 4) {
421 return HalfT::template HSum<N>(lo);
422 } else {
423 auto const hi = GetHalf<1>(a);
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);
428 } else {
429 return _mm512_mask_reduce_add_pd(LaneMask<N>(), a.raw);
430 }
431 }
432 }
433
434 template <int N>
435 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HProd(Simd a) {
436 static_assert(N >= 2 && N <= kSize, "Unsupported N");
437 using HalfT = Simd<Scalar, 4>;
438 auto const lo = GetHalf<0>(a);
439 if constexpr (N <= 4) {
440 return HalfT::template HProd<N>(lo);
441 } else {
442 auto const hi = GetHalf<1>(a);
443 if constexpr (N == 5) {
444 return HalfT::template HProd<4>(lo) * HalfT::template Get<0>(hi);
445 } else if constexpr (N == kSize) {
446 return HalfT::template HProd<4>(lo * hi);
447 } else {
448 return _mm512_mask_reduce_mul_pd(LaneMask<N>(), a.raw);
449 }
450 }
451 }
452
453 template <int N>
454 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Dot(Simd a, Simd b) {
455 static_assert(N >= 2 && N <= kSize, "Unsupported N");
456 return Simd{HSum<N>(a * b)};
457 }
458
459 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
460 return FromMask(_mm512_cmp_pd_mask(raw, rhs.raw, _CMP_LT_OQ)); // AVX512F
461 }
462
463 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
464 return FromMask(_mm512_cmp_pd_mask(raw, rhs.raw, _CMP_GT_OQ)); // AVX512F
465 }
466
467 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
468 return FromMask(_mm512_cmp_pd_mask(raw, rhs.raw, _CMP_LE_OQ)); // AVX512F
469 }
470
471 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
472 return FromMask(_mm512_cmp_pd_mask(raw, rhs.raw, _CMP_GE_OQ)); // AVX512F
473 }
474
475 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
476 return FromMask(_mm512_cmp_pd_mask(a.raw, b.raw, _CMP_EQ_OQ)); // AVX512F
477 }
478
479 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
480 return FromMask(_mm512_cmp_pd_mask(a.raw, b.raw, _CMP_NEQ_UQ)); // AVX512F
481 }
482
483 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
484 return _mm512_setzero_pd(); // AVX512F
485 }
486
487 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
488 return _mm512_cmp_pd_mask(raw, rhs.raw, _CMP_EQ_OQ) == static_cast<__mmask8>(0xFFu);
489 }
490
491 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
492 return _mm512_cmp_pd_mask(raw, rhs.raw, _CMP_NEQ_UQ) != 0;
493 }
494
495 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
496 return _mm512_castsi512_pd(
497 _mm512_xor_si512(_mm512_castpd_si512(raw), _mm512_set1_epi64(-1))); // AVX512F
498 }
499
500 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
501 return _mm512_xor_pd(raw, SignBitMask().raw); // AVX512DQ
502 }
503
504 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
505 return _mm512_add_pd(raw, rhs.raw); // AVX512F
506 }
507
508 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
509 return _mm512_sub_pd(raw, rhs.raw); // AVX512F
510 }
511
512 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
513 return _mm512_mul_pd(raw, rhs.raw); // AVX512F
514 }
515
516 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
517 return _mm512_div_pd(raw, rhs.raw); // AVX512F
518 }
519
520 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
521 return _mm512_and_pd(raw, rhs.raw); // AVX512DQ
522 }
523
524 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
525 return _mm512_or_pd(raw, rhs.raw); // AVX512DQ
526 }
527
528 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
529 return _mm512_xor_pd(raw, rhs.raw); // AVX512DQ
530 }
531
532 private:
533 template <int kTupleCount, int kComponent>
534 [[nodiscard]] static MOCHI_FORCE_INLINE __m512i LoadTransposeIndices() {
535 constexpr int kZeroIndex = kTupleCount * 3;
536 return _mm512_setr_epi64(
537 kComponent,
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);
545 }
546
547 // Returns a mask selecting the lowest N lanes.
548 template <int N>
549 [[nodiscard]] static constexpr __mmask8 LaneMask() {
550 static_assert(N >= 0 && N <= kSize);
551 return LaneMask(N);
552 }
553
554 // Returns a mask selecting the lowest n lanes.
555 [[nodiscard]] static constexpr __mmask8 LaneMask(int n) {
556 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid lane count");
557 return static_cast<__mmask8>((uint32_t{1} << n) - 1);
558 }
559
560 // Converts a canonical logical vector (all-zero or all-one lanes) to a mask.
561 [[nodiscard]] static MOCHI_FORCE_INLINE __mmask8 ToMask(Simd a) {
562 auto const bits = _mm512_castpd_si512(a.raw);
563 auto const mask = _mm512_movepi64_mask(bits); // AVX512DQ
565 _mm512_cmpeq_epi64_mask(bits, _mm512_movm_epi64(mask)) == LaneMask<kSize>(),
566 "Expected a canonical logical mask");
567 return mask;
568 }
569
570 // Expands a mask into a canonical logical vector (all-zero or all-one lanes).
571 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FromMask(__mmask8 mask) {
572 return _mm512_castsi512_pd(_mm512_movm_epi64(mask)); // AVX512DQ
573 }
574};
575
576} // namespace superdex
577
578#endif // MOCHI_USE_SIMD && MOCHI_ARCH_X64_AVX512
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_FORCE_INLINE
T Dot(Simd< T, N > a, Simd< T, N > b)
Definition simd.h:673
constexpr T ACos(T a)
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)
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)