Eigen  5.0.1
 
Loading...
Searching...
No Matches
TypeCasting.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2015 Benoit Steiner <benoit.steiner.goog@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_TYPE_CASTING_AVX_H
12#define EIGEN_TYPE_CASTING_AVX_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21#ifndef EIGEN_VECTORIZE_AVX512
22template <>
23struct type_casting_traits<float, bool> : vectorized_type_casting_traits<float, bool> {};
24template <>
25struct type_casting_traits<bool, float> : vectorized_type_casting_traits<bool, float> {};
26
27template <>
28struct type_casting_traits<float, int> : vectorized_type_casting_traits<float, int> {};
29template <>
30struct type_casting_traits<int, float> : vectorized_type_casting_traits<int, float> {};
31
32template <>
33struct type_casting_traits<float, double> : vectorized_type_casting_traits<float, double> {};
34template <>
35struct type_casting_traits<double, float> : vectorized_type_casting_traits<double, float> {};
36
37template <>
38struct type_casting_traits<double, int> : vectorized_type_casting_traits<double, int> {};
39template <>
40struct type_casting_traits<int, double> : vectorized_type_casting_traits<int, double> {};
41
42template <>
43struct type_casting_traits<half, float> : vectorized_type_casting_traits<half, float> {};
44template <>
45struct type_casting_traits<float, half> : vectorized_type_casting_traits<float, half> {};
46
47template <>
48struct type_casting_traits<bfloat16, float> : vectorized_type_casting_traits<bfloat16, float> {};
49template <>
50struct type_casting_traits<float, bfloat16> : vectorized_type_casting_traits<float, bfloat16> {};
51
52#ifdef EIGEN_VECTORIZE_AVX2
53template <>
54struct type_casting_traits<double, int64_t> : vectorized_type_casting_traits<double, int64_t> {};
55template <>
56struct type_casting_traits<int64_t, double> : vectorized_type_casting_traits<int64_t, double> {};
57#endif
58#endif
59
60EIGEN_STRONG_INLINE __m256 _eigen_mm256_set_m128(__m128 hi, __m128 lo) {
61#if EIGEN_COMP_GNUC && (EIGEN_COMP_CLANG < 1000 || EIGEN_COMP_GNUC < 810)
62 __m256 result = _mm256_castps128_ps256(lo);
63 return _mm256_insertf128_ps(result, hi, 1);
64#else
65 return _mm256_set_m128(hi, lo);
66#endif
67}
68
69EIGEN_STRONG_INLINE __m256d _eigen_mm256_set_m128d(__m128d hi, __m128d lo) {
70#if EIGEN_COMP_GNUC && (EIGEN_COMP_CLANG < 1000 || EIGEN_COMP_GNUC < 810)
71 __m256d result = _mm256_castpd128_pd256(lo);
72 return _mm256_insertf128_pd(result, hi, 1);
73#else
74 return _mm256_set_m128d(hi, lo);
75#endif
76}
77
78EIGEN_STRONG_INLINE __m256i _eigen_mm256_set_m128i(__m128i hi, __m128i lo) {
79#if EIGEN_COMP_GNUC && (EIGEN_COMP_CLANG < 1000 || EIGEN_COMP_GNUC < 810)
80#if defined(EIGEN_VECTORIZE_AVX2)
81 __m256i result = _mm256_castsi128_si256(lo);
82 return _mm256_inserti128_si256(result, hi, 1);
83#else
84 EIGEN_ALIGN32 int32_t tmp[8];
85 _mm_storeu_si128(reinterpret_cast<__m128i*>(tmp), lo);
86 _mm_storeu_si128(reinterpret_cast<__m128i*>(tmp + 4), hi);
87 return _mm256_loadu_si256(reinterpret_cast<const __m256i*>(tmp));
88#endif
89#else
90 return _mm256_set_m128i(hi, lo);
91#endif
92}
93
94template <>
95EIGEN_STRONG_INLINE Packet16b pcast<Packet8f, Packet16b>(const Packet8f& a, const Packet8f& b) {
96 __m256 nonzero_a = _mm256_cmp_ps(a, pzero(a), _CMP_NEQ_UQ);
97 __m256 nonzero_b = _mm256_cmp_ps(b, pzero(b), _CMP_NEQ_UQ);
98 constexpr char kFF = '\255';
99#ifndef EIGEN_VECTORIZE_AVX2
100 __m128i shuffle_mask128_a_lo = _mm_set_epi8(kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, 12, 8, 4, 0);
101 __m128i shuffle_mask128_a_hi = _mm_set_epi8(kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, 12, 8, 4, 0, kFF, kFF, kFF, kFF);
102 __m128i shuffle_mask128_b_lo = _mm_set_epi8(kFF, kFF, kFF, kFF, 12, 8, 4, 0, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF);
103 __m128i shuffle_mask128_b_hi = _mm_set_epi8(12, 8, 4, 0, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF);
104 __m128i a_hi = _mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castps_si256(nonzero_a), 1), shuffle_mask128_a_hi);
105 __m128i a_lo = _mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castps_si256(nonzero_a), 0), shuffle_mask128_a_lo);
106 __m128i b_hi = _mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castps_si256(nonzero_b), 1), shuffle_mask128_b_hi);
107 __m128i b_lo = _mm_shuffle_epi8(_mm256_extractf128_si256(_mm256_castps_si256(nonzero_b), 0), shuffle_mask128_b_lo);
108 __m128i merged = _mm_or_si128(_mm_or_si128(b_lo, b_hi), _mm_or_si128(a_lo, a_hi));
109 return _mm_and_si128(merged, _mm_set1_epi8(1));
110#else
111 __m256i a_shuffle_mask = _mm256_set_epi8(kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, 12, 8, 4, 0, kFF, kFF, kFF, kFF, kFF,
112 kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, 12, 8, 4, 0);
113 __m256i b_shuffle_mask = _mm256_set_epi8(12, 8, 4, 0, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF,
114 kFF, kFF, kFF, 12, 8, 4, 0, kFF, kFF, kFF, kFF, kFF, kFF, kFF, kFF);
115 __m256i a_shuff = _mm256_shuffle_epi8(_mm256_castps_si256(nonzero_a), a_shuffle_mask);
116 __m256i b_shuff = _mm256_shuffle_epi8(_mm256_castps_si256(nonzero_b), b_shuffle_mask);
117 __m256i a_or_b = _mm256_or_si256(a_shuff, b_shuff);
118 __m256i merged = _mm256_or_si256(a_or_b, _mm256_castsi128_si256(_mm256_extractf128_si256(a_or_b, 1)));
119 return _mm256_castsi256_si128(_mm256_and_si256(merged, _mm256_set1_epi8(1)));
120#endif
121}
122
123template <>
124EIGEN_STRONG_INLINE Packet8f pcast<Packet16b, Packet8f>(const Packet16b& a) {
125 const __m256 cst_one = _mm256_set1_ps(1.0f);
126#ifdef EIGEN_VECTORIZE_AVX2
127 __m256i a_extended = _mm256_cvtepi8_epi32(a);
128 __m256i abcd_efgh = _mm256_cmpeq_epi32(a_extended, _mm256_setzero_si256());
129#else
130 __m128i abcd_efhg_ijkl_mnop = _mm_cmpeq_epi8(a, _mm_setzero_si128());
131 __m128i aabb_ccdd_eeff_gghh = _mm_unpacklo_epi8(abcd_efhg_ijkl_mnop, abcd_efhg_ijkl_mnop);
132 __m128i aaaa_bbbb_cccc_dddd = _mm_unpacklo_epi8(aabb_ccdd_eeff_gghh, aabb_ccdd_eeff_gghh);
133 __m128i eeee_ffff_gggg_hhhh = _mm_unpackhi_epi8(aabb_ccdd_eeff_gghh, aabb_ccdd_eeff_gghh);
134 __m256i abcd_efgh = _mm256_setr_m128i(aaaa_bbbb_cccc_dddd, eeee_ffff_gggg_hhhh);
135#endif
136 __m256 result = _mm256_andnot_ps(_mm256_castsi256_ps(abcd_efgh), cst_one);
137 return result;
138}
139
140template <>
141EIGEN_STRONG_INLINE Packet8i pcast<Packet8f, Packet8i>(const Packet8f& a) {
142 return _mm256_cvttps_epi32(a);
143}
144
145template <>
146EIGEN_STRONG_INLINE Packet8i pcast<Packet4d, Packet8i>(const Packet4d& a, const Packet4d& b) {
147 return _eigen_mm256_set_m128i(_mm256_cvttpd_epi32(b), _mm256_cvttpd_epi32(a));
148}
149
150template <>
151EIGEN_STRONG_INLINE Packet4i pcast<Packet4d, Packet4i>(const Packet4d& a) {
152 return _mm256_cvttpd_epi32(a);
153}
154
155template <>
156EIGEN_STRONG_INLINE Packet8f pcast<Packet8i, Packet8f>(const Packet8i& a) {
157 return _mm256_cvtepi32_ps(a);
158}
159
160template <>
161EIGEN_STRONG_INLINE Packet8f pcast<Packet4d, Packet8f>(const Packet4d& a, const Packet4d& b) {
162 return _eigen_mm256_set_m128(_mm256_cvtpd_ps(b), _mm256_cvtpd_ps(a));
163}
164
165template <>
166EIGEN_STRONG_INLINE Packet4f pcast<Packet4d, Packet4f>(const Packet4d& a) {
167 return _mm256_cvtpd_ps(a);
168}
169
170template <>
171EIGEN_STRONG_INLINE Packet4d pcast<Packet8i, Packet4d>(const Packet8i& a) {
172 return _mm256_cvtepi32_pd(_mm256_castsi256_si128(a));
173}
174
175template <>
176EIGEN_STRONG_INLINE Packet4d pcast<Packet4i, Packet4d>(const Packet4i& a) {
177 return _mm256_cvtepi32_pd(a);
178}
179
180template <>
181EIGEN_STRONG_INLINE Packet4d pcast<Packet8f, Packet4d>(const Packet8f& a) {
182 return _mm256_cvtps_pd(_mm256_castps256_ps128(a));
183}
184
185template <>
186EIGEN_STRONG_INLINE Packet4d pcast<Packet4f, Packet4d>(const Packet4f& a) {
187 return _mm256_cvtps_pd(a);
188}
189
190template <>
191EIGEN_STRONG_INLINE Packet8i preinterpret<Packet8i, Packet8f>(const Packet8f& a) {
192 return _mm256_castps_si256(a);
193}
194
195template <>
196EIGEN_STRONG_INLINE Packet8f preinterpret<Packet8f, Packet8i>(const Packet8i& a) {
197 return _mm256_castsi256_ps(a);
198}
199
200template <>
201EIGEN_STRONG_INLINE Packet8ui preinterpret<Packet8ui, Packet8i>(const Packet8i& a) {
202 return Packet8ui(a);
203}
204
205template <>
206EIGEN_STRONG_INLINE Packet8i preinterpret<Packet8i, Packet8ui>(const Packet8ui& a) {
207 return Packet8i(a);
208}
209
210// truncation operations
211
212template <>
213EIGEN_STRONG_INLINE Packet4f preinterpret<Packet4f, Packet8f>(const Packet8f& a) {
214 return _mm256_castps256_ps128(a);
215}
216
217template <>
218EIGEN_STRONG_INLINE Packet2d preinterpret<Packet2d, Packet4d>(const Packet4d& a) {
219 return _mm256_castpd256_pd128(a);
220}
221
222template <>
223EIGEN_STRONG_INLINE Packet4i preinterpret<Packet4i, Packet8i>(const Packet8i& a) {
224 return _mm256_castsi256_si128(a);
225}
226
227template <>
228EIGEN_STRONG_INLINE Packet4ui preinterpret<Packet4ui, Packet8ui>(const Packet8ui& a) {
229 return _mm256_castsi256_si128(a);
230}
231
232#ifdef EIGEN_VECTORIZE_AVX2
233template <>
234EIGEN_STRONG_INLINE Packet4l pcast<Packet4d, Packet4l>(const Packet4d& a) {
235#if defined(EIGEN_VECTORIZE_AVX512DQ) && defined(EIGEN_VECTORIZE_AVS512VL)
236 return _mm256_cvttpd_epi64(a);
237#else
238
239 // if 'a' exceeds the numerical limits of int64_t, the behavior is undefined
240
241 // e <= 0 corresponds to |a| < 1, which should result in zero. incidentally, intel intrinsics with shift arguments
242 // greater than or equal to 64 produce zero. furthermore, negative shifts appear to be interpreted as large positive
243 // shifts (two's complement), which also result in zero. therefore, e does not need to be clamped to [0, 64)
244
245 constexpr int kTotalBits = sizeof(double) * CHAR_BIT, kMantissaBits = std::numeric_limits<double>::digits - 1,
246 kExponentBits = kTotalBits - kMantissaBits - 1, kBias = (1 << (kExponentBits - 1)) - 1;
247
248 const __m256i cst_one = _mm256_set1_epi64x(1);
249 const __m256i cst_total_bits = _mm256_set1_epi64x(kTotalBits);
250 const __m256i cst_bias = _mm256_set1_epi64x(kBias);
251
252 __m256i a_bits = _mm256_castpd_si256(a);
253 // shift left by 1 to clear the sign bit, and shift right by kMantissaBits + 1 to recover biased exponent
254 __m256i biased_e = _mm256_srli_epi64(_mm256_slli_epi64(a_bits, 1), kMantissaBits + 1);
255 __m256i e = _mm256_sub_epi64(biased_e, cst_bias);
256
257 // shift to the left by kExponentBits + 1 to clear the sign and exponent bits
258 __m256i shifted_mantissa = _mm256_slli_epi64(a_bits, kExponentBits + 1);
259 // shift to the right by kTotalBits - e to convert the significand to an integer
260 __m256i result_significand = _mm256_srlv_epi64(shifted_mantissa, _mm256_sub_epi64(cst_total_bits, e));
261
262 // add the implied bit
263 __m256i result_exponent = _mm256_sllv_epi64(cst_one, e);
264 // e <= 0 is interpreted as a large positive shift (2's complement), which also conveniently results in zero
265 __m256i result = _mm256_add_epi64(result_significand, result_exponent);
266 // handle negative arguments
267 __m256i sign_mask = _mm256_cmpgt_epi64(_mm256_setzero_si256(), a_bits);
268 result = _mm256_sub_epi64(_mm256_xor_si256(result, sign_mask), sign_mask);
269 return result;
270#endif
271}
272
273template <>
274EIGEN_STRONG_INLINE Packet4d pcast<Packet4l, Packet4d>(const Packet4l& a) {
275#if defined(EIGEN_VECTORIZE_AVX512DQ) && defined(EIGEN_VECTORIZE_AVS512VL)
276 return _mm256_cvtepi64_pd(a);
277#else
278 int64_t aux[4];
279 pstoreu(aux, a);
280 return _mm256_set_pd(static_cast<double>(aux[3]), static_cast<double>(aux[2]), static_cast<double>(aux[1]),
281 static_cast<double>(aux[0]));
282#endif
283}
284
285template <>
286EIGEN_STRONG_INLINE Packet4d pcast<Packet2l, Packet4d>(const Packet2l& a, const Packet2l& b) {
287 return _eigen_mm256_set_m128d((pcast<Packet2l, Packet2d>(b)), (pcast<Packet2l, Packet2d>(a)));
288}
289
290template <>
291EIGEN_STRONG_INLINE Packet4ul preinterpret<Packet4ul, Packet4l>(const Packet4l& a) {
292 return Packet4ul(a);
293}
294
295template <>
296EIGEN_STRONG_INLINE Packet4l preinterpret<Packet4l, Packet4ul>(const Packet4ul& a) {
297 return Packet4l(a);
298}
299
300template <>
301EIGEN_STRONG_INLINE Packet4l preinterpret<Packet4l, Packet4d>(const Packet4d& a) {
302 return _mm256_castpd_si256(a);
303}
304
305template <>
306EIGEN_STRONG_INLINE Packet4d preinterpret<Packet4d, Packet4l>(const Packet4l& a) {
307 return _mm256_castsi256_pd(a);
308}
309
310// truncation operations
311template <>
312EIGEN_STRONG_INLINE Packet2l preinterpret<Packet2l, Packet4l>(const Packet4l& a) {
313 return _mm256_castsi256_si128(a);
314}
315#endif
316
317#ifndef EIGEN_VECTORIZE_AVX512FP16
318template <>
319EIGEN_STRONG_INLINE Packet8f pcast<Packet8h, Packet8f>(const Packet8h& a) {
320 return half2float(a);
321}
322
323template <>
324EIGEN_STRONG_INLINE Packet8h pcast<Packet8f, Packet8h>(const Packet8f& a) {
325 return float2half(a);
326}
327#endif
328
329template <>
330EIGEN_STRONG_INLINE Packet8f pcast<Packet8bf, Packet8f>(const Packet8bf& a) {
331 return Bf16ToF32(a);
332}
333
334template <>
335EIGEN_STRONG_INLINE Packet8bf pcast<Packet8f, Packet8bf>(const Packet8f& a) {
336 return F32ToBf16(a);
337}
338
339} // end namespace internal
340
341} // end namespace Eigen
342
343#endif // EIGEN_TYPE_CASTING_AVX_H