SuperDex Physics C++ API
Loading...
Searching...
No Matches
x64_simd_float_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_AVX2
22
23namespace superdex {
24
25/***********************************************************************************************
26 Simd<float, 8>
27*/
28template <>
29class Simd<float, 8> {
30 public:
31 MOCHI_NATIVE_SIMD_IMPL_BOILERPLATE(float, 8, __m256);
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 : raw(_mm256_set_ps(h, g, f, e, d, c, b, a)) {} // AVX
42
43 template <class U, MOCHI_REQUIRES_NON_BOOL_SCALAR(U, Scalar)>
44 Simd(U a) : raw(_mm256_set1_ps(a)) {} // AVX
45
46 // Joint two Vec4f into a single Vec8f
47 Simd(Simd<float, 4> const& low, Simd<float, 4> const& high)
48 : raw(_mm256_set_m128(high.raw, low.raw)) {} // AVX
49
50 template <int i>
51 [[nodiscard]] static MOCHI_FORCE_INLINE float Get(Simd v) {
52 static_assert(i >= 0 && i < 8, "Index out of range");
53 return v[i];
54 }
55
56 [[nodiscard]] MOCHI_FORCE_INLINE Scalar operator[](int i) const {
57 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
58#if MOCHI_COMPILER_MSVC
59 return raw.m256_f32[i];
60#else
61 return raw[i];
62#endif
63 }
64
65 template <int iHalf>
66 [[nodiscard]] static MOCHI_FORCE_INLINE Simd<float, 4> GetHalf(Simd a) {
67 static_assert(iHalf == 0 || iHalf == 1);
68 return _mm256_extractf128_ps(a.raw, iHalf); // AVX
69 }
70
71 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, int i, Scalar value) {
72 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
73#if MOCHI_COMPILER_MSVC
74 auto result = v;
75 result.raw.m256_f32[i] = value;
76 return result;
77#else
78 static constexpr __m256i kMasks[] = {
79 // clang-format off
80 {static_cast<long long>(0x00000000FFFFFFFFLL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL)},
81 {static_cast<long long>(0xFFFFFFFF00000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL)},
82 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x00000000FFFFFFFFLL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL)},
83 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0xFFFFFFFF00000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL)},
84 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x00000000FFFFFFFFLL), static_cast<long long>(0x0000000000000000LL)},
85 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0xFFFFFFFF00000000LL), static_cast<long long>(0x0000000000000000LL)},
86 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x00000000FFFFFFFFLL)},
87 {static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0x0000000000000000LL), static_cast<long long>(0xFFFFFFFF00000000LL)}
88 }; // clang-format on
89 return _mm256_blendv_ps(v.raw, _mm256_set1_ps(value), _mm256_castsi256_ps(kMasks[i])); // AVX
90#endif
91 }
92
93 template <int i>
94 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Set(Simd v, Scalar value) {
95 MOCHI_ASSERT_VERBOSE(i >= 0 && i < kSize, "Index out of range");
96 return Set(v, i, value);
97 }
98
99 template <int N>
100 [[nodiscard]] static MOCHI_FORCE_INLINE bool AllTrue(Simd v) {
101 static_assert(N >= 1 && N <= kSize, "Unsupported N");
102 int mask = GetMSBitMask(v); // One bit for each byte in the vector
103 if constexpr (N == kSize) {
104 return mask == 0xFFFFFFFF;
105 } else {
106 int constexpr kNumBits = N * sizeof(Scalar);
107 auto constexpr kMustBeSet = (1UL << kNumBits) - 1;
108 return (mask & kMustBeSet) == kMustBeSet;
109 }
110 }
111
112 template <int N>
113 [[nodiscard]] static MOCHI_FORCE_INLINE bool AnyTrue(Simd v) {
114 static_assert(N >= 1 && N <= kSize, "Unsupported N");
115 int mask = GetMSBitMask(v); // One bit for each byte in the vector
116 if constexpr (N == kSize) {
117 return mask != 0;
118 } else {
119 int constexpr kNumBits = N * sizeof(Scalar);
120 auto constexpr kMayBeSet = (1UL << kNumBits) - 1;
121 return (mask & kMayBeSet) != 0;
122 }
123 }
124
125 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Scalar const* p) {
126 return _mm256_broadcast_ss(p); // AVX
127 }
128
129 template <int i>
130 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Broadcast(Simd v) {
131 return Simd{Get<i>(v)}; // TODO: There is probably a faster way
132 }
133
134 template <int N = kSize>
135 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load([[maybe_unused]] Scalar const* ptr) {
136 static_assert(N >= 0 && N <= 8);
137 if constexpr (N == 0) {
138 return Simd::Zero();
139 } else if constexpr (N == 1) {
140 return Simd{*ptr, 0.0f};
141 } else if constexpr (N == 2) {
142 __m256i mask = _mm256_set_epi32(0, 0, 0, 0, 0, 0, -1, -1); // AVX
143 return _mm256_maskload_ps(ptr, mask); // AVX2
144 } else if constexpr (N == 3) {
145 __m256i mask = _mm256_set_epi32(0, 0, 0, 0, 0, -1, -1, -1); // AVX
146 return _mm256_maskload_ps(ptr, mask); // AVX2
147 } else if constexpr (N == 4) {
148 __m256i mask = _mm256_set_epi32(0, 0, 0, 0, -1, -1, -1, -1); // AVX
149 return _mm256_maskload_ps(ptr, mask); // AVX2
150 } else if constexpr (N == 5) {
151 __m256i mask = _mm256_set_epi32(0, 0, 0, -1, -1, -1, -1, -1); // AVX
152 return _mm256_maskload_ps(ptr, mask); // AVX2
153 } else if constexpr (N == 6) {
154 __m256i mask = _mm256_set_epi32(0, 0, -1, -1, -1, -1, -1, -1); // AVX
155 return _mm256_maskload_ps(ptr, mask); // AVX2
156 } else if constexpr (N == 7) {
157 __m256i mask = _mm256_set_epi32(0, -1, -1, -1, -1, -1, -1, -1); // AVX
158 return _mm256_maskload_ps(ptr, mask); // AVX2
159 } else if constexpr (N == 8) {
160 return _mm256_loadu_ps(ptr); // AVX
161 }
162 }
163
164 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Load(Scalar const* ptr, int n) {
165 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
166#if MOCHI_ARCH_X64_AVX512
167 return _mm256_maskz_loadu_ps(x64_simd::kLaneMasksS8[n], ptr); // AVX512VL
168#else
169 return _mm256_maskload_ps(ptr, x64_simd::kLoadMasksS8[n]); // AVX
170#endif
171 }
172
173 [[nodiscard]] static Simd LoadIndexed(Scalar const* ptr, Simd<int, 8> const& indices) {
174 return _mm256_i32gather_ps(ptr, indices.raw, sizeof(float)); // AVX2
175 }
176
177 template <int N = kSize>
178 static MOCHI_FORCE_INLINE void Store([[maybe_unused]] Scalar* ptr, [[maybe_unused]] Simd v) {
179 static_assert(N >= 0 && N <= kSize);
180 if constexpr (N == 0) {
181 } else if constexpr (N < kSize) {
182 // About 3X faster than a masked store on older AMD CPUs. About the same on others.
183 memcpy(ptr, &v, sizeof(Scalar) * N);
184 } else {
185 _mm256_storeu_ps(ptr, v.raw); // AVX
186 }
187 }
188
189 template <int kTupleCount = kSize>
190 MOCHI_FORCE_INLINE static void
191 LoadTransposed(Scalar const* ptr, Simd& out0, Simd& out1, Simd& out2) {
192 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
193 constexpr int kCount0 = Clamp(kTupleCount * 3 - 0, 0, 8);
194 constexpr int kCount1 = Clamp(kTupleCount * 3 - 8, 0, 8);
195 constexpr int kCount2 = Clamp(kTupleCount * 3 - 16, 0, 8);
196
197 // [ 0, 1, 2, 3, 4, 5, 6, 7]
198 auto a = Simd::Load<kCount0>(ptr).raw;
199 // [ 8, 9, 10, 11, 12, 13, 14, 15]
200 auto b = Simd::Load<kCount1>(kCount1 == 0 ? ptr : ptr + 8).raw;
201 // [16, 17, 18, 19, 20, 21, 22, 23]
202 auto c = Simd::Load<kCount2>(kCount2 == 0 ? ptr : ptr + 16).raw;
203
204 auto d = _mm256_blend_ps(a, b, 0b10010010); // [0,9,_,3, 12,_,6,15]
205 auto e = _mm256_blend_ps(d, c, 0b00100100); // [0,9,18,3, 12,21,6,15]
206 auto f = _mm256_permute2f128_ps(e, e, 0x01); // [12,21,6,15, 0,9,18,3]
207 auto g = _mm256_blend_ps(e, f, 0b01000100); // [0,9,6,3, 12,21,18,15]
208 out0.raw = _mm256_shuffle_ps(g, g, _MM_SHUFFLE(1, 2, 3, 0)); // [0,3,6,9, 12,15,18,21]
209
210 d = _mm256_blend_ps(a, b, 0b00100100); // [_,1,10,_, 4,13,_,7]
211 e = _mm256_blend_ps(d, c, 0b01001001); // [16,1,10,19 4,13,22,7]
212 f = _mm256_permute2f128_ps(e, e, 0x01); // [4,13,22,7 16,1,10,19]
213 g = _mm256_blend_ps(e, f, 0b10011001); // [4,1,10,7, 16,13,22,19]
214 out1.raw = _mm256_shuffle_ps(g, g, _MM_SHUFFLE(2, 3, 0, 1)); // [1,4,7,10, 13,16,19,22]
215
216 d = _mm256_blend_ps(a, b, 0b01001001); // [8,_,2,11, _,5,14,_]
217 e = _mm256_blend_ps(d, c, 0b10010010); // [8,17,2,11, 20,5,14,23]
218 f = _mm256_permute2f128_ps(e, e, 0x01); // [20,5,14,23, 8,17,2,11]
219 g = _mm256_blend_ps(e, f, 0b00100010); // [8,5,2,11, 20,17,14,23]
220 out2.raw = _mm256_shuffle_ps(g, g, _MM_SHUFFLE(3, 0, 1, 2)); // [2,5,8,11, 14,17,20,23]
221 }
222
223 static MOCHI_FORCE_INLINE void Store(Scalar* ptr, Simd v, int n) {
224 MOCHI_ASSERT_VERBOSE(n >= 0 && n <= kSize, "Invalid size parameter");
225#if MOCHI_ARCH_X64_AVX512
226 _mm256_mask_storeu_ps(ptr, x64_simd::kLaneMasksS8[n], v.raw); // AVX512VL
227#else
228 // With AVX2, this is faster than masked store for a predictable value of n.
229 // It is much slower for a random value of n.
230 switch (n) { // clang-format off
231 case 1: Store<1>(ptr, v); break;
232 case 2: Store<2>(ptr, v); break;
233 case 3: Store<3>(ptr, v); break;
234 case 4: Store<4>(ptr, v); break;
235 case 5: Store<5>(ptr, v); break;
236 case 6: Store<6>(ptr, v); break;
237 case 7: Store<7>(ptr, v); break;
238 case 8: Store<8>(ptr, v); break;
239 MOCHI_UNLIKELY default: break;
240 } // clang-format on
241#endif
242 }
243
244 MOCHI_FORCE_INLINE static int StoreSelected(Scalar* ptr, Simd condition, Simd values) {
245 auto mask = _mm256_movemask_ps(condition.raw);
246 // Load 8 bytes from the table, then zero-exend to get the shuffle pattern.
247 auto const* tableRow =
248 reinterpret_cast<__m128i const*>(x64_simd::kStoreSelectedShuffleTableS8[mask]);
249 auto pattern = _mm256_cvtepu8_epi32(_mm_loadl_epi64(tableRow));
250 auto packed = _mm256_permutevar8x32_ps(values.raw, pattern);
251 _mm256_storeu_ps(ptr, packed);
252 return _mm_popcnt_u32(mask);
253 }
254
255 template <int kTupleCount = kSize>
256 MOCHI_FORCE_INLINE static void StoreTransposed(Scalar* ptr, Simd a, Simd b, Simd c) {
257 static_assert(kTupleCount >= 1 && kTupleCount <= kSize, "Invalid kTupleCount");
258 // a = [0,3,6,9, 12,15,18,21]
259 // b = [1,4,7,10, 13,16,19,22]
260 // c = [2,5,8,11, 14,17,20,23]
261 auto d = _mm256_shuffle_ps(a.raw, a.raw, _MM_SHUFFLE(1, 2, 3, 0)); // [0,9,6,3, 12,21,18,15]
262 auto e = _mm256_shuffle_ps(b.raw, b.raw, _MM_SHUFFLE(2, 3, 0, 1)); // [4,1,10,7, 16,13,22,19]
263 auto f = _mm256_shuffle_ps(c.raw, c.raw, _MM_SHUFFLE(3, 0, 1, 2)); // [8,5,2,11, 20,17,14,23]
264 auto g = _mm256_blend_ps(d, e, 0b00100010); // [0,1,_,3, 12,13,_,15]
265 g = _mm256_blend_ps(g, f, 0b01000100); // [0,1,2,3, 12,13,14,15]
266 auto h = _mm256_blend_ps(d, e, 0b10011001); // [4,_,6,7, 16,_,18,19]
267 h = _mm256_blend_ps(h, f, 0b00100010); // [4,5,6,7, 16,17,18,19]
268 h = _mm256_permute2f128_ps(h, h, 0x01); // [16,17,18,19, 4,5,6,7]
269 auto i = _mm256_blend_ps(d, e, 0b01000100); // [_,9,10,_, _,21,22,_]
270 i = _mm256_blend_ps(i, f, 0b10011001); // [8,9,10,11, 20,21,22,23]
271 d = _mm256_blend_ps(g, h, 0b11110000); // [0,1,2,3, 4,5,6,7]
272 e = _mm256_blend_ps(i, g, 0b11110000); // [8,9,10,11, 12,13,14,15]
273 f = _mm256_blend_ps(h, i, 0b11110000); // [16,17,18,19, 20,21,22,23]
274 constexpr int kCount0 = Clamp(kTupleCount * 3 - 0, 0, 8);
275 constexpr int kCount1 = Clamp(kTupleCount * 3 - 8, 0, 8);
276 constexpr int kCount2 = Clamp(kTupleCount * 3 - 16, 0, 8);
277 Simd::Store<kCount0>(ptr, d);
278 if constexpr (kCount1 > 0) {
279 Simd::Store<kCount1>(ptr + 8, e);
280 }
281 if constexpr (kCount2 > 0) {
282 Simd::Store<kCount2>(ptr + 16, f);
283 }
284 }
285
286 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Select(Simd mask, Simd a, Simd b) {
287 return _mm256_blendv_ps(b.raw, a.raw, mask.raw); // AVX
288 }
289
290 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sqrt(Simd v) {
291 return _mm256_sqrt_ps(v.raw); // AVX
292 }
293
294 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpApprox(Simd v) {
295 return _mm256_rcp_ps(v.raw); // AVX
296 }
297
298 [[nodiscard]] static MOCHI_FORCE_INLINE Simd RcpSqrtApprox(Simd v) {
299 return _mm256_rsqrt_ps(v.raw); // AVX
300 }
301
302 // Broadcast the value -0.0. Use this in bitwise operations to affect just the sign bit.
303 [[nodiscard]] static MOCHI_FORCE_INLINE Simd SignBitMask() {
304 return _mm256_castsi256_ps(_mm256_set1_epi32(0x80000000)); // AVX, AVX
305 }
306
307 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Abs(Simd v) {
308 return _mm256_andnot_ps(SignBitMask().raw, v.raw); // AVX
309 }
310
311 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Min(Simd a, Simd b) {
312 return _mm256_min_ps(a.raw, b.raw); // AVX
313 }
314
315 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Max(Simd a, Simd b) {
316 return _mm256_max_ps(a.raw, b.raw); // AVX
317 }
318
319 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Floor(Simd a) {
320 return _mm256_floor_ps(a.raw); // AVX
321 }
322
323 [[nodiscard]] static MOCHI_FORCE_INLINE Simd FastRound(Simd v) {
324 return _mm256_round_ps(v.raw, _MM_FROUND_TO_NEAREST_INT); // AVX
325 }
326
327#if MOCHI_ARCH_X64_SVML
328 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Cos(Simd a) {
329 return _mm256_cos_ps(a.raw); // AVX
330 }
331
332 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Sin(Simd a) {
333 return _mm256_sin_ps(a.raw); // AVX
334 }
335
336 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tan(Simd a) {
337 return _mm256_tan_ps(a.raw); // AVX
338 }
339
340 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ACos(Simd a) {
341 return _mm256_acos_ps(a.raw); // AVX
342 }
343
344 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ASin(Simd a) {
345 return _mm256_asin_ps(a.raw); // AVX
346 }
347
348 [[nodiscard]] static MOCHI_FORCE_INLINE Simd ATan(Simd a) {
349 return _mm256_atan_ps(a.raw); // AVX
350 }
351
352 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Exp(Simd a) {
353 return _mm256_exp_ps(a.raw); // AVX
354 }
355
356 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Ln(Simd a) {
357 return _mm256_log_ps(a.raw); // AVX
358 }
359
360 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Tanh(Simd a) {
361 return _mm256_tanh_ps(a.raw); // AVX
362 }
363#endif // MOCHI_ARCH_X64_SVML
364
365 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulAdd(Simd a, Simd b, Simd c) {
366#if MOCHI_ARCH_X64_FMA
367 return _mm256_fmadd_ps(a.raw, b.raw, c.raw); // FMA
368#else
369 return (a * b) + c;
370#endif
371 }
372
373 [[nodiscard]] static MOCHI_FORCE_INLINE Simd MulSub(Simd a, Simd b, Simd c) {
374#if MOCHI_ARCH_X64_FMA
375 return _mm256_fmsub_ps(a.raw, b.raw, c.raw); // FMA
376#else
377 return (a * b) - c;
378#endif
379 }
380
381 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulAdd(Simd a, Simd b, Simd c) {
382#if MOCHI_ARCH_X64_FMA
383 return _mm256_fnmadd_ps(a.raw, b.raw, c.raw); // FMA
384#else
385 return -(a * b) + c;
386#endif
387 }
388
389 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NegMulSub(Simd a, Simd b, Simd c) {
390#if MOCHI_ARCH_X64_FMA
391 return _mm256_fnmsub_ps(a.raw, b.raw, c.raw); // FMA
392#else
393 return -(a * b) - c;
394#endif
395 }
396
397 template <int N = 8>
398 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMin(Simd a) {
399 static_assert(N >= 2 && N <= 8, "Unsupported N");
400 using HalfT = Simd<Scalar, 4>;
401 if constexpr (N >= 2 && N <= 4) {
402 return HalfT::HMin<N>(GetHalf<0>(a));
403 } else {
404 auto lo = GetHalf<0>(a);
405 auto hi = GetHalf<1>(a);
406 if constexpr (N != 8) {
407 // Set the hi values we don't want to compare to infinity
408 auto inf = HalfT{std::numeric_limits<Scalar>::infinity()};
409 hi = HalfT::Blend<1, N >= 6, N >= 7, 0>(inf, hi);
410 }
411 return HalfT::HMin<4>(HalfT::Min(lo, hi));
412 }
413 }
414
415 template <int N = 8>
416 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HMax(Simd a) {
417 static_assert(N >= 2 && N <= 8, "Unsupported N");
418 using HalfT = Simd<Scalar, 4>;
419 if constexpr (N >= 2 && N <= 4) {
420 return HalfT::HMax<N>(GetHalf<0>(a));
421 } else {
422 auto lo = GetHalf<0>(a);
423 auto hi = GetHalf<1>(a);
424 if constexpr (N != 8) {
425 // Set the hi values we don't want to compare to -infinity
426 auto inf = HalfT{std::numeric_limits<Scalar>::infinity()};
427 hi = HalfT::Blend<1, N >= 6, N >= 7, 0>(-inf, hi);
428 }
429 return HalfT::HMax<4>(HalfT::Max(lo, hi));
430 }
431 }
432
433 template <int N>
434 [[nodiscard]] static MOCHI_FORCE_INLINE Scalar HSum(Simd a) {
435 static_assert(N >= 2 && N <= 8, "Unsupported N");
436 using HalfT = Simd<float, 4>;
437 if constexpr (N >= 2 && N <= 4) {
438 return HalfT::HSum<N>(GetHalf<0>(a));
439 } else if constexpr (N == 5) {
440 auto tmp = GetHalf<0>(a) + GetHalf<1>(a);
441 return HalfT::Get<0>(tmp) + Get<1>(a) + Get<2>(a) + Get<3>(a);
442 } else if constexpr (N == 6) {
443 auto tmp = GetHalf<0>(a) + GetHalf<1>(a);
444 return HalfT::Get<0>(tmp) + HalfT::Get<1>(tmp) + Get<2>(a) + Get<3>(a);
445 } else if constexpr (N == 7) {
446 auto tmp = GetHalf<0>(a) + GetHalf<1>(a);
447 return HalfT::Get<0>(tmp) + HalfT::Get<1>(tmp) + HalfT::Get<2>(tmp) + Get<3>(a);
448 } else if constexpr (N == 8) {
449 // PERF NOTE: Alternatively _mm256_dp_ps could be used to compute the dot product with
450 // Simd{1.0f}. In comparison, this implementation takes 2 extra instructions, but it had ~29%
451 // higher throughput and ~25% lower latency, when benchmarked on an AMD CPU.
452 return HalfT::HSum<4>(GetHalf<0>(a) + GetHalf<1>(a));
453 }
454 }
455
456 template <int N>
457 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Dot(Simd a, Simd b) {
458 static_assert(N == 8, "smaller dot products not yet supported on Vec8f");
459#if MOCHI_COMPILER_CLANG
460 auto ab_half = _mm256_dp_ps(a.raw, b.raw, -1); // AVX
461#else
462 auto ab_half = _mm256_dp_ps(a.raw, b.raw, 0xFF); // AVX
463#endif
464 auto lo_ab = _mm256_extractf128_ps(ab_half, 0); // AVX
465 auto hi_ab = _mm256_extractf128_ps(ab_half, 1); // AVX
466 auto my_dot = _mm_add_ps(lo_ab, hi_ab); // SSE
467 return _mm256_set_m128(my_dot, my_dot); // AVX
468 }
469
470 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<(Simd rhs) const {
471 return _mm256_cmp_ps(this->raw, rhs.raw, _CMP_LT_OQ); // AVX
472 }
473
474 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>(Simd rhs) const {
475 return _mm256_cmp_ps(this->raw, rhs.raw, _CMP_GT_OQ); // AVX
476 }
477
478 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator<=(Simd rhs) const {
479 return _mm256_cmp_ps(this->raw, rhs.raw, _CMP_LE_OQ); // AVX
480 }
481
482 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator>=(Simd rhs) const {
483 return _mm256_cmp_ps(this->raw, rhs.raw, _CMP_GE_OQ); // AVX
484 }
485
486 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Equal(Simd a, Simd b) {
487 return _mm256_cmp_ps(a.raw, b.raw, _CMP_EQ_OQ); // AVX
488 }
489
490 [[nodiscard]] static MOCHI_FORCE_INLINE Simd NotEqual(Simd a, Simd b) {
491 return _mm256_cmp_ps(a.raw, b.raw, _CMP_NEQ_UQ); // AVX
492 }
493
494 [[nodiscard]] static MOCHI_FORCE_INLINE Simd Zero() {
495 return _mm256_setzero_ps(); // AVX
496 }
497
498 [[nodiscard]] MOCHI_FORCE_INLINE bool operator==(Simd rhs) const {
499 auto mask = GetMSBitMask(Equal(*this, rhs));
500 return mask == 0xFFFFFFFF; // All values equal
501 }
502
503 [[nodiscard]] MOCHI_FORCE_INLINE bool operator!=(Simd rhs) const {
504 auto mask = GetMSBitMask(NotEqual(*this, rhs));
505 return mask != 0; // All values equal
506 }
507
508 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator~() const {
509 auto ones = _mm256_castsi256_ps(_mm256_set1_epi32(-1)); // AVX, AVX
510 return _mm256_xor_ps(raw, ones); // AVX
511 }
512
513 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-() const {
514 return Zero() - *this;
515 }
516
517 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator+(Simd rhs) const {
518 return _mm256_add_ps(raw, rhs.raw); // AVX
519 }
520
521 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator-(Simd rhs) const {
522 return _mm256_sub_ps(raw, rhs.raw); // AVX
523 }
524
525 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator*(Simd rhs) const {
526 return _mm256_mul_ps(raw, rhs.raw); // AVX
527 }
528
529 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator/(Simd rhs) const {
530 return _mm256_div_ps(raw, rhs.raw); // AVX
531 }
532
533 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator&(Simd rhs) const {
534 return _mm256_and_ps(raw, rhs.raw); // AVX
535 }
536
537 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator|(Simd rhs) const {
538 return _mm256_or_ps(raw, rhs.raw); // AVX
539 }
540
541 [[nodiscard]] MOCHI_FORCE_INLINE Simd operator^(Simd rhs) const {
542 return _mm256_xor_ps(raw, rhs.raw); // AVX
543 }
544
545 private:
546 // Integer mask with with the most significant bit of each byte in the vector
547 [[nodiscard]] static MOCHI_FORCE_INLINE int GetMSBitMask(Simd a) {
548 return _mm256_movemask_epi8(_mm256_castps_si256(a.raw)); // AVX, AVX
549 }
550};
551
552} // namespace superdex
553
554#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)
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)