Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
SpecialFunctionsImpl.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2015 Eugene Brevdo <ebrevdo@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_SPECIAL_FUNCTIONS_H
12#define EIGEN_SPECIAL_FUNCTIONS_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18namespace internal {
19
20// Parts of this code are based on the Cephes Math Library.
21//
22// Cephes Math Library Release 2.8: June, 2000
23// Copyright 1984, 1987, 1992, 2000 by Stephen L. Moshier
24//
25// Permission has been kindly provided by the original author
26// to incorporate the Cephes software into the Eigen codebase:
27//
28// From: Stephen Moshier
29// To: Eugene Brevdo
30// Subject: Re: Permission to wrap several cephes functions in Eigen
31//
32// Hello Eugene,
33//
34// Thank you for writing.
35//
36// If your licensing is similar to BSD, the formal way that has been
37// handled is simply to add a statement to the effect that you are incorporating
38// the Cephes software by permission of the author.
39//
40// Good luck with your project,
41// Steve
42
43/****************************************************************************
44 * Implementation of lgamma *
45 ****************************************************************************/
46
47template <typename Scalar>
48struct lgamma_impl {
49 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(const Scalar) {
50 EIGEN_STATIC_ASSERT((!std::is_same<Scalar, Scalar>::value), THIS_TYPE_IS_NOT_SUPPORTED)
51 return Scalar(0);
52 }
53};
54
55// Since glibc 2.19
56#if defined(__GLIBC__) && ((__GLIBC__ >= 2 && __GLIBC_MINOR__ >= 19) || __GLIBC__ > 2) && \
57 (defined(_DEFAULT_SOURCE) || defined(_BSD_SOURCE) || defined(_SVID_SOURCE))
58#define EIGEN_HAS_LGAMMA_R
59#endif
60
61// Glibc versions before 2.19
62#if defined(__GLIBC__) && ((__GLIBC__ == 2 && __GLIBC_MINOR__ < 19) || __GLIBC__ < 2) && \
63 (defined(_BSD_SOURCE) || defined(_SVID_SOURCE))
64#define EIGEN_HAS_LGAMMA_R
65#endif
66
67template <>
68struct lgamma_impl<float> {
69 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float run(float x) {
70#if !defined(EIGEN_GPU_COMPILE_PHASE) && defined(EIGEN_HAS_LGAMMA_R) && !defined(__APPLE__)
71 int dummy;
72 return ::lgammaf_r(x, &dummy);
73#elif defined(SYCL_DEVICE_ONLY)
74 return cl::sycl::lgamma(x);
75#else
76 return ::lgammaf(x);
77#endif
78 }
79};
80
81template <>
82struct lgamma_impl<double> {
83 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double run(double x) {
84#if !defined(EIGEN_GPU_COMPILE_PHASE) && defined(EIGEN_HAS_LGAMMA_R) && !defined(__APPLE__)
85 int dummy;
86 return ::lgamma_r(x, &dummy);
87#elif defined(SYCL_DEVICE_ONLY)
88 return cl::sycl::lgamma(x);
89#else
90 return ::lgamma(x);
91#endif
92 }
93};
94
95#undef EIGEN_HAS_LGAMMA_R
96
97/****************************************************************************
98 * Implementation of digamma (psi), based on Cephes *
99 ****************************************************************************/
100
101/*
102 *
103 * Polynomial evaluation helper for the Psi (digamma) function.
104 *
105 * digamma_impl_maybe_poly::run(s) evaluates the asymptotic Psi expansion for
106 * input Scalar s, assuming s is above 10.0.
107 *
108 * If s is above a certain threshold for the given Scalar type, zero
109 * is returned. Otherwise the polynomial is evaluated with enough
110 * coefficients for results matching Scalar machine precision.
111 *
112 *
113 */
114template <typename Scalar>
115struct digamma_impl_maybe_poly {
116 EIGEN_STATIC_ASSERT((std::is_same<Scalar, Scalar>::value == false), THIS_TYPE_IS_NOT_SUPPORTED)
117
118 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(const Scalar) { return Scalar(0); }
119};
120
121template <>
122struct digamma_impl_maybe_poly<float> {
123 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float run(const float s) {
124 constexpr float A[] = {-4.16666666666666666667E-3f, 3.96825396825396825397E-3f, -8.33333333333333333333E-3f,
125 8.33333333333333333333E-2f};
126
127 float z;
128 if (s < 1.0e8f) {
129 z = 1.0f / (s * s);
130 return z * internal::ppolevl<float, 3>::run(z, A);
131 } else
132 return 0.0f;
133 }
134};
135
136template <>
137struct digamma_impl_maybe_poly<double> {
138 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double run(const double s) {
139 constexpr double A[] = {8.33333333333333333333E-2, -2.10927960927960927961E-2, 7.57575757575757575758E-3,
140 -4.16666666666666666667E-3, 3.96825396825396825397E-3, -8.33333333333333333333E-3,
141 8.33333333333333333333E-2};
142
143 double z;
144 if (s < 1.0e17) {
145 z = 1.0 / (s * s);
146 return z * internal::ppolevl<double, 6>::run(z, A);
147 } else
148 return 0.0;
149 }
150};
151
152template <typename Scalar>
153struct digamma_impl {
154 EIGEN_DEVICE_FUNC static Scalar run(Scalar x) {
155 /*
156 *
157 * Psi (digamma) function (modified for Eigen)
158 *
159 *
160 * SYNOPSIS:
161 *
162 * double x, y, psi();
163 *
164 * y = psi( x );
165 *
166 *
167 * DESCRIPTION:
168 *
169 * d -
170 * psi(x) = -- ln | (x)
171 * dx
172 *
173 * is the logarithmic derivative of the gamma function.
174 * For integer x,
175 * n-1
176 * -
177 * psi(n) = -EUL + > 1/k.
178 * -
179 * k=1
180 *
181 * If x is negative, it is transformed to a positive argument by the
182 * reflection formula psi(1-x) = psi(x) + pi cot(pi x).
183 * For general positive x, the argument is made greater than 10
184 * using the recurrence psi(x+1) = psi(x) + 1/x.
185 * Then the following asymptotic expansion is applied:
186 *
187 * inf. B
188 * - 2k
189 * psi(x) = log(x) - 1/2x - > -------
190 * - 2k
191 * k=1 2k x
192 *
193 * where the B2k are Bernoulli numbers.
194 *
195 * ACCURACY (float):
196 * Relative error (except absolute when |psi| < 1):
197 * arithmetic domain # trials peak rms
198 * IEEE 0,30 30000 1.3e-15 1.4e-16
199 * IEEE -30,0 40000 1.5e-15 2.2e-16
200 *
201 * ACCURACY (double):
202 * Absolute error, relative when |psi| > 1 :
203 * arithmetic domain # trials peak rms
204 * IEEE -33,0 30000 8.2e-7 1.2e-7
205 * IEEE 0,33 100000 7.3e-7 7.7e-8
206 *
207 * ERROR MESSAGES:
208 * message condition value returned
209 * psi singularity x integer <=0 INFINITY
210 */
211
212 Scalar p, q, nz, s, w, y;
213 bool negative = false;
214
215 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
216 const Scalar m_pi = Scalar(EIGEN_PI);
217
218 const Scalar zero = Scalar(0);
219 const Scalar one = Scalar(1);
220 const Scalar half = Scalar(0.5);
221 nz = zero;
222
223 // Near 0, psi(x) ~ -1/x - gamma.
224 // For x = +0.0: psi(+0.0) -> -infinity.
225 // For x = -0.0: psi(-0.0) -> +infinity.
226 if (x == zero) {
227 return (std::signbit(x)) ? NumTraits<Scalar>::infinity() : -NumTraits<Scalar>::infinity();
228 }
229 if (x < zero) {
230 negative = true;
231 q = x;
232 p = numext::floor(q);
233 if (p == q) {
234 return nan;
235 }
236 /* Remove the zeros of tan(m_pi x)
237 * by subtracting the nearest integer from x
238 */
239 nz = q - p;
240 if (nz != half) {
241 if (nz > half) {
242 p += one;
243 nz = q - p;
244 }
245 nz = m_pi / numext::tan(m_pi * nz);
246 } else {
247 nz = zero;
248 }
249 x = one - x;
250 }
251
252 /* use the recurrence psi(x+1) = psi(x) + 1/x. */
253 s = x;
254 w = zero;
255 while (s < Scalar(10)) {
256 w += one / s;
257 s += one;
258 }
259
260 y = digamma_impl_maybe_poly<Scalar>::run(s);
261
262 y = numext::log(s) - (half / s) - y - w;
263
264 return (negative) ? y - nz : y;
265 }
266};
267
268// Does unqualified lookup of erf/erfc succeed for T? The lookup below mirrors
269// the one at the call sites in erf_impl/erfc_impl, so it finds std::erf/std::erfc
270// and any overload visible through argument-dependent lookup.
271namespace unqualified_erf {
272EIGEN_USING_STD(erf)
273EIGEN_USING_STD(erfc)
274template <typename T>
275auto test_erf(int) -> decltype(void(erf(std::declval<const T&>())), std::true_type{});
276template <typename T>
277std::false_type test_erf(...);
278template <typename T>
279auto test_erfc(int) -> decltype(void(erfc(std::declval<const T&>())), std::true_type{});
280template <typename T>
281std::false_type test_erfc(...);
282} // namespace unqualified_erf
283
284template <typename T>
285struct has_erf : decltype(unqualified_erf::test_erf<T>(0)) {};
286template <typename T>
287struct has_erfc : decltype(unqualified_erf::test_erfc<T>(0)) {};
288
289/***************************************************************************
290 * Implementation of erfc.
291 ****************************************************************************/
292template <typename Scalar>
293struct generic_fast_erfc {
294 template <typename T>
295 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& x_in);
296};
297
298template <>
299template <typename T>
300EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erfc<float>::run(const T& x_in) {
301 constexpr float kClamp = 11.0f;
302 const T x = pmin<PropagateNaN>(pmax<PropagateNaN>(x_in, pset1<T>(-kClamp)), pset1<T>(kClamp));
303
304 // erfc(x) = 1 + x * S(x^2), |x| <= 1.
305 //
306 // Coefficients for S and T generated with Rminimax command:
307 // ./ratapprox --function="erfc(x)-1" --dom='[-1,1]' --type=[11,0] --num="odd"
308 // --numF="[SG]" --denF="[SG]" --log --dispCoeff="dec"
309 constexpr float alpha[] = {5.61802298761904239654541015625e-04, -4.91381669417023658752441406250e-03,
310 2.67075151205062866210937500000e-02, -1.12800106406211853027343750000e-01,
311 3.76122951507568359375000000000e-01, -1.12837910652160644531250000000e+00};
312 const T x2 = pmul(x, x);
313 const T one = pset1<T>(1.0f);
314 const T erfc_small = pmadd(x, ppolevl<T, 5>::run(x2, alpha), one);
315
316 // Return early if we don't need the more expensive approximation for any
317 // entry in a.
318 const T x_abs_gt_one_mask = pcmp_lt(one, x2);
319 if (!predux_any(x_abs_gt_one_mask)) return erfc_small;
320
321 // erfc(x) = exp(-x^2) * 1/x * P(1/x^2) / Q(1/x^2), 1 < x < 9.
322 //
323 // Coefficients for P and Q generated with Rminimax command:
324 // ./ratapprox --function="erfc(1/sqrt(x))*exp(1/x)/sqrt(x)"
325 // --dom='[0.01,1]' --type=[3,4] --numF="[SG]" --denF="[SG]" --log
326 // --dispCoeff="dec"
327 constexpr float gamma[] = {1.0208116471767425537109375e-01f, 4.2920666933059692382812500e-01f,
328 3.2379078865051269531250000e-01f, 5.3971976041793823242187500e-02f};
329 constexpr float delta[] = {1.7251677811145782470703125e-02f, 3.9137163758277893066406250e-01f,
330 1.0000000000000000000000000e+00f, 6.2173241376876831054687500e-01f,
331 9.5662862062454223632812500e-02f};
332 const T x2_lo = twoprod_low(x, x, x2);
333 // Here we use that
334 // exp(-x^2) = exp(-(x2+x2_lo)^2) ~= exp(-x2)*exp(-x2_lo) ~= exp(-x2)*(1-x2_lo)
335 // since x2_lo < kClamp * eps << 1 in the region we care about. This trick reduces the max error
336 // from 34 ulps to below 5 ulps.
337 const T exp2_hi = pexp(pnegate(x2));
338 const T z = pnmadd(exp2_hi, x2_lo, exp2_hi);
339 const T q2 = preciprocal(x2);
340 const T num = ppolevl<T, 3>::run(q2, gamma);
341 const T denom = pmul(x, ppolevl<T, 4>::run(q2, delta));
342 const T r = pdiv(num, denom);
343 const T maybe_two = pselect(pcmp_lt(x, pset1<T>(0.0f)), pset1<T>(2.0f), pset1<T>(0.0f));
344 const T erfc_large = pmadd(z, r, maybe_two);
345 return pselect(x_abs_gt_one_mask, erfc_large, erfc_small);
346}
347
348// Computes erf(x)/x for |x| <= 1. Used by both erf and erfc implementations.
349// Takes x2 = x^2 as input.
350//
351// PRECONDITION: x2 <= 1.
352template <typename T>
353EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T erf_over_x_double_small(const T& x2) {
354 // erf(x)/x = S(x^2) / T(x^2), x^2 <= 1.
355 //
356 // Coefficients for S and T generated with Rminimax command:
357 // ./ratapprox --function="erf(x)" --dom='[-1,1]' --type=[9,10]
358 // --num="odd" --numF="[D]" --den="even" --denF="[D]" --log --dispCoeff="dec"
359 constexpr double alpha[] = {1.9493725660006057018823477644531294572516344487667083740234375e-04,
360 1.8272566210022942682217328425053892715368419885635375976562500e-03,
361 4.5303363351690106863856044583371840417385101318359375000000000e-02,
362 1.4215015503619179981775744181504705920815467834472656250000000e-01,
363 1.1283791670955125585606992899556644260883331298828125000000000e+00};
364 constexpr double beta[] = {2.0294484101083099089526257108317963684385176748037338256835938e-05,
365 6.8117805899186819641732970609382391558028757572174072265625000e-04,
366 1.0582026056098614921752165685120417037978768348693847656250000e-02,
367 9.3252603143757495374188692949246615171432495117187500000000000e-02,
368 4.5931062818368939559832142549566924571990966796875000000000000e-01,
369 1.0};
370 const T num_small = ppolevl<T, 4>::run(x2, alpha);
371 const T denom_small = ppolevl<T, 5>::run(x2, beta);
372 return pdiv(num_small, denom_small);
373}
374
375// erfc(x) = exp(-x^2) * 1/x * P(1/x^2) / Q(1/x^2), 1 < x < 28.
376//
377// Coefficients for P and Q generated with Rminimax command:
378// ./ratapprox --function="erfc(1/sqrt(x))*exp(1/x)/sqrt(x)" --dom='[0.0013717,1]' --type=[9,9] --numF="[D]"
379// --denF="[D]" --log --dispCoeff="dec"
380//
381// PRECONDITION: 1 < x < 28.
382template <typename T>
383EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T erfc_double_large(const T& x, const T& x2) {
384 constexpr double gamma[] = {1.5252844933226974316088642158462107545346952974796295166015625e-04,
385 1.0909912393738931124520519233556115068495273590087890625000000e-02,
386 1.0628604636755033252537572252549580298364162445068359375000000e-01,
387 3.3492472973137982217295416376146022230386734008789062500000000e-01,
388 4.5065776215933289750026347064704168587923049926757812500000000e-01,
389 2.9433039130294824659017649537418037652969360351562500000000000e-01,
390 9.8792676360600226170838311645638896152377128601074218750000000e-02,
391 1.7095935395503719655962981960328761488199234008789062500000000e-02,
392 1.4249109729504577659398023570247460156679153442382812500000000e-03,
393 4.4567378313647954771875570045835956989321857690811157226562500e-05};
394 constexpr double delta[] = {2.041985103115789845773520028160419315099716186523437500000000e-03,
395 5.316030659946043707142493417450168635696172714233398437500000e-02,
396 3.426242193784684864077405563875799998641014099121093750000000e-01,
397 8.565637124308049799026321124983951449394226074218750000000000e-01,
398 1.000000000000000000000000000000000000000000000000000000000000e+00,
399 5.968805280570776972126623149961233139038085937500000000000000e-01,
400 1.890922854723317836356244470152887515723705291748046875000000e-01,
401 3.152505418656005586885981983868987299501895904541015625000000e-02,
402 2.565085751861882583380047861965067568235099315643310546875000e-03,
403 7.899362131678837697403017248376499992446042597293853759765625e-05};
404 // Compute exp(-x^2).
405 const T x2_lo = twoprod_low(x, x, x2);
406 // Here we use that
407 // exp(-x^2) = exp(-(x2+x2_lo)^2) ~= exp(-x2)*exp(-x2_lo) ~= exp(-x2)*(1-x2_lo)
408 // since x2_lo < kClamp * eps << 1 in the region we care about. This trick reduces the max error
409 // from 258 ulps to below 7 ulps.
410 const T exp2_hi = pexp(pnegate(x2));
411 const T z = pnmadd(exp2_hi, x2_lo, exp2_hi);
412 // Compute r = P / Q.
413 const T q2 = preciprocal(x2);
414 const T num_large = ppolevl<T, 9>::run(q2, gamma);
415 const T denom_large = pmul(x, ppolevl<T, 9>::run(q2, delta));
416 const T r = pdiv(num_large, denom_large);
417 const T maybe_two = pselect(pcmp_lt(x, pset1<T>(0.0)), pset1<T>(2.0), pset1<T>(0.0));
418 return pmadd(z, r, maybe_two);
419}
420
421template <>
422template <typename T>
423EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erfc<double>::run(const T& x_in) {
424 // Clamp x to [-28:28] beyond which erfc(x) is either two or zero (below the underflow threshold).
425 // This avoids having to deal with twoprod(x,x) producing NaN for sufficiently large x.
426 constexpr double kClamp = 28.0;
427 const T x = pmin<PropagateNaN>(pmax<PropagateNaN>(x_in, pset1<T>(-kClamp)), pset1<T>(kClamp));
428
429 // For |x| < 1, we use erfc(x) = 1 - erf(x).
430 const T x2 = pmul(x, x);
431 const T one = pset1<T>(1.0);
432 const T erfc_small = pnmadd(x, erf_over_x_double_small(x2), one);
433
434 // Return early if we don't need the more expensive approximation for any
435 // entry in a.
436 const T x_abs_gt_one_mask = pcmp_lt(one, x2);
437 if (!predux_any(x_abs_gt_one_mask)) return erfc_small;
438
439 const T erfc_large = erfc_double_large(x, x2);
440 return pselect(x_abs_gt_one_mask, erfc_large, erfc_small);
441}
442
443template <typename T>
444struct erfc_impl {
445 typedef typename unpacket_traits<T>::type Scalar;
446 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run(const T& x) { return run_impl(x, std::is_same<T, Scalar>()); }
447
448 private:
449 // Packets of float/double: vectorized rational approximation.
450 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_impl(const T& x, std::false_type) {
451 return generic_fast_erfc<Scalar>::run(x);
452 }
453 // Any other scalar type: defer to an erfc found by argument-dependent lookup
454 // (or std::erfc), keeping custom scalars on their own implementation instead
455 // of the float/double-tuned polynomials.
456 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_impl(const T& x, std::true_type) {
457 EIGEN_STATIC_ASSERT_NON_INTEGER(T)
458 return run_scalar(x, has_erfc<T>());
459 }
460 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_scalar(const T& x, std::true_type) {
461 EIGEN_USING_STD(erfc);
462 return erfc(x);
463 }
464 // Reject the type here instead of letting overload resolution fail inside the
465 // call above: the approximations in this file are tuned for float and double
466 // and are not valid for an arbitrary scalar, so a scalar type that wants erfc
467 // has to supply it.
468 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_scalar(const T& x, std::false_type) {
469 EIGEN_STATIC_ASSERT(has_erfc<T>::value, SCALAR_TYPE_MUST_PROVIDE_AN_ERFC_OVERLOAD_FOUND_BY_ADL_OR_IN_NAMESPACE_STD)
470 return x;
471 }
472};
473
474template <>
475struct erfc_impl<float> {
476 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float run(const float x) {
477#if defined(SYCL_DEVICE_ONLY)
478 return cl::sycl::erfc(x);
479#else
480 return generic_fast_erfc<float>::run(x);
481#endif
482 }
483};
484
485template <>
486struct erfc_impl<double> {
487 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double run(const double x) {
488#if defined(SYCL_DEVICE_ONLY)
489 return cl::sycl::erfc(x);
490#else
491 return generic_fast_erfc<double>::run(x);
492#endif
493 }
494};
495
496/****************************************************************************
497 * Implementation of erf.
498 ****************************************************************************/
499
500template <typename Scalar>
501struct generic_fast_erf {
502 template <typename T>
503 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& x_in);
504};
505
512template <>
513template <typename T>
514EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erf<float>::run(const T& x) {
515 // The monomial coefficients of the numerator polynomial (odd).
516 constexpr float alpha[] = {2.123732201653183437883853912353515625e-06f, 2.861979592125862836837768554687500000e-04f,
517 3.658048342913389205932617187500000000e-03f, 5.243302136659622192382812500000000000e-02f,
518 1.874160766601562500000000000000000000e-01f, 1.128379106521606445312500000000000000e+00f};
519
520 // The monomial coefficients of the denominator polynomial (even).
521 constexpr float beta[] = {3.89185734093189239501953125000e-05f, 1.14329601638019084930419921875e-03f,
522 1.47520881146192550659179687500e-02f, 1.12945675849914550781250000000e-01f,
523 4.99425798654556274414062500000e-01f, 1.0f};
524
525 // Since the polynomials are odd/even, we need x^2.
526 // Since erf(4) == 1 in float, we clamp x^2 to 16 to avoid computing Inf/Inf below.
527 // NaN need not survive this clamp: multiplying by x below restores it.
528 const T x2 = pmin(pset1<T>(16.0f), pmul(x, x));
529
530 // Evaluate the numerator polynomial p.
531 T p = ppolevl<T, 5>::run(x2, alpha);
532 p = pmul(x, p);
533
534 // Evaluate the denominator polynomial q.
535 T q = ppolevl<T, 5>::run(x2, beta);
536 const T r = pdiv(p, q);
537
538 // Clamp to [-1:1].
539 return pmax<PropagateNaN>(pmin<PropagateNaN>(r, pset1<T>(1.0f)), pset1<T>(-1.0f));
540}
541
542template <>
543template <typename T>
544EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erf<double>::run(const T& x_in) {
545 // Clamp x to [-28:28] beyond which erf(x) is ±1 within double precision.
546 // This avoids NaN from twoprod and exp operations for infinite inputs.
547 constexpr double kClamp = 28.0;
548 const T x = pmin<PropagateNaN>(pmax<PropagateNaN>(x_in, pset1<T>(-kClamp)), pset1<T>(kClamp));
549 T x2 = pmul(x, x);
550 T erf_small = pmul(x, erf_over_x_double_small(x2));
551
552 // Return early if we don't need the more expensive approximation for any
553 // entry in a.
554 const T one = pset1<T>(1.0);
555 const T x_abs_gt_one_mask = pcmp_lt(one, x2);
556 if (!predux_any(x_abs_gt_one_mask)) return erf_small;
557
558 // For |x| > 1, use erf(x) = 1 - erfc(x).
559 const T erf_large = psub(one, erfc_double_large(x, x2));
560 return pselect(x_abs_gt_one_mask, erf_large, erf_small);
561}
562
563template <typename T>
564struct erf_impl {
565 typedef typename unpacket_traits<T>::type Scalar;
566 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run(const T& x) { return run_impl(x, std::is_same<T, Scalar>()); }
567
568 private:
569 // Packets of float/double: vectorized rational approximation.
570 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_impl(const T& x, std::false_type) {
571 return generic_fast_erf<Scalar>::run(x);
572 }
573 // Any other scalar type: defer to an erf found by argument-dependent lookup
574 // (or std::erf), keeping custom scalars on their own implementation instead
575 // of the float/double-tuned polynomials.
576 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_impl(const T& x, std::true_type) {
577 EIGEN_STATIC_ASSERT_NON_INTEGER(T)
578 return run_scalar(x, has_erf<T>());
579 }
580 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_scalar(const T& x, std::true_type) {
581 EIGEN_USING_STD(erf);
582 return erf(x);
583 }
584 // Reject the type here instead of letting overload resolution fail inside the
585 // call above: the approximations in this file are tuned for float and double
586 // and are not valid for an arbitrary scalar, so a scalar type that wants erf
587 // has to supply it.
588 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE T run_scalar(const T& x, std::false_type) {
589 EIGEN_STATIC_ASSERT(has_erf<T>::value, SCALAR_TYPE_MUST_PROVIDE_AN_ERF_OVERLOAD_FOUND_BY_ADL_OR_IN_NAMESPACE_STD)
590 return x;
591 }
592};
593
594template <>
595struct erf_impl<float> {
596 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float run(const float x) {
597#if defined(SYCL_DEVICE_ONLY)
598 return cl::sycl::erf(x);
599#else
600 return generic_fast_erf<float>::run(x);
601#endif
602 }
603};
604
605template <>
606struct erf_impl<double> {
607 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double run(const double x) {
608#if defined(SYCL_DEVICE_ONLY)
609 return cl::sycl::erf(x);
610#else
611 return generic_fast_erf<double>::run(x);
612#endif
613 }
614};
615
616/***************************************************************************
617 * Implementation of ndtri. *
618 ****************************************************************************/
619
620/* Inverse of Normal distribution function (modified for Eigen).
621 *
622 *
623 * SYNOPSIS:
624 *
625 * double x, y, ndtri();
626 *
627 * x = ndtri( y );
628 *
629 *
630 *
631 * DESCRIPTION:
632 *
633 * Returns the argument, x, for which the area under the
634 * Gaussian probability density function (integrated from
635 * minus infinity to x) is equal to y.
636 *
637 *
638 * For small arguments 0 < y < exp(-2), the program computes
639 * z = sqrt( -2.0 * log(y) ); then the approximation is
640 * x = z - log(z)/z - (1/z) P(1/z) / Q(1/z).
641 * There are two rational functions P/Q, one for 0 < y < exp(-32)
642 * and the other for y up to exp(-2). For larger arguments,
643 * w = y - 0.5, and x/sqrt(2pi) = w + w**3 R(w**2)/S(w**2)).
644 *
645 *
646 * ACCURACY:
647 *
648 * Relative error:
649 * arithmetic domain # trials peak rms
650 * DEC 0.125, 1 5500 9.5e-17 2.1e-17
651 * DEC 6e-39, 0.135 3500 5.7e-17 1.3e-17
652 * IEEE 0.125, 1 20000 7.2e-16 1.3e-16
653 * IEEE 3e-308, 0.135 50000 4.6e-16 9.8e-17
654 *
655 *
656 * ERROR MESSAGES:
657 *
658 * message condition value returned
659 * ndtri domain x == 0 -INF
660 * ndtri domain x == 1 INF
661 * ndtri domain x < 0, x > 1 NAN
662 */
663/*
664 Cephes Math Library Release 2.2: June, 1992
665 Copyright 1985, 1987, 1992 by Stephen L. Moshier
666 Direct inquiries to 30 Frost Street, Cambridge, MA 02140
667*/
668
669// TODO: Add a cheaper approximation for float.
670
671template <typename T, bool IsScalar = is_scalar<T>::value>
672struct flipsign_impl;
673
674template <typename T>
675struct flipsign_impl<T, false> {
676 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& should_flipsign, const T& x) {
677 const T sign_mask = psignmask<T>();
678 const T sign_bit = pand<T>(should_flipsign, sign_mask);
679 return pxor<T>(sign_bit, x);
680 }
681};
682
683template <typename T>
684struct flipsign_impl<T, true> {
685 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& should_flipsign, const T& x) {
686 return should_flipsign == T(0) ? x : -x;
687 }
688};
689
690template <typename T>
691EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T flipsign(const T& should_flipsign, const T& x) {
692 return flipsign_impl<T>::run(should_flipsign, x);
693}
694
695template <typename T, bool IsScalar = is_scalar<T>::value>
696struct ndtri_negative_infinity_impl;
697
698template <typename T>
699struct ndtri_negative_infinity_impl<T, false> {
700 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& positive_infinity) {
701 return por(psignmask<T>(), positive_infinity);
702 }
703};
704
705template <typename T>
706struct ndtri_negative_infinity_impl<T, true> {
707 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(const T& positive_infinity) { return -positive_infinity; }
708};
709
710// We split this computation in to two so that in the scalar path
711// only one branch is evaluated (due to our template specialization of pselect
712// being an if statement.)
713
714template <typename T, typename ScalarType>
715EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_ndtri_gt_exp_neg_two(const T& b) {
716 const ScalarType p0[] = {ScalarType(-5.99633501014107895267e1), ScalarType(9.80010754185999661536e1),
717 ScalarType(-5.66762857469070293439e1), ScalarType(1.39312609387279679503e1),
718 ScalarType(-1.23916583867381258016e0)};
719 const ScalarType q0[] = {ScalarType(1.0),
720 ScalarType(1.95448858338141759834e0),
721 ScalarType(4.67627912898881538453e0),
722 ScalarType(8.63602421390890590575e1),
723 ScalarType(-2.25462687854119370527e2),
724 ScalarType(2.00260212380060660359e2),
725 ScalarType(-8.20372256168333339912e1),
726 ScalarType(1.59056225126211695515e1),
727 ScalarType(-1.18331621121330003142e0)};
728 const T sqrt2pi = pset1<T>(ScalarType(2.50662827463100050242e0));
729 const T half = pset1<T>(ScalarType(0.5));
730 T c, c2, ndtri_gt_exp_neg_two;
731
732 c = psub(b, half);
733 c2 = pmul(c, c);
734 ndtri_gt_exp_neg_two =
735 pmadd(c, pmul(c2, pdiv(internal::ppolevl<T, 4>::run(c2, p0), internal::ppolevl<T, 8>::run(c2, q0))), c);
736 return pmul(ndtri_gt_exp_neg_two, sqrt2pi);
737}
738
739template <typename T, typename ScalarType>
740EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_ndtri_lt_exp_neg_two(const T& b, const T& should_flipsign) {
741 /* Approximation for interval z = sqrt(-2 log a ) between 2 and 8
742 * i.e., a between exp(-2) = .135 and exp(-32) = 1.27e-14.
743 */
744 const ScalarType p1[] = {ScalarType(4.05544892305962419923e0), ScalarType(3.15251094599893866154e1),
745 ScalarType(5.71628192246421288162e1), ScalarType(4.40805073893200834700e1),
746 ScalarType(1.46849561928858024014e1), ScalarType(2.18663306850790267539e0),
747 ScalarType(-1.40256079171354495875e-1), ScalarType(-3.50424626827848203418e-2),
748 ScalarType(-8.57456785154685413611e-4)};
749 const ScalarType q1[] = {ScalarType(1.0),
750 ScalarType(1.57799883256466749731e1),
751 ScalarType(4.53907635128879210584e1),
752 ScalarType(4.13172038254672030440e1),
753 ScalarType(1.50425385692907503408e1),
754 ScalarType(2.50464946208309415979e0),
755 ScalarType(-1.42182922854787788574e-1),
756 ScalarType(-3.80806407691578277194e-2),
757 ScalarType(-9.33259480895457427372e-4)};
758 /* Approximation for interval z = sqrt(-2 log a ) between 8 and 64
759 * i.e., a between exp(-32) = 1.27e-14 and exp(-2048) = 3.67e-890.
760 */
761 const ScalarType p2[] = {ScalarType(3.23774891776946035970e0), ScalarType(6.91522889068984211695e0),
762 ScalarType(3.93881025292474443415e0), ScalarType(1.33303460815807542389e0),
763 ScalarType(2.01485389549179081538e-1), ScalarType(1.23716634817820021358e-2),
764 ScalarType(3.01581553508235416007e-4), ScalarType(2.65806974686737550832e-6),
765 ScalarType(6.23974539184983293730e-9)};
766 const ScalarType q2[] = {ScalarType(1.0),
767 ScalarType(6.02427039364742014255e0),
768 ScalarType(3.67983563856160859403e0),
769 ScalarType(1.37702099489081330271e0),
770 ScalarType(2.16236993594496635890e-1),
771 ScalarType(1.34204006088543189037e-2),
772 ScalarType(3.28014464682127739104e-4),
773 ScalarType(2.89247864745380683936e-6),
774 ScalarType(6.79019408009981274425e-9)};
775 const T eight = pset1<T>(ScalarType(8.0));
776 const T neg_two = pset1<T>(ScalarType(-2));
777 T x, x0, x1, z;
778
779 x = psqrt(pmul(neg_two, plog(b)));
780 x0 = psub(x, pdiv(plog(x), x));
781 z = preciprocal(x);
782 x1 =
783 pmul(z, pselect(pcmp_lt(x, eight), pdiv(internal::ppolevl<T, 8>::run(z, p1), internal::ppolevl<T, 8>::run(z, q1)),
784 pdiv(internal::ppolevl<T, 8>::run(z, p2), internal::ppolevl<T, 8>::run(z, q2))));
785 return flipsign(should_flipsign, psub(x0, x1));
786}
787
788template <typename T, typename ScalarType>
789EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE T generic_ndtri(const T& a) {
790 const T maxnum = pinf<T>();
791 const T neg_maxnum = ndtri_negative_infinity_impl<T>::run(maxnum);
792
793 const T zero = pset1<T>(ScalarType(0));
794 const T one = pset1<T>(ScalarType(1));
795 // exp(-2)
796 const T exp_neg_two = pset1<T>(ScalarType(0.13533528323661269189));
797 T b, ndtri, should_flipsign;
798
799 should_flipsign = pcmp_le(a, psub(one, exp_neg_two));
800 b = pselect(should_flipsign, a, psub(one, a));
801
802 ndtri = pselect(pcmp_lt(exp_neg_two, b), generic_ndtri_gt_exp_neg_two<T, ScalarType>(b),
803 generic_ndtri_lt_exp_neg_two<T, ScalarType>(b, should_flipsign));
804
805 return pselect(pcmp_eq(a, zero), neg_maxnum, pselect(pcmp_eq(one, a), maxnum, ndtri));
806}
807
808template <typename Scalar>
809struct ndtri_impl {
810 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(const Scalar x) { return generic_ndtri<Scalar, Scalar>(x); }
811};
812
813/**************************************************************************************************************
814 * Implementation of igammac (complemented incomplete gamma integral), based on Cephes *
815 **************************************************************************************************************/
816
817// NOTE: cephes_helper is also used to implement zeta
818template <typename Scalar>
819struct cephes_helper {
820 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar machep() {
821 eigen_assert(false && "machep not supported for this type");
822 return 0.0;
823 }
824 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar big() {
825 eigen_assert(false && "big not supported for this type");
826 return 0.0;
827 }
828 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar biginv() {
829 eigen_assert(false && "biginv not supported for this type");
830 return 0.0;
831 }
832};
833
834template <>
835struct cephes_helper<float> {
836 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float machep() {
837 return NumTraits<float>::epsilon() / 2; // 1.0 - machep == 1.0
838 }
839 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float big() {
840 // use epsneg (1.0 - epsneg == 1.0)
841 return 1.0f / (NumTraits<float>::epsilon() / 2);
842 }
843 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float biginv() {
844 // epsneg
845 return machep();
846 }
847};
848
849template <>
850struct cephes_helper<double> {
851 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double machep() {
852 return NumTraits<double>::epsilon() / 2; // 1.0 - machep == 1.0
853 }
854 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double big() { return 1.0 / NumTraits<double>::epsilon(); }
855 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double biginv() {
856 // inverse of eps
857 return NumTraits<double>::epsilon();
858 }
859};
860
861enum IgammaComputationMode { VALUE, DERIVATIVE, SAMPLE_DERIVATIVE };
862
863template <typename Scalar>
864EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar main_igamma_term(Scalar a, Scalar x) {
865 /* Compute x**a * exp(-x) / gamma(a) */
866 Scalar logax = a * numext::log(x) - x - lgamma_impl<Scalar>::run(a);
867 if (logax < -numext::log(NumTraits<Scalar>::highest()) ||
868 // Assuming x and a aren't Nan.
869 (numext::isnan)(logax)) {
870 return Scalar(0);
871 }
872 return numext::exp(logax);
873}
874
875template <typename Scalar, IgammaComputationMode mode>
876EIGEN_DEVICE_FUNC constexpr int igamma_num_iterations() {
877 /* Returns the maximum number of internal iterations for igamma computation.
878 */
879 return mode == VALUE ? 2000
880 : std::is_same<Scalar, float>::value ? 200
881 : std::is_same<Scalar, double>::value ? 500
882 : 2000;
883}
884
885template <typename Scalar, IgammaComputationMode mode>
886struct igammac_cf_impl {
887 /* Computes igamc(a, x) or derivative (depending on the mode)
888 * using the continued fraction expansion of the complementary
889 * incomplete Gamma function.
890 *
891 * Preconditions:
892 * a > 0
893 * x >= 1
894 * x >= a
895 */
896 EIGEN_DEVICE_FUNC static Scalar run(Scalar a, Scalar x) {
897 const Scalar zero = 0;
898 const Scalar one = 1;
899 const Scalar two = 2;
900 const Scalar machep = cephes_helper<Scalar>::machep();
901 const Scalar big = cephes_helper<Scalar>::big();
902 const Scalar biginv = cephes_helper<Scalar>::biginv();
903
904 if ((numext::isinf)(x)) {
905 return zero;
906 }
907
908 Scalar ax = main_igamma_term<Scalar>(a, x);
909 // This is independent of mode. If this value is zero,
910 // then the function value is zero. If the function value is zero,
911 // then we are in a neighborhood where the function value evaluates to zero,
912 // so the derivative is zero.
913 if (ax == zero) {
914 return zero;
915 }
916
917 // continued fraction
918 Scalar y = one - a;
919 Scalar z = x + y + one;
920 Scalar c = zero;
921 Scalar pkm2 = one;
922 Scalar qkm2 = x;
923 Scalar pkm1 = x + one;
924 Scalar qkm1 = z * x;
925 Scalar ans = pkm1 / qkm1;
926
927 Scalar dpkm2_da = zero;
928 Scalar dqkm2_da = zero;
929 Scalar dpkm1_da = zero;
930 Scalar dqkm1_da = -x;
931 Scalar dans_da = (dpkm1_da - ans * dqkm1_da) / qkm1;
932
933 for (int i = 0; i < igamma_num_iterations<Scalar, mode>(); i++) {
934 c += one;
935 y += one;
936 z += two;
937
938 Scalar yc = y * c;
939 Scalar pk = pkm1 * z - pkm2 * yc;
940 Scalar qk = qkm1 * z - qkm2 * yc;
941
942 Scalar dpk_da = dpkm1_da * z - pkm1 - dpkm2_da * yc + pkm2 * c;
943 Scalar dqk_da = dqkm1_da * z - qkm1 - dqkm2_da * yc + qkm2 * c;
944
945 if (qk != zero) {
946 Scalar ans_prev = ans;
947 ans = pk / qk;
948
949 Scalar dans_da_prev = dans_da;
950 dans_da = (dpk_da - ans * dqk_da) / qk;
951
952 if (mode == VALUE) {
953 if (numext::abs(ans_prev - ans) <= machep * numext::abs(ans)) {
954 break;
955 }
956 } else {
957 if (numext::abs(dans_da - dans_da_prev) <= machep) {
958 break;
959 }
960 }
961 }
962
963 pkm2 = pkm1;
964 pkm1 = pk;
965 qkm2 = qkm1;
966 qkm1 = qk;
967
968 dpkm2_da = dpkm1_da;
969 dpkm1_da = dpk_da;
970 dqkm2_da = dqkm1_da;
971 dqkm1_da = dqk_da;
972
973 if (numext::abs(pk) > big) {
974 pkm2 *= biginv;
975 pkm1 *= biginv;
976 qkm2 *= biginv;
977 qkm1 *= biginv;
978
979 dpkm2_da *= biginv;
980 dpkm1_da *= biginv;
981 dqkm2_da *= biginv;
982 dqkm1_da *= biginv;
983 }
984 }
985
986 /* Compute x**a * exp(-x) / gamma(a) */
987 Scalar dlogax_da = numext::log(x) - digamma_impl<Scalar>::run(a);
988 Scalar dax_da = ax * dlogax_da;
989
990 switch (mode) {
991 case VALUE:
992 return ans * ax;
993 case DERIVATIVE:
994 return ans * dax_da + dans_da * ax;
995 case SAMPLE_DERIVATIVE:
996 default: // this is needed to suppress clang warning
997 return -(dans_da + ans * dlogax_da) * x;
998 }
999 }
1000};
1001
1002template <typename Scalar, IgammaComputationMode mode>
1003struct igamma_series_impl {
1004 /* Computes igam(a, x) or its derivative (depending on the mode)
1005 * using the series expansion of the incomplete Gamma function.
1006 *
1007 * Preconditions:
1008 * x > 0
1009 * a > 0
1010 * !(x > 1 && x > a)
1011 */
1012 EIGEN_DEVICE_FUNC static Scalar run(Scalar a, Scalar x) {
1013 const Scalar zero = 0;
1014 const Scalar one = 1;
1015 const Scalar machep = cephes_helper<Scalar>::machep();
1016
1017 Scalar ax = main_igamma_term<Scalar>(a, x);
1018
1019 // This is independent of mode. If this value is zero,
1020 // then the function value is zero. If the function value is zero,
1021 // then we are in a neighborhood where the function value evaluates to zero,
1022 // so the derivative is zero.
1023 if (ax == zero) {
1024 return zero;
1025 }
1026
1027 ax /= a;
1028
1029 /* power series */
1030 Scalar r = a;
1031 Scalar c = one;
1032 Scalar ans = one;
1033
1034 Scalar dc_da = zero;
1035 Scalar dans_da = zero;
1036
1037 for (int i = 0; i < igamma_num_iterations<Scalar, mode>(); i++) {
1038 r += one;
1039 Scalar term = x / r;
1040 Scalar dterm_da = -x / (r * r);
1041 dc_da = term * dc_da + dterm_da * c;
1042 dans_da += dc_da;
1043 c *= term;
1044 ans += c;
1045
1046 if (mode == VALUE) {
1047 if (c <= machep * ans) {
1048 break;
1049 }
1050 } else {
1051 if (numext::abs(dc_da) <= machep * numext::abs(dans_da)) {
1052 break;
1053 }
1054 }
1055 }
1056
1057 Scalar dlogax_da = numext::log(x) - digamma_impl<Scalar>::run(a + one);
1058 Scalar dax_da = ax * dlogax_da;
1059
1060 switch (mode) {
1061 case VALUE:
1062 return ans * ax;
1063 case DERIVATIVE:
1064 return ans * dax_da + dans_da * ax;
1065 case SAMPLE_DERIVATIVE:
1066 default: // this is needed to suppress clang warning
1067 return -(dans_da + ans * dlogax_da) * x / a;
1068 }
1069 }
1070};
1071
1072template <typename Scalar>
1073struct igammac_impl {
1074 EIGEN_DEVICE_FUNC static Scalar run(Scalar a, Scalar x) {
1075 /* igamc()
1076 *
1077 * Incomplete gamma integral (modified for Eigen)
1078 *
1079 *
1080 *
1081 * SYNOPSIS:
1082 *
1083 * double a, x, y, igamc();
1084 *
1085 * y = igamc( a, x );
1086 *
1087 * DESCRIPTION:
1088 *
1089 * The function is defined by
1090 *
1091 *
1092 * igamc(a,x) = 1 - igam(a,x)
1093 *
1094 * inf.
1095 * -
1096 * 1 | | -t a-1
1097 * = ----- | e t dt.
1098 * - | |
1099 * | (a) -
1100 * x
1101 *
1102 *
1103 * In this implementation both arguments must be positive.
1104 * The integral is evaluated by either a power series or
1105 * continued fraction expansion, depending on the relative
1106 * values of a and x.
1107 *
1108 * ACCURACY (float):
1109 *
1110 * Relative error:
1111 * arithmetic domain # trials peak rms
1112 * IEEE 0,30 30000 7.8e-6 5.9e-7
1113 *
1114 *
1115 * ACCURACY (double):
1116 *
1117 * Tested at random a, x.
1118 * a x Relative error:
1119 * arithmetic domain domain # trials peak rms
1120 * IEEE 0.5,100 0,100 200000 1.9e-14 1.7e-15
1121 * IEEE 0.01,0.5 0,100 200000 1.4e-13 1.6e-15
1122 *
1123 */
1124 /*
1125 Cephes Math Library Release 2.2: June, 1992
1126 Copyright 1985, 1987, 1992 by Stephen L. Moshier
1127 Direct inquiries to 30 Frost Street, Cambridge, MA 02140
1128 */
1129 const Scalar zero = 0;
1130 const Scalar one = 1;
1131 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1132
1133 if ((x < zero) || (a <= zero)) {
1134 // domain error
1135 return nan;
1136 }
1137
1138 if ((numext::isnan)(a) || (numext::isnan)(x)) { // propagate nans
1139 return nan;
1140 }
1141
1142 if ((x < one) || (x < a)) {
1143 // Clamp to [0,1] since 1-igamma can produce tiny negative values
1144 // due to floating-point cancellation for extreme arguments.
1145 return numext::mini(one, numext::maxi(zero, one - igamma_series_impl<Scalar, VALUE>::run(a, x)));
1146 }
1147
1148 return igammac_cf_impl<Scalar, VALUE>::run(a, x);
1149 }
1150};
1151
1152/************************************************************************************************
1153 * Implementation of igamma (incomplete gamma integral), based on Cephes *
1154 ************************************************************************************************/
1155
1156template <typename Scalar, IgammaComputationMode mode>
1157struct igamma_generic_impl {
1158 EIGEN_DEVICE_FUNC static Scalar run(Scalar a, Scalar x) {
1159 /* Depending on the mode, returns
1160 * - VALUE: incomplete Gamma function igamma(a, x)
1161 * - DERIVATIVE: derivative of incomplete Gamma function d/da igamma(a, x)
1162 * - SAMPLE_DERIVATIVE: implicit derivative of a Gamma random variable
1163 * x ~ Gamma(x | a, 1), dx/da = -1 / Gamma(x | a, 1) * d igamma(a, x) / dx
1164 *
1165 * Derivatives are implemented by forward-mode differentiation.
1166 */
1167 const Scalar zero = 0;
1168 const Scalar one = 1;
1169 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1170
1171 if (x == zero) return zero;
1172
1173 if ((x < zero) || (a <= zero)) { // domain error
1174 return nan;
1175 }
1176
1177 if ((numext::isnan)(a) || (numext::isnan)(x)) { // propagate nans
1178 return nan;
1179 }
1180
1181 if ((x > one) && (x > a)) {
1182 Scalar ret = igammac_cf_impl<Scalar, mode>::run(a, x);
1183 if (mode == VALUE) {
1184 // Clamp to [0,1] since 1-igammac can produce tiny negative values
1185 // due to floating-point cancellation for extreme arguments.
1186 return numext::mini(one, numext::maxi(zero, one - ret));
1187 } else {
1188 return -ret;
1189 }
1190 }
1191
1192 Scalar ret = igamma_series_impl<Scalar, mode>::run(a, x);
1193 if (mode == VALUE) {
1194 // Clamp to [0,1] since accumulated series terms can slightly exceed 1.0
1195 // due to floating-point rounding for extreme arguments.
1196 return numext::mini(one, numext::maxi(zero, ret));
1197 }
1198 return ret;
1199 }
1200};
1201
1202template <typename Scalar>
1203struct igamma_impl : igamma_generic_impl<Scalar, VALUE> {
1204 /* igam()
1205 * Incomplete gamma integral.
1206 *
1207 * The CDF of Gamma(a, 1) random variable at the point x.
1208 *
1209 * Accuracy estimation. For each a in [10^-2, 10^-1...10^3] we sample
1210 * 50 Gamma random variables x ~ Gamma(x | a, 1), a total of 300 points.
1211 * The ground truth is computed by mpmath. Mean absolute error:
1212 * float: 1.26713e-05
1213 * double: 2.33606e-12
1214 *
1215 * Cephes documentation below.
1216 *
1217 * SYNOPSIS:
1218 *
1219 * double a, x, y, igam();
1220 *
1221 * y = igam( a, x );
1222 *
1223 * DESCRIPTION:
1224 *
1225 * The function is defined by
1226 *
1227 * x
1228 * -
1229 * 1 | | -t a-1
1230 * igam(a,x) = ----- | e t dt.
1231 * - | |
1232 * | (a) -
1233 * 0
1234 *
1235 *
1236 * In this implementation both arguments must be positive.
1237 * The integral is evaluated by either a power series or
1238 * continued fraction expansion, depending on the relative
1239 * values of a and x.
1240 *
1241 * ACCURACY (double):
1242 *
1243 * Relative error:
1244 * arithmetic domain # trials peak rms
1245 * IEEE 0,30 200000 3.6e-14 2.9e-15
1246 * IEEE 0,100 300000 9.9e-14 1.5e-14
1247 *
1248 *
1249 * ACCURACY (float):
1250 *
1251 * Relative error:
1252 * arithmetic domain # trials peak rms
1253 * IEEE 0,30 20000 7.8e-6 5.9e-7
1254 *
1255 */
1256 /*
1257 Cephes Math Library Release 2.2: June, 1992
1258 Copyright 1985, 1987, 1992 by Stephen L. Moshier
1259 Direct inquiries to 30 Frost Street, Cambridge, MA 02140
1260 */
1261
1262 /* left tail of incomplete gamma function:
1263 *
1264 * inf. k
1265 * a -x - x
1266 * x e > ----------
1267 * - -
1268 * k=0 | (a+k+1)
1269 *
1270 */
1271};
1272
1273template <typename Scalar>
1274struct igamma_der_a_impl : igamma_generic_impl<Scalar, DERIVATIVE> {
1275 /* Derivative of the incomplete Gamma function with respect to a.
1276 *
1277 * Computes d/da igamma(a, x) by forward differentiation of the igamma code.
1278 *
1279 * Accuracy estimation. For each a in [10^-2, 10^-1...10^3] we sample
1280 * 50 Gamma random variables x ~ Gamma(x | a, 1), a total of 300 points.
1281 * The ground truth is computed by mpmath. Mean absolute error:
1282 * float: 6.17992e-07
1283 * double: 4.60453e-12
1284 *
1285 * Reference:
1286 * R. Moore. "Algorithm AS 187: Derivatives of the incomplete gamma
1287 * integral". Journal of the Royal Statistical Society. 1982
1288 */
1289};
1290
1291template <typename Scalar>
1292struct gamma_sample_der_alpha_impl : igamma_generic_impl<Scalar, SAMPLE_DERIVATIVE> {
1293 /* Derivative of a Gamma random variable sample with respect to alpha.
1294 *
1295 * Consider a sample of a Gamma random variable with the concentration
1296 * parameter alpha: sample ~ Gamma(alpha, 1). The reparameterization
1297 * derivative that we want to compute is dsample / dalpha =
1298 * d igammainv(alpha, u) / dalpha, where u = igamma(alpha, sample).
1299 * However, this formula is numerically unstable and expensive, so instead
1300 * we use implicit differentiation:
1301 *
1302 * igamma(alpha, sample) = u, where u ~ Uniform(0, 1).
1303 * Apply d / dalpha to both sides:
1304 * d igamma(alpha, sample) / dalpha
1305 * + d igamma(alpha, sample) / dsample * dsample/dalpha = 0
1306 * d igamma(alpha, sample) / dalpha
1307 * + Gamma(sample | alpha, 1) dsample / dalpha = 0
1308 * dsample/dalpha = - (d igamma(alpha, sample) / dalpha)
1309 * / Gamma(sample | alpha, 1)
1310 *
1311 * Here Gamma(sample | alpha, 1) is the PDF of the Gamma distribution
1312 * (note that the derivative of the CDF w.r.t. sample is the PDF).
1313 * See the reference below for more details.
1314 *
1315 * The derivative of igamma(alpha, sample) is computed by forward
1316 * differentiation of the igamma code. Division by the Gamma PDF is performed
1317 * in the same code, increasing the accuracy and speed due to cancellation
1318 * of some terms.
1319 *
1320 * Accuracy estimation. For each alpha in [10^-2, 10^-1...10^3] we sample
1321 * 50 Gamma random variables sample ~ Gamma(sample | alpha, 1), a total of 300
1322 * points. The ground truth is computed by mpmath. Mean absolute error:
1323 * float: 2.1686e-06
1324 * double: 1.4774e-12
1325 *
1326 * Reference:
1327 * M. Figurnov, S. Mohamed, A. Mnih "Implicit Reparameterization Gradients".
1328 * 2018
1329 */
1330};
1331
1332/*****************************************************************************
1333 * Implementation of Riemann zeta function of two arguments, based on Cephes *
1334 *****************************************************************************/
1335
1336template <typename Scalar>
1337struct zeta_impl_series {
1338 EIGEN_STATIC_ASSERT((std::is_same<Scalar, Scalar>::value == false), THIS_TYPE_IS_NOT_SUPPORTED)
1339
1340 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(const Scalar) { return Scalar(0); }
1341};
1342
1343template <>
1344struct zeta_impl_series<float> {
1345 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool run(float& a, float& b, float& s, const float x,
1346 const float machep) {
1347 int i = 0;
1348 while (i < 9) {
1349 i += 1;
1350 a += 1.0f;
1351 b = numext::pow(a, -x);
1352 s += b;
1353 if (numext::abs(b / s) < machep) return true;
1354 }
1355
1356 // Return whether we are done
1357 return false;
1358 }
1359};
1360
1361template <>
1362struct zeta_impl_series<double> {
1363 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool run(double& a, double& b, double& s, const double x,
1364 const double machep) {
1365 int i = 0;
1366 while ((i < 9) || (a <= 9.0)) {
1367 i += 1;
1368 a += 1.0;
1369 b = numext::pow(a, -x);
1370 s += b;
1371 if (numext::abs(b / s) < machep) return true;
1372 }
1373
1374 // Return whether we are done
1375 return false;
1376 }
1377};
1378
1379template <typename Scalar>
1380struct zeta_impl {
1381 EIGEN_DEVICE_FUNC static Scalar run(Scalar x, Scalar q) {
1382 /* zeta.c
1383 *
1384 * Riemann zeta function of two arguments
1385 *
1386 *
1387 *
1388 * SYNOPSIS:
1389 *
1390 * double x, q, y, zeta();
1391 *
1392 * y = zeta( x, q );
1393 *
1394 *
1395 *
1396 * DESCRIPTION:
1397 *
1398 *
1399 *
1400 * inf.
1401 * - -x
1402 * zeta(x,q) = > (k+q)
1403 * -
1404 * k=0
1405 *
1406 * where x > 1 and q is not a negative integer or zero.
1407 * The Euler-Maclaurin summation formula is used to obtain
1408 * the expansion
1409 *
1410 * n
1411 * - -x
1412 * zeta(x,q) = > (k+q)
1413 * -
1414 * k=1
1415 *
1416 * 1-x inf. B x(x+1)...(x+2j)
1417 * (n+q) 1 - 2j
1418 * + --------- - ------- + > --------------------
1419 * x-1 x - x+2j+1
1420 * 2(n+q) j=1 (2j)! (n+q)
1421 *
1422 * where the B2j are Bernoulli numbers. Note that (see zetac.c)
1423 * zeta(x,1) = zetac(x) + 1.
1424 *
1425 *
1426 *
1427 * ACCURACY:
1428 *
1429 * Relative error for single precision:
1430 * arithmetic domain # trials peak rms
1431 * IEEE 0,25 10000 6.9e-7 1.0e-7
1432 *
1433 * Large arguments may produce underflow in powf(), in which
1434 * case the results are inaccurate.
1435 *
1436 * REFERENCE:
1437 *
1438 * Gradshteyn, I. S., and I. M. Ryzhik, Tables of Integrals,
1439 * Series, and Products, p. 1073; Academic Press, 1980.
1440 *
1441 */
1442
1443 int i;
1444 Scalar p, r, a, b, k, s, t, w;
1445
1446 const Scalar A[] = {
1447 Scalar(12.0),
1448 Scalar(-720.0),
1449 Scalar(30240.0),
1450 Scalar(-1209600.0),
1451 Scalar(47900160.0),
1452 Scalar(-1.8924375803183791606e9), /*1.307674368e12/691*/
1453 Scalar(7.47242496e10),
1454 Scalar(-2.950130727918164224e12), /*1.067062284288e16/3617*/
1455 Scalar(1.1646782814350067249e14), /*5.109094217170944e18/43867*/
1456 Scalar(-4.5979787224074726105e15), /*8.028576626982912e20/174611*/
1457 Scalar(1.8152105401943546773e17), /*1.5511210043330985984e23/854513*/
1458 Scalar(-7.1661652561756670113e18) /*1.6938241367317436694528e27/236364091*/
1459 };
1460
1461 const Scalar maxnum = NumTraits<Scalar>::infinity();
1462 const Scalar zero = Scalar(0.0), half = Scalar(0.5), one = Scalar(1.0);
1463 const Scalar machep = cephes_helper<Scalar>::machep();
1464 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1465
1466 if (x == one) return maxnum;
1467
1468 if (x < one) {
1469 return nan;
1470 }
1471
1472 if (q <= zero) {
1473 if (q == numext::floor(q)) {
1474 if (numext::rint(Scalar(0.5) * x) == Scalar(0.5) * x) {
1475 return maxnum;
1476 } else {
1477 return nan;
1478 }
1479 }
1480 p = x;
1481 r = numext::floor(p);
1482 if (p != r) return nan;
1483 }
1484
1485 /* Permit negative q but continue sum until n+q > +9 .
1486 * This case should be handled by a reflection formula.
1487 * If q<0 and x is an integer, there is a relation to
1488 * the polygamma function.
1489 */
1490 s = numext::pow(q, -x);
1491 a = q;
1492 b = zero;
1493 // Run the summation in a helper function that is specific to the floating precision
1494 if (zeta_impl_series<Scalar>::run(a, b, s, x, machep)) {
1495 return s;
1496 }
1497
1498 // If b is zero, then the tail sum will also end up being zero.
1499 // Exiting early here can prevent NaNs for some large inputs, where
1500 // the tail sum computed below has term `a` which can overflow to `inf`.
1501 if (numext::equal_strict(b, zero)) {
1502 return s;
1503 }
1504
1505 w = a;
1506 s += b * w / (x - one);
1507 s -= half * b;
1508 a = one;
1509 k = zero;
1510
1511 for (i = 0; i < 12; i++) {
1512 a *= x + k;
1513 b /= w;
1514 t = a * b / A[i];
1515 s = s + t;
1516 t = numext::abs(t / s);
1517 if (t < machep) {
1518 break;
1519 }
1520 k += one;
1521 a *= x + k;
1522 b /= w;
1523 k += one;
1524 }
1525 return s;
1526 }
1527};
1528
1529/****************************************************************************
1530 * Implementation of polygamma function *
1531 ****************************************************************************/
1532
1533template <typename Scalar>
1534struct polygamma_impl {
1535 EIGEN_DEVICE_FUNC static Scalar run(Scalar n, Scalar x) {
1536 Scalar zero = 0.0, one = 1.0;
1537 Scalar nplus = n + one;
1538 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1539
1540 // Check that n is a non-negative integer
1541 if (numext::floor(n) != n || n < zero) {
1542 return nan;
1543 }
1544 // Just return the digamma function for n = 0
1545 else if (n == zero) {
1546 return digamma_impl<Scalar>::run(x);
1547 }
1548 // Use the same implementation as scipy
1549 else {
1550 Scalar factorial = numext::exp(lgamma_impl<Scalar>::run(nplus));
1551 return numext::pow(-one, nplus) * factorial * zeta_impl<Scalar>::run(nplus, x);
1552 }
1553 }
1554};
1555
1556/************************************************************************************************
1557 * Implementation of betainc (incomplete beta integral), based on Cephes *
1558 ************************************************************************************************/
1559
1560template <typename Scalar>
1561struct betainc_impl {
1562 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(Scalar, Scalar, Scalar) {
1563 EIGEN_STATIC_ASSERT((!std::is_same<Scalar, Scalar>::value), THIS_TYPE_IS_NOT_SUPPORTED)
1564 /* betaincf.c
1565 *
1566 * Incomplete beta integral
1567 *
1568 *
1569 * SYNOPSIS:
1570 *
1571 * float a, b, x, y, betaincf();
1572 *
1573 * y = betaincf( a, b, x );
1574 *
1575 *
1576 * DESCRIPTION:
1577 *
1578 * Returns incomplete beta integral of the arguments, evaluated
1579 * from zero to x. The function is defined as
1580 *
1581 * x
1582 * - -
1583 * | (a+b) | | a-1 b-1
1584 * ----------- | t (1-t) dt.
1585 * - - | |
1586 * | (a) | (b) -
1587 * 0
1588 *
1589 * The domain of definition is 0 <= x <= 1. In this
1590 * implementation a and b are restricted to positive values.
1591 * The integral from x to 1 may be obtained by the symmetry
1592 * relation
1593 *
1594 * 1 - betainc( a, b, x ) = betainc( b, a, 1-x ).
1595 *
1596 * The integral is evaluated by a continued fraction expansion.
1597 * If a < 1, the function calls itself recursively after a
1598 * transformation to increase a to a+1.
1599 *
1600 * ACCURACY (float):
1601 *
1602 * Tested at random points (a,b,x) with a and b in the indicated
1603 * interval and x between 0 and 1.
1604 *
1605 * arithmetic domain # trials peak rms
1606 * Relative error:
1607 * IEEE 0,30 10000 3.7e-5 5.1e-6
1608 * IEEE 0,100 10000 1.7e-4 2.5e-5
1609 * The useful domain for relative error is limited by underflow
1610 * of the single precision exponential function.
1611 * Absolute error:
1612 * IEEE 0,30 100000 2.2e-5 9.6e-7
1613 * IEEE 0,100 10000 6.5e-5 3.7e-6
1614 *
1615 * Larger errors may occur for extreme ratios of a and b.
1616 *
1617 * ACCURACY (double):
1618 * arithmetic domain # trials peak rms
1619 * IEEE 0,5 10000 6.9e-15 4.5e-16
1620 * IEEE 0,85 250000 2.2e-13 1.7e-14
1621 * IEEE 0,1000 30000 5.3e-12 6.3e-13
1622 * IEEE 0,10000 250000 9.3e-11 7.1e-12
1623 * IEEE 0,100000 10000 8.7e-10 4.8e-11
1624 * Outputs smaller than the IEEE gradual underflow threshold
1625 * were excluded from these statistics.
1626 *
1627 * ERROR MESSAGES:
1628 * message condition value returned
1629 * incbet domain x<0, x>1 nan
1630 * incbet underflow nan
1631 */
1632 return Scalar(0);
1633 }
1634};
1635
1636/* Continued fraction expansion #1 for incomplete beta integral (small_branch = True)
1637 * Continued fraction expansion #2 for incomplete beta integral (small_branch = False)
1638 */
1639template <typename Scalar>
1640struct incbeta_cfe {
1641 EIGEN_STATIC_ASSERT((std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value),
1642 THIS_TYPE_IS_NOT_SUPPORTED)
1643
1644 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar run(Scalar a, Scalar b, Scalar x, bool small_branch) {
1645 const Scalar big = cephes_helper<Scalar>::big();
1646 const Scalar machep = cephes_helper<Scalar>::machep();
1647 const Scalar biginv = cephes_helper<Scalar>::biginv();
1648
1649 const Scalar zero = 0;
1650 const Scalar one = 1;
1651 const Scalar two = 2;
1652
1653 Scalar xk, pk, pkm1, pkm2, qk, qkm1, qkm2;
1654 Scalar k1, k2, k3, k4, k5, k6, k7, k8, k26update;
1655 Scalar ans;
1656 int n;
1657
1658 constexpr int num_iters = (std::is_same<Scalar, float>::value) ? 100 : 300;
1659 const Scalar thresh = (std::is_same<Scalar, float>::value) ? machep : Scalar(3) * machep;
1660 Scalar r = (std::is_same<Scalar, float>::value) ? zero : one;
1661
1662 if (small_branch) {
1663 k1 = a;
1664 k2 = a + b;
1665 k3 = a;
1666 k4 = a + one;
1667 k5 = one;
1668 k6 = b - one;
1669 k7 = k4;
1670 k8 = a + two;
1671 k26update = one;
1672 } else {
1673 k1 = a;
1674 k2 = b - one;
1675 k3 = a;
1676 k4 = a + one;
1677 k5 = one;
1678 k6 = a + b;
1679 k7 = a + one;
1680 k8 = a + two;
1681 k26update = -one;
1682 x = x / (one - x);
1683 }
1684
1685 pkm2 = zero;
1686 qkm2 = one;
1687 pkm1 = one;
1688 qkm1 = one;
1689 ans = one;
1690 n = 0;
1691
1692 do {
1693 xk = -(x * k1 * k2) / (k3 * k4);
1694 pk = pkm1 + pkm2 * xk;
1695 qk = qkm1 + qkm2 * xk;
1696 pkm2 = pkm1;
1697 pkm1 = pk;
1698 qkm2 = qkm1;
1699 qkm1 = qk;
1700
1701 xk = (x * k5 * k6) / (k7 * k8);
1702 pk = pkm1 + pkm2 * xk;
1703 qk = qkm1 + qkm2 * xk;
1704 pkm2 = pkm1;
1705 pkm1 = pk;
1706 qkm2 = qkm1;
1707 qkm1 = qk;
1708
1709 if (qk != zero) {
1710 r = pk / qk;
1711 if (numext::abs(ans - r) < numext::abs(r) * thresh) {
1712 return r;
1713 }
1714 ans = r;
1715 }
1716
1717 k1 += one;
1718 k2 += k26update;
1719 k3 += two;
1720 k4 += two;
1721 k5 += one;
1722 k6 -= k26update;
1723 k7 += two;
1724 k8 += two;
1725
1726 if ((numext::abs(qk) + numext::abs(pk)) > big) {
1727 pkm2 *= biginv;
1728 pkm1 *= biginv;
1729 qkm2 *= biginv;
1730 qkm1 *= biginv;
1731 }
1732 if ((numext::abs(qk) < biginv) || (numext::abs(pk) < biginv)) {
1733 pkm2 *= big;
1734 pkm1 *= big;
1735 qkm2 *= big;
1736 qkm1 *= big;
1737 }
1738 } while (++n < num_iters);
1739
1740 return ans;
1741 }
1742};
1743
1744/* Helper functions depending on the Scalar type */
1745template <typename Scalar>
1746struct betainc_helper {};
1747
1748template <>
1749struct betainc_helper<float> {
1750 /* Core implementation, assumes a large (> 1.0) */
1751 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float incbsa(float aa, float bb, float xx) {
1752 float ans, a, b, t, x, onemx;
1753 bool reversed_a_b = false;
1754
1755 onemx = 1.0f - xx;
1756
1757 /* see if x is greater than the mean */
1758 if (xx > (aa / (aa + bb))) {
1759 reversed_a_b = true;
1760 a = bb;
1761 b = aa;
1762 t = xx;
1763 x = onemx;
1764 } else {
1765 a = aa;
1766 b = bb;
1767 t = onemx;
1768 x = xx;
1769 }
1770
1771 /* Choose expansion for optimal convergence */
1772 if (b > 10.0f) {
1773 if (numext::abs(b * x / a) < 0.3f) {
1774 t = betainc_helper<float>::incbps(a, b, x);
1775 if (reversed_a_b) t = 1.0f - t;
1776 return t;
1777 }
1778 }
1779
1780 ans = x * (a + b - 2.0f) / (a - 1.0f);
1781 if (ans < 1.0f) {
1782 ans = incbeta_cfe<float>::run(a, b, x, true /* small_branch */);
1783 t = b * numext::log(t);
1784 } else {
1785 ans = incbeta_cfe<float>::run(a, b, x, false /* small_branch */);
1786 t = (b - 1.0f) * numext::log(t);
1787 }
1788
1789 t += a * numext::log(x) + lgamma_impl<float>::run(a + b) - lgamma_impl<float>::run(a) - lgamma_impl<float>::run(b);
1790 t += numext::log(ans / a);
1791 t = numext::exp(t);
1792
1793 if (reversed_a_b) t = 1.0f - t;
1794 return t;
1795 }
1796
1797 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE float incbps(float a, float b, float x) {
1798 float t, u, y, s;
1799 const float machep = cephes_helper<float>::machep();
1800
1801 y = a * numext::log(x) + (b - 1.0f) * numext::log1p(-x) - numext::log(a);
1802 y -= lgamma_impl<float>::run(a) + lgamma_impl<float>::run(b);
1803 y += lgamma_impl<float>::run(a + b);
1804
1805 t = x / (1.0f - x);
1806 s = 0.0f;
1807 u = 1.0f;
1808 do {
1809 b -= 1.0f;
1810 if (b == 0.0f) {
1811 break;
1812 }
1813 a += 1.0f;
1814 u *= t * b / a;
1815 s += u;
1816 } while (numext::abs(u) > machep);
1817
1818 return numext::exp(y) * (1.0f + s);
1819 }
1820};
1821
1822template <>
1823struct betainc_impl<float> {
1824 EIGEN_DEVICE_FUNC static float run(float a, float b, float x) {
1825 const float nan = NumTraits<float>::quiet_NaN();
1826 float ans, t;
1827
1828 if (a == 0.0f && b == 0.0f) return nan;
1829 if (x < 0.0f || x > 1.0f) return nan;
1830 if (a < 0.0f) return nan;
1831 if (b < 0.0f) return nan;
1832 if (a == 0.0f) return 1.0f;
1833 if (b == 0.0f) return 0.0f;
1834 if (x == 0.0f) return 0.0f;
1835 if (x == 1.0f) return 1.0f;
1836 // mtherr("betaincf", DOMAIN);
1837
1838 /* transformation for small aa */
1839 if (a <= 1.0f) {
1840 ans = betainc_helper<float>::incbsa(a + 1.0f, b, x);
1841 t = a * numext::log(x) + b * numext::log1p(-x) + lgamma_impl<float>::run(a + b) -
1842 lgamma_impl<float>::run(a + 1.0f) - lgamma_impl<float>::run(b);
1843 return (ans + numext::exp(t));
1844 } else {
1845 return betainc_helper<float>::incbsa(a, b, x);
1846 }
1847 }
1848};
1849
1850template <>
1851struct betainc_helper<double> {
1852 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE double incbps(double a, double b, double x) {
1853 const double machep = cephes_helper<double>::machep();
1854
1855 double s, t, u, v, n, t1, z, ai;
1856
1857 ai = 1.0 / a;
1858 u = (1.0 - b) * x;
1859 v = u / (a + 1.0);
1860 t1 = v;
1861 t = u;
1862 n = 2.0;
1863 s = 0.0;
1864 z = machep * ai;
1865 while (numext::abs(v) > z) {
1866 u = (n - b) * x / n;
1867 t *= u;
1868 v = t / (a + n);
1869 s += v;
1870 n += 1.0;
1871 }
1872 s += t1;
1873 s += ai;
1874
1875 u = a * numext::log(x);
1876 // TODO: gamma() is not directly implemented in Eigen.
1877 /*
1878 if ((a + b) < maxgam && numext::abs(u) < maxlog) {
1879 t = gamma(a + b) / (gamma(a) * gamma(b));
1880 s = s * t * pow(x, a);
1881 }
1882 */
1883 t = lgamma_impl<double>::run(a + b) - lgamma_impl<double>::run(a) - lgamma_impl<double>::run(b) + u +
1884 numext::log(s);
1885 return numext::exp(t);
1886 }
1887};
1888
1889template <>
1890struct betainc_impl<double> {
1891 EIGEN_DEVICE_FUNC static double run(double aa, double bb, double xx) {
1892 const double nan = NumTraits<double>::quiet_NaN();
1893 const double machep = cephes_helper<double>::machep();
1894 double a, b, t, x, xc, w, y;
1895 bool reversed_a_b = false;
1896
1897 if (aa == 0.0 && bb == 0.0) return nan;
1898 if (xx < 0.0 || xx > 1.0) return nan;
1899 if (aa < 0.0) return nan;
1900 if (bb < 0.0) return nan;
1901 if (aa == 0.0) return 1.0;
1902 if (bb == 0.0) return 0.0;
1903 if (xx == 0.0) return 0.0;
1904 if (xx == 1.0) return 1.0;
1905 // mtherr("incbet", DOMAIN);
1906
1907 if ((bb * xx) <= 1.0 && xx <= 0.95) {
1908 return betainc_helper<double>::incbps(aa, bb, xx);
1909 }
1910
1911 w = 1.0 - xx;
1912
1913 /* Reverse a and b if x is greater than the mean. */
1914 if (xx > (aa / (aa + bb))) {
1915 reversed_a_b = true;
1916 a = bb;
1917 b = aa;
1918 xc = xx;
1919 x = w;
1920 } else {
1921 a = aa;
1922 b = bb;
1923 xc = w;
1924 x = xx;
1925 }
1926
1927 if (reversed_a_b && (b * x) <= 1.0 && x <= 0.95) {
1928 t = betainc_helper<double>::incbps(a, b, x);
1929 if (t <= machep) {
1930 t = 1.0 - machep;
1931 } else {
1932 t = 1.0 - t;
1933 }
1934 return t;
1935 }
1936
1937 /* Choose expansion for better convergence. */
1938 y = x * (a + b - 2.0) - (a - 1.0);
1939 if (y < 0.0) {
1940 w = incbeta_cfe<double>::run(a, b, x, true /* small_branch */);
1941 } else {
1942 w = incbeta_cfe<double>::run(a, b, x, false /* small_branch */) / xc;
1943 }
1944
1945 /* Multiply w by the factor
1946 a b _ _ _
1947 x (1-x) | (a+b) / ( a | (a) | (b) ) . */
1948
1949 y = a * numext::log(x);
1950 t = b * numext::log(xc);
1951 // TODO: gamma is not directly implemented in Eigen.
1952 /*
1953 if ((a + b) < maxgam && numext::abs(y) < maxlog && numext::abs(t) < maxlog)
1954 {
1955 t = pow(xc, b);
1956 t *= pow(x, a);
1957 t /= a;
1958 t *= w;
1959 t *= gamma(a + b) / (gamma(a) * gamma(b));
1960 } else {
1961 */
1962 /* Resort to logarithms. */
1963 y += t + lgamma_impl<double>::run(a + b) - lgamma_impl<double>::run(a) - lgamma_impl<double>::run(b);
1964 y += numext::log(w / a);
1965 t = numext::exp(y);
1966
1967 /* } */
1968 // done:
1969
1970 if (reversed_a_b) {
1971 if (t <= machep) {
1972 t = 1.0 - machep;
1973 } else {
1974 t = 1.0 - t;
1975 }
1976 }
1977 return t;
1978 }
1979};
1980
1981} // end namespace internal
1982
1983namespace numext {
1984
1985template <typename Scalar>
1986EIGEN_DEVICE_FUNC inline auto lgamma(const Scalar& x) -> decltype(EIGEN_MATHFUNC_IMPL(lgamma, Scalar)::run(x)) {
1987 return EIGEN_MATHFUNC_IMPL(lgamma, Scalar)::run(x);
1988}
1989
1990template <typename Scalar>
1991EIGEN_DEVICE_FUNC inline auto digamma(const Scalar& x) -> decltype(EIGEN_MATHFUNC_IMPL(digamma, Scalar)::run(x)) {
1992 return EIGEN_MATHFUNC_IMPL(digamma, Scalar)::run(x);
1993}
1994
1995template <typename Scalar>
1996EIGEN_DEVICE_FUNC inline auto zeta(const Scalar& x, const Scalar& q)
1997 -> decltype(EIGEN_MATHFUNC_IMPL(zeta, Scalar)::run(x, q)) {
1998 return EIGEN_MATHFUNC_IMPL(zeta, Scalar)::run(x, q);
1999}
2000
2001template <typename Scalar>
2002EIGEN_DEVICE_FUNC inline auto polygamma(const Scalar& n, const Scalar& x)
2003 -> decltype(EIGEN_MATHFUNC_IMPL(polygamma, Scalar)::run(n, x)) {
2004 return EIGEN_MATHFUNC_IMPL(polygamma, Scalar)::run(n, x);
2005}
2006
2007template <typename Scalar>
2008EIGEN_DEVICE_FUNC inline auto erf(const Scalar& x) -> decltype(EIGEN_MATHFUNC_IMPL(erf, Scalar)::run(x)) {
2009 return EIGEN_MATHFUNC_IMPL(erf, Scalar)::run(x);
2010}
2011
2012template <typename Scalar>
2013EIGEN_DEVICE_FUNC inline auto erfc(const Scalar& x) -> decltype(EIGEN_MATHFUNC_IMPL(erfc, Scalar)::run(x)) {
2014 return EIGEN_MATHFUNC_IMPL(erfc, Scalar)::run(x);
2015}
2016
2017template <typename Scalar>
2018EIGEN_DEVICE_FUNC inline auto ndtri(const Scalar& x) -> decltype(EIGEN_MATHFUNC_IMPL(ndtri, Scalar)::run(x)) {
2019 return EIGEN_MATHFUNC_IMPL(ndtri, Scalar)::run(x);
2020}
2021
2022template <typename Scalar>
2023EIGEN_DEVICE_FUNC inline auto igamma(const Scalar& a, const Scalar& x)
2024 -> decltype(EIGEN_MATHFUNC_IMPL(igamma, Scalar)::run(a, x)) {
2025 return EIGEN_MATHFUNC_IMPL(igamma, Scalar)::run(a, x);
2026}
2027
2028template <typename Scalar>
2029EIGEN_DEVICE_FUNC inline auto igamma_der_a(const Scalar& a, const Scalar& x)
2030 -> decltype(EIGEN_MATHFUNC_IMPL(igamma_der_a, Scalar)::run(a, x)) {
2031 return EIGEN_MATHFUNC_IMPL(igamma_der_a, Scalar)::run(a, x);
2032}
2033
2034template <typename Scalar>
2035EIGEN_DEVICE_FUNC inline auto gamma_sample_der_alpha(const Scalar& a, const Scalar& x)
2036 -> decltype(EIGEN_MATHFUNC_IMPL(gamma_sample_der_alpha, Scalar)::run(a, x)) {
2037 return EIGEN_MATHFUNC_IMPL(gamma_sample_der_alpha, Scalar)::run(a, x);
2038}
2039
2040template <typename Scalar>
2041EIGEN_DEVICE_FUNC inline auto igammac(const Scalar& a, const Scalar& x)
2042 -> decltype(EIGEN_MATHFUNC_IMPL(igammac, Scalar)::run(a, x)) {
2043 return EIGEN_MATHFUNC_IMPL(igammac, Scalar)::run(a, x);
2044}
2045
2046template <typename Scalar>
2047EIGEN_DEVICE_FUNC inline auto betainc(const Scalar& a, const Scalar& b, const Scalar& x)
2048 -> decltype(EIGEN_MATHFUNC_IMPL(betainc, Scalar)::run(a, b, x)) {
2049 return EIGEN_MATHFUNC_IMPL(betainc, Scalar)::run(a, b, x);
2050}
2051
2052} // end namespace numext
2053} // end namespace Eigen
2054
2055#endif // EIGEN_SPECIAL_FUNCTIONS_H
Namespace containing all symbols from the Eigen library.