SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_float_16_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<float, 16>
27*/
28template <>
29class Simd<float, 16> {
30 public:
31 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(float, 16, __m512);
32 Simd(
33 float a,
34 float b,
35 float c = 0.0f,
36 float d = 0.0f,
37 float e = 0.0f,
38 float f = 0.0f,
39 float g = 0.0f,
40 float h = 0.0f,
41 float i = 0.0f,
42 float j = 0.0f,
43 float k = 0.0f,
44 float l = 0.0f,
45 float m = 0.0f,
46 float n = 0.0f,
47 float o = 0.0f,
48 float p = 0.0f)
49 : raw(_mm512_set_ps(p, o, n, m, l, k, j, i, h, g, f, e, d, c, b, a)) {} // AVX512F
50
51 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
52 Simd(U a) : raw(_mm512_set1_ps(a)) {} // AVX512F
53
54 Simd(Simd<float, 8> const& low, Simd<float, 8> const& high)
55 : raw(_mm512_insertf32x8(_mm512_castps256_ps512(low.raw), high.raw, 1)) {} // AVX512DQ
56
57 template <int i>
58 [[nodiscard]] static MOCHI_FORCE_INLINE float Get(Simd v) {
59 static_assert(i >= 0 && i < kSize, "Index out of range");
60 return v[i];
61 }
62
63 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
64 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
65#if MOCHI_COMPILER_MSVC
66 return raw.m512_f32[i];
67#else
68 return raw[i];
69#endif
70 }
71
72 template <int iHalf>
73 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<float, 8> GetHalf(Simd a) {
74 static_assert(iHalf == 0 || iHalf == 1);
75 if constexpr (iHalf == 0) {
76 return _mm512_castps512_ps256(a.raw); // AVX512F
77 } else {
78 return _mm512_extractf32x8_ps(a.raw, 1); // AVX512DQ
79 }
80 }
81
82 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, int i, Scalar value) {
83 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
84 auto const mask = static_cast<__mmask16>(1u << i);
85 return _mm512_mask_broadcastss_ps(v.raw, mask, _mm_set_ss(value));
86 }
87
88 template <int i>
89 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, Scalar value) {
90 static_assert(i >= 0 && i < kSize, "Index out of range");
91 constexpr auto kMask = static_cast<__mmask16>(1u << i);
92 return _mm512_mask_broadcastss_ps(v.raw, kMask, _mm_set_ss(value));
93 }
94
95 template <int N>
96 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
97 static_assert(N >= 1 && N <= kSize, "Unsupported N");
98 auto const mask = ToMask(v);
99 if constexpr (N == kSize) {
100 return _kortestc_mask16_u8(mask, mask) != 0;
101 } else {
102 constexpr auto kLanes = LaneMask<N>();
103 return (mask & kLanes) == kLanes;
104 }
105 }
106
107 template <int N>
108 [[nodiscard]] static MOCHI_FORCE_INLINE bool AnyTrue(Simd v) {
109 static_assert(N >= 1 && N <= kSize, "Unsupported N");
110 auto const mask = ToMask(v);
111 if constexpr (N == kSize) {
112 return _kortestz_mask16_u8(mask, mask) == 0;
113 } else {
114 return (mask & LaneMask<N>()) != 0;
115 }
116 }
117
118 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
119 return _mm512_set1_ps(*p); // AVX512F
120 }
121
122 template <int i>
123 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
124 static_assert(i >= 0 && i < kSize, "Index out of range");
125 if constexpr (i == 0) {
126 return _mm512_broadcastss_ps(_mm512_castps512_ps128(v.raw));
127 } else {
128 constexpr int kLane = i % 4;
129 constexpr int kGroup = i / 4;
130 auto const group =
131 _mm512_shuffle_f32x4(v.raw, v.raw, _MM_SHUFFLE(kGroup, kGroup, kGroup, kGroup));
132 return _mm512_permute_ps(group, _MM_SHUFFLE(kLane, kLane, kLane, kLane));
133 }
134 }
135
136 template <int N = kSize>
137 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
138 static_assert(N >= 0 && N <= kSize);
139 if constexpr (N == 0) {
140 return Zero();
141 } else if constexpr (N == 1) {
142 return _mm512_zextps128_ps512(_mm_load_ss(ptr));
143 } else if constexpr (N == 2) {
144 auto const low = _mm_castsi128_ps(_mm_loadl_epi64(reinterpret_cast<__m128i const*>(ptr)));
145 return _mm512_zextps128_ps512(low);
146 } else if constexpr (N == 3) {
147 return _mm512_zextps128_ps512(
148 _mm_maskz_loadu_ps(static_cast<__mmask8>((uint32_t{1} << N) - 1), ptr));
149 } else if constexpr (N == 4) {
150 return _mm512_zextps128_ps512(_mm_loadu_ps(ptr));
151 } else if constexpr (N < 8) {
152 return _mm512_zextps256_ps512(
153 _mm256_maskz_loadu_ps(static_cast<__mmask8>((uint32_t{1} << N) - 1), ptr));
154 } else if constexpr (N == 8) {
155 return _mm512_zextps256_ps512(_mm256_loadu_ps(ptr));
156 } else if constexpr (N < kSize) {
157 return _mm512_maskz_loadu_ps(LaneMask<N>(), ptr); // AVX512F
158 } else {
159 return _mm512_loadu_ps(ptr); // AVX512F
160 }
161 }
162
163 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
164 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
165 return _mm512_maskz_loadu_ps(LaneMask(n), ptr); // AVX512F
166 }
167
168 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int, 16> const& indices) {
169 return _mm512_i32gather_ps(indices.raw, ptr, sizeof(float)); // AVX512F
170 }
171
172 template <int kTupleCount = kSize>
173 MOCHI_FORCE_INLINE static void
174 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
175 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
176 if constexpr (kTupleCount == 1) {
177 out0 = Load<1>(ptr);
178 out1 = Load<1>(ptr + 1);
179 out2 = Load<1>(ptr + 2);
180 return;
181 }
182 constexpr int kTotalCount = kTupleCount * 3;
183 auto const x0 = Load<Clamp(kTotalCount, 0, kSize)>(ptr).raw;
184 if constexpr (kTupleCount <= 10) {
185 auto const index0 = LoadTransposeIndices<kTupleCount, 0>();
186 auto const index1 = LoadTransposeIndices<kTupleCount, 1>();
187 auto const index2 = LoadTransposeIndices<kTupleCount, 2>();
188 if constexpr (kTupleCount <= 5) {
189 out0.raw = _mm512_permutexvar_ps(index0, x0);
190 out1.raw = _mm512_permutexvar_ps(index1, x0);
191 out2.raw = _mm512_permutexvar_ps(index2, x0);
192 } else {
193 auto const x1 = Load<kTotalCount - kSize>(ptr + kSize).raw;
194 out0.raw = _mm512_permutex2var_ps(x0, index0, x1);
195 out1.raw = _mm512_permutex2var_ps(x0, index1, x1);
196 out2.raw = _mm512_permutex2var_ps(x0, index2, x1);
197 }
198 } else {
199 // clang-format off
200 auto const x1 = Load<kSize>(ptr + kSize).raw;
201 constexpr int kCount2 = kTotalCount - 2 * kSize;
202 auto const x2 = Load<kCount2>(ptr + 2 * kSize).raw;
203 auto const index0 = _mm512_setr_epi32(0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 0, 0, 0, 0, 0);
204 auto const index1 = _mm512_setr_epi32(1, 4, 7, 10, 13, 16, 19, 22, 25, 28, 31, 0, 0, 0, 0, 0);
205 auto const index2 = _mm512_setr_epi32(2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 0, 0, 0, 0, 0, 0);
206 if constexpr (kTupleCount == 11) {
207 out0.raw = _mm512_maskz_permutex2var_ps(LaneMask<kTupleCount>(), x0, index0, x1);
208 out1.raw = _mm512_maskz_permutex2var_ps(LaneMask<kTupleCount>(), x0, index1, x1);
209 constexpr int kZeroIndex = kSize + kCount2;
210 auto const partial2 = _mm512_permutex2var_ps(x0, index2, x1);
211 auto const finalIndex2 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 16, kZeroIndex, kZeroIndex, kZeroIndex, kZeroIndex, kZeroIndex);
212 out2.raw = _mm512_permutex2var_ps(partial2, finalIndex2, x2);
213 } else {
214 constexpr int kZeroIndex = kSize + kCount2;
215 auto const partial0 = _mm512_permutex2var_ps(x0, index0, x1);
216 auto const finalIndex0 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, kTupleCount > 11 ? 17 : kZeroIndex, kTupleCount > 12 ? 20 : kZeroIndex, kTupleCount > 13 ? 23 : kZeroIndex, kTupleCount > 14 ? 26 : kZeroIndex, kTupleCount > 15 ? 29 : kZeroIndex);
217 out0.raw = _mm512_permutex2var_ps(partial0, finalIndex0, x2);
218 auto const partial1 = _mm512_permutex2var_ps(x0, index1, x1);
219 auto const finalIndex1 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, kTupleCount > 11 ? 18 : kZeroIndex, kTupleCount > 12 ? 21 : kZeroIndex, kTupleCount > 13 ? 24 : kZeroIndex, kTupleCount > 14 ? 27 : kZeroIndex, kTupleCount > 15 ? 30 : kZeroIndex);
220 out1.raw = _mm512_permutex2var_ps(partial1, finalIndex1, x2);
221 auto const partial2 = _mm512_permutex2var_ps(x0, index2, x1);
222 auto const finalIndex2 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 16, kTupleCount > 11 ? 19 : kZeroIndex, kTupleCount > 12 ? 22 : kZeroIndex, kTupleCount > 13 ? 25 : kZeroIndex, kTupleCount > 14 ? 28 : kZeroIndex, kTupleCount > 15 ? 31 : kZeroIndex);
223 out2.raw = _mm512_permutex2var_ps(partial2, finalIndex2, x2);
224 }
225 // clang-format on
226 }
227 }
228
229 template <int N = kSize>
230 static MOCHI_FORCE_INLINE void Store([[maybe_unused]] Scalar* ptr, [[maybe_unused]] Simd v) {
231 static_assert(N >= 0 && N <= kSize);
232 if constexpr (N == 0) {
233 } else if constexpr (N == 1) {
234 _mm_store_ss(ptr, _mm512_castps512_ps128(v.raw));
235 } else if constexpr (N == 2) {
236 _mm_storel_epi64(
237 reinterpret_cast<__m128i*>(ptr), _mm_castps_si128(_mm512_castps512_ps128(v.raw)));
238 } else if constexpr (N == 3) {
239 _mm_mask_storeu_ps(
240 ptr, static_cast<__mmask8>((uint32_t{1} << N) - 1), _mm512_castps512_ps128(v.raw));
241 } else if constexpr (N == 4) {
242 _mm_storeu_ps(ptr, _mm512_castps512_ps128(v.raw));
243 } else if constexpr (N < 8) {
244 _mm256_mask_storeu_ps(
245 ptr, static_cast<__mmask8>((uint32_t{1} << N) - 1), _mm512_castps512_ps256(v.raw));
246 } else if constexpr (N == 8) {
247 _mm256_storeu_ps(ptr, _mm512_castps512_ps256(v.raw));
248 } else if constexpr (N < kSize) {
249 _mm512_mask_storeu_ps(ptr, LaneMask<N>(), v.raw); // AVX512F
250 } else {
251 _mm512_storeu_ps(ptr, v.raw); // AVX512F
252 }
253 }
254
255 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
256 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
257 _mm512_mask_storeu_ps(ptr, LaneMask(n), v.raw); // AVX512F
258 }
259
260 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
261 __mmask16 const mask = ToMask(condition);
262 _mm512_mask_compressstoreu_ps(ptr, mask, values.raw); // AVX512F
263 return _mm_popcnt_u32(mask);
264 }
265
266 template <int kTupleCount = kSize>
267 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
268 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
269 if constexpr (kTupleCount == 1) {
270 Store<1>(ptr, a);
271 Store<1>(ptr + 1, b);
272 Store<1>(ptr + 2, c);
273 return;
274 }
275 constexpr int kTotalCount = kTupleCount * 3;
276 auto const ab0 = _mm512_permutex2var_ps(
277 a.raw, _mm512_setr_epi32(0, 16, 0, 1, 17, 0, 2, 18, 0, 3, 19, 0, 4, 20, 0, 5), b.raw);
278 auto const x0 = _mm512_permutex2var_ps(
279 ab0, _mm512_setr_epi32(0, 1, 16, 3, 4, 17, 6, 7, 18, 9, 10, 19, 12, 13, 20, 15), c.raw);
281 if constexpr (kTotalCount > kSize) {
282 auto const ab1 = _mm512_permutex2var_ps(
283 a.raw, _mm512_setr_epi32(21, 0, 6, 22, 0, 7, 23, 0, 8, 24, 0, 9, 25, 0, 10, 26), b.raw);
284 auto const x1 = _mm512_permutex2var_ps(
285 ab1, _mm512_setr_epi32(0, 21, 2, 3, 22, 5, 6, 23, 8, 9, 24, 11, 12, 25, 14, 15), c.raw);
286 Store<Clamp(kTotalCount - kSize, 0, kSize)>(ptr + kSize, x1);
287 }
288 if constexpr (kTotalCount > 2 * kSize) {
289 auto const ab2 = _mm512_permutex2var_ps(
290 a.raw,
291 _mm512_setr_epi32(0, 11, 27, 0, 12, 28, 0, 13, 29, 0, 14, 30, 0, 15, 31, 0),
292 b.raw);
293 auto const x2 = _mm512_permutex2var_ps(
294 ab2, _mm512_setr_epi32(26, 1, 2, 27, 4, 5, 28, 7, 8, 29, 10, 11, 30, 13, 14, 31), c.raw);
295 Store<kTotalCount - 2 * kSize>(ptr + 2 * kSize, x2);
296 }
297 }
298
299 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
300 return _mm512_mask_blend_ps(ToMask(mask), b.raw, a.raw); // AVX512F
301 }
302
303 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sqrt(Simd v) {
304 return _mm512_sqrt_ps(v.raw); // AVX512F
305 }
306
307 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpApprox(Simd v) {
308 return _mm512_rcp14_ps(v.raw); // AVX512F
309 }
310
311 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpSqrtApprox(Simd v) {
312 return _mm512_rsqrt14_ps(v.raw); // AVX512F
313 }
314
315 [[nodiscard]] static MOCHI_FORCE_INLINE Simd SignBitMask() {
316 return _mm512_castsi512_ps(_mm512_set1_epi32(static_cast<int>(0x80000000u))); // AVX512F
317 }
318
319 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Abs(Simd v) {
320 return _mm512_abs_ps(v.raw); // AVX512F
321 }
322
323 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
324 return _mm512_min_ps(a.raw, b.raw); // AVX512F
325 }
326
327 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
328 return _mm512_max_ps(a.raw, b.raw); // AVX512F
329 }
330
331 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Floor(Simd a) {
332 return _mm512_floor_ps(a.raw); // AVX512F
333 }
334
335 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FastRound(Simd v) {
336 return _mm512_roundscale_ps(v.raw, _MM_FROUND_TO_NEAREST_INT); // AVX512F
337 }
338
339#if MOCHI_ARCH_X64_SVML
340 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Cos(Simd a) {
341 return _mm512_cos_ps(a.raw);
342 }
343
344 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sin(Simd a) {
345 return _mm512_sin_ps(a.raw);
346 }
347
348 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tan(Simd a) {
349 return _mm512_tan_ps(a.raw);
350 }
351
352 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ACos(Simd a) {
353 return _mm512_acos_ps(a.raw);
354 }
355
356 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ASin(Simd a) {
357 return _mm512_asin_ps(a.raw);
358 }
359
360 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ATan(Simd a) {
361 return _mm512_atan_ps(a.raw);
362 }
363
364 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Exp(Simd a) {
365 return _mm512_exp_ps(a.raw);
366 }
367
368 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Ln(Simd a) {
369 return _mm512_log_ps(a.raw);
370 }
371
372 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tanh(Simd a) {
373 return _mm512_tanh_ps(a.raw);
374 }
375#endif // MOCHI_ARCH_X64_SVML
376
377 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulAdd(Simd a, Simd b, Simd c) {
378 return _mm512_fmadd_ps(a.raw, b.raw, c.raw); // AVX512F
379 }
380
381 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulSub(Simd a, Simd b, Simd c) {
382 return _mm512_fmsub_ps(a.raw, b.raw, c.raw); // AVX512F
383 }
384
385 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulAdd(Simd a, Simd b, Simd c) {
386 return _mm512_fnmadd_ps(a.raw, b.raw, c.raw); // AVX512F
387 }
388
389 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulSub(Simd a, Simd b, Simd c) {
390 return _mm512_fnmsub_ps(a.raw, b.raw, c.raw); // AVX512F
391 }
392
393 template <int N = kSize>
394 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
395 static_assert(N >= 2 && N <= kSize, "Unsupported N");
396 using HalfT = Simd<Scalar, 8>;
397 auto const lo = GetHalf<0>(a);
398 if constexpr (N <= 8) {
399 return HalfT::template HMin<N>(lo);
400 } else {
401 auto const hi = GetHalf<1>(a);
402 if constexpr (N == 9) {
403 return superdex::Min(HalfT::template HMin<8>(lo), HalfT::template Get<0>(hi));
404 } else if constexpr (N == kSize) {
405 return HalfT::template HMin<8>(HalfT::Min(lo, hi));
406 } else {
407 return _mm512_mask_reduce_min_ps(LaneMask<N>(), a.raw);
408 }
409 }
410 }
411
412 template <int N = kSize>
413 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
414 static_assert(N >= 2 && N <= kSize, "Unsupported N");
415 using HalfT = Simd<Scalar, 8>;
416 auto const lo = GetHalf<0>(a);
417 if constexpr (N <= 8) {
418 return HalfT::template HMax<N>(lo);
419 } else {
420 auto const hi = GetHalf<1>(a);
421 if constexpr (N == 9) {
422 return superdex::Max(HalfT::template HMax<8>(lo), HalfT::template Get<0>(hi));
423 } else if constexpr (N == kSize) {
424 return HalfT::template HMax<8>(HalfT::Max(lo, hi));
425 } else {
426 return _mm512_mask_reduce_max_ps(LaneMask<N>(), a.raw);
427 }
428 }
429 }
430
431 template <int N>
432 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
433 static_assert(N >= 2 && N <= kSize, "Unsupported N");
434 using HalfT = Simd<Scalar, 8>;
435 auto const lo = GetHalf<0>(a);
436 if constexpr (N <= 8) {
437 return HalfT::template HSum<N>(lo);
438 } else {
439 auto const hi = GetHalf<1>(a);
440 if constexpr (N == 9) {
441 return HalfT::template HSum<8>(lo) + HalfT::template Get<0>(hi);
442 } else if constexpr (N == kSize) {
443 return HalfT::template HSum<8>(lo + hi);
444 } else {
445 return _mm512_mask_reduce_add_ps(LaneMask<N>(), a.raw);
446 }
447 }
448 }
449
450 template <int N>
451 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Dot(Simd a, Simd b) {
452 static_assert(N >= 2 && N <= kSize, "Unsupported N");
453 return Simd{HSum<N>(a * b)};
454 }
455
456 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
457 return FromMask(_mm512_cmp_ps_mask(raw, rhs.raw, _CMP_LT_OQ)); // AVX512F
458 }
459
460 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
461 return FromMask(_mm512_cmp_ps_mask(raw, rhs.raw, _CMP_GT_OQ)); // AVX512F
462 }
463
464 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
465 return FromMask(_mm512_cmp_ps_mask(raw, rhs.raw, _CMP_LE_OQ)); // AVX512F
466 }
467
468 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
469 return FromMask(_mm512_cmp_ps_mask(raw, rhs.raw, _CMP_GE_OQ)); // AVX512F
470 }
471
472 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
473 return FromMask(_mm512_cmp_ps_mask(a.raw, b.raw, _CMP_EQ_OQ)); // AVX512F
474 }
475
476 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
477 return FromMask(_mm512_cmp_ps_mask(a.raw, b.raw, _CMP_NEQ_UQ)); // AVX512F
478 }
479
480 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
481 return _mm512_setzero_ps(); // AVX512F
482 }
483
484 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
485 return _mm512_cmp_ps_mask(raw, rhs.raw, _CMP_EQ_OQ) == static_cast<__mmask16>(0xFFFFu);
486 }
487
488 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
489 return _mm512_cmp_ps_mask(raw, rhs.raw, _CMP_NEQ_UQ) != 0;
490 }
491
492 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
493 return _mm512_castsi512_ps(
494 _mm512_xor_si512(_mm512_castps_si512(raw), _mm512_set1_epi32(-1))); // AVX512F
495 }
496
497 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
498 return _mm512_xor_ps(raw, SignBitMask().raw); // AVX512DQ
499 }
500
501 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
502 return _mm512_add_ps(raw, rhs.raw); // AVX512F
503 }
504
505 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
506 return _mm512_sub_ps(raw, rhs.raw); // AVX512F
507 }
508
509 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
510 return _mm512_mul_ps(raw, rhs.raw); // AVX512F
511 }
512
513 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
514 return _mm512_div_ps(raw, rhs.raw); // AVX512F
515 }
516
517 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
518 return _mm512_and_ps(raw, rhs.raw); // AVX512DQ
519 }
520
521 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
522 return _mm512_or_ps(raw, rhs.raw); // AVX512DQ
523 }
524
525 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
526 return _mm512_xor_ps(raw, rhs.raw); // AVX512DQ
527 }
528
529 private:
530 template <int kTupleCount, int kComponent>
531 [[nodiscard]] static MOCHI_FORCE_INLINE __m512i LoadTransposeIndices() {
532 constexpr int kZeroIndex = kTupleCount * 3;
533 return _mm512_setr_epi32(
534 kComponent,
535 kTupleCount > 1 ? 3 + kComponent : kZeroIndex,
536 kTupleCount > 2 ? 6 + kComponent : kZeroIndex,
537 kTupleCount > 3 ? 9 + kComponent : kZeroIndex,
538 kTupleCount > 4 ? 12 + kComponent : kZeroIndex,
539 kTupleCount > 5 ? 15 + kComponent : kZeroIndex,
540 kTupleCount > 6 ? 18 + kComponent : kZeroIndex,
541 kTupleCount > 7 ? 21 + kComponent : kZeroIndex,
542 kTupleCount > 8 ? 24 + kComponent : kZeroIndex,
543 kTupleCount > 9 ? 27 + kComponent : kZeroIndex,
544 kTupleCount > 10 ? 30 + kComponent : kZeroIndex,
545 kTupleCount > 11 ? 33 + kComponent : kZeroIndex,
546 kTupleCount > 12 ? 36 + kComponent : kZeroIndex,
547 kTupleCount > 13 ? 39 + kComponent : kZeroIndex,
548 kTupleCount > 14 ? 42 + kComponent : kZeroIndex,
549 kTupleCount > 15 ? 45 + kComponent : kZeroIndex);
550 }
551
552 // Returns a mask selecting the lowest N lanes.
553 template <int N>
554 [[nodiscard]] static constexpr __mmask16 LaneMask() {
555 static_assert(N >= 0 && N <= kSize);
556 return static_cast<__mmask16>((uint32_t{1} << N) - 1);
557 }
558
559 // Returns a mask selecting the lowest n lanes.
560 [[nodiscard]] static constexpr __mmask16 LaneMask(int n) {
561 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid lane count");
562 return static_cast<__mmask16>((uint32_t{1} << n) - 1);
563 }
564
565 // Converts a canonical logical vector (all-zero or all-one lanes) to a mask.
566 [[nodiscard]] static MOCHI_FORCE_INLINE __mmask16 ToMask(Simd a) {
567 auto const bits = _mm512_castps_si512(a.raw);
568 auto const mask = _mm512_movepi32_mask(bits); // AVX512DQ
570 _mm512_cmpeq_epi32_mask(bits, _mm512_movm_epi32(mask)) == LaneMask<kSize>(),
571 "Expected a canonical logical mask");
572 return mask;
573 }
574
575 // Expands a mask into a canonical logical vector (all-zero or all-one lanes).
576 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FromMask(__mmask16 mask) {
577 return _mm512_castsi512_ps(_mm512_movm_epi32(mask)); // AVX512DQ
578 }
579};
580
581} // namespace superdex
582
583#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
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)