Eigen  5.0.1
 
Loading...
Searching...
No Matches
GenericPacketMathPow.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2018-2025 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_ARCH_GENERIC_PACKET_MATH_POW_H
12#define EIGEN_ARCH_GENERIC_PACKET_MATH_POW_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18namespace internal {
19
20//----------------------------------------------------------------------
21// Cubic Root Functions
22//----------------------------------------------------------------------
23
24// This function implements a single step of Halley's iteration for
25// computing x = y^(1/3):
26// x_{k+1} = x_k - (x_k^3 - y) x_k / (2x_k^3 + y)
27template <typename Packet>
28EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet cbrt_halley_iteration_step(const Packet& x_k,
29 const Packet& y) {
30 using Scalar = typename unpacket_traits<Packet>::type;
31 Packet x_k_cb = pmul(x_k, pmul(x_k, x_k));
32 Packet denom = pmadd(pset1<Packet>(Scalar(2)), x_k_cb, y);
33 Packet num = psub(x_k_cb, y);
34 Packet r = pdiv(num, denom);
35 return pnmadd(x_k, r, x_k);
36}
37
38// Decompose the input such that x^(1/3) = y^(1/3) * 2^e_div3, and y is in the
39// interval [0.125,1].
40template <typename Packet>
41EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet cbrt_decompose(const Packet& x, Packet& e_div3) {
42 using Scalar = typename unpacket_traits<Packet>::type;
43 // Extract the significand s in the range [0.5,1) and exponent e, such that
44 // x = 2^e * s.
45 Packet e, s;
46 s = pfrexp(x, e);
47
48 // Split the exponent into a part divisible by 3 and the remainder.
49 // e = 3*e_div3 + e_mod3.
50 constexpr Scalar kOneThird = Scalar(1) / 3;
51 e_div3 = pceil(pmul(e, pset1<Packet>(kOneThird)));
52 Packet e_mod3 = pnmadd(pset1<Packet>(Scalar(3)), e_div3, e);
53
54 // Replace s by y = (s * 2^e_mod3).
55 return pldexp_fast(s, e_mod3);
56}
57
58template <typename Packet>
59EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet cbrt_special_cases_and_sign(const Packet& x,
60 const Packet& abs_root) {
61 // Set sign.
62 const Packet sign_mask = psignmask<Packet>();
63 const Packet x_sign = pand(sign_mask, x);
64 Packet root = por(x_sign, abs_root);
65
66 // Pass non-finite and zero values of x straight through.
67 const Packet is_not_finite = por(pisinf(x), pisnan(x));
68 const Packet is_zero = pcmp_eq(pzero(x), x);
69 const Packet use_x = por(is_not_finite, is_zero);
70 return pselect(use_x, x, root);
71}
72
73// Generic implementation of cbrt(x) for float.
74//
75// The algorithm computes the cubic root of the input by first
76// decomposing it into an exponent and significand
77// x = s * 2^e.
78//
79// We can then write the cube root as
80//
81// x^(1/3) = 2^(e/3) * s^(1/3)
82// = 2^((3*e_div3 + e_mod3)/3) * s^(1/3)
83// = 2^(e_div3) * 2^(e_mod3/3) * s^(1/3)
84// = 2^(e_div3) * (s * 2^e_mod3)^(1/3)
85//
86// where e_div3 = ceil(e/3) and e_mod3 = e - 3*e_div3.
87//
88// The cube root of the second term y = (s * 2^e_mod3)^(1/3) is coarsely
89// approximated using a cubic polynomial and subsequently refined using a
90// single step of Halley's iteration, and finally the two terms are combined
91// using pldexp_fast.
92//
93// Note: Many alternatives exist for implementing cbrt. See, for example,
94// the excellent discussion in Kahan's note:
95// https://csclub.uwaterloo.ca/~pbarfuss/qbrt.pdf
96// This particular implementation was found to be very fast and accurate
97// among several alternatives tried, but is probably not "optimal" on all
98// platforms.
99//
100// This is accurate to 2 ULP.
101template <typename Packet>
102EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet pcbrt_float(const Packet& x) {
103 using Scalar = typename unpacket_traits<Packet>::type;
104 static_assert(std::is_same<Scalar, float>::value, "Scalar type must be float");
105
106 // Decompose the input such that x^(1/3) = y^(1/3) * 2^e_div3, and y is in the
107 // interval [0.125,1].
108 Packet e_div3;
109 const Packet y = cbrt_decompose(pabs(x), e_div3);
110
111 // Compute initial approximation accurate to 5.22e-3.
112 // The polynomial was computed using Rminimax.
113 constexpr float alpha[] = {5.9220016002655029296875e-01f, -1.3859539031982421875e+00f, 1.4581282138824462890625e+00f,
114 3.408401906490325927734375e-01f};
115 Packet r = ppolevl<Packet, 3>::run(y, alpha);
116
117 // Take one step of Halley's iteration.
118 r = cbrt_halley_iteration_step(r, y);
119
120 // Finally multiply by 2^(e_div3)
121 r = pldexp_fast(r, e_div3);
122
123 return cbrt_special_cases_and_sign(x, r);
124}
125
126// Generic implementation of cbrt(x) for double.
127//
128// The algorithm is identical to the one for float except that a different initial
129// approximation is used for y^(1/3) and two Halley iteration steps are performed.
130//
131// This is accurate to 1 ULP.
132template <typename Packet>
133EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS Packet pcbrt_double(const Packet& x) {
134 using Scalar = typename unpacket_traits<Packet>::type;
135 static_assert(std::is_same<Scalar, double>::value, "Scalar type must be double");
136
137 // Decompose the input such that x^(1/3) = y^(1/3) * 2^e_div3, and y is in the
138 // interval [0.125,1].
139 Packet e_div3;
140 const Packet y = cbrt_decompose(pabs(x), e_div3);
141
142 // Compute initial approximation accurate to 0.016.
143 // The polynomial was computed using Rminimax.
144 constexpr double alpha[] = {-4.69470621553356115551736138513660989701747894287109375e-01,
145 1.072314636518546304699839311069808900356292724609375e+00,
146 3.81249427609571867048288140722434036433696746826171875e-01};
147 Packet r = ppolevl<Packet, 2>::run(y, alpha);
148
149 // Take two steps of Halley's iteration.
150 r = cbrt_halley_iteration_step(r, y);
151 r = cbrt_halley_iteration_step(r, y);
152
153 // Finally multiply by 2^(e_div3).
154 r = pldexp_fast(r, e_div3);
155 return cbrt_special_cases_and_sign(x, r);
156}
157
158//----------------------------------------------------------------------
159// Power Functions (accurate_log2, generic_pow, unary_pow)
160//----------------------------------------------------------------------
161
162// This function computes log2(x) and returns the result as a double word.
163template <typename Scalar>
164struct accurate_log2 {
165 template <typename Packet>
166 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(const Packet& x, Packet& log2_x_hi, Packet& log2_x_lo) const {
167 log2_x_hi = plog2(x);
168 log2_x_lo = pzero(x);
169 }
170};
171
172// This specialization uses a more accurate algorithm to compute log2(x) for
173// floats in [1/sqrt(2);sqrt(2)] with a relative accuracy of ~6.56508e-10.
174// This additional accuracy is needed to counter the error-magnification
175// inherent in multiplying by a potentially large exponent in pow(x,y).
176// The minimax polynomial used was calculated using the Rminimax tool,
177// see https://gitlab.inria.fr/sfilip/rminimax.
178// Command line:
179// $ ratapprox --function="log2(1+x)/x" --dom='[-0.2929,0.41422]'
180// --type=[10,0]
181// --numF="[D,D,SG]" --denF="[SG]" --log --dispCoeff="dec"
182//
183// The resulting implementation of pow(x,y) is accurate to 3 ulps.
184template <>
185struct accurate_log2<float> {
186 template <typename Packet>
187 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(const Packet& z, Packet& log2_x_hi, Packet& log2_x_lo) const {
188 // Split the two lowest order constant coefficient into double-word representation.
189 constexpr double kC0 = 1.442695041742110273474963832995854318141937255859375e+00;
190 constexpr float kC0_hi = static_cast<float>(kC0);
191 constexpr float kC0_lo = static_cast<float>(kC0 - static_cast<double>(kC0_hi));
192 const Packet c0_hi = pset1<Packet>(kC0_hi);
193 const Packet c0_lo = pset1<Packet>(kC0_lo);
194
195 constexpr double kC1 = -7.2134751588268664068692714863573201000690460205078125e-01;
196 constexpr float kC1_hi = static_cast<float>(kC1);
197 constexpr float kC1_lo = static_cast<float>(kC1 - static_cast<double>(kC1_hi));
198 const Packet c1_hi = pset1<Packet>(kC1_hi);
199 const Packet c1_lo = pset1<Packet>(kC1_lo);
200
201 constexpr float c[] = {
202 9.7010828554630279541015625e-02, -1.6896486282348632812500000e-01, 1.7200836539268493652343750e-01,
203 -1.7892272770404815673828125e-01, 2.0505344867706298828125000e-01, -2.4046677350997924804687500e-01,
204 2.8857553005218505859375000e-01, -3.6067414283752441406250000e-01, 4.8089790344238281250000000e-01};
205
206 // Evaluate the higher order terms in the polynomial using
207 // standard arithmetic.
208 const Packet one = pset1<Packet>(1.0f);
209 const Packet x = psub(z, one);
210 Packet p = ppolevl<Packet, 8>::run(x, c);
211 // Evaluate the final two steps in Horner's rule using double-word
212 // arithmetic.
213 Packet p_hi, p_lo;
214 twoprod(x, p, p_hi, p_lo);
215 fast_twosum(c1_hi, c1_lo, p_hi, p_lo, p_hi, p_lo);
216 twoprod(p_hi, p_lo, x, p_hi, p_lo);
217 fast_twosum(c0_hi, c0_lo, p_hi, p_lo, p_hi, p_lo);
218 // Multiply by x to recover log2(z).
219 twoprod(p_hi, p_lo, x, log2_x_hi, log2_x_lo);
220 }
221};
222
223// This specialization uses a more accurate algorithm to compute log2(x) for
224// floats in [1/sqrt(2);sqrt(2)] with a relative accuracy of ~1.27e-18.
225// This additional accuracy is needed to counter the error-magnification
226// inherent in multiplying by a potentially large exponent in pow(x,y).
227// The minimax polynomial used was calculated using the Sollya tool.
228// See sollya.org.
229
230template <>
231struct accurate_log2<double> {
232 template <typename Packet>
233 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(const Packet& x, Packet& log2_x_hi, Packet& log2_x_lo) const {
234 // We use a transformation of variables:
235 // r = c * (x-1) / (x+1),
236 // such that
237 // log2(x) = log2((1 + r/c) / (1 - r/c)) = f(r).
238 // The function f(r) can be approximated well using an odd polynomial
239 // of the form
240 // P(r) = ((Q(r^2) * r^2 + C) * r^2 + 1) * r,
241 // For the implementation of log2<double> here, Q is of degree 6 with
242 // coefficient represented in working precision (double), while C is a
243 // constant represented in extra precision as a double word to achieve
244 // full accuracy.
245 //
246 // The polynomial coefficients were computed by the Sollya script:
247 //
248 // c = 2 / log(2);
249 // trans = c * (x-1)/(x+1);
250 // itrans = (1+x/c)/(1-x/c);
251 // interval=[trans(sqrt(0.5)); trans(sqrt(2))];
252 // print(interval);
253 // f = log2(itrans(x));
254 // p=fpminimax(f,[|1,3,5,7,9,11,13,15,17|],[|1,DD,double...|],interval,relative,floating);
255 const Packet q12 = pset1<Packet>(2.87074255468000586e-9);
256 const Packet q10 = pset1<Packet>(2.38957980901884082e-8);
257 const Packet q8 = pset1<Packet>(2.31032094540014656e-7);
258 const Packet q6 = pset1<Packet>(2.27279857398537278e-6);
259 const Packet q4 = pset1<Packet>(2.31271023278625638e-5);
260 const Packet q2 = pset1<Packet>(2.47556738444535513e-4);
261 const Packet q0 = pset1<Packet>(2.88543873228900172e-3);
262 const Packet C_hi = pset1<Packet>(0.0400377511598501157);
263 const Packet C_lo = pset1<Packet>(-4.77726582251425391e-19);
264 const Packet one = pset1<Packet>(1.0);
265
266 const Packet cst_2_log2e_hi = pset1<Packet>(2.88539008177792677);
267 const Packet cst_2_log2e_lo = pset1<Packet>(4.07660016854549667e-17);
268 // c * (x - 1)
269 Packet t_hi, t_lo;
270 // t = c * (x-1)
271 twoprod(cst_2_log2e_hi, cst_2_log2e_lo, psub(x, one), t_hi, t_lo);
272 // r = c * (x-1) / (x+1),
273 Packet r_hi, r_lo;
274 doubleword_div_fp(t_hi, t_lo, padd(x, one), r_hi, r_lo);
275
276 // r2 = r * r
277 Packet r2_hi, r2_lo;
278 twoprod(r_hi, r_lo, r_hi, r_lo, r2_hi, r2_lo);
279 // r4 = r2 * r2
280 Packet r4_hi, r4_lo;
281 twoprod(r2_hi, r2_lo, r2_hi, r2_lo, r4_hi, r4_lo);
282
283 // Evaluate Q(r^2) in working precision. We evaluate it in two parts
284 // (even and odd in r^2) to improve instruction level parallelism.
285 Packet q_even = pmadd(q12, r4_hi, q8);
286 Packet q_odd = pmadd(q10, r4_hi, q6);
287 q_even = pmadd(q_even, r4_hi, q4);
288 q_odd = pmadd(q_odd, r4_hi, q2);
289 q_even = pmadd(q_even, r4_hi, q0);
290 Packet q = pmadd(q_odd, r2_hi, q_even);
291
292 // Now evaluate the low order terms of P(x) in double word precision.
293 // In the following, due to the increasing magnitude of the coefficients
294 // and r being constrained to [-0.5, 0.5] we can use fast_twosum instead
295 // of the slower twosum.
296 // Q(r^2) * r^2
297 Packet p_hi, p_lo;
298 twoprod(r2_hi, r2_lo, q, p_hi, p_lo);
299 // Q(r^2) * r^2 + C
300 Packet p1_hi, p1_lo;
301 fast_twosum(C_hi, C_lo, p_hi, p_lo, p1_hi, p1_lo);
302 // (Q(r^2) * r^2 + C) * r^2
303 Packet p2_hi, p2_lo;
304 twoprod(r2_hi, r2_lo, p1_hi, p1_lo, p2_hi, p2_lo);
305 // ((Q(r^2) * r^2 + C) * r^2 + 1)
306 Packet p3_hi, p3_lo;
307 fast_twosum(one, p2_hi, p2_lo, p3_hi, p3_lo);
308
309 // log2(x) ~= ((Q(r^2) * r^2 + C) * r^2 + 1) * r
310 twoprod(p3_hi, p3_lo, r_hi, r_lo, log2_x_hi, log2_x_lo);
311 }
312};
313
314// This function implements the non-trivial case of pow(x,y) where x is
315// positive and y is (possibly) non-integer.
316// Formally, pow(x,y) = exp2(y * log2(x)), where exp2(x) is shorthand for 2^x.
317// TODO(rmlarsen): We should probably add this as a packet op 'ppow', to make it
318// easier to specialize or turn off for specific types and/or backends.
319template <typename Packet>
320EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet generic_pow_impl(const Packet& x, const Packet& y) {
321 using Scalar = typename unpacket_traits<Packet>::type;
322 // Split x into exponent e_x and mantissa m_x.
323 Packet e_x;
324 Packet m_x = pfrexp(x, e_x);
325
326 // Adjust m_x to lie in [1/sqrt(2):sqrt(2)] to minimize absolute error in log2(m_x).
327 constexpr Scalar sqrt_half = Scalar(0.70710678118654752440);
328 const Packet m_x_scale_mask = pcmp_lt(m_x, pset1<Packet>(sqrt_half));
329 m_x = pselect(m_x_scale_mask, pmul(pset1<Packet>(Scalar(2)), m_x), m_x);
330 e_x = pselect(m_x_scale_mask, psub(e_x, pset1<Packet>(Scalar(1))), e_x);
331
332 // Compute log2(m_x) with 6 extra bits of accuracy.
333 Packet rx_hi, rx_lo;
334 accurate_log2<Scalar>()(m_x, rx_hi, rx_lo);
335
336 // Compute the two terms {y * e_x, y * r_x} in f = y * log2(x) with doubled
337 // precision using double word arithmetic.
338 Packet f1_hi, f1_lo, f2_hi, f2_lo;
339 twoprod(e_x, y, f1_hi, f1_lo);
340 twoprod(rx_hi, rx_lo, y, f2_hi, f2_lo);
341 // Sum the two terms in f using double word arithmetic. We know
342 // that |e_x| > |log2(m_x)|, except for the case where e_x==0.
343 // This means that we can use fast_twosum(f1,f2).
344 // In the case e_x == 0, e_x * y = f1 = 0, so we don't lose any
345 // accuracy by violating the assumption of fast_twosum, because
346 // it's a no-op.
347 Packet f_hi, f_lo;
348 fast_twosum(f1_hi, f1_lo, f2_hi, f2_lo, f_hi, f_lo);
349
350 // Split f into integer and fractional parts.
351 Packet n_z, r_z;
352 absolute_split(f_hi, n_z, r_z);
353 r_z = padd(r_z, f_lo);
354 Packet n_r;
355 absolute_split(r_z, n_r, r_z);
356 n_z = padd(n_z, n_r);
357
358 // We now have an accurate split of f = n_z + r_z and can compute
359 // x^y = 2**{n_z + r_z) = exp2(r_z) * 2**{n_z}.
360 // Multiplication by the second factor can be done exactly using pldexp(), since
361 // it is an integer power of 2.
362 const Packet e_r = generic_exp2_reduced(r_z);
363
364 // Since we know that e_r is in [1/sqrt(2); sqrt(2)], we can use the fast version
365 // of pldexp to multiply by 2**{n_z} when |n_z| is sufficiently small.
366 constexpr Scalar kPldExpThresh = std::numeric_limits<Scalar>::max_exponent - 2;
367 const Packet pldexp_fast_unsafe = pcmp_lt(pset1<Packet>(kPldExpThresh), pabs(n_z));
368 if (predux_any(pldexp_fast_unsafe)) {
369 return pldexp(e_r, n_z);
370 }
371 return pldexp_fast(e_r, n_z);
372}
373
374// Generic implementation of pow(x,y).
375template <typename Packet>
376EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS std::enable_if_t<!is_scalar<Packet>::value, Packet> generic_pow(
377 const Packet& x, const Packet& y) {
378 using Scalar = typename unpacket_traits<Packet>::type;
379
380 const Packet cst_inf = pinf<Packet>();
381 const Packet cst_zero = pset1<Packet>(Scalar(0));
382 const Packet cst_one = pset1<Packet>(Scalar(1));
383
384 const Packet x_abs = pabs(x);
385 Packet result = generic_pow_impl(x_abs, y);
386
387 // In the following we enforce the special case handling prescribed in
388 // https://en.cppreference.com/w/cpp/numeric/math/pow.
389
390 // Predicates for sign and magnitude of x.
391 const Packet x_is_negative = pcmp_lt(x, cst_zero);
392 const Packet x_is_zero = pcmp_eq(x, cst_zero);
393 const Packet x_is_one = pcmp_eq(x, cst_one);
394 const Packet x_has_signbit = psignbit(x);
395 const Packet x_abs_gt_one = pcmp_lt(cst_one, x_abs);
396 const Packet x_abs_is_inf = pcmp_eq(x_abs, cst_inf);
397
398 // Predicates for sign and magnitude of y.
399 const Packet y_abs = pabs(y);
400 const Packet y_abs_is_inf = pcmp_eq(y_abs, cst_inf);
401 const Packet y_is_negative = pcmp_lt(y, cst_zero);
402 const Packet y_is_zero = pcmp_eq(y, cst_zero);
403 const Packet y_is_one = pcmp_eq(y, cst_one);
404 // Predicates for whether y is integer and odd/even.
405 const Packet y_is_int = pandnot(pcmp_eq(pfloor(y), y), y_abs_is_inf);
406 const Packet y_div_2 = pmul(y, pset1<Packet>(Scalar(0.5)));
407 const Packet y_is_even = pcmp_eq(pround(y_div_2), y_div_2);
408 const Packet y_is_odd_int = pandnot(y_is_int, y_is_even);
409 // Smallest exponent for which (1 + epsilon) overflows to infinity.
410 constexpr Scalar huge_exponent =
411 (NumTraits<Scalar>::max_exponent() * Scalar(EIGEN_LN2)) / NumTraits<Scalar>::epsilon();
412 const Packet y_abs_is_huge = pcmp_le(pset1<Packet>(huge_exponent), y_abs);
413
414 // * pow(base, exp) returns NaN if base is finite and negative
415 // and exp is finite and non-integer.
416 // A packet comparison mask is all ones, a NaN bit pattern, so por selects NaN and pand selects a constant
417 // against zero. Scalar masks hold the value one instead, but this overload takes packets only.
418 result = por(pandnot(x_is_negative, y_is_int), result);
419
420 // * pow(±0, exp), where exp is negative, finite, and is an even integer or
421 // a non-integer, returns +∞
422 // * pow(±0, exp), where exp is positive non-integer or a positive even
423 // integer, returns +0
424 // * pow(+0, exp), where exp is a negative odd integer, returns +∞
425 // * pow(-0, exp), where exp is a negative odd integer, returns -∞
426 // * pow(+0, exp), where exp is a positive odd integer, returns +0
427 // * pow(-0, exp), where exp is a positive odd integer, returns -0
428 // Sign is flipped by the rule below.
429 result = pselect(x_is_zero, pand(y_is_negative, cst_inf), result);
430
431 // pow(base, exp) returns -pow(abs(base), exp) if base has the sign bit set,
432 // and exp is an odd integer exponent.
433 result = pselect(pand(x_has_signbit, y_is_odd_int), pnegate(result), result);
434
435 // * pow(base, -∞) returns +∞ for any |base|<1
436 // * pow(base, -∞) returns +0 for any |base|>1
437 // * pow(base, +∞) returns +0 for any |base|<1
438 // * pow(base, +∞) returns +∞ for any |base|>1
439 // * pow(±0, -∞) returns +∞
440 // * pow(-1, +-∞) = 1
441 Packet inf_y_val = pand(pxor(y_is_negative, x_abs_gt_one), cst_inf);
442 inf_y_val = pselect(pcmp_eq(x, pset1<Packet>(Scalar(-1.0))), cst_one, inf_y_val);
443 result = pselect(y_abs_is_huge, inf_y_val, result);
444
445 // * pow(+∞, exp) returns +0 for any negative exp
446 // * pow(+∞, exp) returns +∞ for any positive exp
447 // * pow(-∞, exp) returns -0 if exp is a negative odd integer.
448 // * pow(-∞, exp) returns +0 if exp is a negative non-integer or negative
449 // even integer.
450 // * pow(-∞, exp) returns -∞ if exp is a positive odd integer.
451 // * pow(-∞, exp) returns +∞ if exp is a positive non-integer or positive
452 // even integer.
453 auto x_pos_inf_value = pandnot(cst_inf, y_is_negative);
454 auto x_neg_inf_value = pselect(y_is_odd_int, pnegate(x_pos_inf_value), x_pos_inf_value);
455 result = pselect(x_abs_is_inf, pselect(x_is_negative, x_neg_inf_value, x_pos_inf_value), result);
456
457 // All cases of NaN inputs return NaN, except the two below.
458 result = por(por(pisnan(x), pisnan(y)), result);
459
460 // * pow(base, 1) returns base.
461 // * pow(base, +/-0) returns 1, regardless of base, even NaN.
462 // * pow(+1, exp) returns 1, regardless of exponent, even NaN.
463 result = pselect(y_is_one, x, pselect(por(x_is_one, y_is_zero), cst_one, result));
464
465 return result;
466}
467
468template <typename Scalar>
469EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS std::enable_if_t<is_scalar<Scalar>::value, Scalar> generic_pow(
470 const Scalar& x, const Scalar& y) {
471 return numext::pow(x, y);
472}
473
474namespace unary_pow {
475
476// Integer exponents up to this magnitude use repeated squaring; larger ones take the log/exp path of generic_pow
477// where that can represent them. The crossover is where a squaring step per exponent bit stops being cheaper
478// than generic_pow, whose cost per element is about four times higher for double than for float (half the
479// lanes, longer polynomials): measured at 2^12 for float and 2^20 for double on AVX2 with FMA. Complex bases
480// take their real scalar's value; it only limits their scalar path, since no vectorized complex pow exists.
481template <typename Scalar>
482constexpr numext::uint64_t max_squaring_exponent() {
483 return numext::uint64_t(1) << (std::is_same<typename NumTraits<Scalar>::Real, double>::value ? 20 : 12);
484}
485
486template <typename ScalarExponent, bool IsInteger = NumTraits<ScalarExponent>::IsInteger>
487struct exponent_helper {
488 using safe_abs_type = numext::uint64_t;
489 // this routine assumes that exp is an integer of magnitude at most max_squaring_exponent() stored as a floating
490 // point type
491 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE safe_abs_type safe_abs(const ScalarExponent& exp) {
492 eigen_assert(((numext::isfinite)(exp) && exp == numext::floor(exp)) && "exp must be an integer");
493 return static_cast<safe_abs_type>(numext::abs(exp));
494 }
495};
496
497template <typename ScalarExponent>
498struct exponent_helper<ScalarExponent, true> {
499 // if `exp` is a signed integer type, cast it to its unsigned counterpart to safely store its absolute value
500 // consider the (rare) case where `exp` is an int32_t: abs(-2147483648) != 2147483648
501 using safe_abs_type = typename numext::get_integer_by_size<sizeof(ScalarExponent)>::unsigned_type;
502 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE safe_abs_type safe_abs(const ScalarExponent& exp) {
503 ScalarExponent mask = numext::signbit(exp);
504 safe_abs_type result = safe_abs_type(exp ^ mask);
505 return result + safe_abs_type(ScalarExponent(1) & mask);
506 }
507};
508
509template <typename ScalarExponent, bool IsSigned = NumTraits<ScalarExponent>::IsSigned>
510struct exponent_is_negative {
511 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const ScalarExponent& exponent) {
512 return exponent < ScalarExponent(0);
513 }
514};
515
516template <typename ScalarExponent>
517struct exponent_is_negative<ScalarExponent, false> {
518 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const ScalarExponent&) { return false; }
519};
520
521// Left-to-right binary exponentiation, so that the base is multiplied in as is: each step squares the running
522// power and multiplies by the base when the next exponent bit is set.
523template <typename AbsExponentType>
524EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE AbsExponentType highest_set_bit(AbsExponentType m) {
525 AbsExponentType bit = AbsExponentType(1);
526 while ((m >> 1) >= bit) bit <<= 1;
527 return bit;
528}
529
530// Repeated squaring for integer bases, wrapping on overflow like the underlying multiplication.
531template <typename Packet, typename ScalarExponent>
532EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet int_pow_wrapping(const Packet& x, const ScalarExponent& exponent) {
533 using Scalar = typename unpacket_traits<Packet>::type;
534 using ExponentHelper = exponent_helper<ScalarExponent>;
535 using AbsExponentType = typename ExponentHelper::safe_abs_type;
536 Packet cst_pos_one = pset1<Packet>(Scalar(1));
537 if (exponent == ScalarExponent(0)) return cst_pos_one;
538 eigen_assert(!exponent_is_negative<ScalarExponent>::run(exponent));
539
540 AbsExponentType m = ExponentHelper::safe_abs(exponent);
541 Packet y = x;
542 for (AbsExponentType bit = highest_set_bit(m) >> 1; bit != 0; bit >>= 1) {
543 y = pmul(y, y);
544 if ((m & bit) != 0) y = pmul(y, x);
545 }
546 return y;
547}
548
549// Double-word repeated squaring needs an exact two-product and the binary layout below; it serves binary32 and
550// binary64 bases, real or complex. Other bases keep plain repeated squaring.
551template <typename Scalar>
552struct is_double_word_base
553 : bool_constant<(std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value) &&
554 std::numeric_limits<Scalar>::is_iec559> {};
555template <typename RealScalar>
556struct is_double_word_base<std::complex<RealScalar>> : is_double_word_base<RealScalar> {};
557
558// Whether exponent_bits_shift_right/exponent_bits_sub are available, not whether integer_packet itself is:
559// AVX/PacketMath.h supplies all three for Packet4d without AVX2, where there is no Packet4l.
560template <typename Packet, typename>
561struct has_exponent_bit_ops : false_type {};
562template <typename Packet>
563struct has_exponent_bit_ops<Packet, void_t<typename unpacket_traits<Packet>::integer_packet>> : true_type {};
564
565template <typename Packet, bool IsComplex = NumTraits<typename unpacket_traits<Packet>::type>::IsComplex,
566 bool IsScalar = is_scalar<Packet>::value>
567struct real_view {
568 using type = Packet;
569};
570template <typename Packet>
571struct real_view<Packet, true, false> {
572 using type = typename unpacket_traits<Packet>::as_real;
573};
574template <typename Scalar>
575struct real_view<Scalar, true, true> {
576 using type = typename NumTraits<Scalar>::Real;
577};
578
579template <typename Packet>
580struct use_double_word : bool_constant<is_double_word_base<typename unpacket_traits<Packet>::type>::value &&
581 has_exponent_bit_ops<typename real_view<Packet>::type>::value> {};
582
583// ARMv7 NEON flushes subnormal float operands and results to zero in its arithmetic and comparisons; the scalar VFP
584// unit, which runs the scalar path, does not. Its dispatch is also under #if EIGEN_ARCH_ARM: elsewhere, although it
585// selects the unchanged path, GCC inlines the pow kernels differently.
586template <typename Packet>
587struct flushes_subnormals
588 : bool_constant<EIGEN_ARCH_ARM && !is_scalar<Packet>::value &&
589 std::is_same<typename NumTraits<typename unpacket_traits<Packet>::type>::Real, float>::value> {};
590
591// Declared in GenericPacketMathFunctionsFwd.h so that a backend header can override them.
592template <typename Packet>
593EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet exponent_bits_shift_right(const Packet& bits) {
594 using PacketI = typename unpacket_traits<Packet>::integer_packet;
595 constexpr int kMantissaBits = numext::numeric_limits<typename unpacket_traits<Packet>::type>::digits - 1;
596 return preinterpret<Packet>(plogical_shift_right<kMantissaBits>(preinterpret<PacketI>(bits)));
597}
598template <typename Packet>
599EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet exponent_bits_sub(const Packet& a_bits, const Packet& b_bits) {
600 using PacketI = typename unpacket_traits<Packet>::integer_packet;
601 return preinterpret<Packet>(psub(preinterpret<PacketI>(a_bits), preinterpret<PacketI>(b_bits)));
602}
603
604// Power-of-two scaling of a real packet through its exponent bits. A finite value's magnitude is never let far
605// from one, so the residuals of the double-word arithmetic stay normal (a subnormal residual costs a microcode
606// assist on x86 for every operation that touches it) and no intermediate overflows or underflows.
607template <typename Packet>
608struct binary_exponent_scaling {
609 using Scalar = typename unpacket_traits<Packet>::type;
610 using Bits = std::make_unsigned_t<typename make_integer<Scalar>::type>;
611 static constexpr int kMantissaBits = numext::numeric_limits<Scalar>::digits - 1;
612 static constexpr int kBias = numext::numeric_limits<Scalar>::max_exponent - 1;
613 static constexpr Bits kExponentMask = ((Bits(1) << (CHAR_BIT * sizeof(Scalar) - kMantissaBits - 1)) - Bits(1))
614 << kMantissaBits;
615 // 2^kMantissaBits: or-ing an integer below 2^kMantissaBits into its low bits adds that integer to it.
616 static constexpr Bits kMagicBits = Bits(kBias + kMantissaBits) << kMantissaBits;
617
618 // For a normal x = m * 2^e with 2 <= |m| < 4, returns 2^-e and sets e (as a floating-point value). The
619 // target [2, 4) rather than [1, 2) keeps 2^-e a normal number for every x.
620 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet inverse_scale(const Packet& x, Packet& e) {
621 Packet exponent_bits = pand(x, pset1frombits<Packet>(kExponentMask));
622 Packet biased = exponent_bits_shift_right(exponent_bits);
623 Packet magic = pset1frombits<Packet>(kMagicBits);
624 e = psub(por(biased, magic), padd(magic, pset1<Packet>(Scalar(kBias + 1))));
625 // 2^(bias + 1 - E) has biased exponent 2 * bias + 1 - E, in [1, 2 * bias] for a normal x.
626 Packet two_bias_plus_one = pset1frombits<Packet>(Bits(2 * kBias + 1) << kMantissaBits);
627 return exponent_bits_sub(two_bias_plus_one, exponent_bits);
628 }
629
630 // x * 2^e for the scaled power, whose |x| lies within 2^(+-62): beyond the clamp the result is infinite or zero
631 // either way, and the scalar pldexp converts the exponent to int. Below kZeroBelow the result rounds to zero,
632 // which multiplying by zero gives exactly, whereas pldexp would underflow through several subnormal
633 // intermediates, each a microcode assist on x86.
634 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet scale_result(const Packet& x, const Packet& e) {
635 constexpr int kLimit = 4 * numext::numeric_limits<Scalar>::max_exponent;
636 constexpr int kZeroBelow =
637 numext::numeric_limits<Scalar>::min_exponent - numext::numeric_limits<Scalar>::digits - 64;
638 Packet zero_below = pset1<Packet>(Scalar(kZeroBelow));
639 Packet value = pselect(pcmp_lt(e, zero_below), pmul(x, pzero(x)), x);
640 Packet exponent = pmin(pmax(e, zero_below), pset1<Packet>(Scalar(kLimit)));
641#if EIGEN_ARCH_ARM
642 return with_subnormals(value, exponent, pldexp(value, exponent), flushes_subnormals<Packet>());
643#else
644 return pldexp(value, exponent);
645#endif
646 }
647
648 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet with_subnormals(const Packet&, const Packet&, const Packet& r,
649 false_type) {
650 return r;
651 }
652 // A flushing pldexp returns zero where x * 2^e lies below the smallest normal, i.e. t = |x| * 2^(e - min_exponent +
653 // digits) < 2^kMantissaBits. Rounded once, as pldexp rounds, that is k * 2^(min_exponent - digits) for k = rint(t) <=
654 // 2^kMantissaBits, the bit pattern of its magnitude. Only lanes with |r| < min are scaled, as elsewhere the scaling
655 // can overflow and its conversion raise FE_INVALID, and of those only lanes with t < 2^kMantissaBits are rebuilt:
656 // pldexp also zeroes some normal x * 2^e through a flushed partial product (2^-124 * 2^-1 in float).
657 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet with_subnormals(const Packet& x, const Packet& e, const Packet& r,
658 true_type) {
659 using PacketI = typename unpacket_traits<Packet>::integer_packet;
660 constexpr int kShift = numext::numeric_limits<Scalar>::digits - numext::numeric_limits<Scalar>::min_exponent;
661 Packet below = pcmp_lt(pabs(r), pset1<Packet>((numext::numeric_limits<Scalar>::min)()));
662 if (!predux_any(below)) return r;
663 Packet t = pldexp(pand(below, pabs(x)), padd(e, pset1<Packet>(Scalar(kShift))));
664 Packet k = preinterpret<Packet>(pcast<Packet, PacketI>(print(t)));
665 Packet rebuilt = pand(below, pcmp_lt(t, pset1<Packet>(Scalar(Bits(1) << kMantissaBits))));
666 return pselect(rebuilt, por(k, pand(x, pset1<Packet>(Scalar(-0.0)))), r);
667 }
668
669 // Whether a lane is subnormal, tested on its bits, as a flushing comparison reads a subnormal as zero.
670 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_subnormal(const Packet& x) {
671 using PacketI = typename unpacket_traits<Packet>::integer_packet;
672 using Int = typename unpacket_traits<PacketI>::type;
673 PacketI magnitude = preinterpret<PacketI>(pandnot(x, pset1<Packet>(Scalar(-0.0))));
674 PacketI subnormal =
675 pand(pcmp_lt(pzero(magnitude), magnitude), pcmp_lt(magnitude, pset1<PacketI>(Int(1) << kMantissaBits)));
676 return predux_any(preinterpret<Packet>(subnormal));
677 }
678
679 // Factors x = m * 2^e with 2 <= |m| < 4 for a finite, nonzero x, where m = (x * lift) * scale: lift is 2^digits
680 // for a subnormal x and one otherwise, and both products are exact. Zero and infinity make m NaN (their scale
681 // is infinite, respectively zero), which the callers resolve at the end; NaN stays NaN.
682 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void input_scale(const Packet& x, Packet& lift, Packet& scale,
683 Packet& e) {
684 constexpr int kDigits = numext::numeric_limits<Scalar>::digits;
685 Packet is_subnormal = pcmp_lt(pabs(x), pset1<Packet>((numext::numeric_limits<Scalar>::min)()));
686 lift = pselect(is_subnormal, pset1<Packet>(Scalar(Bits(1) << kDigits)), pset1<Packet>(Scalar(1)));
687 scale = inverse_scale(pmul(x, lift), e);
688 e = psub(e, pselect(is_subnormal, pset1<Packet>(Scalar(kDigits)), pzero(x)));
689 }
690};
691
692// The running power of repeated squaring: a double word {hi, lo} times 2^exponent, with the exponent kept as a
693// floating-point value that is exact while the result is finite, plus what a zero or infinite base turns into:
694// itself, or its reciprocal, with the sign dropped for an even exponent. Such a base makes hi NaN from the input
695// scaling on, and a finite nonzero base cannot make hi NaN or zero since the scaled power stays within 2^(+-62),
696// so result() substitutes the special value exactly where hi is NaN or zero. (A zero hi would also lose its sign
697// in the double-word sums.) NaN needs nothing: it propagates.
698template <typename Packet, bool IsComplex = NumTraits<typename unpacket_traits<Packet>::type>::IsComplex>
699struct repeated_squaring_ops {
700 using Scalar = typename unpacket_traits<Packet>::type;
701 using Scaling = binary_exponent_scaling<Packet>;
702 using Bound = Packet;
703 struct State {
704 Packet hi, lo, exponent, special;
705 };
706 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_subnormal(const Packet& x) { return Scaling::any_subnormal(x); }
707 // Whether a finite nonzero base has a power r below the smallest normal.
708 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_below_normal(const Packet& x, const Packet& r) {
709 Packet abs_x = pabs(x);
710 Packet regular = pand(pcmp_lt(pzero(x), abs_x), pcmp_lt(abs_x, pset1<Packet>(NumTraits<Scalar>::infinity())));
711 return predux_any(pand(regular, pcmp_lt(pabs(r), pset1<Packet>((numext::numeric_limits<Scalar>::min)()))));
712 }
713 // Whether every lane lies within [1/bound, bound], where the power and its residuals stay normal without scaling.
714 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool in_range(const Packet& x, const Packet& bound) {
715 Packet abs_x = pabs(x);
716 Packet out = por(pcmp_lt(pmul(abs_x, bound), pset1<Packet>(Scalar(1))), pcmp_lt_or_nan(bound, abs_x));
717 return !predux_any(out);
718 }
719 // The double word of m or of 1/m = q + (1 - q*m)/m, where 1 - q*m is formed from the exact product
720 // q*m = p_hi + p_lo (1 - p_hi is exact as p_hi is within rounding of 1); renormalizing makes hi the correctly
721 // rounded reciprocal.
722 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void power_base(const Packet& m, bool reciprocal, Packet& hi,
723 Packet& lo) {
724 if (!reciprocal) {
725 hi = m;
726 lo = pzero(m);
727 return;
728 }
729 Packet cst_pos_one = pset1<Packet>(Scalar(1));
730 Packet q = pdiv(cst_pos_one, m);
731 Packet p_hi, p_lo;
732 twoprod(q, m, p_hi, p_lo);
733 fast_twosum(q, pdiv(psub(psub(cst_pos_one, p_hi), p_lo), m), hi, lo);
734 }
735 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE State base(const Packet& x, bool reciprocal, bool scaled,
736 const Scalar&) {
737 State b;
738 if (!scaled) {
739 power_base(x, reciprocal, b.hi, b.lo);
740 b.exponent = pzero(x);
741 b.special = x;
742 return b;
743 }
744 Packet lift, scale;
745 Scaling::input_scale(x, lift, scale, b.exponent);
746 power_base(pmul(pmul(x, lift), scale), reciprocal, b.hi, b.lo);
747 if (reciprocal) {
748 b.exponent = pnegate(b.exponent);
749 b.special = pdiv(pset1<Packet>(Scalar(1)), x);
750 } else {
751 b.special = x;
752 }
753 return b;
754 }
755 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void multiply(State& y, const State& b) {
756 fast_twoprod(y.hi, y.lo, b.hi, b.lo, y.hi, y.lo);
757 y.exponent = padd(y.exponent, b.exponent);
758 }
759 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void square(State& y) { multiply(y, y); }
760 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void renormalize(State& y) {
761 Packet e;
762 Packet scale = Scaling::inverse_scale(y.hi, e);
763 y.hi = pmul(y.hi, scale);
764 y.lo = pmul(y.lo, scale);
765 y.exponent = padd(y.exponent, e);
766 }
767 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet result(const Packet&, const State& y, const State& b, bool odd,
768 bool scaled) {
769 if (!scaled) return y.hi;
770 Packet use_special = por(pisnan(y.hi), pcmp_eq(y.hi, pzero(y.hi)));
771 return pselect(use_special, odd ? b.special : pabs(b.special), Scaling::scale_result(y.hi, y.exponent));
772 }
773 // With a scaled |base| in [1/4, 4) a step at most cubes the magnitude bound, so four steps keep it within 2^(+-62)
774 // and the residuals, u times smaller, normal.
775 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE int renormalization_steps(const Scalar&) { return 4; }
776 // Zero, infinite and NaN bases are resolved by the special value in result().
777 struct SpecialBases {
778 static constexpr bool any = false;
779 };
780 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE SpecialBases special_bases(const Packet&, bool) {
781 return SpecialBases();
782 }
783 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet without_special_bases(const Packet& x, const SpecialBases&) {
784 return x;
785 }
786 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet with_special_bases(const Packet& r, const SpecialBases&, bool) {
787 return r;
788 }
789};
790
791// The real and imaginary parts of a complex value or packet as two values of a real representation R, on which
792// the complex algorithm below runs component-wise. For a complex packet R is its interleaved real view with both
793// lanes of a pair holding the same component, so every lane operation applies to the pair at once.
794template <typename Packet, bool IsScalar = is_scalar<Packet>::value>
795struct complex_components {
796 using Scalar = typename unpacket_traits<Packet>::type;
797 using RealScalar = typename NumTraits<Scalar>::Real;
798 using R = typename unpacket_traits<Packet>::as_real;
799 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R flip(const R& x) { return pcplxflip(Packet(x)).v; }
800 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R odd_lanes() {
801 return pcmp_eq(pset1<Packet>(Scalar(0, 1)).v, pset1<R>(RealScalar(1)));
802 }
803 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void split(const Packet& z, R& re, R& im) {
804 R odd = odd_lanes();
805 re = pselect(odd, flip(z.v), z.v);
806 im = pselect(odd, z.v, flip(z.v));
807 }
808 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet join(const R& re, const R& im) {
809 return Packet(pselect(odd_lanes(), im, re));
810 }
811 // max(|re|, |im|) in both lanes of each pair.
812 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R magnitude(const Packet& z) {
813 R abs_z = pabs(z.v);
814 return pmax(abs_z, flip(abs_z));
815 }
816 // min(|re|, |im|) in both lanes of each pair.
817 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R minor_magnitude(const Packet& z) {
818 R abs_z = pabs(z.v);
819 return pmin(abs_z, flip(abs_z));
820 }
821 // Whether a component exceeds bound or is NaN, in its own lane: pmax need not propagate NaN.
822 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R exceeds(const Packet& z, const R& bound) {
823 return pcmp_lt_or_nan(bound, pabs(z.v));
824 }
825 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_zero(const Packet& z) {
826 return predux_any(pcmp_eq(z.v, pzero(z.v)));
827 }
828 // In both lanes of each pair: both components zero; a component infinite; a component NaN and none infinite.
829 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void special_masks(const Packet& z, R& zero, R& inf, R& nan) {
830 R zero_lane = pcmp_eq(z.v, pzero(z.v));
831 R inf_lane = pcmp_eq(pabs(z.v), pset1<R>(NumTraits<RealScalar>::infinity()));
832 R nan_lane = pisnan(z.v);
833 zero = pand(zero_lane, flip(zero_lane));
834 inf = por(inf_lane, flip(inf_lane));
835 nan = pandnot(por(nan_lane, flip(nan_lane)), inf);
836 }
837 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet replace(const R& mask, const Scalar& value, const Packet& z) {
838 return Packet(pselect(mask, pset1<Packet>(value).v, z.v));
839 }
840};
841
842template <typename Scalar>
843struct complex_components<Scalar, true> {
844 using R = typename NumTraits<Scalar>::Real;
845 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void split(const Scalar& z, R& re, R& im) {
846 re = numext::real(z);
847 im = numext::imag(z);
848 }
849 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Scalar join(const R& re, const R& im) { return Scalar(re, im); }
850 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R magnitude(const Scalar& z) {
851 return numext::maxi(numext::abs(numext::real(z)), numext::abs(numext::imag(z)));
852 }
853 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R minor_magnitude(const Scalar& z) {
854 return numext::mini(numext::abs(numext::real(z)), numext::abs(numext::imag(z)));
855 }
856 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE R exceeds(const Scalar& z, const R& bound) {
857 return por(pcmp_lt_or_nan(bound, numext::abs(numext::real(z))),
858 pcmp_lt_or_nan(bound, numext::abs(numext::imag(z))));
859 }
860 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_zero(const Scalar& z) {
861 return numext::real(z) == R(0) || numext::imag(z) == R(0);
862 }
863 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void special_masks(const Scalar& z, R& zero, R& inf, R& nan) {
864 R re = numext::real(z), im = numext::imag(z);
865 bool is_inf = (numext::isinf)(re) || (numext::isinf)(im);
866 zero = re == R(0) && im == R(0) ? R(1) : R(0);
867 inf = is_inf ? R(1) : R(0);
868 nan = !is_inf && ((numext::isnan)(re) || (numext::isnan)(im)) ? R(1) : R(0);
869 }
870 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Scalar replace(const R& mask, const Scalar& value, const Scalar& z) {
871 return mask != R(0) ? value : z;
872 }
873};
874
875// Complex bases keep the real and imaginary parts as separate double words sharing one exponent, scaled by the
876// larger component: z^2 = (a^2 - b^2) + 2ab i is three products and a product (a + bi)(c + di) four. Either
877// component may legitimately be zero. Zero and infinite bases are the two points 0 and infinity of the extended
878// complex plane, as std::proj sees it: every value with an infinite component, a NaN one included, is the same
879// infinity, +inf + 0i. So a zero base gives +0 + 0i for n > 0 and +inf + 0i for n < 0, an infinite base the reverse,
880// whatever the signs of their zeros; a base with a NaN component and no infinite one gives NaN + NaN i. std::pow
881// leaves all of these to the implementation, and libstdc++'s follow its evaluation order: 1/0 = inf + nan i,
882// (-inf + 2i)^3 = nan + inf i.
883template <typename Packet>
884struct repeated_squaring_ops<Packet, true> {
885 using Scalar = typename unpacket_traits<Packet>::type;
886 using Real = typename NumTraits<Scalar>::Real;
887 using Components = complex_components<Packet>;
888 using R = typename Components::R;
889 using Scaling = binary_exponent_scaling<R>;
890 using Bound = R;
891 // Lanes marked `separated` hold only the larger component L (a, or bi) of z = L (1 + delta); the power is
892 // L^n (1 + i c) with i c = n delta, c = (coefficient_hi + coefficient_lo) 2^coefficient_exponent (see base()).
893 struct State {
894 R re_hi, re_lo, im_hi, im_lo, exponent;
895 R separated, coefficient_hi, coefficient_lo, coefficient_exponent;
896 bool any_separated;
897 bool negative, any_zero;
898 };
899 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_subnormal(const Packet& x) {
900 return Scaling::any_subnormal(x.v);
901 }
902 // A complex power has no single-operation form for the scalar path to round differently.
903 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool any_below_normal(const Packet&, const Packet&) { return false; }
904 // Four steps between renormalizations let the power fall to 2^-62 of the scale when the base is a reciprocal in
905 // [1/4, 1/2), and the smaller component, of order r = |s/L| times the power, keeps normal residuals only for
906 // r >= 2^(min_exponent - 1 + digits + 62). Below that the base is separated (see base()), and first order is
907 // exact to within (n r)^2 / 2 <= u/4 for |n| <= 2^kFirstOrderLog2Limit: every n for double, 2^25 for float.
908 // Larger n renormalize at every step instead, which keeps the power above 2^-6 (a step from [2, 4), or from the
909 // base itself) and lets the separation drop 56 binades, where first order holds for any n.
910 static constexpr int kMinExponent = numext::numeric_limits<Real>::min_exponent;
911 static constexpr int kDigits = numext::numeric_limits<Real>::digits;
912 static constexpr int kSeparationExponent = kMinExponent + kDigits + 64;
913 static constexpr int kFirstOrderLog2Limit = (1 - kDigits - 2 * kSeparationExponent) / 2;
914 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool every_step(const Real& count) {
915 return kFirstOrderLog2Limit < 64 &&
916 count > Real(numext::uint64_t(1) << (kFirstOrderLog2Limit < 64 ? kFirstOrderLog2Limit : 0));
917 }
918 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE int renormalization_steps(const Real& count) {
919 return every_step(count) ? 1 : 4;
920 }
921 // Whether every lane's max(|re|, |im|), and min(|re|, |im|) unless zero, lie within [1/bound, bound], where the
922 // power and its residuals stay normal without scaling. A zero component stays exactly zero in every power.
923 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool in_range(const Packet& x, const R& bound) {
924 R magnitude = Components::magnitude(x);
925 R minor = Components::minor_magnitude(x);
926 R one = pset1<R>(Real(1));
927 R out = por(pcmp_lt(pmul(magnitude, bound), one), Components::exceeds(x, bound));
928 out = por(out, pandnot(pcmp_lt(pmul(minor, bound), one), pcmp_eq(minor, pzero(minor))));
929 return predux_any(out) == false;
930 }
931 template <typename Count>
932 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE State base(const Packet& x, bool reciprocal, bool scaled,
933 const Count& m) {
934 State b;
935 b.any_separated = false;
936 R re, im;
937 Components::split(x, re, im);
938 b.negative = reciprocal;
939 b.any_zero = Components::any_zero(x);
940 if (!scaled) {
941 power_base(re, im, reciprocal, b);
942 b.exponent = pzero(re);
943 return b;
944 }
945 R lift, scale;
946 Scaling::input_scale(Components::magnitude(x), lift, scale, b.exponent);
947 R wr = pmul(pmul(re, lift), scale), wi = pmul(pmul(im, lift), scale);
948 // Both components take the larger one's scale, under which a nonzero smaller component s far below the larger
949 // one L would lose its residuals (see kSeparationExponent). Then z = L (1 + delta) with |delta| tiny and
950 // z^n = L^n (1 + n delta) to working precision: the power runs on L alone and result() adds L^n n delta, with
951 // delta = i b/a, or a/(bi) = -i a/b, and -delta for the reciprocal base.
952 int separation_exponent = every_step(Real(m)) ? kSeparationExponent - 56 : kSeparationExponent;
953 R separation =
954 pset1frombits<R>(typename Scaling::Bits(Scaling::kBias + separation_exponent) << Scaling::kMantissaBits);
955 R zero = pzero(wr);
956 R re_separated = pandnot(pcmp_lt(pabs(wr), separation), pcmp_eq(re, zero));
957 R im_separated = pandnot(pcmp_lt(pabs(wi), separation), pcmp_eq(im, zero));
958 b.separated = por(re_separated, im_separated);
959 b.any_separated = predux_any(b.separated);
960 if (b.any_separated) {
961 R s_lift, s_scale, s_exponent, count_exponent;
962 Scaling::input_scale(pselect(im_separated, im, re), s_lift, s_scale, s_exponent);
963 R s = pmul(pmul(pselect(im_separated, im, re), s_lift), s_scale);
964 // The count n is a double word too, as Real(n) is inexact beyond 2^digits: its leading digits, which Real
965 // holds exactly, and the rest.
966 numext::uint64_t n = numext::uint64_t(m);
967 numext::uint64_t unit = highest_set_bit(n) >> (kDigits - 1);
968 numext::uint64_t n_hi = unit > 1 ? n & ~(unit - 1) : n;
969 R count_hi = pset1<R>(Real(n_hi)), count_lo = pset1<R>(Real(n - n_hi));
970 R count_scale = Scaling::inverse_scale(count_hi, count_exponent);
971 // The coefficient n s / L is a double word so that result() rounds the added component once: the ratio
972 // q + (s - q L) / L, where s - q L = (s - p_hi) - p_lo is exact, times the count. The scaled |s / L| is in
973 // (1/2, 2) and the scaled count in [2, 4); an eighth keeps |c hi| within 2^62.
974 R l = pselect(im_separated, wr, wi);
975 R q = pdiv(s, l);
976 R p_hi, p_lo, c_hi, c_lo;
977 twoprod(q, l, p_hi, p_lo);
978 R q_lo = pdiv(psub(psub(s, p_hi), p_lo), l);
979 R scaled_count = pmul(count_hi, count_scale);
980 R scaled_count_lo = pmul(count_lo, count_scale);
981 twoprod(scaled_count, q, c_hi, c_lo);
982 fast_twosum(c_hi, pmadd(scaled_count, q_lo, pmadd(scaled_count_lo, q, c_lo)), p_hi, p_lo);
983 R sign = pselect(im_separated, pset1<R>(Real(reciprocal ? -0.125 : 0.125)),
984 pset1<R>(Real(reciprocal ? 0.125 : -0.125)));
985 b.coefficient_hi = pmul(p_hi, sign);
986 b.coefficient_lo = pmul(p_lo, sign);
987 b.coefficient_exponent = padd(padd(psub(s_exponent, b.exponent), count_exponent), pset1<R>(Real(3)));
988 wr = pselect(re_separated, zero, wr);
989 wi = pselect(im_separated, zero, wi);
990 }
991 power_base(wr, wi, reciprocal, b);
992 if (reciprocal) b.exponent = pnegate(b.exponent);
993 return b;
994 }
995 // The double words of w or of 1/w = q + e*q with q = conj(w)/|w|^2 and e = 1 - q*w = er - t i. The norm cannot
996 // over- or underflow for max(|wr|, |wi|) in [2, 4) or within the unscaled range, and the few ulps of error in q
997 // are what the residual corrects: e is of order u and is formed from exact products so that e*q carries the
998 // correction to order u^2. 1 - s_hi is exact as s_hi is within rounding of 1.
999 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void power_base(const R& wr, const R& wi, bool reciprocal, State& b) {
1000 b.re_lo = b.im_lo = pzero(wr);
1001 if (!reciprocal) {
1002 b.re_hi = wr;
1003 b.im_hi = wi;
1004 return;
1005 }
1006 R inv = pdiv(pset1<R>(typename NumTraits<Scalar>::Real(1)), pmadd(wr, wr, pmul(wi, wi)));
1007 R qr = pmul(wr, inv), qi = pnegate(pmul(wi, inv));
1008 R a_hi, a_lo, c_hi, c_lo, s_hi, s_lo, t_hi, t_lo;
1009 twoprod(qr, wr, a_hi, a_lo);
1010 twoprod(qi, wi, c_hi, c_lo);
1011 twodiff(a_hi, a_lo, c_hi, c_lo, s_hi, s_lo);
1012 twoprod(qr, wi, a_hi, a_lo);
1013 twoprod(qi, wr, c_hi, c_lo);
1014 twosum(a_hi, a_lo, c_hi, c_lo, t_hi, t_lo);
1015 R er = psub(psub(pset1<R>(typename NumTraits<Scalar>::Real(1)), s_hi), s_lo);
1016 R t = padd(t_hi, t_lo);
1017 fast_twosum(qr, pmadd(er, qr, pmul(t, qi)), b.re_hi, b.re_lo);
1018 fast_twosum(qi, pmsub(er, qi, pmul(t, qr)), b.im_hi, b.im_lo);
1019 }
1020 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void square(State& y) {
1021 R a_hi, a_lo, c_hi, c_lo, p_hi, p_lo;
1022 fast_twoprod(y.re_hi, y.re_lo, y.re_hi, y.re_lo, a_hi, a_lo);
1023 fast_twoprod(y.im_hi, y.im_lo, y.im_hi, y.im_lo, c_hi, c_lo);
1024 fast_twoprod(y.re_hi, y.re_lo, y.im_hi, y.im_lo, p_hi, p_lo);
1025 twodiff(a_hi, a_lo, c_hi, c_lo, y.re_hi, y.re_lo);
1026 y.im_hi = padd(p_hi, p_hi);
1027 y.im_lo = padd(p_lo, p_lo);
1028 y.exponent = padd(y.exponent, y.exponent);
1029 }
1030 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void multiply(State& y, const State& b) {
1031 R ac_hi, ac_lo, bd_hi, bd_lo, ad_hi, ad_lo, bc_hi, bc_lo;
1032 fast_twoprod(y.re_hi, y.re_lo, b.re_hi, b.re_lo, ac_hi, ac_lo);
1033 fast_twoprod(y.im_hi, y.im_lo, b.im_hi, b.im_lo, bd_hi, bd_lo);
1034 fast_twoprod(y.re_hi, y.re_lo, b.im_hi, b.im_lo, ad_hi, ad_lo);
1035 fast_twoprod(y.im_hi, y.im_lo, b.re_hi, b.re_lo, bc_hi, bc_lo);
1036 twodiff(ac_hi, ac_lo, bd_hi, bd_lo, y.re_hi, y.re_lo);
1037 twosum(ad_hi, ad_lo, bc_hi, bc_lo, y.im_hi, y.im_lo);
1038 y.exponent = padd(y.exponent, b.exponent);
1039 }
1040 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void renormalize(State& y) {
1041 R e;
1042 R scale = Scaling::inverse_scale(pmax(pabs(y.re_hi), pabs(y.im_hi)), e);
1043 y.re_hi = pmul(y.re_hi, scale);
1044 y.re_lo = pmul(y.re_lo, scale);
1045 y.im_hi = pmul(y.im_hi, scale);
1046 y.im_lo = pmul(y.im_lo, scale);
1047 y.exponent = padd(y.exponent, e);
1048 }
1049 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet result(const Packet& x, const State& y, const State& b, bool odd,
1050 bool scaled) {
1051 if (!scaled) {
1052 R re = y.re_hi, im = y.im_hi;
1053 if (b.any_zero) zero_component_signs(x, b.negative, odd, re, im);
1054 return Components::join(re, im);
1055 }
1056 R re = Scaling::scale_result(y.re_hi, y.exponent);
1057 R im = Scaling::scale_result(y.im_hi, y.exponent);
1058 if (b.any_separated) {
1059 // L^n is real or imaginary, and i c L^n fills the component it leaves exactly zero.
1060 R e = padd(y.exponent, b.coefficient_exponent);
1061 R zero = pzero(re);
1062 R re_hi, re_lo, im_hi, im_lo;
1063 fast_twoprod(b.coefficient_hi, b.coefficient_lo, y.im_hi, y.im_lo, re_hi, re_lo);
1064 fast_twoprod(b.coefficient_hi, b.coefficient_lo, y.re_hi, y.re_lo, im_hi, im_lo);
1065 re = pselect(pand(b.separated, pcmp_eq(y.re_hi, zero)), Scaling::scale_result(pnegate(re_hi), e), re);
1066 im = pselect(pand(b.separated, pcmp_eq(y.im_hi, zero)), Scaling::scale_result(im_hi, e), im);
1067 }
1068 if (b.any_zero) zero_component_signs(x, b.negative, odd, re, im);
1069 return Components::join(re, im);
1070 }
1071 // Zero, infinite and NaN bases are replaced by 1 + i before the loop, so that they neither force the scaled loop on
1072 // their packet nor come back NaN, and their results are set at the end (see the comment above the struct). 1 + i
1073 // has no zero component, which would take the packet through zero_component_signs. Every one of them fails
1074 // in_range, so only a packet bound for the scaled loop is checked, and only whether there are any is carried through
1075 // the loop.
1076 struct SpecialBases {
1077 Packet x;
1078 bool any;
1079 };
1080 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE SpecialBases special_bases(const Packet& x, bool scaled) {
1081 SpecialBases special;
1082 special.x = x;
1083 special.any = false;
1084 if (scaled) {
1085 R zero, inf, nan;
1086 Components::special_masks(x, zero, inf, nan);
1087 special.any = predux_any(por(por(zero, inf), nan));
1088 }
1089 return special;
1090 }
1091 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet without_special_bases(const Packet& x,
1092 const SpecialBases& special) {
1093 if (!special.any) return x;
1094 R zero, inf, nan;
1095 Components::special_masks(x, zero, inf, nan);
1096 return Components::replace(por(por(zero, inf), nan), Scalar(1, 1), x);
1097 }
1098 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet with_special_bases(const Packet& r, const SpecialBases& special,
1099 bool negative) {
1100 if (!special.any) return r;
1101 R zero, inf, nan;
1102 Components::special_masks(special.x, zero, inf, nan);
1103 Real real_inf = NumTraits<Real>::infinity(), real_nan = NumTraits<Real>::quiet_NaN();
1104 Packet out = Components::replace(negative ? zero : inf, Scalar(real_inf, Real(0)), r);
1105 out = Components::replace(negative ? inf : zero, Scalar(0), out);
1106 return Components::replace(nan, Scalar(real_nan, real_nan), out);
1107 }
1108 // A base with one component s = +-0 is the limit of L (1 + i t), with t = s/L for a real L and -s/b for L = bi,
1109 // so z^n = L^n (1 + i n t): the zero component of the power has the sign of i n t L^n, which the double-word sums
1110 // lose (-0 + 0 = +0). Sign bits combine by xor; a scalar comparison mask is one rather than all ones, so masks
1111 // select.
1112 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void zero_component_signs(const Packet& x, bool negative, bool odd,
1113 R& re, R& im) {
1114 R zero = pzero(re);
1115 R x_re, x_im;
1116 Components::split(x, x_re, x_im);
1117 R imaginary = pcmp_eq(x_re, zero);
1118 R one_zero = pxor(imaginary, pcmp_eq(x_im, zero));
1119 R minus_zero = pset1<R>(Real(-0.0));
1120 // i^2 = -1 from t for L = bi and from i t L^n for an imaginary power cancel for odd n.
1121 R power_imaginary = odd ? imaginary : zero;
1122 R n_sign = negative ? minus_zero : zero;
1123 R extra = odd ? n_sign : pselect(imaginary, pxor(n_sign, minus_zero), n_sign);
1124 R sign = pxor(pand(pxor(pxor(x_re, x_im), pselect(power_imaginary, im, re)), minus_zero), extra);
1125 re = pselect(pand(one_zero, power_imaginary), sign, re);
1126 im = pselect(pandnot(one_zero, power_imaginary), sign, im);
1127 }
1128};
1129
1130// Plain repeated squaring for the remaining floating-point bases.
1131template <typename Packet, typename ScalarExponent>
1132EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow_plain(const Packet& x, const ScalarExponent& exponent) {
1133 using Scalar = typename unpacket_traits<Packet>::type;
1134 using ExponentHelper = exponent_helper<ScalarExponent>;
1135 using AbsExponentType = typename ExponentHelper::safe_abs_type;
1136 if (exponent == ScalarExponent(0)) return pset1<Packet>(Scalar(1));
1137
1138 Packet base = exponent_is_negative<ScalarExponent>::run(exponent) ? pdiv(pset1<Packet>(Scalar(1)), x) : x;
1139 AbsExponentType m = ExponentHelper::safe_abs(exponent);
1140 Packet y = base;
1141 for (AbsExponentType bit = highest_set_bit(m) >> 1; bit != 0; bit >>= 1) {
1142 y = pmul(y, y);
1143 if ((m & bit) != 0) y = pmul(y, base);
1144 }
1145 return y;
1146}
1147
1148// Repeated squaring for a binary32 or binary64 base, real or complex, carried in double-word arithmetic. Plain
1149// repeated squaring doubles the accumulated rounding error at every squaring, so its error grows like n * u for
1150// x^n. Keeping the running power as an unevaluated sum {hi, lo} bounds each step's relative error by about
1151// 7 * u^2 (fast_twoprod), so the result is correctly rounded unless the exact power lies within about 7 * n * u^2
1152// of a rounding boundary, and stays within 1 ulp up to n = max_squaring_exponent() for float. For a complex base
1153// the bound is normwise, relative to |z^n|, as a^2 - b^2 and ac - bd can cancel, and the n u^2 term passes u/2
1154// as |n| approaches 1/u: (0.6 + 0.8i)^n in complex<float> is 0.2 u off at n = 4096, 0.5 u at 2^24 and 4.2 u at
1155// 2^31 - 1. The power is scaled by powers of two throughout and only the final pldexp can overflow or underflow.
1156template <typename Packet, typename ScalarExponent>
1157EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow_double_word(const Packet& x, const ScalarExponent& exponent) {
1158 using Scalar = typename unpacket_traits<Packet>::type;
1159 using ExponentHelper = exponent_helper<ScalarExponent>;
1160 using AbsExponentType = typename ExponentHelper::safe_abs_type;
1161 using Ops = repeated_squaring_ops<Packet>;
1162 if (exponent == ScalarExponent(0)) return pset1<Packet>(Scalar(1));
1163
1164 bool negative = exponent_is_negative<ScalarExponent>::run(exponent);
1165 AbsExponentType m = ExponentHelper::safe_abs(exponent);
1166 bool odd = (m & AbsExponentType(1)) != 0;
1167 if (m == AbsExponentType(1) && !negative) return x;
1168 // A real x^-1 and x^2 are a single correctly rounded operation, which is what the double word would produce. On
1169 // ARMv7 NEON pdiv is a refined reciprocal estimate instead.
1170 EIGEN_IF_CONSTEXPR (!NumTraits<Scalar>::IsComplex) {
1171 if (m == AbsExponentType(1) && !flushes_subnormals<Packet>::value) return pdiv(pset1<Packet>(Scalar(1)), x);
1172 if (m == AbsExponentType(2) && !negative) return pmul(x, x);
1173 }
1174 AbsExponentType top = highest_set_bit(m);
1175
1176 // Bases whose magnitude lies within 2^(+-B) for B = budget / |n| keep every power up to x^n and its residuals
1177 // normal: |x^n| <= 2^budget < max and u * |x^n| >= 2^(min_exponent - 1), the smallest normal, so the loop
1178 // needs neither scaling nor exponents and the result is hi itself. A complex base is tested on both nonzero
1179 // components, as every term a^k b^(n-k) of its power must stay normal too; its magnitude exceeds the larger one
1180 // by up to sqrt(2), so B is one less to absorb sqrt(2)^n (an exhausted budget therefore admits nothing), and its
1181 // reciprocal forms |w|^2, which bounds B by half the exponent range. Zero, infinity and NaN fail the test.
1182 using Real = typename NumTraits<Scalar>::Real;
1183 using RealBits = std::make_unsigned_t<typename make_integer<Real>::type>;
1184 constexpr int kBudget = -(numext::numeric_limits<Real>::min_exponent + numext::numeric_limits<Real>::digits);
1185 constexpr int kMantissaBits = numext::numeric_limits<Real>::digits - 1;
1186 constexpr RealBits kBiasBits = RealBits(numext::numeric_limits<Real>::max_exponent - 1);
1187 int b = numext::uint64_t(m) > numext::uint64_t(kBudget) ? 0 : kBudget / int(m);
1188 EIGEN_IF_CONSTEXPR (NumTraits<Scalar>::IsComplex)
1189 b = numext::mini(b - 1, (numext::numeric_limits<Real>::max_exponent - 1) / 2);
1190 // A flushing packet always scales: the budget keeps the unscaled power normal, but not all of its residuals, nor the
1191 // components and the norm of a complex reciprocal.
1192 EIGEN_IF_CONSTEXPR (flushes_subnormals<Packet>::value) b = -1;
1193 typename Ops::Bound bound =
1194 pset1frombits<typename Ops::Bound>(RealBits(kBiasBits + RealBits(b < 0 ? 0 : b)) << kMantissaBits);
1195 bool scaled = b < 0 || !Ops::in_range(x, bound);
1196 // Zero, infinite and NaN bases are never in range, so only a packet bound for the scaled loop can hold one.
1197 typename Ops::SpecialBases special = Ops::special_bases(x, scaled);
1198 Packet x_regular = Ops::without_special_bases(x, special);
1199 if (special.any) scaled = b < 0 || !Ops::in_range(x_regular, bound);
1200
1201 typename Ops::State base = Ops::base(x_regular, negative, scaled, m);
1202 typename Ops::State y = base;
1203 int renormalization_steps = Ops::renormalization_steps(Real(m));
1204 int steps_since_renormalization = 0;
1205 for (AbsExponentType bit = top >> 1; bit != 0; bit >>= 1) {
1206 Ops::square(y);
1207 if ((m & bit) != 0) Ops::multiply(y, base);
1208 if (scaled && ++steps_since_renormalization == renormalization_steps) {
1209 Ops::renormalize(y);
1210 steps_since_renormalization = 0;
1211 }
1212 }
1213 return Ops::with_special_bases(Ops::result(x_regular, y, base, odd, scaled), special, negative);
1214}
1215
1216template <typename Packet, typename ScalarExponent>
1217EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE Packet int_pow_lanewise(const Packet& x, const ScalarExponent& exponent) {
1218 using Scalar = typename unpacket_traits<Packet>::type;
1219 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<Packet>::alignment) Scalar values[unpacket_traits<Packet>::size];
1220 pstore(values, x);
1221 for (Scalar& value : values) value = int_pow_double_word(value, exponent);
1222 return pload<Packet>(values);
1223}
1224
1225template <typename Packet, typename ScalarExponent>
1226EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow_double_word(const Packet& x, const ScalarExponent& exponent,
1227 false_type) {
1228 return int_pow_double_word(x, exponent);
1229}
1230
1231// A flushing packet is computed lane by lane on scalars, whose arithmetic does not flush, where it holds a subnormal
1232// base, which the comparisons, magnitudes and scalings of the base would all read as zero, or where a real x^-1 or x^2
1233// falls below the smallest normal: the scalar path rounds those once, as 1/x and x * x do.
1234template <typename Packet, typename ScalarExponent>
1235EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow_double_word(const Packet& x, const ScalarExponent& exponent,
1236 true_type) {
1237 using Ops = repeated_squaring_ops<Packet>;
1238 if (Ops::any_subnormal(x)) return int_pow_lanewise(x, exponent);
1239 Packet r = int_pow_double_word(x, exponent);
1240 if (exponent_helper<ScalarExponent>::safe_abs(exponent) <= 2 && Ops::any_below_normal(x, r))
1241 return int_pow_lanewise(x, exponent);
1242 return r;
1243}
1244
1245template <typename Packet, typename ScalarExponent>
1246EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow(const Packet& x, const ScalarExponent& exponent, true_type) {
1247#if EIGEN_ARCH_ARM
1248 return int_pow_double_word(x, exponent, flushes_subnormals<Packet>());
1249#else
1250 return int_pow_double_word(x, exponent);
1251#endif
1252}
1253
1254template <typename Packet, typename ScalarExponent>
1255EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow(const Packet& x, const ScalarExponent& exponent, false_type) {
1256 return int_pow_plain(x, exponent);
1257}
1258
1259template <typename Packet, typename ScalarExponent>
1260EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet int_pow(const Packet& x, const ScalarExponent& exponent) {
1261 return int_pow(x, exponent, use_double_word<Packet>());
1262}
1263
1264// The largest integer-valued floating-point exponent that repeated squaring handles; beyond it generic_pow takes
1265// over. Plain squaring is accurate to a few ulps only for small exponents.
1266template <typename Packet, typename ScalarExponent>
1267EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool use_repeated_squaring(const ScalarExponent& exponent) {
1268 using Scalar = typename unpacket_traits<Packet>::type;
1269 return use_double_word<Packet>::value ? numext::abs(exponent) <= ScalarExponent(max_squaring_exponent<Scalar>())
1270 : (exponent <= ScalarExponent(7) && exponent >= ScalarExponent(-3));
1271}
1272
1273// The same for an exponent of integer type, which repeated squaring handles exactly whatever its size: generic_pow
1274// only takes over where the exponent converts exactly to the base type.
1275template <typename Packet, typename ScalarExponent>
1276EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool use_repeated_squaring_for_integer(const ScalarExponent& exponent) {
1277 using Scalar = typename unpacket_traits<Packet>::type;
1278 constexpr numext::uint64_t kExactLimit = numext::uint64_t(1) << numext::numeric_limits<Scalar>::digits;
1279 return exponent_helper<ScalarExponent>::safe_abs(exponent) > kExactLimit ||
1280 use_repeated_squaring<Packet>(static_cast<Scalar>(exponent));
1281}
1282
1283template <typename Packet>
1284EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<!is_scalar<Packet>::value, Packet> gen_pow(
1285 const Packet& x, const typename unpacket_traits<Packet>::type& exponent) {
1286 const Packet exponent_packet = pset1<Packet>(exponent);
1287 // generic_pow_impl requires positive x; sign/error handling is done by the caller.
1288 return generic_pow_impl(pabs(x), exponent_packet);
1289}
1290
1291template <typename Scalar>
1292EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<is_scalar<Scalar>::value, Scalar> gen_pow(
1293 const Scalar& x, const Scalar& exponent) {
1294 return numext::pow(x, exponent);
1295}
1296
1297// Handle special cases for pow(x, exponent) where both base and exponent are
1298// floating point and the exponent is a non-integer scalar (uniform across all
1299// SIMD lanes). This allows us to use scalar branches on exponent properties.
1300// Reached from packetOp only, so the comparison masks are all ones and select a NaN or an infinity bitwise.
1301template <typename Packet, typename ScalarExponent>
1302EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet handle_nonint_nonint_errors(const Packet& x, const Packet& powx,
1303 const ScalarExponent& exponent) {
1304 using Scalar = typename unpacket_traits<Packet>::type;
1305 const Packet cst_zero = pzero(x);
1306 const Packet cst_one = pset1<Packet>(Scalar(1));
1307 const Packet cst_inf = pinf<Packet>();
1308 const Packet cst_nan = pnan<Packet>();
1309
1310 const Packet abs_x = pabs(x);
1311
1312 // x < 0 with non-integer exponent -> NaN.
1313 Packet result = por(pcmp_lt(x, cst_zero), powx);
1314
1315 if (!(numext::isfinite)(exponent)) {
1316 if (exponent != exponent) {
1317 // pow(x, NaN) = NaN, except pow(+1, NaN) = 1.
1318 result = pselect(pcmp_eq(x, cst_one), cst_one, cst_nan);
1319 } else {
1320 // Exponent is +inf or -inf.
1321 const Packet abs_x_is_one = pcmp_eq(abs_x, cst_one);
1322 if (exponent > ScalarExponent(0)) {
1323 // pow(x, +inf): |x| > 1 -> +inf, |x| < 1 -> 0, |x| == 1 -> 1.
1324 result = pand(pcmp_lt(cst_one, abs_x), cst_inf);
1325 } else {
1326 // pow(x, -inf): |x| < 1 -> +inf, |x| > 1 -> 0, |x| == 1 -> 1.
1327 result = pand(pcmp_lt(abs_x, cst_one), cst_inf);
1328 }
1329 // pow(+-1, +-inf) = 1.
1330 result = pselect(abs_x_is_one, cst_one, result);
1331 }
1332 } else {
1333 // Finite non-integer exponent.
1334 const Packet x_is_zero = pcmp_eq(x, cst_zero);
1335 const Packet abs_x_is_inf = pcmp_eq(abs_x, cst_inf);
1336 if (exponent < ScalarExponent(0)) {
1337 // pow(+-0, negative non-integer) = +inf. pow(+-inf, negative) = +0.
1338 result = pselect(x_is_zero, cst_inf, result);
1339 result = pselect(abs_x_is_inf, cst_zero, result);
1340 } else {
1341 // pow(+-0, positive non-integer) = +0. pow(+-inf, positive) = +inf.
1342 result = pselect(x_is_zero, cst_zero, result);
1343 result = pselect(abs_x_is_inf, cst_inf, result);
1344 }
1345 }
1346
1347 // NaN base produces NaN. This overrides all cases above, but pow(NaN, 0) = 1
1348 // and pow(NaN, integer) are handled by the integer exponent path and never
1349 // reach this function.
1350 result = por(pisnan(x), result);
1351
1352 return result;
1353}
1354
1355template <typename Packet, typename ScalarExponent,
1356 std::enable_if_t<NumTraits<typename unpacket_traits<Packet>::type>::IsSigned, bool> = true>
1357EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet handle_negative_exponent(const Packet& x, const ScalarExponent& exponent) {
1358 using Scalar = typename unpacket_traits<Packet>::type;
1359
1360 // signed integer base, signed integer exponent case
1361
1362 // This routine handles negative exponents.
1363 // The return value is either 0, 1, or -1.
1364 Packet cst_pos_one = pset1<Packet>(Scalar(1));
1365 const bool exponent_is_odd = exponent % ScalarExponent(2) != ScalarExponent(0);
1366 const Packet exp_is_odd = exponent_is_odd ? ptrue<Packet>(x) : pzero<Packet>(x);
1367
1368 const Packet abs_x = pabs(x);
1369 const Packet abs_x_is_one = pcmp_eq(abs_x, cst_pos_one);
1370
1371 Packet result = pselect(exp_is_odd, x, abs_x);
1372 result = pselect(abs_x_is_one, result, pzero<Packet>(x));
1373 return result;
1374}
1375
1376template <typename Packet, typename ScalarExponent,
1377 std::enable_if_t<!NumTraits<typename unpacket_traits<Packet>::type>::IsSigned, bool> = true>
1378EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet handle_negative_exponent(const Packet& x, const ScalarExponent&) {
1379 using Scalar = typename unpacket_traits<Packet>::type;
1380
1381 // unsigned integer base, signed integer exponent case
1382
1383 // This routine handles negative exponents.
1384 // The return value is either 0 or 1
1385
1386 const Scalar pos_one = Scalar(1);
1387
1388 const Packet cst_pos_one = pset1<Packet>(pos_one);
1389
1390 const Packet x_is_one = pcmp_eq(x, cst_pos_one);
1391
1392 return pand(x_is_one, x);
1393}
1394
1395} // end namespace unary_pow
1396
1397template <typename Packet, typename ScalarExponent,
1398 bool BaseIsIntegerType = NumTraits<typename unpacket_traits<Packet>::type>::IsInteger,
1399 bool ExponentIsIntegerType = NumTraits<ScalarExponent>::IsInteger,
1400 bool ExponentIsSigned = NumTraits<ScalarExponent>::IsSigned>
1401struct unary_pow_impl;
1402
1403template <typename Packet, typename ScalarExponent, bool ExponentIsSigned>
1404struct unary_pow_impl<Packet, ScalarExponent, false, false, ExponentIsSigned> {
1405 using Scalar = typename unpacket_traits<Packet>::type;
1406 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet run(const Packet& x, const ScalarExponent& exponent) {
1407 const bool exponent_is_integer = (numext::isfinite)(exponent) && numext::round(exponent) == exponent;
1408 if (exponent_is_integer) {
1409 return unary_pow::use_repeated_squaring<Packet>(exponent) ? unary_pow::int_pow(x, exponent)
1410 : generic_pow(x, pset1<Packet>(exponent));
1411 } else {
1412 Packet result = unary_pow::gen_pow(x, exponent);
1413 result = unary_pow::handle_nonint_nonint_errors(x, result, exponent);
1414 return result;
1415 }
1416 }
1417};
1418
1419template <typename Packet, typename ScalarExponent, bool ExponentIsSigned>
1420struct unary_pow_impl<Packet, ScalarExponent, false, true, ExponentIsSigned> {
1421 using Scalar = typename unpacket_traits<Packet>::type;
1422 // Only real float and double bases with double-word support fall back to generic_pow for large exponents:
1423 // complex bases have no vectorized generic_pow, and the bases left without double-word support (half,
1424 // bfloat16, the GPU packets, the MSA and HVX float packets) have no packet generic_pow either, so their
1425 // squaring loop -- at most 64 steps -- is all there is.
1426 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet run(const Packet& x, const ScalarExponent& exponent) {
1427 return run(x, exponent,
1428 bool_constant < unary_pow::use_double_word<Packet>::value && !NumTraits<Scalar>::IsComplex > ());
1429 }
1430 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet run(const Packet& x, const ScalarExponent& exponent, true_type) {
1431 return unary_pow::use_repeated_squaring_for_integer<Packet>(exponent)
1432 ? unary_pow::int_pow(x, exponent)
1433 : generic_pow(x, pset1<Packet>(static_cast<Scalar>(exponent)));
1434 }
1435 static EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet run(const Packet& x, const ScalarExponent& exponent, false_type) {
1436 return unary_pow::int_pow(x, exponent);
1437 }
1438};
1439
1440template <typename Packet, typename ScalarExponent>
1441struct unary_pow_impl<Packet, ScalarExponent, true, true, true> {
1442 using Scalar = typename unpacket_traits<Packet>::type;
1443 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet run(const Packet& x, const ScalarExponent& exponent) {
1444 if (exponent < ScalarExponent(0)) {
1445 return unary_pow::handle_negative_exponent(x, exponent);
1446 } else {
1447 return unary_pow::int_pow_wrapping(x, exponent);
1448 }
1449 }
1450};
1451
1452template <typename Packet, typename ScalarExponent>
1453struct unary_pow_impl<Packet, ScalarExponent, true, true, false> {
1454 using Scalar = typename unpacket_traits<Packet>::type;
1455 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet run(const Packet& x, const ScalarExponent& exponent) {
1456 return unary_pow::int_pow_wrapping(x, exponent);
1457 }
1458};
1459
1460} // end namespace internal
1461} // end namespace Eigen
1462
1463#endif // EIGEN_ARCH_GENERIC_PACKET_MATH_POW_H