12#ifndef EIGEN_MATHFUNCTIONSIMPL_H
13#define EIGEN_MATHFUNCTIONSIMPL_H
16#include "./InternalHeaderCheck.h"
37template <
typename Packet,
int Steps>
38struct generic_reciprocal_newton_step {
39 static_assert(Steps > 0,
"Steps must be at least 1.");
40 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Packet run(
const Packet& a,
const Packet& approx_a_recip) {
41 using Scalar =
typename unpacket_traits<Packet>::type;
42 const Packet one = pset1<Packet>(Scalar(1));
45 const Packet x = generic_reciprocal_newton_step<Packet, Steps - 1>::run(a, approx_a_recip);
46 const Packet tmp = pnmadd(a, x, one);
48 const Packet refined = pmadd(x, tmp, x);
51 const Packet redo = pcmp_lt_or_nan(pabs(refined), pset1<Packet>((std::numeric_limits<Scalar>::denorm_min)()));
52 return predux_any(redo) ? pselect(redo, pdiv(one, a), refined) : refined;
56template <
typename Packet>
57struct generic_reciprocal_newton_step<Packet, 0> {
58 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Packet run(
const Packet& ,
const Packet& approx_rsqrt) {
78template <
typename Packet,
int Steps>
79struct generic_rsqrt_newton_step {
80 static_assert(Steps > 0,
"Steps must be at least 1.");
81 using Scalar =
typename unpacket_traits<Packet>::type;
82 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Packet run(
const Packet& a,
const Packet& approx_rsqrt) {
83 const Scalar kMinusHalf = Scalar(-1) / Scalar(2);
84 const Packet cst_minus_half = pset1<Packet>(kMinusHalf);
85 const Packet cst_minus_one = pset1<Packet>(Scalar(-1));
87 Packet inv_sqrt = approx_rsqrt;
88 for (
int step = 0; step < Steps; ++step) {
92 Packet r2 = pmul(a, inv_sqrt);
93 Packet half_r = pmul(inv_sqrt, cst_minus_half);
94 Packet h_n = pmadd(r2, inv_sqrt, cst_minus_one);
95 inv_sqrt = pmadd(half_r, h_n, inv_sqrt);
102 return pselect(pisnan(inv_sqrt), approx_rsqrt, inv_sqrt);
106template <
typename Packet>
107struct generic_rsqrt_newton_step<Packet, 0> {
108 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Packet run(
const Packet& ,
const Packet& approx_rsqrt) {
128template <
typename Packet,
int Steps = 1>
129struct generic_sqrt_newton_step {
130 static_assert(Steps > 0,
"Steps must be at least 1.");
132 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Packet run(
const Packet& a,
const Packet& approx_rsqrt) {
133 using Scalar =
typename unpacket_traits<Packet>::type;
134 const Packet one_point_five = pset1<Packet>(Scalar(1.5));
135 const Packet minus_half = pset1<Packet>(Scalar(-0.5));
137 const Packet inf_mask = pcmp_eq(a, pset1<Packet>(NumTraits<Scalar>::infinity()));
138 const Packet return_a = por(pcmp_eq(a, pzero(a)), inf_mask);
142 Packet rsqrt = pmul(approx_rsqrt, pmadd(pmul(minus_half, approx_rsqrt), pmul(a, approx_rsqrt), one_point_five));
143 for (
int step = 1; step < Steps; ++step) {
144 rsqrt = pmul(rsqrt, pmadd(pmul(minus_half, rsqrt), pmul(a, rsqrt), one_point_five));
149 return pselect(return_a, a, pmul(a, rsqrt));
153template <
typename RealScalar>
154EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE RealScalar positive_real_hypot(
const RealScalar& x,
155 const RealScalar& y) {
157 if ((numext::isinf)(x) || (numext::isinf)(y))
return NumTraits<RealScalar>::infinity();
158 if ((numext::isnan)(x) || (numext::isnan)(y))
return NumTraits<RealScalar>::quiet_NaN();
160 EIGEN_USING_STD(sqrt);
161 RealScalar p = numext::maxi(x, y);
162 if (numext::is_exactly_zero(p))
return RealScalar(0);
163 RealScalar qp = numext::mini(y, x) / p;
164 return p * sqrt(RealScalar(1) + qp * qp);
167template <
typename Scalar>
169 using RealScalar =
typename NumTraits<Scalar>::Real;
170 static EIGEN_DEVICE_FUNC
inline RealScalar run(
const Scalar& x,
const Scalar& y) {
171 return positive_real_hypot<RealScalar>(numext::abs(x), numext::abs(y));
175template <
typename ComplexT,
bool Reciprocal>
176EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE ComplexT complex_sqrt_extreme(
const ComplexT& z) {
177 using T =
typename NumTraits<ComplexT>::Real;
178 const T x = numext::real(z);
179 const T y = numext::imag(z);
181 const T inf = NumTraits<T>::infinity();
182 EIGEN_IF_CONSTEXPR (Reciprocal) {
183 if ((numext::isinf)(x) || (numext::isinf)(y))
return ComplexT(zero, numext::copysign(zero, -y));
185 if ((numext::isinf)(y))
return ComplexT(inf, y);
186 if ((numext::isinf)(x)) {
187 const T other = (numext::isnan)(y) ? y : zero;
188 return x > zero ? ComplexT(inf, numext::copysign(other, y))
189 : ComplexT(numext::abs(other), numext::copysign(inf, y));
192 if ((numext::isnan)(x) || (numext::isnan)(y)) {
193 return ComplexT(NumTraits<T>::quiet_NaN(), NumTraits<T>::quiet_NaN());
195 if (numext::is_exactly_zero(x) && numext::is_exactly_zero(y)) {
196 return Reciprocal ? ComplexT(inf, NumTraits<T>::quiet_NaN()) : ComplexT(zero, y);
199 T ax = numext::abs(x);
200 T ay = numext::abs(y);
201 const T p = numext::maxi(ax, ay);
202 const T r = numext::mini(ax, ay) / p;
203 const T h = numext::sqrt(T(1) + r * r);
207 const T scale = numext::sqrt(p);
208 const T sum = ax + h;
209 const T w = numext::sqrt(T(0.5) * sum);
212 EIGEN_IF_CONSTEXPR (Reciprocal) {
213 major = (w / h) / scale;
214 minor = (ay / sum) * major;
217 minor = numext::abs(y) / (T(2) * major);
219 if (numext::is_exactly_zero(x)) minor = major;
220 const T imag_sign = Reciprocal ? -y : y;
221 return x < zero ? ComplexT(minor, numext::copysign(major, imag_sign))
222 : ComplexT(major, numext::copysign(minor, imag_sign));
225template <
typename ComplexT,
bool Reciprocal>
226EIGEN_DEVICE_FUNC
constexpr ComplexT complex_sqrt_impl(
const ComplexT& z) {
227 using T =
typename NumTraits<ComplexT>::Real;
228 const T x = numext::real(z);
229 const T y = numext::imag(z);
230 const T ax = numext::abs(x);
231 const T ay = numext::abs(y);
232 const bool real_larger = ax > ay;
233 const T p = real_larger ? ax : ay;
234 const T q = real_larger ? ay : ax;
237 if (EIGEN_PREDICT_FALSE(!(p > T(2) * (numext::numeric_limits<T>::min)() && p <= NumTraits<T>::highest() / T(4)))) {
238 return complex_sqrt_extreme<ComplexT, Reciprocal>(z);
241 const T abs_z = p * numext::sqrt(T(1) + r * r);
242 const T sum = ax + abs_z;
243 const T w = numext::sqrt(T(0.5) * sum);
246 EIGEN_IF_CONSTEXPR (Reciprocal) {
249 minor = (ay / sum) * major;
251 minor = ay / (T(2) * w);
253 if (numext::is_exactly_zero(x)) minor = major;
254 const T imag_sign = Reciprocal ? -y : y;
255 return x < T(0) ? ComplexT(minor, numext::copysign(major, imag_sign))
256 : ComplexT(major, numext::copysign(minor, imag_sign));
260template <
typename ComplexT>
261EIGEN_DEVICE_FUNC
constexpr ComplexT complex_sqrt(
const ComplexT& z) {
262 return complex_sqrt_impl<ComplexT, false>(z);
265template <
typename ComplexT>
266EIGEN_DEVICE_FUNC
constexpr ComplexT complex_rsqrt(
const ComplexT& z) {
267 return complex_sqrt_impl<ComplexT, true>(z);
270template <
typename ComplexT>
271EIGEN_DEVICE_FUNC
constexpr ComplexT complex_log(
const ComplexT& z) {
273 using T =
typename NumTraits<ComplexT>::Real;
274 T a = numext::abs(z);
275 EIGEN_USING_STD(atan2);
276 T b = atan2(z.imag(), z.real());
277 return ComplexT(numext::log(a), b);