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_AVX512_H
12#define EIGEN_REDUCTIONS_AVX512_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21// Preserve this backend's any-bit semantics by testing 32-bit chunks for every scalar width.
22EIGEN_STRONG_INLINE bool avx512_predux_any(const Packet16i& bits) { return _mm512_test_epi32_mask(bits, bits) != 0; }
23
24/* -- -- -- -- -- -- -- -- -- -- -- -- Packet16i -- -- -- -- -- -- -- -- -- -- -- -- */
25
26template <>
27EIGEN_STRONG_INLINE int predux(const Packet16i& a) {
28 return _mm512_reduce_add_epi32(a);
29}
30
31template <>
32EIGEN_STRONG_INLINE int predux_mul(const Packet16i& a) {
33 return _mm512_reduce_mul_epi32(a);
34}
35
36template <>
37EIGEN_STRONG_INLINE int predux_min(const Packet16i& a) {
38 return _mm512_reduce_min_epi32(a);
39}
40
41template <>
42EIGEN_STRONG_INLINE int predux_max(const Packet16i& a) {
43 return _mm512_reduce_max_epi32(a);
44}
45
46template <>
47EIGEN_STRONG_INLINE bool predux_any(const Packet16i& a) {
48 return avx512_predux_any(a);
49}
50
51template <>
52EIGEN_STRONG_INLINE bool predux_all(const Packet16i& a) {
53 return _mm512_cmp_epi32_mask(a, _mm512_setzero_epi32(), _MM_CMPINT_EQ) == 0;
54}
55
56/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8l -- -- -- -- -- -- -- -- -- -- -- -- */
57
58template <>
59EIGEN_STRONG_INLINE int64_t predux(const Packet8l& a) {
60 return _mm512_reduce_add_epi64(a);
61}
62
63#if EIGEN_COMP_MSVC
64// MSVC's _mm512_reduce_mul_epi64 is borked, at least up to and including 1939.
65// alignas(64) int64_t data[] = { 1,1,-1,-1,1,-1,-1,-1 };
66// int64_t out = _mm512_reduce_mul_epi64(_mm512_load_epi64(data));
67// produces garbage: 4294967295. This occurs when the result should be negative.
68// Fall back to a manual approach:
69template <>
70EIGEN_STRONG_INLINE int64_t predux_mul(const Packet8l& a) {
71 Packet4l lane0 = _mm512_extracti64x4_epi64(a, 0);
72 Packet4l lane1 = _mm512_extracti64x4_epi64(a, 1);
73 return predux_mul(pmul(lane0, lane1));
74}
75#else
76template <>
77EIGEN_STRONG_INLINE int64_t predux_mul<Packet8l>(const Packet8l& a) {
78 return _mm512_reduce_mul_epi64(a);
79}
80#endif
81
82template <>
83EIGEN_STRONG_INLINE int64_t predux_min(const Packet8l& a) {
84 return _mm512_reduce_min_epi64(a);
85}
86
87template <>
88EIGEN_STRONG_INLINE int64_t predux_max(const Packet8l& a) {
89 return _mm512_reduce_max_epi64(a);
90}
91
92template <>
93EIGEN_STRONG_INLINE bool predux_any(const Packet8l& a) {
94 return avx512_predux_any(a);
95}
96
97template <>
98EIGEN_STRONG_INLINE bool predux_all(const Packet8l& a) {
99 return _mm512_cmp_epi64_mask(a, _mm512_setzero_si512(), _MM_CMPINT_EQ) == 0;
100}
101
102/* -- -- -- -- -- -- -- -- -- -- -- -- Packet16f -- -- -- -- -- -- -- -- -- -- -- -- */
103
104template <>
105EIGEN_STRONG_INLINE float predux(const Packet16f& a) {
106 return _mm512_reduce_add_ps(a);
107}
108
109template <>
110EIGEN_STRONG_INLINE float predux_mul(const Packet16f& a) {
111 return _mm512_reduce_mul_ps(a);
112}
113
114template <>
115EIGEN_STRONG_INLINE float predux_min(const Packet16f& a) {
116 return _mm512_reduce_min_ps(a);
117}
118
119template <>
120EIGEN_STRONG_INLINE float predux_min<PropagateNumbers>(const Packet16f& a) {
121 Packet8f lane0 = _mm512_extractf32x8_ps(a, 0);
122 Packet8f lane1 = _mm512_extractf32x8_ps(a, 1);
123 return predux_min<PropagateNumbers>(pmin<PropagateNumbers>(lane0, lane1));
124}
125
126template <>
127EIGEN_STRONG_INLINE float predux_min<PropagateNaN>(const Packet16f& a) {
128 Packet8f lane0 = _mm512_extractf32x8_ps(a, 0);
129 Packet8f lane1 = _mm512_extractf32x8_ps(a, 1);
130 return predux_min<PropagateNaN>(pmin<PropagateNaN>(lane0, lane1));
131}
132
133template <>
134EIGEN_STRONG_INLINE float predux_max(const Packet16f& a) {
135 return _mm512_reduce_max_ps(a);
136}
137
138template <>
139EIGEN_STRONG_INLINE float predux_max<PropagateNumbers>(const Packet16f& a) {
140 Packet8f lane0 = _mm512_extractf32x8_ps(a, 0);
141 Packet8f lane1 = _mm512_extractf32x8_ps(a, 1);
142 return predux_max<PropagateNumbers>(pmax<PropagateNumbers>(lane0, lane1));
143}
144
145template <>
146EIGEN_STRONG_INLINE float predux_max<PropagateNaN>(const Packet16f& a) {
147 Packet8f lane0 = _mm512_extractf32x8_ps(a, 0);
148 Packet8f lane1 = _mm512_extractf32x8_ps(a, 1);
149 return predux_max<PropagateNaN>(pmax<PropagateNaN>(lane0, lane1));
150}
151
152template <>
153EIGEN_STRONG_INLINE bool predux_any(const Packet16f& a) {
154 return avx512_predux_any(_mm512_castps_si512(a));
155}
156
157template <>
158EIGEN_STRONG_INLINE bool predux_all(const Packet16f& a) {
159 return _mm512_cmp_ps_mask(a, _mm512_setzero_ps(), _CMP_EQ_OQ) == 0;
160}
161
162template <>
163EIGEN_STRONG_INLINE Index predux_count(const Packet16f& a) {
164 return Index(popcount(static_cast<unsigned int>(_mm512_cmp_ps_mask(a, _mm512_setzero_ps(), _CMP_NEQ_UQ))));
165}
166
167/* -- -- -- -- -- -- -- -- -- -- -- -- Packet8d -- -- -- -- -- -- -- -- -- -- -- -- */
168
169template <>
170EIGEN_STRONG_INLINE double predux(const Packet8d& a) {
171 return _mm512_reduce_add_pd(a);
172}
173
174template <>
175EIGEN_STRONG_INLINE double predux_mul(const Packet8d& a) {
176 return _mm512_reduce_mul_pd(a);
177}
178
179template <>
180EIGEN_STRONG_INLINE double predux_min(const Packet8d& a) {
181 return _mm512_reduce_min_pd(a);
182}
183
184template <>
185EIGEN_STRONG_INLINE double predux_min<PropagateNumbers>(const Packet8d& a) {
186 Packet4d lane0 = _mm512_extractf64x4_pd(a, 0);
187 Packet4d lane1 = _mm512_extractf64x4_pd(a, 1);
188 return predux_min<PropagateNumbers>(pmin<PropagateNumbers>(lane0, lane1));
189}
190
191template <>
192EIGEN_STRONG_INLINE double predux_min<PropagateNaN>(const Packet8d& a) {
193 Packet4d lane0 = _mm512_extractf64x4_pd(a, 0);
194 Packet4d lane1 = _mm512_extractf64x4_pd(a, 1);
195 return predux_min<PropagateNaN>(pmin<PropagateNaN>(lane0, lane1));
196}
197
198template <>
199EIGEN_STRONG_INLINE double predux_max(const Packet8d& a) {
200 return _mm512_reduce_max_pd(a);
201}
202
203template <>
204EIGEN_STRONG_INLINE double predux_max<PropagateNumbers>(const Packet8d& a) {
205 Packet4d lane0 = _mm512_extractf64x4_pd(a, 0);
206 Packet4d lane1 = _mm512_extractf64x4_pd(a, 1);
207 return predux_max<PropagateNumbers>(pmax<PropagateNumbers>(lane0, lane1));
208}
209
210template <>
211EIGEN_STRONG_INLINE double predux_max<PropagateNaN>(const Packet8d& a) {
212 Packet4d lane0 = _mm512_extractf64x4_pd(a, 0);
213 Packet4d lane1 = _mm512_extractf64x4_pd(a, 1);
214 return predux_max<PropagateNaN>(pmax<PropagateNaN>(lane0, lane1));
215}
216
217template <>
218EIGEN_STRONG_INLINE bool predux_any(const Packet8d& a) {
219 return avx512_predux_any(_mm512_castpd_si512(a));
220}
221
222template <>
223EIGEN_STRONG_INLINE bool predux_all(const Packet8d& a) {
224 return _mm512_cmp_pd_mask(a, _mm512_setzero_pd(), _CMP_EQ_OQ) == 0;
225}
226
227template <>
228EIGEN_STRONG_INLINE Index predux_count(const Packet8d& a) {
229 return Index(popcount(static_cast<unsigned int>(_mm512_cmp_pd_mask(a, _mm512_setzero_pd(), _CMP_NEQ_UQ))));
230}
231
232// Count 16-bit floating-point lanes through integer bits so fast-math cannot treat NaN masks as zero.
233EIGEN_STRONG_INLINE Index predux_count_16bit(const __m256i& a) {
234#if defined(EIGEN_VECTORIZE_AVX512VL) && defined(__AVX512BW__)
235 // A cast rather than _cvtmask16_u32, whose nvc++ definition calls a __builtin_ia32_kmovw that nvc++ cannot link.
236 const unsigned int nonzero_lanes = static_cast<unsigned int>(_mm256_test_epi16_mask(a, _mm256_set1_epi16(0x7fff)));
237 return Index(popcount(nonzero_lanes));
238#else
239 const __m256i magnitude = _mm256_and_si256(a, _mm256_set1_epi16(0x7fff));
240 const __m256i zeros = _mm256_cmpeq_epi16(magnitude, _mm256_setzero_si256());
241 const unsigned int zero_bytes = static_cast<unsigned int>(_mm256_movemask_epi8(zeros));
242 return Index(16 - popcount(zero_bytes) / 2);
243#endif
244}
245
246#ifndef EIGEN_VECTORIZE_AVX512FP16
247/* -- -- -- -- -- -- -- -- -- -- -- -- Packet16h -- -- -- -- -- -- -- -- -- -- -- -- */
248
249template <>
250EIGEN_STRONG_INLINE half predux(const Packet16h& from) {
251 return half(predux(half2float(from)));
252}
253
254template <>
255EIGEN_STRONG_INLINE half predux_mul(const Packet16h& from) {
256 return half(predux_mul(half2float(from)));
257}
258
259template <>
260EIGEN_STRONG_INLINE half predux_min(const Packet16h& from) {
261 return half(predux_min(half2float(from)));
262}
263
264template <>
265EIGEN_STRONG_INLINE half predux_min<PropagateNumbers>(const Packet16h& from) {
266 return half(predux_min<PropagateNumbers>(half2float(from)));
267}
268
269template <>
270EIGEN_STRONG_INLINE half predux_min<PropagateNaN>(const Packet16h& from) {
271 return half(predux_min<PropagateNaN>(half2float(from)));
272}
273
274template <>
275EIGEN_STRONG_INLINE half predux_max(const Packet16h& from) {
276 return half(predux_max(half2float(from)));
277}
278
279template <>
280EIGEN_STRONG_INLINE half predux_max<PropagateNumbers>(const Packet16h& from) {
281 return half(predux_max<PropagateNumbers>(half2float(from)));
282}
283
284template <>
285EIGEN_STRONG_INLINE half predux_max<PropagateNaN>(const Packet16h& from) {
286 return half(predux_max<PropagateNaN>(half2float(from)));
287}
288
289template <>
290EIGEN_STRONG_INLINE bool predux_any(const Packet16h& a) {
291 return predux_any<Packet8i>(a.m_val);
292}
293
294template <>
295EIGEN_STRONG_INLINE bool predux_all(const Packet16h& a) {
296 return predux_count_16bit(a.m_val) == 16;
297}
298
299template <>
300EIGEN_STRONG_INLINE Index predux_count(const Packet16h& a) {
301 return predux_count_16bit(a.m_val);
302}
303#endif
304
305/* -- -- -- -- -- -- -- -- -- -- -- -- Packet16bf -- -- -- -- -- -- -- -- -- -- -- -- */
306
307template <>
308EIGEN_STRONG_INLINE bfloat16 predux(const Packet16bf& from) {
309 return static_cast<bfloat16>(predux<Packet16f>(Bf16ToF32(from)));
310}
311
312template <>
313EIGEN_STRONG_INLINE bfloat16 predux_mul(const Packet16bf& from) {
314 return static_cast<bfloat16>(predux_mul<Packet16f>(Bf16ToF32(from)));
315}
316
317template <>
318EIGEN_STRONG_INLINE bfloat16 predux_min(const Packet16bf& from) {
319 return static_cast<bfloat16>(predux_min<Packet16f>(Bf16ToF32(from)));
320}
321
322template <>
323EIGEN_STRONG_INLINE bfloat16 predux_min<PropagateNumbers>(const Packet16bf& from) {
324 return static_cast<bfloat16>(predux_min<PropagateNumbers>(Bf16ToF32(from)));
325}
326
327template <>
328EIGEN_STRONG_INLINE bfloat16 predux_min<PropagateNaN>(const Packet16bf& from) {
329 return static_cast<bfloat16>(predux_min<PropagateNaN>(Bf16ToF32(from)));
330}
331
332template <>
333EIGEN_STRONG_INLINE bfloat16 predux_max(const Packet16bf& from) {
334 return static_cast<bfloat16>(predux_max(Bf16ToF32(from)));
335}
336
337template <>
338EIGEN_STRONG_INLINE bfloat16 predux_max<PropagateNumbers>(const Packet16bf& from) {
339 return static_cast<bfloat16>(predux_max<PropagateNumbers>(Bf16ToF32(from)));
340}
341
342template <>
343EIGEN_STRONG_INLINE bfloat16 predux_max<PropagateNaN>(const Packet16bf& from) {
344 return static_cast<bfloat16>(predux_max<PropagateNaN>(Bf16ToF32(from)));
345}
346
347template <>
348EIGEN_STRONG_INLINE bool predux_any(const Packet16bf& a) {
349 return predux_any<Packet8i>(a.m_val);
350}
351
352template <>
353EIGEN_STRONG_INLINE bool predux_all(const Packet16bf& a) {
354 return predux_count_16bit(a.m_val) == 16;
355}
356
357template <>
358EIGEN_STRONG_INLINE Index predux_count(const Packet16bf& a) {
359 return predux_count_16bit(a.m_val);
360}
361
362} // end namespace internal
363} // end namespace Eigen
364
365#endif // EIGEN_REDUCTIONS_AVX512_H