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) 2019 Rasmus Munk Larsen <rmlarsen@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_AVX512_H
12#define EIGEN_TYPE_CASTING_AVX512_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <>
22struct type_casting_traits<float, bool> : vectorized_type_casting_traits<float, bool> {};
23template <>
24struct type_casting_traits<bool, float> : vectorized_type_casting_traits<bool, float> {};
25
26template <>
27struct type_casting_traits<float, int> : vectorized_type_casting_traits<float, int> {};
28template <>
29struct type_casting_traits<int, float> : vectorized_type_casting_traits<int, float> {};
30
31template <>
32struct type_casting_traits<float, double> : vectorized_type_casting_traits<float, double> {};
33template <>
34struct type_casting_traits<double, float> : vectorized_type_casting_traits<double, float> {};
35
36template <>
37struct type_casting_traits<double, int> : vectorized_type_casting_traits<double, int> {};
38template <>
39struct type_casting_traits<int, double> : vectorized_type_casting_traits<int, double> {};
40
41template <>
42struct type_casting_traits<double, int64_t> : vectorized_type_casting_traits<double, int64_t> {};
43template <>
44struct type_casting_traits<int64_t, double> : vectorized_type_casting_traits<int64_t, double> {};
45
46template <>
47struct type_casting_traits<half, float> : vectorized_type_casting_traits<half, float> {};
48template <>
49struct type_casting_traits<float, half> : vectorized_type_casting_traits<float, half> {};
50
51template <>
52struct type_casting_traits<bfloat16, float> : vectorized_type_casting_traits<bfloat16, float> {};
53template <>
54struct type_casting_traits<float, bfloat16> : vectorized_type_casting_traits<float, bfloat16> {};
55
56EIGEN_STRONG_INLINE __mmask16 _eigen_mm512_cmpneq_ps_mask(__m512 a, __m512 b) {
57#if EIGEN_COMP_GNUC && (EIGEN_COMP_CLANG < 1000 || EIGEN_COMP_GNUC < 810)
58 return _mm512_cmp_ps_mask(a, b, _CMP_NEQ_UQ);
59#else
60 return _mm512_cmpneq_ps_mask(a, b);
61#endif
62}
63
64template <>
65EIGEN_STRONG_INLINE Packet16b pcast<Packet16f, Packet16b>(const Packet16f& a) {
66 __mmask16 mask = _eigen_mm512_cmpneq_ps_mask(a, pzero(a));
67 return _mm512_maskz_cvtepi32_epi8(mask, _mm512_set1_epi32(1));
68}
69
70template <>
71EIGEN_STRONG_INLINE Packet16f pcast<Packet16b, Packet16f>(const Packet16b& a) {
72 return _mm512_cvtepi32_ps(_mm512_and_si512(_mm512_cvtepi8_epi32(a), _mm512_set1_epi32(1)));
73}
74
75template <>
76EIGEN_STRONG_INLINE Packet16i pcast<Packet16f, Packet16i>(const Packet16f& a) {
77 return _mm512_cvttps_epi32(a);
78}
79
80template <>
81EIGEN_STRONG_INLINE Packet8d pcast<Packet16f, Packet8d>(const Packet16f& a) {
82 return _mm512_cvtps_pd(_mm512_castps512_ps256(a));
83}
84
85template <>
86EIGEN_STRONG_INLINE Packet8d pcast<Packet8f, Packet8d>(const Packet8f& a) {
87 return _mm512_cvtps_pd(a);
88}
89
90template <>
91EIGEN_STRONG_INLINE Packet8l pcast<Packet8d, Packet8l>(const Packet8d& a) {
92#if defined(EIGEN_VECTORIZE_AVX512DQ) && defined(EIGEN_VECTORIZE_AVX512VL)
93 return _mm512_cvttpd_epi64(a);
94#else
95 constexpr int kTotalBits = sizeof(double) * CHAR_BIT, kMantissaBits = std::numeric_limits<double>::digits - 1,
96 kExponentBits = kTotalBits - kMantissaBits - 1, kBias = (1 << (kExponentBits - 1)) - 1;
97
98 const __m512i cst_one = _mm512_set1_epi64(1);
99 const __m512i cst_total_bits = _mm512_set1_epi64(kTotalBits);
100 const __m512i cst_bias = _mm512_set1_epi64(kBias);
101
102 __m512i a_bits = _mm512_castpd_si512(a);
103 // shift left by 1 to clear the sign bit, and shift right by kMantissaBits + 1 to recover biased exponent
104 __m512i biased_e = _mm512_srli_epi64(_mm512_slli_epi64(a_bits, 1), kMantissaBits + 1);
105 __m512i e = _mm512_sub_epi64(biased_e, cst_bias);
106
107 // shift to the left by kExponentBits + 1 to clear the sign and exponent bits
108 __m512i shifted_mantissa = _mm512_slli_epi64(a_bits, kExponentBits + 1);
109 // shift to the right by kTotalBits - e to convert the significand to an integer
110 __m512i result_significand = _mm512_srlv_epi64(shifted_mantissa, _mm512_sub_epi64(cst_total_bits, e));
111
112 // add the implied bit
113 __m512i result_exponent = _mm512_sllv_epi64(cst_one, e);
114 // e <= 0 is interpreted as a large positive shift (2's complement), which also conveniently results in zero
115 __m512i result = _mm512_add_epi64(result_significand, result_exponent);
116 // handle negative arguments
117 __mmask8 sign_mask = _mm512_cmplt_epi64_mask(a_bits, _mm512_setzero_si512());
118 result = _mm512_mask_sub_epi64(result, sign_mask, _mm512_setzero_si512(), result);
119 return result;
120#endif
121}
122
123template <>
124EIGEN_STRONG_INLINE Packet16f pcast<Packet16i, Packet16f>(const Packet16i& a) {
125 return _mm512_cvtepi32_ps(a);
126}
127
128template <>
129EIGEN_STRONG_INLINE Packet8d pcast<Packet16i, Packet8d>(const Packet16i& a) {
130 return _mm512_cvtepi32_pd(_mm512_castsi512_si256(a));
131}
132
133template <>
134EIGEN_STRONG_INLINE Packet8d pcast<Packet8i, Packet8d>(const Packet8i& a) {
135 return _mm512_cvtepi32_pd(a);
136}
137
138template <>
139EIGEN_STRONG_INLINE Packet8d pcast<Packet8l, Packet8d>(const Packet8l& a) {
140#if defined(EIGEN_VECTORIZE_AVX512DQ) && defined(EIGEN_VECTORIZE_AVX512VL)
141 return _mm512_cvtepi64_pd(a);
142#else
143 EIGEN_ALIGN64 int64_t aux[8];
144 pstore(aux, a);
145 return _mm512_set_pd(static_cast<double>(aux[7]), static_cast<double>(aux[6]), static_cast<double>(aux[5]),
146 static_cast<double>(aux[4]), static_cast<double>(aux[3]), static_cast<double>(aux[2]),
147 static_cast<double>(aux[1]), static_cast<double>(aux[0]));
148#endif
149}
150
151template <>
152EIGEN_STRONG_INLINE Packet16f pcast<Packet8d, Packet16f>(const Packet8d& a, const Packet8d& b) {
153 return cat256(_mm512_cvtpd_ps(a), _mm512_cvtpd_ps(b));
154}
155
156template <>
157EIGEN_STRONG_INLINE Packet16i pcast<Packet8d, Packet16i>(const Packet8d& a, const Packet8d& b) {
158 return cat256i(_mm512_cvttpd_epi32(a), _mm512_cvttpd_epi32(b));
159}
160
161template <>
162EIGEN_STRONG_INLINE Packet8i pcast<Packet8d, Packet8i>(const Packet8d& a) {
163 return _mm512_cvtpd_epi32(a);
164}
165template <>
166EIGEN_STRONG_INLINE Packet8f pcast<Packet8d, Packet8f>(const Packet8d& a) {
167 return _mm512_cvtpd_ps(a);
168}
169
170template <>
171EIGEN_STRONG_INLINE Packet16i preinterpret<Packet16i, Packet16f>(const Packet16f& a) {
172 return _mm512_castps_si512(a);
173}
174
175template <>
176EIGEN_STRONG_INLINE Packet16f preinterpret<Packet16f, Packet16i>(const Packet16i& a) {
177 return _mm512_castsi512_ps(a);
178}
179
180template <>
181EIGEN_STRONG_INLINE Packet8d preinterpret<Packet8d, Packet16f>(const Packet16f& a) {
182 return _mm512_castps_pd(a);
183}
184
185template <>
186EIGEN_STRONG_INLINE Packet8d preinterpret<Packet8d, Packet8l>(const Packet8l& a) {
187 return _mm512_castsi512_pd(a);
188}
189
190template <>
191EIGEN_STRONG_INLINE Packet8l preinterpret<Packet8l, Packet8d>(const Packet8d& a) {
192 return _mm512_castpd_si512(a);
193}
194
195template <>
196EIGEN_STRONG_INLINE Packet16f preinterpret<Packet16f, Packet8d>(const Packet8d& a) {
197 return _mm512_castpd_ps(a);
198}
199
200template <>
201EIGEN_STRONG_INLINE Packet8f preinterpret<Packet8f, Packet16f>(const Packet16f& a) {
202 return _mm512_castps512_ps256(a);
203}
204
205template <>
206EIGEN_STRONG_INLINE Packet4f preinterpret<Packet4f, Packet16f>(const Packet16f& a) {
207 return _mm512_castps512_ps128(a);
208}
209
210template <>
211EIGEN_STRONG_INLINE Packet4d preinterpret<Packet4d, Packet8d>(const Packet8d& a) {
212 return _mm512_castpd512_pd256(a);
213}
214
215template <>
216EIGEN_STRONG_INLINE Packet2d preinterpret<Packet2d, Packet8d>(const Packet8d& a) {
217 return _mm512_castpd512_pd128(a);
218}
219
220template <>
221EIGEN_STRONG_INLINE Packet16f preinterpret<Packet16f, Packet8f>(const Packet8f& a) {
222 return _mm512_castps256_ps512(a);
223}
224
225template <>
226EIGEN_STRONG_INLINE Packet16f preinterpret<Packet16f, Packet4f>(const Packet4f& a) {
227 return _mm512_castps128_ps512(a);
228}
229
230template <>
231EIGEN_STRONG_INLINE Packet8d preinterpret<Packet8d, Packet4d>(const Packet4d& a) {
232 return _mm512_castpd256_pd512(a);
233}
234
235template <>
236EIGEN_STRONG_INLINE Packet8d preinterpret<Packet8d, Packet2d>(const Packet2d& a) {
237 return _mm512_castpd128_pd512(a);
238}
239
240template <>
241EIGEN_STRONG_INLINE Packet8i preinterpret<Packet8i, Packet16i>(const Packet16i& a) {
242 return _mm512_castsi512_si256(a);
243}
244template <>
245EIGEN_STRONG_INLINE Packet4i preinterpret<Packet4i, Packet16i>(const Packet16i& a) {
246 return _mm512_castsi512_si128(a);
247}
248
249#ifndef EIGEN_VECTORIZE_AVX512FP16
250template <>
251EIGEN_STRONG_INLINE Packet8h preinterpret<Packet8h, Packet16h>(const Packet16h& a) {
252 return _mm256_castsi256_si128(a);
253}
254
255template <>
256EIGEN_STRONG_INLINE Packet16f pcast<Packet16h, Packet16f>(const Packet16h& a) {
257 return half2float(a);
258}
259
260template <>
261EIGEN_STRONG_INLINE Packet16h pcast<Packet16f, Packet16h>(const Packet16f& a) {
262 return float2half(a);
263}
264
265#endif
266
267template <>
268EIGEN_STRONG_INLINE Packet8bf preinterpret<Packet8bf, Packet16bf>(const Packet16bf& a) {
269 return _mm256_castsi256_si128(a);
270}
271
272template <>
273EIGEN_STRONG_INLINE Packet16f pcast<Packet16bf, Packet16f>(const Packet16bf& a) {
274 return Bf16ToF32(a);
275}
276
277template <>
278EIGEN_STRONG_INLINE Packet16bf pcast<Packet16f, Packet16bf>(const Packet16f& a) {
279 return F32ToBf16(a);
280}
281
282} // end namespace internal
283
284} // end namespace Eigen
285
286#endif // EIGEN_TYPE_CASTING_AVX512_H