SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_int_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<int, 16>
27*/
28template <>
29class Simd<int, 16> {
30 public:
31 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(int, 16, __m512i);
32
33 Simd(
34 int a,
35 int b,
36 int c = 0,
37 int d = 0,
38 int e = 0,
39 int f = 0,
40 int g = 0,
41 int h = 0,
42 int i = 0,
43 int j = 0,
44 int k = 0,
45 int l = 0,
46 int m = 0,
47 int n = 0,
48 int o = 0,
49 int p = 0)
50 : raw(_mm512_set_epi32(p, o, n, m, l, k, j, i, h, g, f, e, d, c, b, a)) {}
51
52 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
53 Simd(U a) : raw(_mm512_set1_epi32(a)) {}
54
55 Simd(Simd<int, 8> const& low, Simd<int, 8> const& high)
56 : raw(_mm512_inserti64x4(_mm512_castsi256_si512(low.raw), high.raw, 1)) {}
57
58 template <int i>
59 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar Get(Simd v) {
60 static_assert(i >= 0 && i < kSize, "Index out of range");
61 constexpr int kQuarter = i / 4;
62 constexpr int kLane = i % 4;
63 if constexpr (kQuarter == 0) {
64 return Simd<int, 4>::template Get<kLane>(_mm512_castsi512_si128(v.raw));
65 } else {
66 return Simd<int, 4>::template Get<kLane>(_mm512_extracti32x4_epi32(v.raw, kQuarter));
67 }
68 }
69
70 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
71 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
72 auto const indices = _mm512_set1_epi32(i);
73 return _mm_cvtsi128_si32(_mm512_castsi512_si128(_mm512_permutexvar_epi32(indices, raw)));
74 }
75
76 template <int iHalf>
77 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<int, 8> GetHalf(Simd a) {
78 static_assert(iHalf == 0 || iHalf == 1);
79 if constexpr (iHalf == 0) {
80 return _mm512_castsi512_si256(a.raw);
81 } else {
82 return _mm512_extracti64x4_epi64(a.raw, 1);
83 }
84 }
85
86 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, int i, Scalar value) {
87 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
88 auto const mask = static_cast<__mmask16>(uint32_t{1} << i);
89 return _mm512_mask_broadcastd_epi32(v.raw, mask, _mm_cvtsi32_si128(value));
90 }
91
92 template <int i>
93 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, Scalar value) {
94 static_assert(i >= 0 && i < kSize, "Index out of range");
95 constexpr auto kMask = static_cast<__mmask16>(uint32_t{1} << i);
96 return _mm512_mask_broadcastd_epi32(v.raw, kMask, _mm_cvtsi32_si128(value));
97 }
98
99 [[nodiscard]] static MOCHI_FORCE_INLINE Simd
100 SetInt64(int64_t a, int64_t b, int64_t c, int64_t d, int64_t e, int64_t f, int64_t g, int64_t h) {
101 return _mm512_set_epi64(h, g, f, e, d, c, b, a);
102 }
103
104 template <int N>
105 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
106 static_assert(N >= 1 && N <= kSize, "Unsupported N");
107 auto const mask = ToMask(v);
108 if constexpr (N == kSize) {
109 return _kortestc_mask16_u8(mask, mask) != 0;
110 } else {
111 constexpr auto kLanes = LaneMask<N>();
112 return (mask & kLanes) == kLanes;
113 }
114 }
115
116 template <int N>
117 [[nodiscard]] static MOCHI_FORCE_INLINE bool AnyTrue(Simd v) {
118 static_assert(N >= 1 && N <= kSize, "Unsupported N");
119 auto const mask = ToMask(v);
120 if constexpr (N == kSize) {
121 return _kortestz_mask16_u8(mask, mask) == 0;
122 } else {
123 return (mask & LaneMask<N>()) != 0;
124 }
125 }
126
127 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
128 return Simd{*p};
129 }
130
131 template <int i>
132 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
133 static_assert(i >= 0 && i < kSize, "Index out of range");
134 if constexpr (i == 0) {
135 return _mm512_broadcastd_epi32(_mm512_castsi512_si128(v.raw));
136 } else {
137 constexpr int kLane = i % 4;
138 constexpr int kGroup = i / 4;
139 auto const group =
140 _mm512_shuffle_i32x4(v.raw, v.raw, _MM_SHUFFLE(kGroup, kGroup, kGroup, kGroup));
141 return _mm512_shuffle_epi32(
142 group, static_cast<_MM_PERM_ENUM>(_MM_SHUFFLE(kLane, kLane, kLane, kLane)));
143 }
144 }
145
146 template <int N = kSize>
147 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
148 static_assert(N >= 2 && N <= kSize, "Unsupported N");
149 using HalfT = Simd<Scalar, 8>;
150 auto const lo = GetHalf<0>(a);
151 if constexpr (N <= 8) {
152 return HalfT::template HMin<N>(lo);
153 } else {
154 auto const hi = GetHalf<1>(a);
155 if constexpr (N == 9) {
156 return superdex::Min(HalfT::template HMin<8>(lo), HalfT::template Get<0>(hi));
157 } else if constexpr (N == kSize) {
158 return HalfT::template HMin<8>(HalfT::Min(lo, hi));
159 } else {
160 return _mm512_mask_reduce_min_epi32(LaneMask<N>(), a.raw);
161 }
162 }
163 }
164
165 template <int N = kSize>
166 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
167 static_assert(N >= 2 && N <= kSize, "Unsupported N");
168 using HalfT = Simd<Scalar, 8>;
169 auto const lo = GetHalf<0>(a);
170 if constexpr (N <= 8) {
171 return HalfT::template HMax<N>(lo);
172 } else {
173 auto const hi = GetHalf<1>(a);
174 if constexpr (N == 9) {
175 return superdex::Max(HalfT::template HMax<8>(lo), HalfT::template Get<0>(hi));
176 } else if constexpr (N == kSize) {
177 return HalfT::template HMax<8>(HalfT::Max(lo, hi));
178 } else {
179 return _mm512_mask_reduce_max_epi32(LaneMask<N>(), a.raw);
180 }
181 }
182 }
183
184 template <int N>
185 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
186 static_assert(N >= 2 && N <= kSize, "Unsupported N");
187 using HalfT = Simd<Scalar, 8>;
188 auto const lo = GetHalf<0>(a);
189 if constexpr (N <= 8) {
190 return HalfT::template HSum<N>(lo);
191 } else {
192 auto const hi = GetHalf<1>(a);
193 if constexpr (N == 9) {
194 return HalfT::template HSum<8>(lo) + HalfT::template Get<0>(hi);
195 } else if constexpr (N == kSize) {
196 return HalfT::template HSum<8>(lo + hi);
197 } else {
198 return _mm512_mask_reduce_add_epi32(LaneMask<N>(), a.raw);
199 }
200 }
201 }
202
203 template <int N = kSize>
204 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
205 static_assert(N >= 0 && N <= kSize);
206 if constexpr (N == 0) {
207 return Zero();
208 } else if constexpr (N == 1) {
209 return _mm512_zextsi128_si512(_mm_cvtsi32_si128(*ptr));
210 } else if constexpr (N == 2) {
211 return _mm512_zextsi128_si512(_mm_loadl_epi64(reinterpret_cast<__m128i const*>(ptr)));
212 } else if constexpr (N < 4) {
213 return _mm512_zextsi128_si512(
214 _mm_maskz_loadu_epi32(static_cast<__mmask8>((uint32_t{1} << N) - 1), ptr));
215
216 } else if constexpr (N == 4) {
217 return _mm512_zextsi128_si512(_mm_loadu_si128(reinterpret_cast<__m128i const*>(ptr)));
218 } else if constexpr (N < 8) {
219 return _mm512_zextsi256_si512(
220 _mm256_maskz_loadu_epi32(static_cast<__mmask8>((uint32_t{1} << N) - 1), ptr));
221 } else if constexpr (N == 8) {
222 return _mm512_zextsi256_si512(_mm256_loadu_si256(reinterpret_cast<__m256i const*>(ptr)));
223 } else if constexpr (N < kSize) {
224 return _mm512_maskz_loadu_epi32(LaneMask<N>(), ptr);
225 } else {
226 return _mm512_loadu_si512(ptr);
227 }
228 }
229
230 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
231 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
232 auto const mask = static_cast<__mmask16>((uint32_t{1} << n) - 1);
233 return _mm512_maskz_loadu_epi32(mask, ptr);
234 }
235
236 template <int kTupleCount = kSize>
237 MOCHI_FORCE_INLINE static void
238 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
239 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
240 if constexpr (kTupleCount == 1) {
241 out0 = Load<1>(ptr);
242 out1 = Load<1>(ptr + 1);
243 out2 = Load<1>(ptr + 2);
244 return;
245 }
246 constexpr int kTotalCount = kTupleCount * 3;
247 auto const x0 = Load<Clamp(kTotalCount, 0, kSize)>(ptr).raw;
248 if constexpr (kTupleCount <= 10) {
249 auto const index0 = LoadTransposeIndices<kTupleCount, 0>();
250 auto const index1 = LoadTransposeIndices<kTupleCount, 1>();
251 auto const index2 = LoadTransposeIndices<kTupleCount, 2>();
252 if constexpr (kTupleCount <= 5) {
253 out0.raw = _mm512_permutexvar_epi32(index0, x0);
254 out1.raw = _mm512_permutexvar_epi32(index1, x0);
255 out2.raw = _mm512_permutexvar_epi32(index2, x0);
256 } else {
257 auto const x1 = Load<kTotalCount - kSize>(ptr + kSize).raw;
258 out0.raw = _mm512_permutex2var_epi32(x0, index0, x1);
259 out1.raw = _mm512_permutex2var_epi32(x0, index1, x1);
260 out2.raw = _mm512_permutex2var_epi32(x0, index2, x1);
261 }
262 } else {
263 // clang-format off
264 auto const x1 = Load<kSize>(ptr + kSize).raw;
265 constexpr int kCount2 = kTotalCount - 2 * kSize;
266 auto const x2 = Load<kCount2>(ptr + 2 * kSize).raw;
267 auto const index0 = _mm512_setr_epi32(0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 0, 0, 0, 0, 0);
268 auto const index1 = _mm512_setr_epi32(1, 4, 7, 10, 13, 16, 19, 22, 25, 28, 31, 0, 0, 0, 0, 0);
269 auto const index2 = _mm512_setr_epi32(2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 0, 0, 0, 0, 0, 0);
270 if constexpr (kTupleCount == 11) {
271 out0.raw = _mm512_maskz_permutex2var_epi32(LaneMask<kTupleCount>(), x0, index0, x1);
272 out1.raw = _mm512_maskz_permutex2var_epi32(LaneMask<kTupleCount>(), x0, index1, x1);
273 constexpr int kZeroIndex = kSize + kCount2;
274 auto const partial2 = _mm512_permutex2var_epi32(x0, index2, x1);
275 auto const finalIndex2 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 16, kZeroIndex, kZeroIndex, kZeroIndex, kZeroIndex, kZeroIndex);
276 out2.raw = _mm512_permutex2var_epi32(partial2, finalIndex2, x2);
277 } else {
278 constexpr int kZeroIndex = kSize + kCount2;
279 auto const partial0 = _mm512_permutex2var_epi32(x0, index0, x1);
280 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);
281 out0.raw = _mm512_permutex2var_epi32(partial0, finalIndex0, x2);
282 auto const partial1 = _mm512_permutex2var_epi32(x0, index1, x1);
283 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);
284 out1.raw = _mm512_permutex2var_epi32(partial1, finalIndex1, x2);
285 auto const partial2 = _mm512_permutex2var_epi32(x0, index2, x1);
286 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);
287 out2.raw = _mm512_permutex2var_epi32(partial2, finalIndex2, x2);
288 }
289 // clang-format on
290 }
291 }
292
293 template <int N = kSize>
294 static MOCHI_FORCE_INLINE void Store([[maybe_unused]] Scalar* ptr, [[maybe_unused]] Simd v) {
295 static_assert(N >= 0 && N <= kSize);
296 if constexpr (N == 0) {
297 } else if constexpr (N == 1) {
298 _mm_storeu_si32(ptr, _mm512_castsi512_si128(v.raw));
299 } else if constexpr (N == 2) {
300 _mm_storel_epi64(reinterpret_cast<__m128i*>(ptr), _mm512_castsi512_si128(v.raw));
301 } else if constexpr (N < 4) {
302 _mm_mask_storeu_epi32(
303 ptr, static_cast<__mmask8>((uint32_t{1} << N) - 1), _mm512_castsi512_si128(v.raw));
304 } else if constexpr (N == 4) {
305 _mm_storeu_si128(reinterpret_cast<__m128i*>(ptr), _mm512_castsi512_si128(v.raw));
306 } else if constexpr (N < 8) {
307 _mm256_mask_storeu_epi32(
308 ptr, static_cast<__mmask8>((uint32_t{1} << N) - 1), _mm512_castsi512_si256(v.raw));
309 } else if constexpr (N == 8) {
310 _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), _mm512_castsi512_si256(v.raw));
311 } else if constexpr (N < kSize) {
312 _mm512_mask_storeu_epi32(ptr, LaneMask<N>(), v.raw);
313 } else {
314 _mm512_storeu_si512(ptr, v.raw);
315 }
316 }
317
318 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
319 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
320 auto const mask = static_cast<__mmask16>((uint32_t{1} << n) - 1);
321 _mm512_mask_storeu_epi32(ptr, mask, v.raw);
322 }
323
324 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
325 auto const mask = ToMask(condition);
326 _mm512_mask_compressstoreu_epi32(ptr, mask, values.raw);
327 return _mm_popcnt_u32(mask);
328 }
329
330 template <int kTupleCount = kSize>
331 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
332 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
333 if constexpr (kTupleCount == 1) {
334 Store<1>(ptr, a);
335 Store<1>(ptr + 1, b);
336 Store<1>(ptr + 2, c);
337 return;
338 }
339 constexpr int kTotalCount = kTupleCount * 3;
340 auto const ab0 = _mm512_permutex2var_epi32(
341 a.raw, _mm512_setr_epi32(0, 16, 0, 1, 17, 0, 2, 18, 0, 3, 19, 0, 4, 20, 0, 5), b.raw);
342 auto const x0 = _mm512_permutex2var_epi32(
343 ab0, _mm512_setr_epi32(0, 1, 16, 3, 4, 17, 6, 7, 18, 9, 10, 19, 12, 13, 20, 15), c.raw);
345 if constexpr (kTotalCount > kSize) {
346 auto const ab1 = _mm512_permutex2var_epi32(
347 a.raw, _mm512_setr_epi32(21, 0, 6, 22, 0, 7, 23, 0, 8, 24, 0, 9, 25, 0, 10, 26), b.raw);
348 auto const x1 = _mm512_permutex2var_epi32(
349 ab1, _mm512_setr_epi32(0, 21, 2, 3, 22, 5, 6, 23, 8, 9, 24, 11, 12, 25, 14, 15), c.raw);
350 Store<Clamp(kTotalCount - kSize, 0, kSize)>(ptr + kSize, x1);
351 }
352 if constexpr (kTotalCount > 2 * kSize) {
353 auto const ab2 = _mm512_permutex2var_epi32(
354 a.raw,
355 _mm512_setr_epi32(0, 11, 27, 0, 12, 28, 0, 13, 29, 0, 14, 30, 0, 15, 31, 0),
356 b.raw);
357 auto const x2 = _mm512_permutex2var_epi32(
358 ab2, _mm512_setr_epi32(26, 1, 2, 27, 4, 5, 28, 7, 8, 29, 10, 11, 30, 13, 14, 31), c.raw);
359 Store<kTotalCount - 2 * kSize>(ptr + 2 * kSize, x2);
360 }
361 }
362
363 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
364 return _mm512_min_epi32(a.raw, b.raw);
365 }
366
367 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
368 return _mm512_max_epi32(a.raw, b.raw);
369 }
370
371 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
372 return _mm512_mask_blend_epi32(ToMask(mask), b.raw, a.raw);
373 }
374
375 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
376 return _mm512_setzero_si512();
377 }
378
379 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
380 return FromMask(_mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_LT));
381 }
382
383 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
384 return FromMask(_mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_GT));
385 }
386
387 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
388 return FromMask(_mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_LE));
389 }
390
391 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
392 return FromMask(_mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_GE));
393 }
394
395 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
396 return FromMask(_mm512_cmp_epi32_mask(a.raw, b.raw, _MM_CMPINT_EQ));
397 }
398
399 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
400 return FromMask(_mm512_cmp_epi32_mask(a.raw, b.raw, _MM_CMPINT_NE));
401 }
402
403 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
404 return _mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_EQ) == __mmask16{0xFFFF};
405 }
406
407 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
408 return _mm512_cmp_epi32_mask(raw, rhs.raw, _MM_CMPINT_NE) != 0;
409 }
410
411 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
412 return _mm512_xor_si512(raw, _mm512_set1_epi32(-1));
413 }
414
415 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
416 return _mm512_sub_epi32(_mm512_setzero_si512(), raw);
417 }
418
419 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
420 return _mm512_add_epi32(raw, rhs.raw);
421 }
422
423 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
424 return _mm512_sub_epi32(raw, rhs.raw);
425 }
426
427 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
428 return _mm512_mullo_epi32(raw, rhs.raw);
429 }
430
431 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
432#if MOCHI_ARCH_X64_SVML
433 return _mm512_div_epi32(raw, rhs.raw);
434#else
435 return Simd{
436 Get<0>(*this) / Get<0>(rhs),
437 Get<1>(*this) / Get<1>(rhs),
438 Get<2>(*this) / Get<2>(rhs),
439 Get<3>(*this) / Get<3>(rhs),
440 Get<4>(*this) / Get<4>(rhs),
441 Get<5>(*this) / Get<5>(rhs),
442 Get<6>(*this) / Get<6>(rhs),
443 Get<7>(*this) / Get<7>(rhs),
444 Get<8>(*this) / Get<8>(rhs),
445 Get<9>(*this) / Get<9>(rhs),
446 Get<10>(*this) / Get<10>(rhs),
447 Get<11>(*this) / Get<11>(rhs),
448 Get<12>(*this) / Get<12>(rhs),
449 Get<13>(*this) / Get<13>(rhs),
450 Get<14>(*this) / Get<14>(rhs),
451 Get<15>(*this) / Get<15>(rhs)};
452#endif
453 }
454
455 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
456 return _mm512_and_si512(raw, rhs.raw);
457 }
458
459 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
460 return _mm512_or_si512(raw, rhs.raw);
461 }
462
463 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
464 return _mm512_xor_si512(raw, rhs.raw);
465 }
466
467 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<<(int rhs) const {
468 return _mm512_sll_epi32(raw, _mm_cvtsi32_si128(rhs));
469 }
470
471 template <int kShift>
472 [[nodiscard]] MOCHI_FORCE_INLINE static Simd ShiftRight(Simd a) {
473 static_assert(kShift >= 0 && kShift < 32, "Shift amount out-of-range");
474 if constexpr (kShift == 0) {
475 return a;
476 } else {
477 return _mm512_srai_epi32(a.raw, kShift);
478 }
479 }
480
481 private:
482 template <int kTupleCount, int kComponent>
483 [[nodiscard]] static MOCHI_FORCE_INLINE __m512i LoadTransposeIndices() {
484 constexpr int kZeroIndex = kTupleCount * 3;
485 return _mm512_setr_epi32(
486 kComponent,
487 kTupleCount > 1 ? 3 + kComponent : kZeroIndex,
488 kTupleCount > 2 ? 6 + kComponent : kZeroIndex,
489 kTupleCount > 3 ? 9 + kComponent : kZeroIndex,
490 kTupleCount > 4 ? 12 + kComponent : kZeroIndex,
491 kTupleCount > 5 ? 15 + kComponent : kZeroIndex,
492 kTupleCount > 6 ? 18 + kComponent : kZeroIndex,
493 kTupleCount > 7 ? 21 + kComponent : kZeroIndex,
494 kTupleCount > 8 ? 24 + kComponent : kZeroIndex,
495 kTupleCount > 9 ? 27 + kComponent : kZeroIndex,
496 kTupleCount > 10 ? 30 + kComponent : kZeroIndex,
497 kTupleCount > 11 ? 33 + kComponent : kZeroIndex,
498 kTupleCount > 12 ? 36 + kComponent : kZeroIndex,
499 kTupleCount > 13 ? 39 + kComponent : kZeroIndex,
500 kTupleCount > 14 ? 42 + kComponent : kZeroIndex,
501 kTupleCount > 15 ? 45 + kComponent : kZeroIndex);
502 }
503
504 // Returns a mask selecting the lowest N lanes.
505 template <int N>
506 [[nodiscard]] static constexpr __mmask16 LaneMask() {
507 static_assert(N >= 0 && N <= kSize);
508 if constexpr (N == kSize) {
509 return __mmask16{0xFFFF};
510 } else {
511 return static_cast<__mmask16>((uint32_t{1} << N) - 1);
512 }
513 }
514
515 // Converts a canonical logical vector (all-zero or all-one lanes) to a mask.
516 [[nodiscard]] static MOCHI_FORCE_INLINE __mmask16 ToMask(Simd a) {
517 auto const mask = _mm512_movepi32_mask(a.raw);
519 _mm512_cmpeq_epi32_mask(a.raw, _mm512_movm_epi32(mask)) == LaneMask<kSize>(),
520 "Expected a canonical logical mask");
521 return mask;
522 }
523
524 // Expands a mask into a canonical logical vector (all-zero or all-one lanes).
525 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FromMask(__mmask16 mask) {
526 return _mm512_movm_epi32(mask);
527 }
528};
529
530} // namespace superdex
531
532#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<<(int shift) 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
constexpr T const & Min(T const &a, T const &b)
constexpr auto Equal(T const &a, T const &b)
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
Simd< T, N > Set(Simd< T, N > a, T value)
Definition simd_inl.h:313
constexpr auto NotEqual(T const &a, T const &b)
T HMax(Simd< T, N > a)
Definition simd_inl.h:395
V Broadcast(typename V::Scalar a)
Definition simd_inl.h:115
bool AnyTrue(T const &a)
Definition basic_utils.h:66
constexpr T Select(bool condition, T a, T b)
T Get(Simd< T, N > v)
Definition simd_inl.h:303
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
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
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
Simd< T, N > ShiftRight(Simd< T, N > a)
Definition simd_inl.h:260
#define MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(T, N, NativeT)