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