11#ifndef EIGEN_SPECIAL_FUNCTIONS_H
12#define EIGEN_SPECIAL_FUNCTIONS_H
15#include "./InternalHeaderCheck.h"
47template <
typename Scalar>
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)
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
62#if defined(__GLIBC__) && ((__GLIBC__ == 2 && __GLIBC_MINOR__ < 19) || __GLIBC__ < 2) && \
63 (defined(_BSD_SOURCE) || defined(_SVID_SOURCE))
64#define EIGEN_HAS_LGAMMA_R
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__)
72 return ::lgammaf_r(x, &dummy);
73#elif defined(SYCL_DEVICE_ONLY)
74 return cl::sycl::lgamma(x);
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__)
86 return ::lgamma_r(x, &dummy);
87#elif defined(SYCL_DEVICE_ONLY)
88 return cl::sycl::lgamma(x);
95#undef EIGEN_HAS_LGAMMA_R
114template <
typename Scalar>
115struct digamma_impl_maybe_poly {
116 EIGEN_STATIC_ASSERT((std::is_same<Scalar, Scalar>::value ==
false), THIS_TYPE_IS_NOT_SUPPORTED)
118 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar run(
const Scalar) {
return Scalar(0); }
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};
130 return z * internal::ppolevl<float, 3>::run(z, A);
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};
146 return z * internal::ppolevl<double, 6>::run(z, A);
152template <
typename Scalar>
154 EIGEN_DEVICE_FUNC
static Scalar run(Scalar x) {
212 Scalar p, q, nz, s, w, y;
213 bool negative =
false;
215 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
216 const Scalar m_pi = Scalar(EIGEN_PI);
218 const Scalar zero = Scalar(0);
219 const Scalar one = Scalar(1);
220 const Scalar half = Scalar(0.5);
227 return (std::signbit(x)) ? NumTraits<Scalar>::infinity() : -NumTraits<Scalar>::infinity();
232 p = numext::floor(q);
245 nz = m_pi / numext::tan(m_pi * nz);
255 while (s < Scalar(10)) {
260 y = digamma_impl_maybe_poly<Scalar>::run(s);
262 y = numext::log(s) - (half / s) - y - w;
264 return (negative) ? y - nz : y;
271namespace unqualified_erf {
275auto test_erf(
int) ->
decltype(void(erf(std::declval<const T&>())), std::true_type{});
277std::false_type test_erf(...);
279auto test_erfc(
int) ->
decltype(void(erfc(std::declval<const T&>())), std::true_type{});
281std::false_type test_erfc(...);
285struct has_erf : decltype(unqualified_erf::test_erf<T>(0)) {};
287struct has_erfc : decltype(unqualified_erf::test_erfc<T>(0)) {};
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);
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));
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);
318 const T x_abs_gt_one_mask = pcmp_lt(one, x2);
319 if (!predux_any(x_abs_gt_one_mask))
return erfc_small;
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);
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);
353EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T erf_over_x_double_small(
const T& x2) {
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,
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);
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};
405 const T x2_lo = twoprod_low(x, x, x2);
410 const T exp2_hi = pexp(pnegate(x2));
411 const T z = pnmadd(exp2_hi, x2_lo, exp2_hi);
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);
423EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erfc<double>::run(
const T& x_in) {
426 constexpr double kClamp = 28.0;
427 const T x = pmin<PropagateNaN>(pmax<PropagateNaN>(x_in, pset1<T>(-kClamp)), pset1<T>(kClamp));
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);
436 const T x_abs_gt_one_mask = pcmp_lt(one, x2);
437 if (!predux_any(x_abs_gt_one_mask))
return erfc_small;
439 const T erfc_large = erfc_double_large(x, x2);
440 return pselect(x_abs_gt_one_mask, erfc_large, erfc_small);
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>()); }
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);
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>());
460 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE T run_scalar(
const T& x, std::true_type) {
461 EIGEN_USING_STD(erfc);
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)
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);
480 return generic_fast_erfc<float>::run(x);
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);
491 return generic_fast_erfc<double>::run(x);
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);
514EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erf<float>::run(
const T& x) {
516 constexpr float alpha[] = {2.123732201653183437883853912353515625e-06f, 2.861979592125862836837768554687500000e-04f,
517 3.658048342913389205932617187500000000e-03f, 5.243302136659622192382812500000000000e-02f,
518 1.874160766601562500000000000000000000e-01f, 1.128379106521606445312500000000000000e+00f};
521 constexpr float beta[] = {3.89185734093189239501953125000e-05f, 1.14329601638019084930419921875e-03f,
522 1.47520881146192550659179687500e-02f, 1.12945675849914550781250000000e-01f,
523 4.99425798654556274414062500000e-01f, 1.0f};
528 const T x2 = pmin(pset1<T>(16.0f), pmul(x, x));
531 T p = ppolevl<T, 5>::run(x2, alpha);
535 T q = ppolevl<T, 5>::run(x2, beta);
536 const T r = pdiv(p, q);
539 return pmax<PropagateNaN>(pmin<PropagateNaN>(r, pset1<T>(1.0f)), pset1<T>(-1.0f));
544EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T generic_fast_erf<double>::run(
const T& x_in) {
547 constexpr double kClamp = 28.0;
548 const T x = pmin<PropagateNaN>(pmax<PropagateNaN>(x_in, pset1<T>(-kClamp)), pset1<T>(kClamp));
550 T erf_small = pmul(x, erf_over_x_double_small(x2));
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;
559 const T erf_large = psub(one, erfc_double_large(x, x2));
560 return pselect(x_abs_gt_one_mask, erf_large, erf_small);
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>()); }
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);
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>());
580 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE T run_scalar(
const T& x, std::true_type) {
581 EIGEN_USING_STD(erf);
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)
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);
600 return generic_fast_erf<float>::run(x);
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);
611 return generic_fast_erf<double>::run(x);
671template <typename T, bool IsScalar = is_scalar<T>::value>
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);
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;
691EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T flipsign(
const T& should_flipsign,
const T& x) {
692 return flipsign_impl<T>::run(should_flipsign, x);
695template <typename T, bool IsScalar = is_scalar<T>::value>
696struct ndtri_negative_infinity_impl;
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);
706struct ndtri_negative_infinity_impl<T, true> {
707 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T run(
const T& positive_infinity) {
return -positive_infinity; }
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;
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);
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) {
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)};
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));
779 x = psqrt(pmul(neg_two, plog(b)));
780 x0 = psub(x, pdiv(plog(x), x));
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));
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);
793 const T zero = pset1<T>(ScalarType(0));
794 const T one = pset1<T>(ScalarType(1));
796 const T exp_neg_two = pset1<T>(ScalarType(0.13533528323661269189));
797 T b, ndtri, should_flipsign;
799 should_flipsign = pcmp_le(a, psub(one, exp_neg_two));
800 b = pselect(should_flipsign, a, psub(one, a));
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));
805 return pselect(pcmp_eq(a, zero), neg_maxnum, pselect(pcmp_eq(one, a), maxnum, ndtri));
808template <
typename Scalar>
810 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar run(
const Scalar x) {
return generic_ndtri<Scalar, Scalar>(x); }
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");
824 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar big() {
825 eigen_assert(
false &&
"big not supported for this type");
828 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar biginv() {
829 eigen_assert(
false &&
"biginv not supported for this type");
835struct cephes_helper<float> {
836 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
float machep() {
837 return NumTraits<float>::epsilon() / 2;
839 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
float big() {
841 return 1.0f / (NumTraits<float>::epsilon() / 2);
843 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
float biginv() {
850struct cephes_helper<double> {
851 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
double machep() {
852 return NumTraits<double>::epsilon() / 2;
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() {
857 return NumTraits<double>::epsilon();
861enum IgammaComputationMode { VALUE, DERIVATIVE, SAMPLE_DERIVATIVE };
863template <
typename Scalar>
864EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar main_igamma_term(Scalar a, Scalar x) {
866 Scalar logax = a * numext::log(x) - x - lgamma_impl<Scalar>::run(a);
867 if (logax < -numext::log(NumTraits<Scalar>::highest()) ||
869 (numext::isnan)(logax)) {
872 return numext::exp(logax);
875template <
typename Scalar, IgammaComputationMode mode>
876EIGEN_DEVICE_FUNC
constexpr int igamma_num_iterations() {
879 return mode == VALUE ? 2000
880 : std::is_same<Scalar, float>::value ? 200
881 : std::is_same<Scalar, double>::value ? 500
885template <
typename Scalar, IgammaComputationMode mode>
886struct igammac_cf_impl {
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();
904 if ((numext::isinf)(x)) {
908 Scalar ax = main_igamma_term<Scalar>(a, x);
919 Scalar z = x + y + one;
923 Scalar pkm1 = x + one;
925 Scalar ans = pkm1 / qkm1;
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;
933 for (
int i = 0; i < igamma_num_iterations<Scalar, mode>(); i++) {
939 Scalar pk = pkm1 * z - pkm2 * yc;
940 Scalar qk = qkm1 * z - qkm2 * yc;
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;
946 Scalar ans_prev = ans;
949 Scalar dans_da_prev = dans_da;
950 dans_da = (dpk_da - ans * dqk_da) / qk;
953 if (numext::abs(ans_prev - ans) <= machep * numext::abs(ans)) {
957 if (numext::abs(dans_da - dans_da_prev) <= machep) {
973 if (numext::abs(pk) > big) {
987 Scalar dlogax_da = numext::log(x) - digamma_impl<Scalar>::run(a);
988 Scalar dax_da = ax * dlogax_da;
994 return ans * dax_da + dans_da * ax;
995 case SAMPLE_DERIVATIVE:
997 return -(dans_da + ans * dlogax_da) * x;
1002template <
typename Scalar, IgammaComputationMode mode>
1003struct igamma_series_impl {
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();
1017 Scalar ax = main_igamma_term<Scalar>(a, x);
1034 Scalar dc_da = zero;
1035 Scalar dans_da = zero;
1037 for (
int i = 0; i < igamma_num_iterations<Scalar, mode>(); i++) {
1039 Scalar term = x / r;
1040 Scalar dterm_da = -x / (r * r);
1041 dc_da = term * dc_da + dterm_da * c;
1046 if (mode == VALUE) {
1047 if (c <= machep * ans) {
1051 if (numext::abs(dc_da) <= machep * numext::abs(dans_da)) {
1057 Scalar dlogax_da = numext::log(x) - digamma_impl<Scalar>::run(a + one);
1058 Scalar dax_da = ax * dlogax_da;
1064 return ans * dax_da + dans_da * ax;
1065 case SAMPLE_DERIVATIVE:
1067 return -(dans_da + ans * dlogax_da) * x / a;
1072template <
typename Scalar>
1073struct igammac_impl {
1074 EIGEN_DEVICE_FUNC
static Scalar run(Scalar a, Scalar x) {
1129 const Scalar zero = 0;
1130 const Scalar one = 1;
1131 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1133 if ((x < zero) || (a <= zero)) {
1138 if ((numext::isnan)(a) || (numext::isnan)(x)) {
1142 if ((x < one) || (x < a)) {
1145 return numext::mini(one, numext::maxi(zero, one - igamma_series_impl<Scalar, VALUE>::run(a, x)));
1148 return igammac_cf_impl<Scalar, VALUE>::run(a, x);
1156template <
typename Scalar, IgammaComputationMode mode>
1157struct igamma_generic_impl {
1158 EIGEN_DEVICE_FUNC
static Scalar run(Scalar a, Scalar x) {
1167 const Scalar zero = 0;
1168 const Scalar one = 1;
1169 const Scalar nan = NumTraits<Scalar>::quiet_NaN();
1171 if (x == zero)
return zero;
1173 if ((x < zero) || (a <= zero)) {
1177 if ((numext::isnan)(a) || (numext::isnan)(x)) {
1181 if ((x > one) && (x > a)) {
1182 Scalar ret = igammac_cf_impl<Scalar, mode>::run(a, x);
1183 if (mode == VALUE) {
1186 return numext::mini(one, numext::maxi(zero, one - ret));
1192 Scalar ret = igamma_series_impl<Scalar, mode>::run(a, x);
1193 if (mode == VALUE) {
1196 return numext::mini(one, numext::maxi(zero, ret));
1202template <
typename Scalar>
1203struct igamma_impl : igamma_generic_impl<Scalar, VALUE> {
1273template <
typename Scalar>
1274struct igamma_der_a_impl : igamma_generic_impl<Scalar, DERIVATIVE> {
1291template <
typename Scalar>
1292struct gamma_sample_der_alpha_impl : igamma_generic_impl<Scalar, SAMPLE_DERIVATIVE> {
1336template <
typename Scalar>
1337struct zeta_impl_series {
1338 EIGEN_STATIC_ASSERT((std::is_same<Scalar, Scalar>::value ==
false), THIS_TYPE_IS_NOT_SUPPORTED)
1340 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE Scalar run(
const Scalar) {
return Scalar(0); }
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) {
1351 b = numext::pow(a, -x);
1353 if (numext::abs(b / s) < machep)
return true;
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) {
1366 while ((i < 9) || (a <= 9.0)) {
1369 b = numext::pow(a, -x);
1371 if (numext::abs(b / s) < machep)
return true;
1379template <
typename Scalar>
1381 EIGEN_DEVICE_FUNC
static Scalar run(Scalar x, Scalar q) {
1444 Scalar p, r, a, b, k, s, t, w;
1446 const Scalar A[] = {
1452 Scalar(-1.8924375803183791606e9),
1453 Scalar(7.47242496e10),
1454 Scalar(-2.950130727918164224e12),
1455 Scalar(1.1646782814350067249e14),
1456 Scalar(-4.5979787224074726105e15),
1457 Scalar(1.8152105401943546773e17),
1458 Scalar(-7.1661652561756670113e18)
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();
1466 if (x == one)
return maxnum;
1473 if (q == numext::floor(q)) {
1474 if (numext::rint(Scalar(0.5) * x) == Scalar(0.5) * x) {
1481 r = numext::floor(p);
1482 if (p != r)
return nan;
1490 s = numext::pow(q, -x);
1494 if (zeta_impl_series<Scalar>::run(a, b, s, x, machep)) {
1501 if (numext::equal_strict(b, zero)) {
1506 s += b * w / (x - one);
1511 for (i = 0; i < 12; i++) {
1516 t = numext::abs(t / s);
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();
1541 if (numext::floor(n) != n || n < zero) {
1545 else if (n == zero) {
1546 return digamma_impl<Scalar>::run(x);
1550 Scalar factorial = numext::exp(lgamma_impl<Scalar>::run(nplus));
1551 return numext::pow(-one, nplus) * factorial * zeta_impl<Scalar>::run(nplus, x);
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)
1639template <
typename Scalar>
1641 EIGEN_STATIC_ASSERT((std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value),
1642 THIS_TYPE_IS_NOT_SUPPORTED)
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();
1649 const Scalar zero = 0;
1650 const Scalar one = 1;
1651 const Scalar two = 2;
1653 Scalar xk, pk, pkm1, pkm2, qk, qkm1, qkm2;
1654 Scalar k1, k2, k3, k4, k5, k6, k7, k8, k26update;
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;
1693 xk = -(x * k1 * k2) / (k3 * k4);
1694 pk = pkm1 + pkm2 * xk;
1695 qk = qkm1 + qkm2 * xk;
1701 xk = (x * k5 * k6) / (k7 * k8);
1702 pk = pkm1 + pkm2 * xk;
1703 qk = qkm1 + qkm2 * xk;
1711 if (numext::abs(ans - r) < numext::abs(r) * thresh) {
1726 if ((numext::abs(qk) + numext::abs(pk)) > big) {
1732 if ((numext::abs(qk) < biginv) || (numext::abs(pk) < biginv)) {
1738 }
while (++n < num_iters);
1745template <
typename Scalar>
1746struct betainc_helper {};
1749struct betainc_helper<float> {
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;
1758 if (xx > (aa / (aa + bb))) {
1759 reversed_a_b =
true;
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;
1780 ans = x * (a + b - 2.0f) / (a - 1.0f);
1782 ans = incbeta_cfe<float>::run(a, b, x,
true );
1783 t = b * numext::log(t);
1785 ans = incbeta_cfe<float>::run(a, b, x,
false );
1786 t = (b - 1.0f) * numext::log(t);
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);
1793 if (reversed_a_b) t = 1.0f - t;
1797 EIGEN_DEVICE_FUNC
static EIGEN_STRONG_INLINE
float incbps(
float a,
float b,
float x) {
1799 const float machep = cephes_helper<float>::machep();
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);
1816 }
while (numext::abs(u) > machep);
1818 return numext::exp(y) * (1.0f + s);
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();
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;
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));
1845 return betainc_helper<float>::incbsa(a, b, x);
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();
1855 double s, t, u, v, n, t1, z, ai;
1865 while (numext::abs(v) > z) {
1866 u = (n - b) * x / n;
1875 u = a * numext::log(x);
1883 t = lgamma_impl<double>::run(a + b) - lgamma_impl<double>::run(a) - lgamma_impl<double>::run(b) + u +
1885 return numext::exp(t);
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;
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;
1907 if ((bb * xx) <= 1.0 && xx <= 0.95) {
1908 return betainc_helper<double>::incbps(aa, bb, xx);
1914 if (xx > (aa / (aa + bb))) {
1915 reversed_a_b =
true;
1927 if (reversed_a_b && (b * x) <= 1.0 && x <= 0.95) {
1928 t = betainc_helper<double>::incbps(a, b, x);
1938 y = x * (a + b - 2.0) - (a - 1.0);
1940 w = incbeta_cfe<double>::run(a, b, x,
true );
1942 w = incbeta_cfe<double>::run(a, b, x,
false ) / xc;
1949 y = a * numext::log(x);
1950 t = b * numext::log(xc);
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);
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);
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);
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);
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);
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);
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);
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);
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);
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);
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);
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);
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);
Namespace containing all symbols from the Eigen library.