Eigen  5.0.1
 
Loading...
Searching...
No Matches
Reductions.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2025 Charlie Schlosser <cs.schlosser@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_REDUCTIONS_AVX_H
12#define EIGEN_REDUCTIONS_AVX_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8i -- -- -- -- -- -- -- -- -- -- -- -- */
22
23template <>
24EIGEN_STRONG_INLINE int predux(const Packet8i& a) {
25 Packet4i lo = _mm256_castsi256_si128(a);
26 Packet4i hi = _mm256_extractf128_si256(a, 1);
27 return predux(padd(lo, hi));
28}
29
30template <>
31EIGEN_STRONG_INLINE int predux_mul(const Packet8i& a) {
32 Packet4i lo = _mm256_castsi256_si128(a);
33 Packet4i hi = _mm256_extractf128_si256(a, 1);
34 return predux_mul(pmul(lo, hi));
35}
36
37template <>
38EIGEN_STRONG_INLINE int predux_min(const Packet8i& a) {
39 Packet4i lo = _mm256_castsi256_si128(a);
40 Packet4i hi = _mm256_extractf128_si256(a, 1);
41 return predux_min(pmin(lo, hi));
42}
43
44template <>
45EIGEN_STRONG_INLINE int predux_max(const Packet8i& a) {
46 Packet4i lo = _mm256_castsi256_si128(a);
47 Packet4i hi = _mm256_extractf128_si256(a, 1);
48 return predux_max(pmax(lo, hi));
49}
50
51template <>
52EIGEN_STRONG_INLINE bool predux_any(const Packet8i& a) {
53#ifdef EIGEN_VECTORIZE_AVX2
54 return _mm256_movemask_epi8(a) != 0x0;
55#else
56 return _mm256_movemask_ps(_mm256_castsi256_ps(a)) != 0x0;
57#endif
58}
59
60/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8ui -- -- -- -- -- -- -- -- -- -- -- -- */
61
62template <>
63EIGEN_STRONG_INLINE uint32_t predux(const Packet8ui& a) {
64 Packet4ui lo = _mm256_castsi256_si128(a);
65 Packet4ui hi = _mm256_extractf128_si256(a, 1);
66 return predux(padd(lo, hi));
67}
68
69template <>
70EIGEN_STRONG_INLINE uint32_t predux_mul(const Packet8ui& a) {
71 Packet4ui lo = _mm256_castsi256_si128(a);
72 Packet4ui hi = _mm256_extractf128_si256(a, 1);
73 return predux_mul(pmul(lo, hi));
74}
75
76template <>
77EIGEN_STRONG_INLINE uint32_t predux_min(const Packet8ui& a) {
78 Packet4ui lo = _mm256_castsi256_si128(a);
79 Packet4ui hi = _mm256_extractf128_si256(a, 1);
80 return predux_min(pmin(lo, hi));
81}
82
83template <>
84EIGEN_STRONG_INLINE uint32_t predux_max(const Packet8ui& a) {
85 Packet4ui lo = _mm256_castsi256_si128(a);
86 Packet4ui hi = _mm256_extractf128_si256(a, 1);
87 return predux_max(pmax(lo, hi));
88}
89
90template <>
91EIGEN_STRONG_INLINE bool predux_any(const Packet8ui& a) {
92#ifdef EIGEN_VECTORIZE_AVX2
93 return _mm256_movemask_epi8(a) != 0x0;
94#else
95 return _mm256_movemask_ps(_mm256_castsi256_ps(a)) != 0x0;
96#endif
97}
98
99#ifdef EIGEN_VECTORIZE_AVX2
100
101/* -- -- -- -- -- -- -- -- -- -- -- -- Packet4l -- -- -- -- -- -- -- -- -- -- -- -- */
102
103template <>
104EIGEN_STRONG_INLINE int64_t predux(const Packet4l& a) {
105 Packet2l lo = _mm256_castsi256_si128(a);
106 Packet2l hi = _mm256_extractf128_si256(a, 1);
107 return predux(padd(lo, hi));
108}
109
110template <>
111EIGEN_STRONG_INLINE bool predux_any(const Packet4l& a) {
112 return _mm256_movemask_pd(_mm256_castsi256_pd(a)) != 0x0;
113}
114
115/* -- -- -- -- -- -- -- -- -- -- -- -- Packet4ul -- -- -- -- -- -- -- -- -- -- -- -- */
116
117template <>
118EIGEN_STRONG_INLINE uint64_t predux(const Packet4ul& a) {
119 return static_cast<uint64_t>(predux(Packet4l(a)));
120}
121
122template <>
123EIGEN_STRONG_INLINE bool predux_any(const Packet4ul& a) {
124 return _mm256_movemask_pd(_mm256_castsi256_pd(a)) != 0x0;
125}
126
127#endif
128
129/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8f -- -- -- -- -- -- -- -- -- -- -- -- */
130
131template <>
132EIGEN_STRONG_INLINE float predux(const Packet8f& a) {
133 Packet4f lo = _mm256_castps256_ps128(a);
134 Packet4f hi = _mm256_extractf128_ps(a, 1);
135 return predux(padd(lo, hi));
136}
137
138template <>
139EIGEN_STRONG_INLINE float predux_mul(const Packet8f& a) {
140 Packet4f lo = _mm256_castps256_ps128(a);
141 Packet4f hi = _mm256_extractf128_ps(a, 1);
142 return predux_mul(pmul(lo, hi));
143}
144
145template <>
146EIGEN_STRONG_INLINE float predux_min(const Packet8f& a) {
147 Packet4f lo = _mm256_castps256_ps128(a);
148 Packet4f hi = _mm256_extractf128_ps(a, 1);
149 return predux_min(pmin(lo, hi));
150}
151
152template <>
153EIGEN_STRONG_INLINE float predux_min<PropagateNumbers>(const Packet8f& a) {
154 Packet4f lo = _mm256_castps256_ps128(a);
155 Packet4f hi = _mm256_extractf128_ps(a, 1);
156 return predux_min<PropagateNumbers>(pmin<PropagateNumbers>(lo, hi));
157}
158
159template <>
160EIGEN_STRONG_INLINE float predux_min<PropagateNaN>(const Packet8f& a) {
161 Packet4f lo = _mm256_castps256_ps128(a);
162 Packet4f hi = _mm256_extractf128_ps(a, 1);
163 return predux_min<PropagateNaN>(pmin<PropagateNaN>(lo, hi));
164}
165
166template <>
167EIGEN_STRONG_INLINE float predux_max(const Packet8f& a) {
168 Packet4f lo = _mm256_castps256_ps128(a);
169 Packet4f hi = _mm256_extractf128_ps(a, 1);
170 return predux_max(pmax(lo, hi));
171}
172
173template <>
174EIGEN_STRONG_INLINE float predux_max<PropagateNumbers>(const Packet8f& a) {
175 Packet4f lo = _mm256_castps256_ps128(a);
176 Packet4f hi = _mm256_extractf128_ps(a, 1);
177 return predux_max<PropagateNumbers>(pmax<PropagateNumbers>(lo, hi));
178}
179
180template <>
181EIGEN_STRONG_INLINE float predux_max<PropagateNaN>(const Packet8f& a) {
182 Packet4f lo = _mm256_castps256_ps128(a);
183 Packet4f hi = _mm256_extractf128_ps(a, 1);
184 return predux_max<PropagateNaN>(pmax<PropagateNaN>(lo, hi));
185}
186
187template <>
188EIGEN_STRONG_INLINE bool predux_any(const Packet8f& a) {
189 return _mm256_movemask_ps(a) != 0x0;
190}
191
192template <>
193EIGEN_STRONG_INLINE Index predux_count(const Packet8f& a) {
194 const unsigned int mask =
195 static_cast<unsigned int>(_mm256_movemask_ps(_mm256_cmp_ps(a, _mm256_setzero_ps(), _CMP_NEQ_UQ)));
196 return Index(popcount(mask));
197}
198
199/* -- -- -- -- -- -- -- -- -- -- -- -- Packet4d -- -- -- -- -- -- -- -- -- -- -- -- */
200
201template <>
202EIGEN_STRONG_INLINE double predux(const Packet4d& a) {
203 Packet2d lo = _mm256_castpd256_pd128(a);
204 Packet2d hi = _mm256_extractf128_pd(a, 1);
205 return predux(padd(lo, hi));
206}
207
208template <>
209EIGEN_STRONG_INLINE double predux_mul(const Packet4d& a) {
210 Packet2d lo = _mm256_castpd256_pd128(a);
211 Packet2d hi = _mm256_extractf128_pd(a, 1);
212 return predux_mul(pmul(lo, hi));
213}
214
215template <>
216EIGEN_STRONG_INLINE double predux_min(const Packet4d& a) {
217 Packet2d lo = _mm256_castpd256_pd128(a);
218 Packet2d hi = _mm256_extractf128_pd(a, 1);
219 return predux_min(pmin(lo, hi));
220}
221
222template <>
223EIGEN_STRONG_INLINE double predux_min<PropagateNumbers>(const Packet4d& a) {
224 Packet2d lo = _mm256_castpd256_pd128(a);
225 Packet2d hi = _mm256_extractf128_pd(a, 1);
226 return predux_min<PropagateNumbers>(pmin<PropagateNumbers>(lo, hi));
227}
228
229template <>
230EIGEN_STRONG_INLINE double predux_min<PropagateNaN>(const Packet4d& a) {
231 Packet2d lo = _mm256_castpd256_pd128(a);
232 Packet2d hi = _mm256_extractf128_pd(a, 1);
233 return predux_min<PropagateNaN>(pmin<PropagateNaN>(lo, hi));
234}
235
236template <>
237EIGEN_STRONG_INLINE double predux_max(const Packet4d& a) {
238 Packet2d lo = _mm256_castpd256_pd128(a);
239 Packet2d hi = _mm256_extractf128_pd(a, 1);
240 return predux_max(pmax(lo, hi));
241}
242
243template <>
244EIGEN_STRONG_INLINE double predux_max<PropagateNumbers>(const Packet4d& a) {
245 Packet2d lo = _mm256_castpd256_pd128(a);
246 Packet2d hi = _mm256_extractf128_pd(a, 1);
247 return predux_max<PropagateNumbers>(pmax<PropagateNumbers>(lo, hi));
248}
249
250template <>
251EIGEN_STRONG_INLINE double predux_max<PropagateNaN>(const Packet4d& a) {
252 Packet2d lo = _mm256_castpd256_pd128(a);
253 Packet2d hi = _mm256_extractf128_pd(a, 1);
254 return predux_max<PropagateNaN>(pmax<PropagateNaN>(lo, hi));
255}
256
257template <>
258EIGEN_STRONG_INLINE bool predux_any(const Packet4d& a) {
259 return _mm256_movemask_pd(a) != 0x0;
260}
261
262template <>
263EIGEN_STRONG_INLINE Index predux_count(const Packet4d& a) {
264 const unsigned int mask =
265 static_cast<unsigned int>(_mm256_movemask_pd(_mm256_cmp_pd(a, _mm256_setzero_pd(), _CMP_NEQ_UQ)));
266 return Index(popcount(mask));
267}
268
269/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8h -- -- -- -- -- -- -- -- -- -- -- -- */
270#ifndef EIGEN_VECTORIZE_AVX512FP16
271
272template <>
273EIGEN_STRONG_INLINE half predux(const Packet8h& a) {
274 return static_cast<half>(predux(half2float(a)));
275}
276
277template <>
278EIGEN_STRONG_INLINE half predux_mul(const Packet8h& a) {
279 return static_cast<half>(predux_mul(half2float(a)));
280}
281
282template <>
283EIGEN_STRONG_INLINE half predux_min(const Packet8h& a) {
284 return static_cast<half>(predux_min(half2float(a)));
285}
286
287template <>
288EIGEN_STRONG_INLINE half predux_min<PropagateNumbers>(const Packet8h& a) {
289 return static_cast<half>(predux_min<PropagateNumbers>(half2float(a)));
290}
291
292template <>
293EIGEN_STRONG_INLINE half predux_min<PropagateNaN>(const Packet8h& a) {
294 return static_cast<half>(predux_min<PropagateNaN>(half2float(a)));
295}
296
297template <>
298EIGEN_STRONG_INLINE half predux_max(const Packet8h& a) {
299 return static_cast<half>(predux_max(half2float(a)));
300}
301
302template <>
303EIGEN_STRONG_INLINE half predux_max<PropagateNumbers>(const Packet8h& a) {
304 return static_cast<half>(predux_max<PropagateNumbers>(half2float(a)));
305}
306
307template <>
308EIGEN_STRONG_INLINE half predux_max<PropagateNaN>(const Packet8h& a) {
309 return static_cast<half>(predux_max<PropagateNaN>(half2float(a)));
310}
311
312template <>
313EIGEN_STRONG_INLINE bool predux_any(const Packet8h& a) {
314 return _mm_movemask_epi8(a) != 0;
315}
316#endif // EIGEN_VECTORIZE_AVX512FP16
317
318/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8bf -- -- -- -- -- -- -- -- -- -- -- -- */
319
320template <>
321EIGEN_STRONG_INLINE bfloat16 predux(const Packet8bf& a) {
322 return static_cast<bfloat16>(predux<Packet8f>(Bf16ToF32(a)));
323}
324
325template <>
326EIGEN_STRONG_INLINE bfloat16 predux_mul(const Packet8bf& a) {
327 return static_cast<bfloat16>(predux_mul<Packet8f>(Bf16ToF32(a)));
328}
329
330template <>
331EIGEN_STRONG_INLINE bfloat16 predux_min(const Packet8bf& a) {
332 return static_cast<bfloat16>(predux_min(Bf16ToF32(a)));
333}
334
335template <>
336EIGEN_STRONG_INLINE bfloat16 predux_min<PropagateNumbers>(const Packet8bf& a) {
337 return static_cast<bfloat16>(predux_min<PropagateNumbers>(Bf16ToF32(a)));
338}
339
340template <>
341EIGEN_STRONG_INLINE bfloat16 predux_min<PropagateNaN>(const Packet8bf& a) {
342 return static_cast<bfloat16>(predux_min<PropagateNaN>(Bf16ToF32(a)));
343}
344
345template <>
346EIGEN_STRONG_INLINE bfloat16 predux_max(const Packet8bf& a) {
347 return static_cast<bfloat16>(predux_max<Packet8f>(Bf16ToF32(a)));
348}
349
350template <>
351EIGEN_STRONG_INLINE bfloat16 predux_max<PropagateNumbers>(const Packet8bf& a) {
352 return static_cast<bfloat16>(predux_max<PropagateNumbers>(Bf16ToF32(a)));
353}
354
355template <>
356EIGEN_STRONG_INLINE bfloat16 predux_max<PropagateNaN>(const Packet8bf& a) {
357 return static_cast<bfloat16>(predux_max<PropagateNaN>(Bf16ToF32(a)));
358}
359
360template <>
361EIGEN_STRONG_INLINE bool predux_any(const Packet8bf& a) {
362 return _mm_movemask_epi8(a) != 0;
363}
364
365} // end namespace internal
366} // end namespace Eigen
367
368#endif // EIGEN_REDUCTIONS_AVX_H