SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_int64_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<int64_t, 4>
27*/
28template <>
29class Simd<int64_t, 4> {
30 public:
31 static_assert(sizeof(int64_t) == sizeof(long long));
32
33 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(int64_t, 4, __m256i);
34 Simd(int64_t a, int64_t b, int64_t c = 0, int64_t d = 0)
35 : raw(_mm256_set_epi64x(
36 static_cast<long long>(d),
37 static_cast<long long>(c),
38 static_cast<long long>(b),
39 static_cast<long long>(a))) {} // AVX
40 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
41 Simd(U a) : raw(_mm256_set1_epi64x(static_cast<long long>(a))) {} // AVX
42
43 template <int i>
44 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar Get(Simd v) {
45 static_assert(i >= 0 && i < kSize, "Index out of range");
46 return v[i];
47 }
48
49 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
50 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
51#if MOCHI_COMPILER_MSVC
52 return raw.m256i_i64[i];
53#else
54 return static_cast<Scalar>(raw[i]);
55#endif
56 }
57
58 template <int iHalf>
59 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<int64_t, 2> GetHalf(Simd a) {
60 static_assert(iHalf == 0 || iHalf == 1);
61 return _mm256_extracti128_si256(a.raw, iHalf); // AVX2
62 }
63
64 template <int N>
65 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
66 static_assert(N >= 1 && N <= kSize, "Unsupported N");
67 int mask = GetMSBitMask(v); // One bit for each byte in the vector
68 if constexpr (N == kSize) {
69 return mask == 0xFFFFFFFF;
70 } else {
71 int constexpr kNumBits = N * sizeof(Scalar);
72 auto constexpr kMustBeSet = (1UL << kNumBits) - 1;
73 return (mask & kMustBeSet) == kMustBeSet;
74 }
75 }
76
77 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
78 return Simd{*p};
79 }
80
81 template <int i>
82 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
83 return Shuffle<i, i, i, i>(v);
84 }
85
86 template <int N = 4>
87 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
88 static_assert(N >= 2 && N <= 4, "Unsupported N");
89 using HalfT = Simd<Scalar, 2>;
90 if constexpr (N == 2) {
91 return Get<0>(Min(a, Broadcast<1>(a)));
92 } else if constexpr (N == 3) {
93 auto lo = GetHalf<0>(a);
94 auto hi = GetHalf<1>(a);
95 return HalfT::Get<0>(HalfT::Min(HalfT::Min(lo, HalfT::Broadcast<1>(lo)), hi));
96 } else {
97 return HalfT::HMin(HalfT::Min(GetHalf<0>(a), GetHalf<1>(a)));
98 }
99 }
100
101 template <int N = 4>
102 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
103 static_assert(N >= 2 && N <= 4, "Unsupported N");
104 using HalfT = Simd<Scalar, 2>;
105 if constexpr (N == 2) {
106 return Get<0>(Max(a, Broadcast<1>(a)));
107 } else if constexpr (N == 3) {
108 auto lo = GetHalf<0>(a);
109 auto hi = GetHalf<1>(a);
110 return HalfT::Get<0>(HalfT::Max(HalfT::Max(lo, HalfT::Broadcast<1>(lo)), hi));
111 } else {
112 return HalfT::HMax(HalfT::Max(GetHalf<0>(a), GetHalf<1>(a)));
113 }
114 }
115
116 template <int N>
117 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
118 static_assert(N >= 2 && N <= 4, "Unsupported N");
119 if constexpr (N == 2) {
120 return Get<0>(a) + Get<1>(a);
121 } else if constexpr (N == 3) {
122 using HalfT = Simd<int64_t, 2>;
123 HalfT tmp = GetHalf<0>(a) + GetHalf<1>(a);
124 return HalfT::Get<0>(tmp) + Get<1>(a); // (a[0] + a[2]) + a[1]
125 } else if constexpr (N == 4) {
126 using HalfT = Simd<int64_t, 2>;
127 HalfT tmp = GetHalf<0>(a) + GetHalf<1>(a);
128 return HalfT::Get<0>(tmp) + HalfT::Get<1>(tmp); // (a[0] + a[2]) + (a[1] + a[3]);
129 }
130 }
131
132 template <int N = kSize>
133 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
134 static_assert(N >= 0 && N <= kSize);
135 if constexpr (N == 0) {
136 return Simd::Zero();
137 } else if constexpr (N == 1) {
138 return Simd{*ptr, 0, 0, 0};
139 } else if constexpr (N == 2) {
140 __m256i mask = _mm256_set_epi32(0, 0, 0, 0, -1, -1, -1, -1); // AVX
141 return _mm256_maskload_epi64(reinterpret_cast<long long const*>(ptr), mask); // AVX2
142 } else if constexpr (N == 3) {
143 __m256i mask = _mm256_set_epi32(0, 0, -1, -1, -1, -1, -1, -1); // AVX
144 return _mm256_maskload_epi64(reinterpret_cast<long long const*>(ptr), mask); // AVX2
145 } else {
146 return _mm256_loadu_si256(reinterpret_cast<__m256i const*>(ptr)); // AVX
147 }
148 }
149
150 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
151 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
152#if MOCHI_ARCH_X64_AVX512
153 return _mm256_maskz_loadu_epi64(x64_simd::kLaneMasksS8[n], ptr); // AVX512VL
154#else
155 switch (n) { // clang-format off
156 case 1: return Load<1>(ptr);
157 case 2: return Load<2>(ptr);
158 case 3: return Load<3>(ptr);
159 case 4: return Load<4>(ptr);
160 MOCHI_UNLIKELY default: return Zero();
161 } // clang-format on
162#endif
163 }
164
165 template <int kTupleCount = kSize>
166 MOCHI_FORCE_INLINE static void
167 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
168 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
169 constexpr int kCount0 = Clamp(kTupleCount * 3 - 0, 0, 4);
170 constexpr int kCount1 = Clamp(kTupleCount * 3 - 4, 0, 4);
171 constexpr int kCount2 = Clamp(kTupleCount * 3 - 8, 0, 4);
172 auto a = _mm256_castsi256_pd(Simd::Load<kCount0>(ptr).raw); // [0,1,2,3]
173 auto b =
174 _mm256_castsi256_pd(Simd::Load<kCount1>(kCount1 == 0 ? ptr : ptr + 4).raw); // [4,5,6,7]
175 auto c =
176 _mm256_castsi256_pd(Simd::Load<kCount2>(kCount2 == 0 ? ptr : ptr + 8).raw); // [8,9,10,11]
177
178 auto d = _mm256_blend_pd(a, b, 0b0100); // [0,_,6,3]
179 d = _mm256_blend_pd(d, c, 0b0010); // [0,9,6,3]
180 auto e = _mm256_permute2f128_pd(d, d, 0x01); // [6,3,0,9]
181 out0.raw = _mm256_castpd_si256(_mm256_blend_pd(d, e, 0b1010)); // [0,3,6,9]
182
183 d = _mm256_blend_pd(a, b, 0b1001); // [4,1,_,7]
184 d = _mm256_blend_pd(d, c, 0b0100); // [4,1,10,7]
185 out1.raw = _mm256_castpd_si256(_mm256_shuffle_pd(d, d, 0b0101)); // [1,4,7,10]
186
187 d = _mm256_blend_pd(a, b, 0b0010); // [_,5,2,_]
188 d = _mm256_blend_pd(d, c, 0b1001); // [8,5,2,11]
189 e = _mm256_permute2f128_pd(d, d, 0x01); // [2,11,8,5]
190 out2.raw = _mm256_castpd_si256(_mm256_blend_pd(d, e, 0b0101)); // [2,5,8,11]
191 }
192
193 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
194#if MOCHI_ARCH_X64_AVX512
195 return _mm256_min_epi64(a.raw, b.raw); // AVX512VL
196#else
197 return Simd{
201 superdex::Min(Get<3>(a), Get<3>(b))};
202#endif
203 }
204
205 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
206#if MOCHI_ARCH_X64_AVX512
207 return _mm256_max_epi64(a.raw, b.raw); // AVX512VL
208#else
209 return Simd{
213 superdex::Max(Get<3>(a), Get<3>(b))};
214#endif
215 }
216
217 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
218 return _mm256_blendv_epi8(b.raw, a.raw, mask.raw); // AVX2
219 }
220
221 template <int x, int y, int z, int w>
222 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Shuffle(Simd a) {
223 if constexpr (x == 0 && y == 1 && z == 2 && w == 3) {
224 return a;
225 } else {
226 return _mm256_permute4x64_epi64(a.raw, _MM_SHUFFLE(w, z, y, x)); // AVX2
227 }
228 }
229
230 template <int x, int y, int z, int w>
231 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Shuffle(Simd a, Simd b) {
232 return Simd{Get<x>(a), Get<y>(a), Get<z>(b), Get<w>(b)};
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 memcpy(ptr, &v.raw, sizeof(Scalar) * N);
241 } else {
242 _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), v.raw); // AVX2
243 }
244 }
245
246 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
247 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
248#if MOCHI_ARCH_X64_AVX512
249 _mm256_mask_storeu_epi64(ptr, x64_simd::kLaneMasksS8[n], v.raw); // AVX512VL
250#else
251 // With AVX2, this is faster than masked store for a predictable value of n.
252 // It is much slower for a random value of n.
253 switch (n) { // clang-format off
254 case 1: Store<1>(ptr, v); break;
255 case 2: Store<2>(ptr, v); break;
256 case 3: Store<3>(ptr, v); break;
257 case 4: Store<4>(ptr, v); break;
258 MOCHI_UNLIKELY default: break;
259 } // clang-format on
260#endif
261 }
262
263 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
264#if MOCHI_ARCH_X64_AVX512
265 auto const mask = _mm256_movepi64_mask(condition.raw);
266 _mm256_mask_compressstoreu_epi64(ptr, mask, values.raw); // AVX512VL
267 return _mm_popcnt_u32(mask);
268#else
269 auto mask = _mm256_movemask_pd(_mm256_castsi256_pd(condition.raw));
270 // Load 8 bytes from the table, then zero-exend to get the shuffle pattern.
271 auto const* tableRow =
272 reinterpret_cast<__m128i const*>(x64_simd::kStoreSelectedShuffleTableD4[mask]);
273 auto pattern = _mm256_cvtepu8_epi32(_mm_loadl_epi64(tableRow));
274 auto packed = _mm256_permutevar8x32_epi32(values.raw, pattern);
275 _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), packed);
276 return _mm_popcnt_u32(mask);
277#endif
278 }
279
280 template <int kTupleCount = kSize>
281 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
282 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
283 // a = [0,3,6,9], b = [1,4,7,10], c = [2,5,8,11]
284 auto a_ = _mm256_castsi256_pd(a.raw);
285 auto b_ = _mm256_castsi256_pd(b.raw);
286 auto c_ = _mm256_castsi256_pd(c.raw);
287 auto d = _mm256_shuffle_pd(b_, b_, 0b0101); // [4,1,10,7]
288 auto e = _mm256_blend_pd(a_, c_, 0b0101); // [2,3,8,9]
289 e = _mm256_permute2f128_pd(e, e, 0x01); // [8,9,2,3]
290 auto f = _mm256_blend_pd(a_, d, 0b1010); // [0,1,6,7]
291 auto g = _mm256_blend_pd(d, c_, 0b1010); // [4,5,10,11]
292 constexpr int kCount0 = Clamp(kTupleCount * 3 - 0, 0, 4);
293 constexpr int kCount1 = Clamp(kTupleCount * 3 - 4, 0, 4);
294 constexpr int kCount2 = Clamp(kTupleCount * 3 - 8, 0, 4);
295 Simd::Store<kCount0>(
296 ptr, Simd{_mm256_castpd_si256(_mm256_blend_pd(f, e, 0b1100))}); // [0,1,2,3]
297 if constexpr (kCount1 > 0) {
298 Simd::Store<kCount1>(
299 ptr + 4, Simd{_mm256_castpd_si256(_mm256_blend_pd(g, f, 0b1100))}); // [4,5,6,7]
300 }
301 if constexpr (kCount2 > 0) {
302 Simd::Store<kCount2>(
303 ptr + 8, Simd{_mm256_castpd_si256(_mm256_blend_pd(e, g, 0b1100))}); // [8,9,10,11]
304 }
305 }
306
307 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
308 return _mm256_setzero_si256(); // AVX
309 }
310
311 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
312 return _mm256_cmpgt_epi64(rhs.raw, this->raw); // AVX2
313 }
314
315 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
316 return _mm256_cmpgt_epi64(this->raw, rhs.raw); // AVX2
317 }
318
319 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
320 return ~(*this > rhs); // No native support
321 }
322
323 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
324 return ~(*this < rhs); // No native support
325 }
326
327 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
328 return _mm256_cmpeq_epi64(a.raw, b.raw); // AVX2
329 }
330
331 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
332 return ~Equal(a, b); // // No native support
333 }
334
335 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
336 auto mask = GetMSBitMask(Equal(*this, rhs));
337 return mask == 0xFFFFFFFF; // All values equal
338 }
339
340 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
341 auto mask = GetMSBitMask(NotEqual(*this, rhs));
342 return mask != 0; // Any values not equal
343 }
344
345 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
346 auto ones = _mm256_cmpeq_epi64(raw, raw); // AVX2
347 return _mm256_xor_si256(raw, ones); // AVX2
348 }
349
350 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
351 return _mm256_sub_epi64(_mm256_setzero_si256(), raw); // AVX2
352 }
353
354 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
355 return _mm256_add_epi64(raw, rhs.raw); // SSE
356 }
357
358 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
359 return _mm256_sub_epi64(raw, rhs.raw); // SSE
360 }
361
362 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
363#if MOCHI_ARCH_X64_AVX512
364 return _mm256_mullo_epi64(raw, rhs.raw); // AVX512VL
365#else
366 return Simd{
367 Get<0>(*this) * Get<0>(rhs),
368 Get<1>(*this) * Get<1>(rhs),
369 Get<2>(*this) * Get<2>(rhs),
370 Get<3>(*this) * Get<3>(rhs)};
371#endif
372 }
373
374 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
375#if MOCHI_ARCH_X64_SVML
376 return _mm256_div_epi64(raw, rhs.raw); // SSE
377#else
378 // Fallback
379 return Simd{
380 Get<0>(*this) / Get<0>(rhs),
381 Get<1>(*this) / Get<1>(rhs),
382 Get<2>(*this) / Get<2>(rhs),
383 Get<3>(*this) / Get<3>(rhs)};
384#endif
385 }
386
387 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
388 return _mm256_and_si256(raw, rhs.raw); // AVX2
389 }
390
391 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
392 return _mm256_or_si256(raw, rhs.raw); // AVX2
393 }
394
395 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
396 return _mm256_xor_si256(raw, rhs.raw); // AVX2
397 }
398
399 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<<(int rhs) const {
400 return _mm256_slli_epi64(raw, rhs); // AVX2
401 }
402
403 template <int kShift>
404 [[nodiscard]] MOCHI_FORCE_INLINE static Simd ShiftRight(Simd a) {
405 static_assert(kShift >= 0 && kShift < 64, "Shift amount out-of-range");
406 if constexpr (kShift == 0) {
407 return a;
408 } else {
409#if MOCHI_ARCH_X64_AVX512
410 return _mm256_srai_epi64(a.raw, kShift); // AVX512VL
411#else
412 auto shifted = _mm256_srli_epi64(a.raw, kShift); // AVX2
413 auto signMask = _mm256_cmpgt_epi64(_mm256_setzero_si256(), a.raw); // AVX2
414 auto signFill = _mm256_slli_epi64(signMask, 64 - kShift); // AVX2
415 return _mm256_or_si256(shifted, signFill); // AVX2
416#endif
417 }
418 }
419
420 private:
421 // Integer mask with the most significant bit of each byte in the vector
422 [[nodiscard]] static MOCHI_FORCE_INLINE int GetMSBitMask(Simd a) {
423 return _mm256_movemask_epi8(a.raw); // AVX2
424 }
425};
426
427} // namespace superdex
428
429#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<<(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_UNLIKELY
#define MOCHI_FORCE_INLINE
Simd< T, 2 > Shuffle(Simd< T, 2 > a)
Definition simd_inl.h:273
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
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
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)