Eigen  5.0.1
 
Loading...
Searching...
No Matches
MathFunctions.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2016 Pedro Gonnet (pedro.gonnet@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 THIRD_PARTY_EIGEN3_EIGEN_SRC_CORE_ARCH_AVX512_MATHFUNCTIONS_H_
12#define THIRD_PARTY_EIGEN3_EIGEN_SRC_CORE_ARCH_AVX512_MATHFUNCTIONS_H_
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20EIGEN_INSTANTIATE_GENERIC_MATH_FUNCS_FLOAT(Packet16f)
21EIGEN_INSTANTIATE_GENERIC_MATH_FUNCS_DOUBLE(Packet8d)
22
23template <>
24EIGEN_STRONG_INLINE Packet16h pfrexp(const Packet16h& a, Packet16h& exponent) {
25 Packet16f fexponent;
26 const Packet16h out = float2half(pfrexp<Packet16f>(half2float(a), fexponent));
27 exponent = float2half(fexponent);
28 return out;
29}
30
31template <>
32EIGEN_STRONG_INLINE Packet16h pldexp(const Packet16h& a, const Packet16h& exponent) {
33 return float2half(pldexp<Packet16f>(half2float(a), half2float(exponent)));
34}
35
36template <>
37EIGEN_STRONG_INLINE Packet16bf pfrexp(const Packet16bf& a, Packet16bf& exponent) {
38 // Both results are exact: the mantissa keeps the input's significand, and the exponent is a
39 // small integer.
40 Packet16f fexponent;
41 const Packet16bf out = F32ToBf16Truncate(pfrexp<Packet16f>(Bf16ToF32(a), fexponent));
42 exponent = F32ToBf16Truncate(fexponent);
43 return out;
44}
45
46template <>
47EIGEN_STRONG_INLINE Packet16bf pldexp(const Packet16bf& a, const Packet16bf& exponent) {
48 return F32ToBf16(pldexp<Packet16f>(Bf16ToF32(a), Bf16ToF32(exponent)));
49}
50
51#if EIGEN_FAST_MATH
52template <>
53EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet16f psqrt<Packet16f>(const Packet16f& x) {
54 return generic_sqrt_newton_step<Packet16f>::run(x, _mm512_rsqrt14_ps(x));
55}
56
57template <>
58EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet8d psqrt<Packet8d>(const Packet8d& x) {
59#ifdef EIGEN_VECTORIZE_AVX512ER
60 return generic_sqrt_newton_step<Packet8d, /*Steps=*/1>::run(x, _mm512_rsqrt28_pd(x));
61#else
62 return generic_sqrt_newton_step<Packet8d, /*Steps=*/2>::run(x, _mm512_rsqrt14_pd(x));
63#endif
64}
65#else
66template <>
67EIGEN_STRONG_INLINE Packet16f psqrt<Packet16f>(const Packet16f& x) {
68 return _mm512_sqrt_ps(x);
69}
70
71template <>
72EIGEN_STRONG_INLINE Packet8d psqrt<Packet8d>(const Packet8d& x) {
73 return _mm512_sqrt_pd(x);
74}
75#endif
76
77// prsqrt for float.
78#if defined(EIGEN_VECTORIZE_AVX512ER)
79template <>
80EIGEN_STRONG_INLINE Packet16f prsqrt<Packet16f>(const Packet16f& x) {
81 return _mm512_rsqrt28_ps(x);
82}
83#elif EIGEN_FAST_MATH
84
85template <>
86EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet16f prsqrt<Packet16f>(const Packet16f& x) {
87 return generic_rsqrt_newton_step<Packet16f, /*Steps=*/1>::run(x, _mm512_rsqrt14_ps(x));
88}
89#endif
90
91// prsqrt for double.
92#if EIGEN_FAST_MATH
93template <>
94EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet8d prsqrt<Packet8d>(const Packet8d& x) {
95#ifdef EIGEN_VECTORIZE_AVX512ER
96 return generic_rsqrt_newton_step<Packet8d, /*Steps=*/1>::run(x, _mm512_rsqrt28_pd(x));
97#else
98 return generic_rsqrt_newton_step<Packet8d, /*Steps=*/2>::run(x, _mm512_rsqrt14_pd(x));
99#endif
100}
101
102template <>
103EIGEN_STRONG_INLINE Packet16f preciprocal<Packet16f>(const Packet16f& a) {
104#ifdef EIGEN_VECTORIZE_AVX512ER
105 return _mm512_rcp28_ps(a);
106#else
107 // generic_reciprocal_newton_step, with its NaN-or-zero test as one compare into a mask register (r == 0 or
108 // unordered), as for AVX: GCC compiles the generic predux_any(pcmp_lt_or_nan(...)) to a compare mask converted to a
109 // vector and back.
110 const Packet16f one = pset1<Packet16f>(1.0f);
111 const Packet16f x = _mm512_rcp14_ps(a);
112 const Packet16f refined = pmadd(x, pnmadd(a, x, one), x);
113 const __mmask16 redo = _mm512_cmp_ps_mask(refined, _mm512_setzero_ps(), _CMP_EQ_UQ);
114 return redo == 0 ? refined : _mm512_mask_div_ps(refined, redo, one, a);
115#endif
116}
117#endif
118
119EIGEN_INSTANTIATE_GENERIC_MATH_FUNCS_BF16(Packet16f, Packet16bf)
120
121#ifndef EIGEN_VECTORIZE_AVX512FP16
122EIGEN_INSTANTIATE_GENERIC_MATH_FUNCS_F16(Packet16f, Packet16h)
123#endif // EIGEN_VECTORIZE_AVX512FP16
124
125} // end namespace internal
126
127} // end namespace Eigen
128
129#endif // THIRD_PARTY_EIGEN3_EIGEN_SRC_CORE_ARCH_AVX512_MATHFUNCTIONS_H_