From 919ae4c5884ec0cfdd9e41c44268befb7f5a7b3a Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:13 +0300 Subject: [PATCH 01/15] Add digamma_diff for digamma differences without cancellation --- stan/math/prim/fun.hpp | 1 + stan/math/prim/fun/digamma_diff.hpp | 199 ++++++++++++++++++ test/unit/math/mix/fun/digamma_diff_test.cpp | 17 ++ test/unit/math/prim/fun/digamma_diff_test.cpp | 104 +++++++++ 4 files changed, 321 insertions(+) create mode 100644 stan/math/prim/fun/digamma_diff.hpp create mode 100644 test/unit/math/mix/fun/digamma_diff_test.cpp create mode 100644 test/unit/math/prim/fun/digamma_diff_test.cpp diff --git a/stan/math/prim/fun.hpp b/stan/math/prim/fun.hpp index 8ef4204c81b..25526122101 100644 --- a/stan/math/prim/fun.hpp +++ b/stan/math/prim/fun.hpp @@ -66,6 +66,7 @@ #include #include #include +#include #include #include #include diff --git a/stan/math/prim/fun/digamma_diff.hpp b/stan/math/prim/fun/digamma_diff.hpp new file mode 100644 index 00000000000..22ea6070efd --- /dev/null +++ b/stan/math/prim/fun/digamma_diff.hpp @@ -0,0 +1,199 @@ +#ifndef STAN_MATH_PRIM_FUN_DIGAMMA_DIFF_HPP +#define STAN_MATH_PRIM_FUN_DIGAMMA_DIFF_HPP + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace stan { +namespace math { + +namespace internal { +/** + * Gradient code that needs digamma(x) - digamma(x + d) can use the plain + * difference for x below this value: its absolute error is a few eps times + * |digamma|, and x times that error (the error of a gradient in log(x)) + * stays at the rounding level. From this value on it must use digamma_diff. + * The derivatives of lbeta use this rule. + */ +constexpr double digamma_diff_min_x = 10.0; +} // namespace internal + +/** + * Return the difference of the digamma function at two arguments that + * differ by a nonnegative offset, + * + \f[ + \mbox{digamma\_diff}(x, d) = \Psi(x + d) - \Psi(x). + \f] + * + * The plain difference of two `digamma` calls loses all accuracy when `x` + * is large and `d` is not: both values are close to \f$\log x\f$ and their + * difference is close to \f$d / x\f$, so the absolute error is about + * \f$\epsilon \log x\f$. This function has a relative error of a few ulp + * for all \f$x > 0\f$ and \f$d \ge 0\f$. Densities with shape or + * dispersion parameters use it for gradients such as + * \f$\Psi(\alpha + n) - \Psi(\alpha)\f$. + * + * Method. Two cases have a cheaper form of the same accuracy. + * + * - If `d` is not an autodiff type and is an integer from 0 to 8 (a count), + * \f$\Psi(x + d) - \Psi(x) = \sum_{j=0}^{d-1} 1 / (x + j)\f$. Every term + * is positive, the terms are added from the smallest, and for \f$d = 1\f$ + * the result is exactly \f$1/x\f$. + * - If \f$x < 10\f$ and \f$d \ge 10\f$, the plain difference + * \f$\Psi(x + d) - \Psi(x)\f$ is used. The result is at least + * \f$\Psi(20) - \Psi(10) \approx 0.72\f$ and \f$\Psi(x) \le \Psi(10) + * \approx 2.25\f$, so the rounding errors of the two values grow by a + * factor of at most about 7 (nothing cancels where \f$\Psi(x) < 0\f$). + * + * Otherwise, for \f$x < 10\f$ the recurrence + * \f$\Psi(y + 1) = \Psi(y) + 1/y\f$ shifts both arguments up by the same + * integer \f$J\f$: + * + \f[ + \Psi(x + d) - \Psi(x) = \sum_{j=0}^{J-1} \frac{d}{(x + j)(x + j + d)} + + \Psi(y + d) - \Psi(y), \quad y = x + J \ge 10. + \f] + * + * Every term of the sum is positive. A term with \f$d \ge x + j\f$ is + * split into \f$1/(x + j) - 1/(x + j + d)\f$, which does not cancel, and + * the parts \f$1/(x + j)\f$ are added last, so that for small \f$x\f$ the + * dominant \f$1/x\f$ keeps its correct rounding; for \f$d = 1\f$ the result + * is exactly \f$1/x\f$. For \f$y \ge 10\f$ the asymptotic + * expansion \f$\Psi(y) = \log y - 1/(2y) - \sum_i B_{2i} / (2i\,y^{2i})\f$ + * gives + * + \f[ + \Psi(y + d) - \Psi(y) = \log\left(1 + \frac{d}{y}\right) + + \frac{d}{2y(y + d)} + + \sum_{i=1}^{8} \frac{B_{2i}}{2i} \left(u^i - v^i\right), + \f] + * + * with \f$u = y^{-2}\f$ and \f$v = (y + d)^{-2}\f$. The differences + * \f$u^i - v^i = (u - v) h_i\f$, \f$h_1 = 1\f$, + * \f$h_{i+1} = u h_i + v^i\f$, and + * \f$u - v = u \frac{d}{y + d} \left(1 + \frac{y}{y + d}\right)\f$ are + * formed without cancellation. The truncation error is below + * \f$|B_{18}| y^{-18} \approx 6 \times 10^{-17}\f$ relative. + * + * @tparam T1 type of the first argument + * @tparam T2 type of the second argument + * @param x first argument, positive + * @param d offset, nonnegative + * @return \f$\Psi(x + d) - \Psi(x)\f$ + * @throw std::domain_error if `x` is not positive or `d` is negative + */ +template * = nullptr> +inline return_type_t digamma_diff(const T1& x, const T2& d) { + using T_ret = return_type_t; + static constexpr const char* function = "digamma_diff"; + if (is_any_nan(x, d)) { + return NOT_A_NUMBER; + } + check_positive(function, "first argument", x); + check_nonnegative(function, "second argument", d); + // d = 0 needs no special case: every term below is then exactly 0, and + // the derivative with respect to d stays correct + if (is_inf(d)) { + return INFTY; + } + if (is_inf(x)) { + return T_ret(0.0); + } + + // A count d from 0 to 8: sum of the positive terms 1 / (x + j), from the + // smallest (exactly 1 / x for d = 1) + if constexpr (std::is_arithmetic::value) { + if (d <= 8 && d == std::floor(d)) { + T_ret sum(0.0); + for (int j = static_cast(d) - 1; j >= 0; --j) { + sum += inv(x + j); + } + return sum; + } + } + // x < 10 and d >= 10: the plain difference does not cancel much + if (value_of_rec(x) < 10.0 && value_of_rec(d) >= 10.0) { + return digamma(x + d) - digamma(x); + } + + // B_{2i} / (2i), i = 1..8 + static constexpr double coeffs[] + = {1.0 / 12.0, -1.0 / 120.0, 1.0 / 252.0, -1.0 / 240.0, + 1.0 / 132.0, -691.0 / 32760.0, 1.0 / 12.0, -3617.0 / 8160.0}; + static constexpr int n_coeffs = 8; + static constexpr double shift_to = 10.0; + + // Each shift term is d / (y (y + d)) = 1 / y - 1 / (y + d). For d < y it + // is formed as (d / (y + d)) / y, which neither cancels nor overflows. For + // d >= y the two parts do not cancel either, and the parts 1 / y are + // summed separately and added last: for small x the result is dominated + // by 1 / x, which then keeps its correct rounding (for d = 1 the result is + // exactly 1 / x, as for digamma(x + 1) - digamma(x)). + T_ret inv_sum(0.0); + T_ret shift_sum(0.0); + T_ret y = x; + while (value_of_rec(y) < shift_to) { + if (value_of_rec(d) >= value_of_rec(y)) { + inv_sum += inv(y); + shift_sum -= inv(y + d); + } else { + shift_sum += (d / (y + d)) / y; + } + y += 1.0; + } + + const T_ret y_plus_d = y + d; + const T_ret d_frac = d / y_plus_d; + // square(inv(.)) underflows to 0 where inv(square(.)) would overflow + const T_ret u = square(inv(y)); + const T_ret v = square(inv(y_plus_d)); + const T_ret u_minus_v = u * d_frac * (1.0 + y / y_plus_d); + T_ret h(1.0); + T_ret v_pow = v; + T_ret series = coeffs[0]; + for (int i = 1; i < n_coeffs; ++i) { + h = u * h + v_pow; + v_pow *= v; + series += coeffs[i] * h; + } + return inv_sum + + (shift_sum + log1p(d / y) + 0.5 * d_frac / y + u_minus_v * series); +} + +/** + * Enables the vectorized application of the digamma_diff function, when the + * first and/or second arguments are containers. + * + * @tparam T1 type of first input + * @tparam T2 type of second input + * @param a First input + * @param b Second input + * @return digamma_diff function applied to the two inputs. + */ +template * = nullptr> +inline auto digamma_diff(T1&& a, T2&& b) { + return apply_scalar_binary( + [](auto&& c, auto&& d) { + return digamma_diff(std::forward(c), + std::forward(d)); + }, + std::forward(a), std::forward(b)); +} + +} // namespace math +} // namespace stan + +#endif diff --git a/test/unit/math/mix/fun/digamma_diff_test.cpp b/test/unit/math/mix/fun/digamma_diff_test.cpp new file mode 100644 index 00000000000..b6c71d57260 --- /dev/null +++ b/test/unit/math/mix/fun/digamma_diff_test.cpp @@ -0,0 +1,17 @@ +#include + +TEST(mathMixScalFun, digamma_diff) { + auto f = [](const auto& x, const auto& d) { + return stan::math::digamma_diff(x, d); + }; + // both sides of the shift threshold x = 10, integer and real offsets + stan::test::expect_ad(f, 0.5, 1.0); + stan::test::expect_ad(f, 2.3, 0.5); + stan::test::expect_ad(f, 9.5, 3.0); + stan::test::expect_ad(f, 10.5, 0.25); + stan::test::expect_ad(f, 25.0, 57.0); + stan::test::expect_ad(f, 1e3, 117.0); + // invalid arguments throw for every type + stan::test::expect_ad(f, -1.0, 2.0); + stan::test::expect_ad(f, 2.0, -1.0); +} diff --git a/test/unit/math/prim/fun/digamma_diff_test.cpp b/test/unit/math/prim/fun/digamma_diff_test.cpp new file mode 100644 index 00000000000..7cdae318efb --- /dev/null +++ b/test/unit/math/prim/fun/digamma_diff_test.cpp @@ -0,0 +1,104 @@ +#include +#include +#include +#include +#include + +TEST(MathFunctions, digamma_diff_special_cases) { + using stan::math::digamma_diff; + const double nan = std::numeric_limits::quiet_NaN(); + const double inf = std::numeric_limits::infinity(); + + EXPECT_TRUE(std::isnan(digamma_diff(nan, 1.0))); + EXPECT_TRUE(std::isnan(digamma_diff(1.0, nan))); + EXPECT_EQ(digamma_diff(2.5, 0.0), 0.0); + EXPECT_EQ(digamma_diff(1e18, 0), 0.0); + EXPECT_EQ(digamma_diff(2.5, inf), inf); + EXPECT_EQ(digamma_diff(inf, 2.5), 0.0); + EXPECT_THROW(digamma_diff(0.0, 1.0), std::domain_error); + EXPECT_THROW(digamma_diff(-1.5, 1.0), std::domain_error); + EXPECT_THROW(digamma_diff(1.0, -0.5), std::domain_error); +} + +TEST(MathFunctions, digamma_diff_integer_offset) { + // psi(x + k) - psi(x) = sum_{j < k} 1 / (x + j) for integer k; the sum in + // double has a rounding error of up to about k eps / 2 + using stan::math::digamma_diff; + for (double x : {0.25, 1.0, 3.5, 9.5, 10.0, 47.0}) { + double sum = 0; + for (int k = 1; k <= 20; ++k) { + sum += 1.0 / (x + k - 1); + EXPECT_NEAR(digamma_diff(x, k), sum, 2e-15 * sum) + << "x = " << x << ", k = " << k; + } + } +} + +TEST(MathFunctions, digamma_diff_count_and_plain_paths) { + // A count d from 0 to 8 is summed directly; for x < 10 and d >= 10 the + // plain difference is used. d = 1 gives exactly 1 / x, as + // digamma(x + 1) - digamma(x) does. References: mpmath at 80 digits. + using stan::math::digamma_diff; + for (double x : {1e-300, 0.1, 3.7, 9.99, 1e8, 1e300}) { + EXPECT_EQ(digamma_diff(x, 1), 1.0 / x) << "x = " << x; + EXPECT_EQ(digamma_diff(x, 1.0), 1.0 / x) << "x = " << x; + } + // count path + EXPECT_NEAR(digamma_diff(2.5, 8), 1.599844393652443188, 1e-15 * 1.6); + EXPECT_NEAR(digamma_diff(0x1.89374bc6a7efap-9, 7), 335.77886985018505786, + 1e-15 * 335.8); + // plain difference, x < 10 and d >= 10 + EXPECT_NEAR(digamma_diff(0x1.ee45a1cac0831p+2, 10.0), 0.86831971186762265543, + 2e-15 * 0.87); + EXPECT_NEAR(digamma_diff(0x1.7ae147ae147aep-2, 57.0), 6.8360822546693604884, + 2e-15 * 6.84); + EXPECT_NEAR(digamma_diff(0x1.3fae147ae147bp+3, 1e6), 11.564819675088084883, + 2e-15 * 11.6); + // d = 9 is neither a count of the first path nor large: the shift and the + // asymptotic series + EXPECT_NEAR(digamma_diff(6.25, 9.0), 0.9409809180725561924, 1e-15 * 0.94); +} + +namespace digamma_diff_test_internal { +struct TestValue { + double x; + double d; + double val; +}; + +// psi(x + d) - psi(x) computed with mpmath at 60 digits plus the digits lost +// to cancellation (log10(x) for large x, log10(1 / d) for small d); the +// arguments are written in hex so that they are exact. Points of large x +// are where the plain difference digamma(x + d) - digamma(x) is wrong: at +// x = 1e12, d = 57 it has a relative error of 4.6e-06, at x = 1e18 of 1. +std::vector testValues = { + {0x1.0624dd2f1a9fcp-10, 0x1.0000000000000p-1, 9.9861698830979385e+2}, + {0x1.0000000000000p-1, 0x1.0000000000000p+0, 2.0000000000000000}, + {0x1.0000000000000p+0, 0x1.8000000000000p+1, 1.8333333333333333}, + {0x1.4000000000000p+1, 0x1.5798ee2308c3ap-27, 4.9035775491921462e-9}, + {0x1.3800000000000p+3, 0x1.0000000000000p-2, 2.6643054022145096e-2}, + {0x1.4000000000000p+3, 0x1.c800000000000p+5, 1.9454587802736456}, + {0x1.5000000000000p+3, 0x1.e848000000000p+19, 1.1512519523616630e+1}, + {0x1.2a00000000000p+5, 0x1.0000000000000p-1, 1.3512896711535966e-2}, + {0x1.f400000000000p+9, 0x1.d400000000000p+6, 1.1069890905638717e-1}, + {0x1.e848000000000p+19, 0x1.8000000000000p+1, 2.9999970000050000e-6}, + {0x1.7d78400000000p+26, 0x1.5798ee2308c3ap-27, 1.0000000050000000e-16}, + {0x1.d1a94a2000000p+39, 0x1.c800000000000p+5, 5.6999999998404000e-11}, + {0x1.b48eb57e00000p+44, 0x1.9000000000000p+8, 1.3333333333244667e-11}, + {0x1.c6bf526340000p+49, 0x1.e848000000000p+19, 9.9999999950000050e-10}, + {0x1.bc16d674ec800p+59, 0x1.0000000000000p+0, 1.0000000000000000e-18}, + {0x1.8232558201159p+59, 0x1.d400000000000p+6, 1.3453882098446929e-16}, + {0x1.0000000000000p-74, 0x1.0000000000000p+1, 1.8889465931478581e+22}, +}; +} // namespace digamma_diff_test_internal + +TEST(MathFunctions, digamma_diff_precomputed) { + using digamma_diff_test_internal::TestValue; + using digamma_diff_test_internal::testValues; + using stan::math::digamma_diff; + + for (const TestValue& t : testValues) { + EXPECT_NEAR(digamma_diff(t.x, t.d), t.val, 1e-15 * std::fabs(t.val)) + << "x = " << t.x << ", d = " << t.d; + } +} From a8251814afe8e73bff1f1af7bf5f032d4e524144 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:19 +0300 Subject: [PATCH 02/15] Add internal helpers for log ratios of rising factorials --- stan/math/prim/fun.hpp | 1 + .../prim/fun/log_rising_factorial_ratio.hpp | 276 ++++++++++++++++++ 2 files changed, 277 insertions(+) create mode 100644 stan/math/prim/fun/log_rising_factorial_ratio.hpp diff --git a/stan/math/prim/fun.hpp b/stan/math/prim/fun.hpp index 25526122101..6e767e25c45 100644 --- a/stan/math/prim/fun.hpp +++ b/stan/math/prim/fun.hpp @@ -187,6 +187,7 @@ #include #include #include +#include #include #include #include diff --git a/stan/math/prim/fun/log_rising_factorial_ratio.hpp b/stan/math/prim/fun/log_rising_factorial_ratio.hpp new file mode 100644 index 00000000000..bb2805f35d0 --- /dev/null +++ b/stan/math/prim/fun/log_rising_factorial_ratio.hpp @@ -0,0 +1,276 @@ +#ifndef STAN_MATH_PRIM_FUN_LOG_RISING_FACTORIAL_RATIO_HPP +#define STAN_MATH_PRIM_FUN_LOG_RISING_FACTORIAL_RATIO_HPP + +#include +#include +#include +#include +#include +#include +#include + +namespace stan { +namespace math { +namespace internal { + +/** + * Return log1p(t) - t for 0 <= t < 1 without cancellation. For t <= 1/4 + * use log1p(t) = 2 atanh(r), r = t / (2 + t), whose series gives + * log1p(t) - t = r (2 r^2 (1/3 + r^2/5 + ...) - t) with r^2 <= 1/81. + * + * @tparam T type of the argument + * @param t argument in [0, 1) + * @return log1p(t) - t + */ +template +inline T log1pmx(const T& t) { + if (t > 0.25) { + return log1p(t) - t; + } + const T r = t / (2.0 + t); + const T r2 = r * r; + T s = 1.0 / 21.0; + s = 1.0 / 19.0 + r2 * s; + s = 1.0 / 17.0 + r2 * s; + s = 1.0 / 15.0 + r2 * s; + s = 1.0 / 13.0 + r2 * s; + s = 1.0 / 11.0 + r2 * s; + s = 1.0 / 9.0 + r2 * s; + s = 1.0 / 7.0 + r2 * s; + s = 1.0 / 5.0 + r2 * s; + s = 1.0 / 3.0 + r2 * s; + return r * (2.0 * r2 * s - t); +} + +/** + * Return the log of the ratio of two rising factorials with the same + * number of factors, + * + * log((x)_k / (x + b)_k) = lgamma(x + k) - lgamma(x) + * - lgamma(x + b + k) + lgamma(x + b), + * + * for x > 0, b >= 0, k >= 0, without the cancellation of the four lgamma + * terms when x is large. The expression is symmetric in b and k; let + * m = min(b, k), M = max(b, k) and y = x + M. + * + * For x >= lgamma_stirling_diff_useful, write every lgamma through its + * Stirling form. The terms linear in the arguments cancel exactly and + * + * result = m log1p(-M / (y + m)) + T(x, m) - T(y, m) + * + lgamma_stirling_diff(x + m) - lgamma_stirling_diff(x) + * - lgamma_stirling_diff(y + m) + lgamma_stirling_diff(y), + * + * with T(z, m) = (z - 1/2) log1p(m / z). For m < z, T(z, m) is about m; + * it is split into m and z (log1p(m / z) - m / z) - log1p(m / z) / 2, and + * the m parts of the two terms cancel exactly. + * + * For x < lgamma_stirling_diff_useful, use + * lbeta(m, y) - lbeta(m, x), from lgamma(z + m) - lgamma(z) + * = lgamma(m) - lbeta(m, z). + * + * @tparam T type of the arguments + * @param x first argument, positive + * @param b shift of the second rising factorial, nonnegative + * @param k number of factors, nonnegative + * @return log((x)_k / (x + b)_k) + */ +template +inline T log_rising_factorial_ratio(const T& x, const T& b, const T& k) { + const T m = (b < k) ? b : k; + const T big = (b < k) ? k : b; + if (m == 0) { + return T(0.0); + } + const T y = x + big; + if (x < lgamma_stirling_diff_useful) { + return lbeta(m, y) - lbeta(m, x); + } + const T t_x = m / x; + const T t_y = m / y; + T t_diff; + if (t_x < 1.0) { + // m < x < y: both split, the two m parts cancel + t_diff + = x * log1pmx(t_x) - y * log1pmx(t_y) - 0.5 * (log1p(t_x) - log1p(t_y)); + } else if (t_y < 1.0) { + // x <= m < y: only the second term is split + t_diff = (x - 0.5) * log1p(t_x) - m - (y * log1pmx(t_y) - 0.5 * log1p(t_y)); + } else { + t_diff = (x - 0.5) * log1p(t_x) - (y - 0.5) * log1p(t_y); + } + // m log((x + m) / (y + m)); log1p form while the ratio is near 1 + const T shift_frac = big / (y + m); + const T log_ratio + = (shift_frac < 0.5) ? log1p(-shift_frac) : log((x + m) / (y + m)); + return m * log_ratio + t_diff + lgamma_stirling_diff(x + m) + - lgamma_stirling_diff(x) - lgamma_stirling_diff(y + m) + + lgamma_stirling_diff(y); +} + +/** + * Return (x - 1/2) log1p(k / x) - c, where c is k when k < x and 0 + * otherwise, and add c to `k_removed`. For k < x the term is about k and is + * formed as x (log1p(k / x) - k / x) - log1p(k / x) / 2, which is of the + * size k^2 / x. The counts k are integers, so their sum is exact. + * + * @tparam T type of the arguments + * @param x shape, at least lgamma_stirling_diff_useful + * @param k count, a nonnegative integer + * @param[in, out] k_removed sum of the removed counts + * @return the term without the removed count + */ +template +inline T stirling_log1p_term(const T& x, const T& k, T& k_removed) { + if (k == 0) { + return T(0.0); + } + const T t = k / x; + if (t < 1.0) { + k_removed += k; + return x * log1pmx(t) - 0.5 * log1p(t); + } + return (x - 0.5) * log1p(t); +} + +/** + * Return an estimate of the size of the terms that `lbeta(a, b)` adds up, + * for a, b >= lgamma_stirling_diff_useful: min(a, b) (1 + log1p(max / min)). + * Only used to choose between two forms; the accuracy is not important. + * + * @param a first argument + * @param b second argument + * @return the size estimate + */ +inline double lbeta_terms_size(double a, double b) { + const double small = std::fmin(a, b); + const double large = std::fmax(a, b); + return small * (1.0 + std::log1p(large / small)); +} + +/** + * Return the part of log_beta_ratio() that depends on the shapes only, so + * that a caller can compute it once per pair of shapes: the three + * lgamma_stirling_diff terms when both shapes are at least + * lgamma_stirling_diff_useful, and lbeta(alpha, beta) otherwise. + * + * @tparam T type of the arguments + * @param alpha first shape, positive + * @param beta second shape, positive + * @return the shape-only part + */ +template +inline T log_beta_ratio_denominator(const T& alpha, const T& beta) { + if (alpha >= lgamma_stirling_diff_useful + && beta >= lgamma_stirling_diff_useful) { + return lgamma_stirling_diff(alpha) + lgamma_stirling_diff(beta) + - lgamma_stirling_diff(alpha + beta); + } + return lbeta(alpha, beta); +} + +/** + * Return lbeta(alpha + n, beta + m) - lbeta(alpha, beta), the log of the + * ratio of rising factorials (alpha)_n (beta)_m / (alpha + beta)_(n + m), + * for shapes alpha, beta > 0 and integer counts n, m >= 0. + * + * When both shapes are large, each lbeta is of the order of the shapes + * while the difference is of the order of n + m, so the plain difference + * keeps no correct digits from shapes near 1e15. For + * shapes of at least lgamma_stirling_diff_useful, write every lgamma as + * (x - 1/2) log(x) - x + log(2 pi) / 2 + lgamma_stirling_diff(x). With + * N = n + m and s = alpha + beta, the terms linear in the arguments cancel + * exactly and the rest is + * + * T(alpha, n) + T(beta, m) - T(s, N) + * + n log((alpha + n) / (s + N)) + m log((beta + m) / (s + N)) + * + the six lgamma_stirling_diff terms, + * + * with T(x, k) = (x - 1/2) log1p(k / x), split by stirling_log1p_term(). + * This form has rounding errors of the size of N; the lbeta form has + * rounding errors of the size of the smaller shape. The function takes the + * form with the smaller terms. + * + * @tparam T type of the arguments + * @param alpha first shape, positive + * @param beta second shape, positive + * @param n first count, a nonnegative integer + * @param m second count, a nonnegative integer + * @param denominator log_beta_ratio_denominator(alpha, beta) + * @return lbeta(alpha + n, beta + m) - lbeta(alpha, beta) + */ +template +inline T log_beta_ratio(const T& alpha, const T& beta, const T& n, const T& m, + const T& denominator) { + const bool large_shapes = alpha >= lgamma_stirling_diff_useful + && beta >= lgamma_stirling_diff_useful; + if (large_shapes) { + const T total_count = n + m; + const T alpha_plus_beta = alpha + beta; + const T total = alpha_plus_beta + total_count; + T removed_plus(0.0); + T removed_minus(0.0); + const T term_alpha = stirling_log1p_term(alpha, n, removed_plus); + const T term_beta = stirling_log1p_term(beta, m, removed_plus); + const T term_total + = stirling_log1p_term(alpha_plus_beta, total_count, removed_minus); + // n log(p) + m log(q) with p + q = 1: take the log1m form for the larger + // of p and q, so that a ratio near 1 keeps its digits + const T p = (alpha + n) / total; + const T q = (beta + m) / total; + const T log_n = (p < q) ? n * log(p) : n * log1p(-q); + const T log_m = (p < q) ? m * log1p(-p) : m * log(q); + const double size_stirling = std::fabs(value_of_rec(term_alpha)) + + std::fabs(value_of_rec(term_beta)) + + std::fabs(value_of_rec(term_total)) + + std::fabs(value_of_rec(log_n)) + + std::fabs(value_of_rec(log_m)); + const double size_lbeta + = lbeta_terms_size(value_of_rec(alpha), value_of_rec(beta)) + + lbeta_terms_size(value_of_rec(alpha + n), value_of_rec(beta + m)); + if (size_stirling <= size_lbeta) { + return term_alpha + term_beta - term_total + + (removed_plus - removed_minus) + log_n + log_m + + lgamma_stirling_diff(alpha + n) + lgamma_stirling_diff(beta + m) + - lgamma_stirling_diff(total) - denominator; + } + // here denominator holds the Stirling remainders, not lbeta + return lbeta(alpha + n, beta + m) - lbeta(alpha, beta); + } + return lbeta(alpha + n, beta + m) - denominator; +} + +/** + * Return the log pmf of the beta negative binomial distribution at the + * integer k >= 0, + * + * lgamma(r + k) - lgamma(k + 1) - lgamma(r) + lbeta(alpha + r, beta + k) + * - lbeta(alpha, beta), + * + * as -log(k) - lbeta(k, beta) + log_rising_factorial_ratio(r, alpha + beta, + * k) + log_rising_factorial_ratio(alpha, beta, r), which does not cancel + * when the parameters are large. For k = 0 only the + * last term remains. + * + * @tparam T type of the arguments + * @param k outcome, a nonnegative integer + * @param r number of successes, positive + * @param alpha first shape, positive + * @param beta second shape, positive + * @return the log pmf + */ +template +inline T beta_neg_binomial_log_pmf(const T& k, const T& r, const T& alpha, + const T& beta) { + T lp = log_rising_factorial_ratio(alpha, beta, r); + if (k > 0) { + lp += log_rising_factorial_ratio(r, T(alpha + beta), k) - log(k) + - lbeta(k, beta); + } + return lp; +} + +} // namespace internal +} // namespace math +} // namespace stan + +#endif From 82900c8554171f2e4bd11f55263a511ea4ceb72d Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:31 +0300 Subject: [PATCH 03/15] Avoid the lbeta cancellation in beta_binomial_lpmf at large shapes --- stan/math/prim/prob/beta_binomial_lpmf.hpp | 102 +++++++----------- .../math/prim/prob/beta_binomial_test.cpp | 72 +++++++++++++ .../math/rev/prob/beta_binomial_lpmf_test.cpp | 93 ++++++++++++++++ 3 files changed, 204 insertions(+), 63 deletions(-) create mode 100644 test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp diff --git a/stan/math/prim/prob/beta_binomial_lpmf.hpp b/stan/math/prim/prob/beta_binomial_lpmf.hpp index 12065ce56aa..650a1a5d116 100644 --- a/stan/math/prim/prob/beta_binomial_lpmf.hpp +++ b/stan/math/prim/prob/beta_binomial_lpmf.hpp @@ -5,16 +5,20 @@ #include #include #include -#include -#include +#include #include -#include +#include +#include +#include +#include +#include #include #include #include #include #include #include +#include namespace stan { namespace math { @@ -77,8 +81,6 @@ inline return_type_t beta_binomial_lpmf(const T_n& n, scalar_seq_view N_vec(N_ref); scalar_seq_view alpha_vec(alpha_ref); scalar_seq_view beta_vec(beta_ref); - size_t size_alpha = stan::math::size(alpha); - size_t size_beta = stan::math::size(beta); size_t size_n_N = max_size(n, N); size_t size_alpha_beta = max_size(alpha, beta); size_t max_size_seq_view = max_size(n, N, alpha, beta); @@ -95,73 +97,47 @@ inline return_type_t beta_binomial_lpmf(const T_n& n, if constexpr (include_summand::value) normalizing_constant[i] = binomial_coefficient_log(N_vec[i], n_vec[i]); + // The value is lbeta(n + alpha, N - n + beta) - lbeta(alpha, beta). For + // large shapes the two lbeta values are of the order of the shapes and + // their plain difference keeps no correct digits from shapes near 1e15; + // internal::log_beta_ratio forms the difference + // without that cancellation. Its shape-only part is computed once per + // pair of shapes. VectorBuilder lbeta_denominator( size_alpha_beta); for (size_t i = 0; i < size_alpha_beta; i++) { - lbeta_denominator[i] = lbeta(alpha_vec.val(i), beta_vec.val(i)); + lbeta_denominator[i] = internal::log_beta_ratio_denominator( + T_partials_return(alpha_vec.val(i)), + T_partials_return(beta_vec.val(i))); } - VectorBuilder lbeta_diff( - max_size_seq_view); for (size_t i = 0; i < max_size_seq_view; i++) { - lbeta_diff[i] = lbeta(n_vec[i] + alpha_vec.val(i), - N_vec[i] - n_vec[i] + beta_vec.val(i)) - - lbeta_denominator[i]; - } - - VectorBuilder, T_partials_return, T_n, T_size1> - digamma_n_plus_alpha(max_size(n, alpha)); - if constexpr (is_autodiff_v) { - for (size_t i = 0; i < max_size(n, alpha); i++) { - digamma_n_plus_alpha[i] = digamma(n_vec.val(i) + alpha_vec.val(i)); - } - } - - VectorBuilder, T_partials_return, T_size1, - T_size2> - digamma_alpha_plus_beta(size_alpha_beta); - if constexpr (is_any_autodiff_v) { - for (size_t i = 0; i < size_alpha_beta; i++) { - digamma_alpha_plus_beta[i] = digamma(alpha_vec.val(i) + beta_vec.val(i)); + const T_partials_return alpha_dbl = alpha_vec.val(i); + const T_partials_return beta_dbl = beta_vec.val(i); + const T_partials_return n_i = n_vec.val(i); + const T_partials_return N_i = N_vec.val(i); + const T_partials_return m_i = N_i - n_i; + if constexpr (include_summand::value) { + logp += normalizing_constant[i]; } - } - - VectorBuilder, T_partials_return, T_N, - T_size1, T_size2> - digamma_diff(max_size(N, alpha, beta)); - if constexpr (is_any_autodiff_v) { - for (size_t i = 0; i < max_size(N, alpha, beta); i++) { - digamma_diff[i] - = digamma_alpha_plus_beta[i] - - digamma(N_vec.val(i) + alpha_vec.val(i) + beta_vec.val(i)); + logp += internal::log_beta_ratio(alpha_dbl, beta_dbl, n_i, m_i, + lbeta_denominator[i]); + + // Each partial is a sum of two digamma differences psi(x + k) - psi(x) + // of size k / x; digamma_diff forms them without cancellation. + if constexpr (is_any_autodiff_v) { + const T_partials_return digamma_diff_total + = digamma_diff(alpha_dbl + beta_dbl, N_i); + if constexpr (is_autodiff_v) { + partials<0>(ops_partials)[i] + += digamma_diff(alpha_dbl, n_i) - digamma_diff_total; + } + if constexpr (is_autodiff_v) { + partials<1>(ops_partials)[i] + += digamma_diff(beta_dbl, m_i) - digamma_diff_total; + } } } - - VectorBuilder, T_partials_return, T_size1> - digamma_alpha(size_alpha); - for (size_t i = 0; i < size_alpha; i++) - if constexpr (is_autodiff_v) - digamma_alpha[i] = digamma(alpha_vec.val(i)); - - VectorBuilder, T_partials_return, T_size2> - digamma_beta(size_beta); - for (size_t i = 0; i < size_beta; i++) - if constexpr (is_autodiff_v) - digamma_beta[i] = digamma(beta_vec.val(i)); - - for (size_t i = 0; i < max_size_seq_view; i++) { - if constexpr (include_summand::value) - logp += normalizing_constant[i]; - logp += lbeta_diff[i]; - - if constexpr (is_autodiff_v) - partials<0>(ops_partials)[i] - += digamma_n_plus_alpha[i] + digamma_diff[i] - digamma_alpha[i]; - if constexpr (is_autodiff_v) - partials<1>(ops_partials)[i] - += digamma(N_vec.val(i) - n_vec.val(i) + beta_vec.val(i)) - + digamma_diff[i] - digamma_beta[i]; - } return ops_partials.build(logp); } diff --git a/test/unit/math/prim/prob/beta_binomial_test.cpp b/test/unit/math/prim/prob/beta_binomial_test.cpp index 9adc12fe794..fc52cf1e1ce 100644 --- a/test/unit/math/prim/prob/beta_binomial_test.cpp +++ b/test/unit/math/prim/prob/beta_binomial_test.cpp @@ -4,6 +4,8 @@ #include #include #include +#include +#include #include #include @@ -56,3 +58,73 @@ TEST(ProbDistributionBetaBinomial, error_check) { 4, 0.6, stan::math::positive_infinity(), rng), std::domain_error); } + +namespace beta_binomial_test_internal { +struct TestValue { + int n; + int N; + double alpha; + double beta; + double value; +}; + +// Log pmf computed with mpmath at 80 digits as lchoose(N, n) +// + lbeta(n + alpha, N - n + beta) - lbeta(alpha, beta) with mp.loggamma, +// and checked at 130 digits. The shapes are written in hex so that they are +// exact. The first three are n = 57, N = 117, alpha = beta = exp(lc) / 2 +// for lc = 32, 36, 42, where the plain lbeta difference was wrong by +// 4.5e-3, 9.8e-2 and 81. The next three are n = 500, N = 1000, +// alpha = beta = 0.5e10, 0.5e13, 0.5e19, where the plain lbeta difference +// was wrong by 5.3e-8, 2.8e-4 and 693. In the last two points one or both +// shapes are small, where the lbeta form is the accurate one. +std::vector testValues = { + {57, 117, 0x1.1f43fcc4b662cp+45, 0x1.1f43fcc4b662cp+45, + -2.6471538352642870}, + {57, 117, 0x1.ea215a1d20d76p+50, 0x1.ea215a1d20d76p+50, + -2.6471538352636157}, + {57, 117, 0x1.8232558201159p+59, 0x1.8232558201159p+59, + -2.6471538352636032}, + {500, 1000, 0x1.2a05f20000000p+32, 0x1.2a05f20000000p+32, + -3.6799190420941268}, + {500, 1000, 0x1.2309ce5400000p+42, 0x1.2309ce5400000p+42, + -3.6799189921441293}, + {500, 1000, 0x1.158e460913d00p+62, 0x1.158e460913d00p+62, + -3.6799189920941293}, + {400, 1000, 0x1.b48eb57e00000p+44, 0x1.977420dc00000p+42, + -4.1354072540579752e+2}, + {0, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -3.3658802476858363e+1}, + {117, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -1.6219644025102715e+2}, + {3, 1000000, 0x1.e848000000000p+19, 0x1.bc16d674ec800p+59, + -4.3238292143128877e+1}, + {1000000, 1000000, 0x1.cfde000000000p+19, 0x1.e000000000000p+3, + -1.0786783324715004e+1}, + {5, 20, 0x1.4000000000000p+3, 0x1.9000000000000p+4, -1.8540068216786033}, +}; +} // namespace beta_binomial_test_internal + +TEST(ProbDistributionsBetaBinomial, large_shapes) { + using beta_binomial_test_internal::TestValue; + using beta_binomial_test_internal::testValues; + using stan::math::beta_binomial_lpmf; + + for (const TestValue& t : testValues) { + const double tol = 1e-13 * std::max(1.0, std::fabs(t.value)); + EXPECT_NEAR(beta_binomial_lpmf(t.n, t.N, t.alpha, t.beta), t.value, tol) + << "n = " << t.n << ", N = " << t.N << ", alpha = " << t.alpha + << ", beta = " << t.beta; + } +} + +TEST(ProbDistributionsBetaBinomial, binomial_limit) { + // The difference to the binomial limit tends to 0 + // from below as the shapes grow; at lc = 42 it is -3.1e-17 + using stan::math::beta_binomial_lpmf; + using stan::math::binomial_lpmf; + const double s = 0x1.8232558201159p+59; // exp(42) / 2 + const double diff + = beta_binomial_lpmf(57, 117, s, s) - binomial_lpmf(57, 117, 0.5); + EXPECT_LE(diff, 1e-13); + EXPECT_GE(diff, -1e-13); +} diff --git a/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp b/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp new file mode 100644 index 00000000000..1ff5e18e742 --- /dev/null +++ b/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp @@ -0,0 +1,93 @@ +#include +#include +#include +#include +#include +#include + +namespace beta_binomial_lpmf_rev_test_internal { +struct TestValue { + int n; + int N; + double alpha; + double beta; + double value; + double grad_log_alpha; // alpha * d/dalpha + double grad_log_beta; // beta * d/dbeta +}; + +// Computed with mpmath at 80 digits: the value from mp.loggamma, the +// partials from mp.digamma, both checked against mp.diff at 130 digits. +// The shapes are in hex so that they are exact. The gradients are given in +// the log-shape parameterization, which is the one a sampler sees when a +// model puts a prior on log(alpha) or log(concentration). +std::vector testValues = { + // n = 57, N = 117, alpha = beta = exp(32) / 2 + {57, 117, 0x1.1f43fcc4b662cp+45, 0x1.1f43fcc4b662cp+45, -2.6471538352642870, + -1.4999999999974545, 1.4999999999981384}, + {400, 1000, 0x1.b48eb57e00000p+44, 0x1.977420dc00000p+42, + -4.1354072540579752e+2, -4.1081081080252486e+2, 4.1081081078769344e+2}, + {0, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -3.3658802476858363e+1, -2.9249999999999996e+1, 2.9249999999999990e+1}, + {117, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -1.6219644025102715e+2, 8.7749999999999936e+1, -8.7749999999999987e+1}, + // one shape large, the other moderate or small + {3, 1000000, 0x1.e848000000000p+19, 0x1.bc16d674ec800p+59, + -4.3238292143128877e+1, 2.9999960000050000, -2.9999989999970000}, + {0, 1000000, 0x1.03caccd133500p+59, 0x1.35c28f5c28f5cp+2, + -2.8094819745857234e+7, -9.9999999999914529e+5, 5.9751970421184601e+1}, + {1000000, 1000000, 0x1.cfde000000000p+19, 0x1.e000000000000p+3, + -1.0786783324715004e+1, 7.6922233997281191, -1.0786722596873199e+1}, + {5, 20, 0x1.4000000000000p+3, 0x1.9000000000000p+4, -1.8540068216786033, + -3.4626330704764243e-1, 5.0911144710080701e-1}, +}; +} // namespace beta_binomial_lpmf_rev_test_internal + +TEST(ProbDistributionsBetaBinomial, log_shape_gradients) { + using beta_binomial_lpmf_rev_test_internal::TestValue; + using beta_binomial_lpmf_rev_test_internal::testValues; + using stan::math::var; + + for (const TestValue& t : testValues) { + var alpha = t.alpha; + var beta = t.beta; + var lp = stan::math::beta_binomial_lpmf(t.n, t.N, alpha, beta); + lp.grad(); + EXPECT_NEAR(lp.val(), t.value, 1e-13 * std::max(1.0, std::fabs(t.value))) + << "n = " << t.n << ", N = " << t.N << ", alpha = " << t.alpha + << ", beta = " << t.beta; + // Each partial is a difference of two digamma differences. Its rounding + // error is a few eps of those terms, which in the log-shape + // parameterization are of size up to N. + const double scale = 1e-12 * std::max(1.0, 1e-3 * t.N); + EXPECT_NEAR(t.alpha * alpha.adj(), t.grad_log_alpha, + scale * std::max(1.0, std::fabs(t.grad_log_alpha))) + << "n = " << t.n << ", N = " << t.N << ", alpha = " << t.alpha + << ", beta = " << t.beta; + EXPECT_NEAR(t.beta * beta.adj(), t.grad_log_beta, + scale * std::max(1.0, std::fabs(t.grad_log_beta))) + << "n = " << t.n << ", N = " << t.N << ", alpha = " << t.alpha + << ", beta = " << t.beta; + stan::math::recover_memory(); + } +} + +TEST(ProbDistributionsBetaBinomial, log_concentration_gradient) { + // alpha = beta = exp(lc) / 2, n = 57, N = 117. The + // gradient in lc goes to 0 like 27 / exp(lc). develop returned 0 or + // noise (-0.14 at lc = 32) from lc = 20 on, which made the log density a + // plateau that warmup could not leave. + using stan::math::var; + const std::vector lcs = {16.0, 18.0, 20.0, 26.0, 30.0, 32.0}; + const std::vector grads = { + 6.0768267180691362e-6, 8.2241757434662373e-7, 1.1130227121763732e-7, + 2.7589080736553729e-10, 5.0531164031234141e-12, 6.8386493965016465e-13}; + for (size_t i = 0; i < lcs.size(); ++i) { + var lc = lcs[i]; + var s = stan::math::exp(lc) / 2; + var lp = stan::math::beta_binomial_lpmf(57, 117, s, s); + lp.grad(); + EXPECT_NEAR(lc.adj(), grads[i], 1e-12) << "lc = " << lcs[i]; + stan::math::recover_memory(); + } +} From 58a339b4909836c07a56a56dd2e774a36fec4c76 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:38 +0300 Subject: [PATCH 04/15] Avoid the cancellation in the beta_binomial cdfs at large shapes --- stan/math/prim/prob/beta_binomial_cdf.hpp | 20 +++++--- stan/math/prim/prob/beta_binomial_lccdf.hpp | 21 +++++--- stan/math/prim/prob/beta_binomial_lcdf.hpp | 54 ++++++++++++++++++--- 3 files changed, 76 insertions(+), 19 deletions(-) diff --git a/stan/math/prim/prob/beta_binomial_cdf.hpp b/stan/math/prim/prob/beta_binomial_cdf.hpp index 7d90678d154..0bab43798d7 100644 --- a/stan/math/prim/prob/beta_binomial_cdf.hpp +++ b/stan/math/prim/prob/beta_binomial_cdf.hpp @@ -5,7 +5,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -103,18 +104,24 @@ inline return_type_t beta_binomial_cdf(const T_n& n, const T_partials_return F = hypergeometric_3F2({one, mu, 1 - N_minus_n}, {n_dbl + 2, 1 - nu}, one); - T_partials_return C = lbeta(nu, mu) - lbeta(alpha_dbl, beta_dbl) - - lbeta(N_minus_n, n_dbl + 2); + // lbeta(nu, mu) - lbeta(alpha, beta) without the cancellation for large + // shapes + T_partials_return C + = internal::log_beta_ratio( + alpha_dbl, beta_dbl, n_dbl + 1, N_minus_n - 1, + internal::log_beta_ratio_denominator(alpha_dbl, beta_dbl)) + - lbeta(N_minus_n, n_dbl + 2); C = F * exp(C) / (N_dbl + 1); const T_partials_return Pi = 1 - C; P *= Pi; + // digamma(alpha + beta) - digamma(mu + nu), mu + nu = alpha + beta + N T_partials_return digammaDiff = is_constant_all::value ? 0 - : digamma(alpha_dbl + beta_dbl) - digamma(mu + nu); + : -digamma_diff(alpha_dbl + beta_dbl, N_dbl); T_partials_return dF[6]; if constexpr (is_any_autodiff_v) { @@ -122,12 +129,13 @@ inline return_type_t beta_binomial_cdf(const T_n& n, } if constexpr (is_autodiff_v) { const T_partials_return g - = -C * (digamma(mu) - digamma(alpha_dbl) + digammaDiff + dF[1] / F); + = -C * (digamma_diff(alpha_dbl, n_dbl + 1) + digammaDiff + dF[1] / F); partials<0>(ops_partials)[i] += g / Pi; } if constexpr (is_autodiff_v) { const T_partials_return g - = -C * (digamma(nu) - digamma(beta_dbl) + digammaDiff - dF[4] / F); + = -C + * (digamma_diff(beta_dbl, N_minus_n - 1) + digammaDiff - dF[4] / F); partials<1>(ops_partials)[i] += g / Pi; } } diff --git a/stan/math/prim/prob/beta_binomial_lccdf.hpp b/stan/math/prim/prob/beta_binomial_lccdf.hpp index d611d23b700..8f17cb47054 100644 --- a/stan/math/prim/prob/beta_binomial_lccdf.hpp +++ b/stan/math/prim/prob/beta_binomial_lccdf.hpp @@ -5,7 +5,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -101,18 +102,24 @@ inline return_type_t beta_binomial_lccdf( const T_partials_return F = hypergeometric_3F2( {one, mu, -N_dbl + n_dbl + 1}, {n_dbl + 2, 1 - nu}, one); - T_partials_return C = lbeta(nu, mu) - lbeta(alpha_dbl, beta_dbl) - - lbeta(N_dbl - n_dbl, n_dbl + 2); + // lbeta(nu, mu) - lbeta(alpha, beta) without the cancellation for large + // shapes + T_partials_return C + = internal::log_beta_ratio( + alpha_dbl, beta_dbl, n_dbl + 1, N_dbl - n_dbl - 1, + internal::log_beta_ratio_denominator(alpha_dbl, beta_dbl)) + - lbeta(N_dbl - n_dbl, n_dbl + 2); C = F * exp(C) / (N_dbl + 1); const T_partials_return Pi = C; P += log(Pi); + // digamma(alpha + beta) - digamma(mu + nu), mu + nu = alpha + beta + N T_partials_return digammaDiff = is_constant_all::value ? 0 - : digamma(alpha_dbl + beta_dbl) - digamma(mu + nu); + : -digamma_diff(alpha_dbl + beta_dbl, N_dbl); T_partials_return dF[6]; if constexpr (is_any_autodiff_v) { @@ -120,11 +127,11 @@ inline return_type_t beta_binomial_lccdf( } if constexpr (is_autodiff_v) { partials<0>(ops_partials)[i] - += digamma(mu) - digamma(alpha_dbl) + digammaDiff + dF[1] / F; + += digamma_diff(alpha_dbl, n_dbl + 1) + digammaDiff + dF[1] / F; } if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] - += digamma(nu) - digamma(beta_dbl) + digammaDiff - dF[4] / F; + partials<1>(ops_partials)[i] += digamma_diff(beta_dbl, N_dbl - n_dbl - 1) + + digammaDiff - dF[4] / F; } } diff --git a/stan/math/prim/prob/beta_binomial_lcdf.hpp b/stan/math/prim/prob/beta_binomial_lcdf.hpp index eea4ceb5e1e..d95431a2960 100644 --- a/stan/math/prim/prob/beta_binomial_lcdf.hpp +++ b/stan/math/prim/prob/beta_binomial_lcdf.hpp @@ -5,7 +5,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -105,18 +106,58 @@ inline return_type_t beta_binomial_lcdf(const T_n& n, const T_partials_return F = hypergeometric_3F2({one, mu, 1 - N_minus_n}, {n_dbl + 2, 1 - nu}, one); - T_partials_return C = lbeta(nu, mu) - lbeta(alpha_dbl, beta_dbl) - - lbeta(N_minus_n, n_dbl + 2); + // lbeta(nu, mu) - lbeta(alpha, beta) without the cancellation for large + // shapes + T_partials_return C + = internal::log_beta_ratio( + alpha_dbl, beta_dbl, n_dbl + 1, N_minus_n - 1, + internal::log_beta_ratio_denominator(alpha_dbl, beta_dbl)) + - lbeta(N_minus_n, n_dbl + 2); C = F * exp(C) / (N_dbl + 1); + if (C > 0.5 && alpha_dbl != 1.0) { + // The cdf is below 1/2, and log(1 - C) loses its digits; C can even + // round to 1 or above. Use the mirror: X <= n if and only if + // N - X > N - n - 1, and N - X is beta_binomial(N, beta, alpha). Its + // ccdf is the same series with the shapes swapped and needs no + // complement. For alpha = 1 the mirrored series has the parameters + // -n and 1 - alpha - n = -n, whose terms are 0 / 0 at k = n + 1, so + // that case keeps log(1 - C). + const T_partials_return n_m = N_minus_n - 1; + const T_partials_return mu_m = beta_dbl + N_minus_n; + const T_partials_return nu_m = alpha_dbl + n_dbl; + const T_partials_return F_m + = hypergeometric_3F2({one, mu_m, -n_dbl}, {n_m + 2, 1 - nu_m}, one); + P += internal::log_beta_ratio( + beta_dbl, alpha_dbl, N_minus_n, n_dbl, + internal::log_beta_ratio_denominator(beta_dbl, alpha_dbl)) + - lbeta(n_dbl + 1, N_minus_n + 1) + log(F_m) - log(N_dbl + 1); + if constexpr (is_any_autodiff_v) { + T_partials_return dF_m[6]; + grad_F32(dF_m, one, mu_m, -n_dbl, n_m + 2, 1 - nu_m, one); + const T_partials_return dpsi_total + = digamma_diff(alpha_dbl + beta_dbl, N_dbl); + if constexpr (is_autodiff_v) { + partials<0>(ops_partials)[i] + += digamma_diff(alpha_dbl, n_dbl) - dpsi_total - dF_m[4] / F_m; + } + if constexpr (is_autodiff_v) { + partials<1>(ops_partials)[i] + += digamma_diff(beta_dbl, N_minus_n) - dpsi_total + dF_m[1] / F_m; + } + } + continue; + } + const T_partials_return Pi = 1 - C; P += log(Pi); + // digamma(alpha + beta) - digamma(mu + nu), mu + nu = alpha + beta + N T_partials_return digammaDiff = is_constant_all::value ? 0 - : digamma(alpha_dbl + beta_dbl) - digamma(mu + nu); + : -digamma_diff(alpha_dbl + beta_dbl, N_dbl); T_partials_return dF[6]; if constexpr (is_any_autodiff_v) { @@ -124,12 +165,13 @@ inline return_type_t beta_binomial_lcdf(const T_n& n, } if constexpr (is_autodiff_v) { const T_partials_return g - = -C * (digamma(mu) - digamma(alpha_dbl) + digammaDiff + dF[1] / F); + = -C * (digamma_diff(alpha_dbl, n_dbl + 1) + digammaDiff + dF[1] / F); partials<0>(ops_partials)[i] += g / Pi; } if constexpr (is_autodiff_v) { const T_partials_return g - = -C * (digamma(nu) - digamma(beta_dbl) + digammaDiff - dF[4] / F); + = -C + * (digamma_diff(beta_dbl, N_minus_n - 1) + digammaDiff - dF[4] / F); partials<1>(ops_partials)[i] += g / Pi; } } From 25f9d3dbaa638b9796984ab933b256a8d35bb79e Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:45 +0300 Subject: [PATCH 05/15] Avoid the cancellation in beta_neg_binomial at large shapes --- stan/math/prim/prob/beta_neg_binomial_cdf.hpp | 51 ++++++------ .../prim/prob/beta_neg_binomial_lccdf.hpp | 51 ++++++------ .../math/prim/prob/beta_neg_binomial_lcdf.hpp | 50 ++++++------ .../math/prim/prob/beta_neg_binomial_lpmf.hpp | 77 +++++++++++-------- 4 files changed, 118 insertions(+), 111 deletions(-) diff --git a/stan/math/prim/prob/beta_neg_binomial_cdf.hpp b/stan/math/prim/prob/beta_neg_binomial_cdf.hpp index 00aa643ad8b..e39c7385e8c 100644 --- a/stan/math/prim/prob/beta_neg_binomial_cdf.hpp +++ b/stan/math/prim/prob/beta_neg_binomial_cdf.hpp @@ -4,7 +4,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -101,46 +102,44 @@ inline return_type_t beta_neg_binomial_cdf( std::initializer_list{1.0, b_plus_n + 1.0, r_plus_n + 1.0}, std::initializer_list{n_dbl + 2.0, a_plus_r + b_plus_n + 1.0}, 1.0); - auto C = lgamma(r_plus_n + 1.0) + lbeta(a_plus_r, b_plus_n + 1.0) - - lgamma(r_dbl) - lbeta(alpha_dbl, beta_dbl) - lgamma(n_dbl + 2.0); + // C is the log pmf at n + 1; formed from lgamma and lbeta directly it + // cancels for large parameters + const T_partials_return k = n_dbl + 1.0; + const T_partials_return C = internal::beta_neg_binomial_log_pmf( + k, T_partials_return(r_dbl), T_partials_return(alpha_dbl), + T_partials_return(beta_dbl)); auto ccdf = stan::math::exp(C + stan::math::log(F)); cdf *= 1.0 - ccdf; if constexpr (is_any_autodiff_v) { auto chain_rule_term = -ccdf / (1.0 - ccdf); - auto digamma_n_r_alpha_beta = digamma(a_plus_r + b_plus_n + 1.0); T_partials_return dF[6]; grad_F32, is_autodiff_v, false, true, false>(dF, 1.0, b_plus_n + 1.0, r_plus_n + 1.0, n_dbl + 2.0, a_plus_r + b_plus_n + 1.0, 1.0, precision, max_steps); - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_r_alpha = digamma(a_plus_r); - if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] - += (digamma(r_plus_n + 1) - + (digamma_r_alpha - digamma_n_r_alpha_beta) - + (dF[2] + dF[4]) / F - digamma(r_dbl)) - * chain_rule_term; - } + // the partials of C as differences digamma(x + d) - digamma(x) + const T_partials_return alpha_plus_beta = alpha_dbl + beta_dbl; + const T_partials_return dpsi_r_plus_ab_k + = digamma_diff(r_dbl + alpha_plus_beta, k); + if constexpr (is_autodiff_v) { + partials<0>(ops_partials)[i] + += (digamma_diff(r_dbl, k) - dpsi_r_plus_ab_k + - digamma_diff(a_plus_r, beta_dbl) + (dF[2] + dF[4]) / F) + * chain_rule_term; + } + if constexpr (is_any_autodiff_v) { + const T_partials_return dpsi_ab_r + = digamma_diff(alpha_plus_beta, r_dbl); if constexpr (is_autodiff_v) { partials<1>(ops_partials)[i] - += (digamma_r_alpha - digamma_n_r_alpha_beta + dF[4] / F - - digamma(alpha_dbl)) + += (digamma_diff(alpha_dbl, r_dbl) - dpsi_ab_r - dpsi_r_plus_ab_k + + dF[4] / F) * chain_rule_term; } - } - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_alpha_beta = digamma(alpha_dbl + beta_dbl); - if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] += digamma_alpha_beta * chain_rule_term; - } if constexpr (is_autodiff_v) { partials<2>(ops_partials)[i] - += (digamma(b_plus_n + 1) - digamma_n_r_alpha_beta - + (dF[1] + dF[4]) / F - - (digamma(beta_dbl) - digamma_alpha_beta)) + += (digamma_diff(beta_dbl, k) - dpsi_r_plus_ab_k - dpsi_ab_r + + (dF[1] + dF[4]) / F) * chain_rule_term; } } diff --git a/stan/math/prim/prob/beta_neg_binomial_lccdf.hpp b/stan/math/prim/prob/beta_neg_binomial_lccdf.hpp index 49dc472e4e5..259fddff2b8 100644 --- a/stan/math/prim/prob/beta_neg_binomial_lccdf.hpp +++ b/stan/math/prim/prob/beta_neg_binomial_lccdf.hpp @@ -4,7 +4,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -99,42 +100,40 @@ inline return_type_t beta_neg_binomial_lccdf( std::initializer_list{1.0, b_plus_n + 1.0, r_plus_n + 1.0}, std::initializer_list{n_dbl + 2.0, a_plus_r + b_plus_n + 1.0}, 1.0); - auto C = lgamma(r_plus_n + 1.0) + lbeta(a_plus_r, b_plus_n + 1.0) - - lgamma(r_dbl) - lbeta(alpha_dbl, beta_dbl) - lgamma(n_dbl + 2.0); + // C is the log pmf at n + 1; formed from lgamma and lbeta directly it + // cancels for large parameters + const T_partials_return k = n_dbl + 1.0; + const T_partials_return C = internal::beta_neg_binomial_log_pmf( + k, T_partials_return(r_dbl), T_partials_return(alpha_dbl), + T_partials_return(beta_dbl)); log_ccdf += C + stan::math::log(F); if constexpr (is_any_autodiff_v) { - auto digamma_n_r_alpha_beta = digamma(a_plus_r + b_plus_n + 1.0); T_partials_return dF[6]; grad_F32, is_autodiff_v, false, true, false>(dF, 1.0, b_plus_n + 1.0, r_plus_n + 1.0, n_dbl + 2.0, a_plus_r + b_plus_n + 1.0, 1.0, precision, max_steps); - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_r_alpha = digamma(a_plus_r); - if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] - += digamma(r_plus_n + 1) - + (digamma_r_alpha - digamma_n_r_alpha_beta) - + (dF[2] + dF[4]) / F - digamma(r_dbl); - } - if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] += digamma_r_alpha - - digamma_n_r_alpha_beta + dF[4] / F - - digamma(alpha_dbl); - } + // the partials of C as differences digamma(x + d) - digamma(x) + const T_partials_return alpha_plus_beta = alpha_dbl + beta_dbl; + const T_partials_return dpsi_r_plus_ab_k + = digamma_diff(r_dbl + alpha_plus_beta, k); + if constexpr (is_autodiff_v) { + partials<0>(ops_partials)[i] + += digamma_diff(r_dbl, k) - dpsi_r_plus_ab_k + - digamma_diff(a_plus_r, beta_dbl) + (dF[2] + dF[4]) / F; } - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_alpha_beta = digamma(alpha_dbl + beta_dbl); + if constexpr (is_any_autodiff_v) { + const T_partials_return dpsi_ab_r + = digamma_diff(alpha_plus_beta, r_dbl); if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] += digamma_alpha_beta; + partials<1>(ops_partials)[i] += digamma_diff(alpha_dbl, r_dbl) + - dpsi_ab_r - dpsi_r_plus_ab_k + + dF[4] / F; } if constexpr (is_autodiff_v) { - partials<2>(ops_partials)[i] - += digamma(b_plus_n + 1) - digamma_n_r_alpha_beta - + (dF[1] + dF[4]) / F - - (digamma(beta_dbl) - digamma_alpha_beta); + partials<2>(ops_partials)[i] += digamma_diff(beta_dbl, k) + - dpsi_r_plus_ab_k - dpsi_ab_r + + (dF[1] + dF[4]) / F; } } } diff --git a/stan/math/prim/prob/beta_neg_binomial_lcdf.hpp b/stan/math/prim/prob/beta_neg_binomial_lcdf.hpp index 897ffa135a4..6bd1cedc1f4 100644 --- a/stan/math/prim/prob/beta_neg_binomial_lcdf.hpp +++ b/stan/math/prim/prob/beta_neg_binomial_lcdf.hpp @@ -4,7 +4,8 @@ #include #include #include -#include +#include +#include #include #include #include @@ -99,43 +100,42 @@ inline return_type_t beta_neg_binomial_lcdf( std::initializer_list{1.0, b_plus_n + 1.0, r_plus_n + 1.0}, std::initializer_list{n_dbl + 2.0, a_plus_r + b_plus_n + 1.0}, 1.0); - auto C = lgamma(r_plus_n + 1.0) + lbeta(a_plus_r, b_plus_n + 1.0) - - lgamma(r_dbl) - lbeta(alpha_dbl, beta_dbl) - lgamma(n_dbl + 2.0); + // C is the log pmf at n + 1; formed from lgamma and lbeta directly it + // cancels for large parameters + const T_partials_return k = n_dbl + 1.0; + const T_partials_return C = internal::beta_neg_binomial_log_pmf( + k, T_partials_return(r_dbl), T_partials_return(alpha_dbl), + T_partials_return(beta_dbl)); auto ccdf = stan::math::exp(C) * F; log_cdf += log1m(ccdf); if constexpr (is_any_autodiff_v) { auto chain_rule_term = -ccdf / (1.0 - ccdf); - auto digamma_n_r_alpha_beta = digamma(a_plus_r + b_plus_n + 1.0); T_partials_return dF[6]; grad_F32, is_autodiff_v, false, true, false>(dF, 1.0, b_plus_n + 1.0, r_plus_n + 1.0, n_dbl + 2.0, a_plus_r + b_plus_n + 1.0, 1.0, precision, max_steps); - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_r_alpha = digamma(a_plus_r); - if constexpr (is_autodiff_v) { - auto partial_lccdf = digamma(r_plus_n + 1.0) - + (digamma_r_alpha - digamma_n_r_alpha_beta) - + (dF[2] + dF[4]) / F - digamma(r_dbl); - partials<0>(ops_partials)[i] += partial_lccdf * chain_rule_term; - } - if constexpr (is_autodiff_v) { - auto partial_lccdf = digamma_r_alpha - digamma_n_r_alpha_beta - + dF[4] / F - digamma(alpha_dbl); - partials<1>(ops_partials)[i] += partial_lccdf * chain_rule_term; - } + // the partials of C as differences digamma(x + d) - digamma(x) + const T_partials_return alpha_plus_beta = alpha_dbl + beta_dbl; + const T_partials_return dpsi_r_plus_ab_k + = digamma_diff(r_dbl + alpha_plus_beta, k); + if constexpr (is_autodiff_v) { + auto partial_lccdf = digamma_diff(r_dbl, k) - dpsi_r_plus_ab_k + - digamma_diff(a_plus_r, beta_dbl) + + (dF[2] + dF[4]) / F; + partials<0>(ops_partials)[i] += partial_lccdf * chain_rule_term; } - - if constexpr (is_autodiff_v || is_autodiff_v) { - auto digamma_alpha_beta = digamma(alpha_dbl + beta_dbl); + if constexpr (is_any_autodiff_v) { + const T_partials_return dpsi_ab_r + = digamma_diff(alpha_plus_beta, r_dbl); if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] += digamma_alpha_beta * chain_rule_term; + auto partial_lccdf = digamma_diff(alpha_dbl, r_dbl) - dpsi_ab_r + - dpsi_r_plus_ab_k + dF[4] / F; + partials<1>(ops_partials)[i] += partial_lccdf * chain_rule_term; } if constexpr (is_autodiff_v) { - auto partial_lccdf = digamma(b_plus_n + 1.0) - digamma_n_r_alpha_beta - + (dF[1] + dF[4]) / F - - (digamma(beta_dbl) - digamma_alpha_beta); + auto partial_lccdf = digamma_diff(beta_dbl, k) - dpsi_r_plus_ab_k + - dpsi_ab_r + (dF[1] + dF[4]) / F; partials<2>(ops_partials)[i] += partial_lccdf * chain_rule_term; } } diff --git a/stan/math/prim/prob/beta_neg_binomial_lpmf.hpp b/stan/math/prim/prob/beta_neg_binomial_lpmf.hpp index 8bb8eb14981..8be05d0cb08 100644 --- a/stan/math/prim/prob/beta_neg_binomial_lpmf.hpp +++ b/stan/math/prim/prob/beta_neg_binomial_lpmf.hpp @@ -4,9 +4,11 @@ #include #include #include -#include +#include #include #include +#include +#include #include #include #include @@ -74,47 +76,54 @@ inline return_type_t beta_neg_binomial_lpmf( scalar_seq_view alpha_vec(alpha_ref); scalar_seq_view beta_vec(beta_ref); const size_t max_size_seq_view = max_size(n, r, alpha, beta); + // With D(x, k) = lgamma(x + k) - lgamma(x) and A = alpha + beta, the log + // pmf is + // + // -lgamma(n + 1) + D(beta, n) + [D(r, n) - D(r + A, n)] + // + [D(alpha, r) - D(A, r)]. + // + // Formed from lbeta and lgamma directly, the terms are of the size of the + // shapes and cancel when the shapes are large. Here + // D(beta, n) = lgamma(n) - lbeta(n, beta) for n > 0, and each bracket is + // internal::log_rising_factorial_ratio, which does not cancel. The + // partials are sums of digamma differences psi(x + k) - psi(x), formed by + // digamma_diff. T_partials_return logp(0.0); for (size_t i = 0; i < max_size_seq_view; i++) { - if constexpr (include_summand::value) { - logp -= lgamma(n_vec[i] + 1); + const T_partials_return r_dbl = r_vec.val(i); + const T_partials_return alpha_dbl = alpha_vec.val(i); + const T_partials_return beta_dbl = beta_vec.val(i); + const T_partials_return n_dbl = n_vec.val(i); + const T_partials_return alpha_plus_beta = alpha_dbl + beta_dbl; + if (n_dbl > 0) { + if constexpr (include_summand::value) { + logp -= log(n_dbl); // -lgamma(n + 1) + lgamma(n) + } else { + logp += lgamma(n_dbl); + } + logp -= lbeta(n_dbl, beta_dbl); } - T_partials_return lbeta_denominator = lbeta(r_vec.val(i), alpha_vec.val(i)); - T_partials_return lgamma_numerator = lgamma(n_vec[i] + beta_vec.val(i)); - T_partials_return lgamma_denominator = lgamma(beta_vec.val(i)); - T_partials_return lbeta_numerator - = lbeta(n_vec[i] + r_vec.val(i), alpha_vec.val(i) + beta_vec.val(i)); - logp += lbeta_numerator + lgamma_numerator - lbeta_denominator - - lgamma_denominator; - if constexpr (is_any_autodiff_v) { - T_partials_return digamma_n_r_alpha_beta = digamma( - n_vec[i] + r_vec.val(i) + alpha_vec.val(i) + beta_vec.val(i)); + logp += internal::log_rising_factorial_ratio(r_dbl, alpha_plus_beta, n_dbl) + + internal::log_rising_factorial_ratio(alpha_dbl, beta_dbl, r_dbl); - if constexpr (is_autodiff_v || is_autodiff_v) { - T_partials_return digamma_r_alpha - = digamma(r_vec.val(i) + alpha_vec.val(i)); - if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] - += digamma(n_vec[i] + r_vec.val(i)) - digamma_n_r_alpha_beta - - (digamma(r_vec.val(i)) - digamma_r_alpha); - } + if constexpr (is_any_autodiff_v) { + const T_partials_return dpsi_r_plus_ab_n + = digamma_diff(r_dbl + alpha_plus_beta, n_dbl); + if constexpr (is_autodiff_v) { + partials<0>(ops_partials)[i] + += digamma_diff(r_dbl, n_dbl) - dpsi_r_plus_ab_n + - digamma_diff(r_dbl + alpha_dbl, beta_dbl); + } + if constexpr (is_any_autodiff_v) { + const T_partials_return dpsi_ab_r + = digamma_diff(alpha_plus_beta, r_dbl); if constexpr (is_autodiff_v) { partials<1>(ops_partials)[i] - += -digamma_n_r_alpha_beta - - (digamma(alpha_vec.val(i)) - digamma_r_alpha); + += digamma_diff(alpha_dbl, r_dbl) - dpsi_ab_r - dpsi_r_plus_ab_n; } - } - if constexpr (is_autodiff_v || is_autodiff_v) { - T_partials_return digamma_alpha_beta - = digamma(alpha_vec.val(i) + beta_vec.val(i)); if constexpr (is_autodiff_v) { - partials<2>(ops_partials)[i] += digamma_alpha_beta - - digamma_n_r_alpha_beta - + digamma(n_vec[i] + beta_vec.val(i)) - - digamma(beta_vec.val(i)); - } - if constexpr (is_autodiff_v) { - partials<1>(ops_partials)[i] += digamma_alpha_beta; + partials<2>(ops_partials)[i] + += digamma_diff(beta_dbl, n_dbl) - dpsi_r_plus_ab_n - dpsi_ab_r; } } } From fa5573904f376a11251833e0fe08b8a9b2bc4311 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:10:53 +0300 Subject: [PATCH 06/15] Use digamma_diff in the negative binomial gradients --- .../prim/prob/neg_binomial_2_log_glm_lpmf.hpp | 76 +++++++++++-------- .../prim/prob/neg_binomial_2_log_lpmf.hpp | 5 +- stan/math/prim/prob/neg_binomial_2_lpmf.hpp | 5 +- stan/math/prim/prob/neg_binomial_lpmf.hpp | 19 ++--- 4 files changed, 53 insertions(+), 52 deletions(-) diff --git a/stan/math/prim/prob/neg_binomial_2_log_glm_lpmf.hpp b/stan/math/prim/prob/neg_binomial_2_log_glm_lpmf.hpp index 3e7254f0262..8f93f6124ec 100644 --- a/stan/math/prim/prob/neg_binomial_2_log_glm_lpmf.hpp +++ b/stan/math/prim/prob/neg_binomial_2_log_glm_lpmf.hpp @@ -5,8 +5,9 @@ #include #include #include +#include #include -#include +#include #include #include #include @@ -147,42 +148,52 @@ neg_binomial_2_log_glm_lpmf(const T_y& y, const T_x& x, const T_alpha& alpha, } check_finite(function, "Matrix of independent variables", theta); T_precision_val log_phi = log(phi_arr); - Array logsumexp_theta_logphi - = (theta > log_phi) - .select(theta + log1p_exp(log_phi - theta), - log_phi + log1p_exp(theta - log_phi)); + // log1p(mu / phi) with mu = exp(theta) + Array log1p_exp_theta_m_log_phi + = log1p_exp(theta - log_phi); T_sum_val y_plus_phi = y_arr + phi_arr; - // Compute the log-density. - T_partials_return logp(0); - if constexpr (include_summand::value) { - if constexpr (is_vector::value) { - logp -= sum(lgamma(y_arr + 1.0)); - } else { - logp -= sum(lgamma(y_arr + 1.0)) * N_instances; - } - } + // Compute the log-density. The log pmf of one instance, with + // mu = exp(theta), is + // + // lchoose(y + phi - 1, y) - y log1p(phi / mu) - phi log1p(mu / phi), + // + // the form of neg_binomial_2_log_lpmf. Formed as phi log(phi) - lgamma(phi) + // + lgamma(y + phi) - (y + phi) log(mu + phi) + y theta, the terms are of + // the size of phi log(phi) and cancel for large phi. + T_partials_return logp = -sum(y_arr * log1p_exp(log_phi - theta) + + phi_arr * log1p_exp_theta_m_log_phi); if constexpr (include_summand::value) { - if constexpr (is_vector::value) { + if constexpr (is_vector::value || is_vector::value) { + scalar_seq_view y_vec(y_val_vec); scalar_seq_view phi_vec(phi_val_vec); for (size_t n = 0; n < N_instances; ++n) { - logp += multiply_log(phi_vec[n], phi_vec[n]) - lgamma(phi_vec[n]); + logp += binomial_coefficient_log(y_vec[n] + phi_vec[n] - 1, y_vec[n]); } } else { - logp += N_instances * (multiply_log(phi_val, phi_val) - lgamma(phi_val)); + logp + += N_instances * binomial_coefficient_log(y_val + phi_val - 1, y_val); } } - logp -= sum(y_plus_phi * logsumexp_theta_logphi); - - if constexpr (include_summand::value) { - logp += sum(y_arr * theta); - } - if constexpr (include_summand::value) { - if constexpr (is_vector::value || is_vector::value) { - logp += sum(lgamma(y_plus_phi)); + // Under propto, remove -lgamma(y + 1) (in lchoose), phi log(phi) for data + // phi, and y theta for data x, alpha and beta (in the log1p_exp terms). + if constexpr (!include_summand::value) { + if constexpr (include_summand::value) { + if constexpr (is_vector::value) { + logp += sum(lgamma(y_arr + 1.0)); + } else { + logp += sum(lgamma(y_arr + 1.0)) * N_instances; + } } else { - logp += sum(lgamma(y_plus_phi)) * N_instances; + if constexpr (is_vector::value) { + logp -= sum(phi_arr * log_phi); + } else { + logp -= N_instances * multiply_log(phi_val, phi_val); + } + } + if constexpr (!include_summand::value) { + logp -= sum(y_arr * theta); } } @@ -220,16 +231,17 @@ neg_binomial_2_log_glm_lpmf(const T_y& y, const T_x& x, const T_alpha& alpha, } } if constexpr (is_autodiff_v) { + // 1 - (y + phi) / (mu + phi) and log(phi) - log(mu + phi) are formed + // without cancellation, and digamma_diff replaces + // digamma(y + phi) - digamma(phi) if constexpr (is_vector::value) { edge<3>(ops_partials).partials_ - = 1 - y_plus_phi / (theta_exp + phi_arr) + log_phi - - logsumexp_theta_logphi + digamma(y_plus_phi) - digamma(phi_arr); + = (theta_exp - y_arr) / (theta_exp + phi_arr) + - log1p_exp_theta_m_log_phi + digamma_diff(phi_arr, y_arr); } else { partials<3>(ops_partials)[0] - = N_instances - + sum(-y_plus_phi / (theta_exp + phi_arr) + log_phi - - logsumexp_theta_logphi + digamma(y_plus_phi) - - digamma(phi_arr)); + = sum((theta_exp - y_arr) / (theta_exp + phi_arr) + - log1p_exp_theta_m_log_phi + digamma_diff(phi_arr, y_arr)); } } } diff --git a/stan/math/prim/prob/neg_binomial_2_log_lpmf.hpp b/stan/math/prim/prob/neg_binomial_2_log_lpmf.hpp index 44564538fbf..3c5d542a78b 100644 --- a/stan/math/prim/prob/neg_binomial_2_log_lpmf.hpp +++ b/stan/math/prim/prob/neg_binomial_2_log_lpmf.hpp @@ -4,7 +4,7 @@ #include #include #include -#include +#include #include #include #include @@ -124,8 +124,7 @@ inline return_type_t neg_binomial_2_log_lpmf( if constexpr (is_autodiff_v) { partials<1>(ops_partials)[i] += exp_eta_over_exp_eta_phi[i] - n_vec[i] / (exp_eta[i] + phi_val[i]) - - log1p_exp_eta_m_logphi[i] - - (digamma(phi_val[i]) - digamma(n_plus_phi[i])); + - log1p_exp_eta_m_logphi[i] + digamma_diff(phi_val[i], n_vec[i]); } } return ops_partials.build(logp); diff --git a/stan/math/prim/prob/neg_binomial_2_lpmf.hpp b/stan/math/prim/prob/neg_binomial_2_lpmf.hpp index b708d4ce3a9..6d522d95143 100644 --- a/stan/math/prim/prob/neg_binomial_2_lpmf.hpp +++ b/stan/math/prim/prob/neg_binomial_2_lpmf.hpp @@ -4,7 +4,7 @@ #include #include #include -#include +#include #include #include #include @@ -87,8 +87,7 @@ inline return_type_t neg_binomial_2_lpmf( auto log_term = select(mu_val < phi_val, log1p(-mu_val / mu_plus_phi), log_phi - log_mu_plus_phi); partials<1>(ops_partials) = (mu_val - value_of(n_vec)) / mu_plus_phi - + log_term - digamma(phi_val) - + digamma(n_plus_phi); + + log_term + digamma_diff(phi_val, n_vec); } return ops_partials.build(logp); } diff --git a/stan/math/prim/prob/neg_binomial_lpmf.hpp b/stan/math/prim/prob/neg_binomial_lpmf.hpp index 61c88160c6b..8ace0337b48 100644 --- a/stan/math/prim/prob/neg_binomial_lpmf.hpp +++ b/stan/math/prim/prob/neg_binomial_lpmf.hpp @@ -5,7 +5,7 @@ #include #include #include -#include +#include #include #include #include @@ -60,19 +60,10 @@ inline return_type_t neg_binomial_lpmf( scalar_seq_view n_vec(n_ref); scalar_seq_view alpha_vec(alpha_ref); scalar_seq_view beta_vec(beta_ref); - size_t size_alpha = stan::math::size(alpha); size_t size_beta = stan::math::size(beta); size_t size_alpha_beta = max_size(alpha, beta); size_t max_size_seq_view = max_size(n, alpha, beta); - VectorBuilder, T_partials_return, T_shape> - digamma_alpha(size_alpha); - if constexpr (is_autodiff_v) { - for (size_t i = 0; i < size_alpha; ++i) { - digamma_alpha[i] = digamma(alpha_vec.val(i)); - } - } - VectorBuilder log1p_inv_beta(size_beta); VectorBuilder log1p_beta(size_beta); for (size_t i = 0; i < size_beta; ++i) { @@ -88,8 +79,8 @@ inline return_type_t neg_binomial_lpmf( for (size_t i = 0; i < size_alpha_beta; ++i) { const T_partials_return alpha_dbl = alpha_vec.val(i); const T_partials_return beta_dbl = beta_vec.val(i); - lambda_m_alpha_over_1p_beta[i] - = alpha_dbl / beta_dbl - alpha_dbl / (1 + beta_dbl); + // alpha / beta - alpha / (1 + beta), without the cancellation + lambda_m_alpha_over_1p_beta[i] = alpha_dbl / beta_dbl / (1 + beta_dbl); } } @@ -106,8 +97,8 @@ inline return_type_t neg_binomial_lpmf( logp -= alpha_dbl * log1p_inv_beta[i] + n_vec[i] * log1p_beta[i]; if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] += digamma(alpha_dbl + n_vec[i]) - - digamma_alpha[i] - log1p_inv_beta[i]; + partials<0>(ops_partials)[i] + += digamma_diff(alpha_dbl, n_vec[i]) - log1p_inv_beta[i]; } if constexpr (is_autodiff_v) { partials<1>(ops_partials)[i] From 38af31d08930c03811fa36c463cc7f7f3eb99c34 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:00 +0300 Subject: [PATCH 07/15] Use digamma_diff in dirichlet_multinomial_lpmf --- stan/math/prim/prob/dirichlet_multinomial_lpmf.hpp | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/stan/math/prim/prob/dirichlet_multinomial_lpmf.hpp b/stan/math/prim/prob/dirichlet_multinomial_lpmf.hpp index b371cfb340d..2f514e61b2b 100644 --- a/stan/math/prim/prob/dirichlet_multinomial_lpmf.hpp +++ b/stan/math/prim/prob/dirichlet_multinomial_lpmf.hpp @@ -5,7 +5,7 @@ #include #include #include -#include +#include #include #include #include @@ -89,10 +89,9 @@ inline return_type_t dirichlet_multinomial_lpmf( auto ops_partials = make_partials_propagator(alpha_ref); if constexpr (is_autodiff_v) { + // digamma_diff(x, 0) is 0, so the categories with n = 0 need no select partials<0>(ops_partials) - = (ns_array > 0) - .select(digamma(alpha_val + ns_array) - digamma(alpha_val), 0.0) - + digamma(a_sum) - digamma(a_sum + n_sum); + = digamma_diff(alpha_val, ns_array) - digamma_diff(a_sum, n_sum); } return ops_partials.build(lp); } From ceec7d3e463baf03fdacbbb6990f813522a6dce1 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:06 +0300 Subject: [PATCH 08/15] Avoid the lgamma cancellation in the lkj_corr constant --- stan/math/prim/prob/lkj_corr_lpdf.hpp | 29 ++++++++++++++++++++++----- 1 file changed, 24 insertions(+), 5 deletions(-) diff --git a/stan/math/prim/prob/lkj_corr_lpdf.hpp b/stan/math/prim/prob/lkj_corr_lpdf.hpp index f10d55ba204..d4fdd145f16 100644 --- a/stan/math/prim/prob/lkj_corr_lpdf.hpp +++ b/stan/math/prim/prob/lkj_corr_lpdf.hpp @@ -4,10 +4,14 @@ #include #include #include +#include +#include #include #include #include #include +#include +#include namespace stan { namespace math { @@ -16,10 +20,12 @@ template inline return_type_t do_lkj_constant(const T_shape& eta, const unsigned int& K) { // Lewandowski, Kurowicka, and Joe (2009) theorem 5 - return_type_t constant; + using T_partials_return = partials_return_t; + const T_partials_return eta_val = value_of(eta); + T_partials_return constant; const int Km1 = K - 1; using stan::math::lgamma; - if (eta == 1.0) { + if (eta_val == 1.0) { // C++ integer division is appropriate in this block Eigen::VectorXd denominator(Km1 / 2); for (int k = 1; k <= denominator.rows(); k++) { @@ -35,12 +41,25 @@ inline return_type_t do_lkj_constant(const T_shape& eta, - Km1 * lgamma(static_cast(K)); } } else { - constant = Km1 * lgamma(eta + 0.5 * Km1); + // Km1 lgamma(eta + Km1 / 2) - sum_k lgamma(eta + (Km1 - k) / 2) is a sum + // of differences lgamma(x + k / 2) - lgamma(x), x = eta + (Km1 - k) / 2, + // which cancel for large eta. Each one equals + // lgamma(k / 2) - lbeta(k / 2, x), which does not cancel. + constant = -0.25 * Km1 * K * LOG_PI; for (int k = 1; k <= Km1; k++) { - constant -= 0.5 * k * LOG_PI + lgamma(eta + 0.5 * (Km1 - k)); + constant += lgamma(0.5 * k) - lbeta(0.5 * k, eta_val + 0.5 * (Km1 - k)); } } - return constant; + // d/deta = sum_k digamma(x + k / 2) - digamma(x), also at eta = 1 + auto ops_partials = make_partials_propagator(eta); + if constexpr (is_autodiff_v) { + T_partials_return d_eta(0.0); + for (int k = 1; k <= Km1; k++) { + d_eta += digamma_diff(eta_val + 0.5 * (Km1 - k), 0.5 * k); + } + partials<0>(ops_partials)[0] = d_eta; + } + return ops_partials.build(constant); } // LKJ_Corr(y|eta) [ y correlation matrix (not covariance matrix) From 0acd8e9c3cd100e5ee6a6a168e616c7396055ac9 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:12 +0300 Subject: [PATCH 09/15] Avoid the lgamma cancellation in the yule_simon cdfs --- stan/math/prim/prob/yule_simon_cdf.hpp | 18 +++++++++++------- stan/math/prim/prob/yule_simon_lccdf.hpp | 15 +++++++++------ stan/math/prim/prob/yule_simon_lcdf.hpp | 19 +++++++++++-------- 3 files changed, 31 insertions(+), 21 deletions(-) diff --git a/stan/math/prim/prob/yule_simon_cdf.hpp b/stan/math/prim/prob/yule_simon_cdf.hpp index 10a5ba2637e..73989e7da2e 100644 --- a/stan/math/prim/prob/yule_simon_cdf.hpp +++ b/stan/math/prim/prob/yule_simon_cdf.hpp @@ -4,9 +4,11 @@ #include #include #include -#include +#include #include +#include #include +#include #include #include #include @@ -66,17 +68,19 @@ inline return_type_t yule_simon_cdf(const T_n& n, T_partials_return cdf(1.0); auto ops_partials = make_partials_propagator(alpha_ref); for (size_t i = 0; i < max_size_seq_view; i++) { - auto np1 = n_vec.val(i) + 1.0; - auto ap1 = alpha_vec.val(i) + 1.0; - auto nap1 = n_vec.val(i) + ap1; + const T_partials_return n_dbl = n_vec.val(i); + const T_partials_return ap1 = alpha_vec.val(i) + 1.0; - auto ccdf = stan::math::exp(lgamma(ap1) + lgamma(np1) - lgamma(nap1)); + // lgamma(alpha + 1) + lgamma(n + 1) - lgamma(n + alpha + 1) without the + // cancellation for large alpha + const T_partials_return ccdf + = stan::math::exp(log(n_dbl) + lbeta(n_dbl, ap1)); cdf *= 1.0 - ccdf; if constexpr (is_autodiff_v) { - auto chain_rule_term = -ccdf / (1.0 - ccdf); + const T_partials_return chain_rule_term = -ccdf / (1.0 - ccdf); partials<0>(ops_partials)[i] - += (digamma(ap1) - digamma(nap1)) * chain_rule_term; + -= digamma_diff(ap1, n_dbl) * chain_rule_term; } } diff --git a/stan/math/prim/prob/yule_simon_lccdf.hpp b/stan/math/prim/prob/yule_simon_lccdf.hpp index 335fa158e8c..96b294ac477 100644 --- a/stan/math/prim/prob/yule_simon_lccdf.hpp +++ b/stan/math/prim/prob/yule_simon_lccdf.hpp @@ -4,9 +4,11 @@ #include #include #include -#include +#include #include +#include #include +#include #include #include #include @@ -66,13 +68,14 @@ inline return_type_t yule_simon_lccdf(const T_n& n, T_partials_return log_ccdf(0.0); auto ops_partials = make_partials_propagator(alpha_ref); for (size_t i = 0; i < max_size_seq_view; i++) { - auto np1 = n_vec.val(i) + 1.0; - auto ap1 = alpha_vec.val(i) + 1.0; - auto nap1 = n_vec.val(i) + ap1; - log_ccdf += lgamma(ap1) + lgamma(np1) - lgamma(nap1); + const T_partials_return n_dbl = n_vec.val(i); + const T_partials_return ap1 = alpha_vec.val(i) + 1.0; + // lgamma(alpha + 1) + lgamma(n + 1) - lgamma(n + alpha + 1) without the + // cancellation for large alpha + log_ccdf += log(n_dbl) + lbeta(n_dbl, ap1); if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] += digamma(ap1) - digamma(nap1); + partials<0>(ops_partials)[i] -= digamma_diff(ap1, n_dbl); } } diff --git a/stan/math/prim/prob/yule_simon_lcdf.hpp b/stan/math/prim/prob/yule_simon_lcdf.hpp index b84e54a6a82..d1a4546d5a0 100644 --- a/stan/math/prim/prob/yule_simon_lcdf.hpp +++ b/stan/math/prim/prob/yule_simon_lcdf.hpp @@ -4,9 +4,11 @@ #include #include #include -#include +#include #include +#include #include +#include #include #include #include @@ -66,18 +68,19 @@ inline return_type_t yule_simon_lcdf(const T_n& n, T_partials_return log_cdf(0.0); auto ops_partials = make_partials_propagator(alpha_ref); for (size_t i = 0; i < max_size_seq_view; i++) { - auto np1 = n_vec.val(i) + 1.0; - auto ap1 = alpha_vec.val(i) + 1.0; - auto nap1 = n_vec.val(i) + ap1; + const T_partials_return n_dbl = n_vec.val(i); + const T_partials_return ap1 = alpha_vec.val(i) + 1.0; - auto log_ccdf = lgamma(ap1) + lgamma(np1) - lgamma(nap1); + // lgamma(alpha + 1) + lgamma(n + 1) - lgamma(n + alpha + 1) without the + // cancellation for large alpha + const T_partials_return log_ccdf = log(n_dbl) + lbeta(n_dbl, ap1); log_cdf += log1m_exp(log_ccdf); if constexpr (is_autodiff_v) { - auto ccdf = stan::math::exp(log_ccdf); - auto chain_rule_term = -ccdf / (1.0 - ccdf); + const T_partials_return ccdf = stan::math::exp(log_ccdf); + const T_partials_return chain_rule_term = -ccdf / (1.0 - ccdf); partials<0>(ops_partials)[i] - += (digamma(ap1) - digamma(nap1)) * chain_rule_term; + -= digamma_diff(ap1, n_dbl) * chain_rule_term; } } From 60ca002dfc9e4df850b1c45c43b095b95dc6f895 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:20 +0300 Subject: [PATCH 10/15] Use digamma_diff in the derivatives of lbeta at large arguments --- stan/math/fwd/fun/lbeta.hpp | 29 ++++++++++++++++++++++++++-- stan/math/rev/fun/lbeta.hpp | 38 +++++++++++++++++++++++++++---------- 2 files changed, 55 insertions(+), 12 deletions(-) diff --git a/stan/math/fwd/fun/lbeta.hpp b/stan/math/fwd/fun/lbeta.hpp index eb2c05b4c1a..509766d97d5 100644 --- a/stan/math/fwd/fun/lbeta.hpp +++ b/stan/math/fwd/fun/lbeta.hpp @@ -5,13 +5,38 @@ #include #include +#include #include +#include namespace stan { namespace math { +namespace internal { +/** + * Return digamma(x) - digamma(x + y), the partial derivative of + * lbeta(x, y) in x. For large x the plain difference loses all digits, + * so from x = digamma_diff_min_x on it is formed by + * digamma_diff. Below that value (and for NaN) the plain difference is + * accurate enough for a gradient and cheaper. + */ +template +inline return_type_t lbeta_partial_fwd(const T1& x, const T2& y) { + if (value_of_rec(x) >= digamma_diff_min_x) { + return -digamma_diff(x, y); + } + return digamma(x) - digamma(x + y); +} +} // namespace internal template inline fvar lbeta(const fvar& x1, const fvar& x2) { + if (value_of_rec(x1.val_) >= internal::digamma_diff_min_x + || value_of_rec(x2.val_) >= internal::digamma_diff_min_x) { + return fvar(lbeta(x1.val_, x2.val_), + x1.d_ * internal::lbeta_partial_fwd(x1.val_, x2.val_) + + x2.d_ * internal::lbeta_partial_fwd(x2.val_, x1.val_)); + } + // both arguments small: the plain differences share digamma(x1 + x2) return fvar(lbeta(x1.val_, x2.val_), x1.d_ * digamma(x1.val_) + x2.d_ * digamma(x2.val_) - (x1.d_ + x2.d_) * digamma(x1.val_ + x2.val_)); @@ -20,13 +45,13 @@ inline fvar lbeta(const fvar& x1, const fvar& x2) { template inline fvar lbeta(double x1, const fvar& x2) { return fvar(lbeta(x1, x2.val_), - x2.d_ * digamma(x2.val_) - x2.d_ * digamma(x1 + x2.val_)); + x2.d_ * internal::lbeta_partial_fwd(x2.val_, x1)); } template inline fvar lbeta(const fvar& x1, double x2) { return fvar(lbeta(x1.val_, x2), - x1.d_ * digamma(x1.val_) - x1.d_ * digamma(x1.val_ + x2)); + x1.d_ * internal::lbeta_partial_fwd(x1.val_, x2)); } } // namespace math } // namespace stan diff --git a/stan/math/rev/fun/lbeta.hpp b/stan/math/rev/fun/lbeta.hpp index e5a8b01f8b8..84c8e421879 100644 --- a/stan/math/rev/fun/lbeta.hpp +++ b/stan/math/rev/fun/lbeta.hpp @@ -4,21 +4,43 @@ #include #include #include +#include #include namespace stan { namespace math { namespace internal { +/** + * Return digamma(x) - digamma(x + y), the partial derivative of + * lbeta(x, y) in x. For large x the plain difference loses all digits, + * so from x = digamma_diff_min_x on it is formed by + * digamma_diff. Below that value (and for NaN) the plain difference is + * accurate enough for a gradient and cheaper. + */ +inline double lbeta_partial(double x, double y) { + if (x >= digamma_diff_min_x) { + return -digamma_diff(x, y); + } + return digamma(x) - digamma(x + y); +} + class lbeta_vv_vari : public op_vv_vari { public: lbeta_vv_vari(vari* avi, vari* bvi) : op_vv_vari(lbeta(avi->val_, bvi->val_), avi, bvi) {} void chain() { - const double digamma_ab = digamma(avi_->val_ + bvi_->val_); - avi_->adj_ += adj_ * (digamma(avi_->val_) - digamma_ab); - - bvi_->adj_ += adj_ * (digamma(bvi_->val_) - digamma_ab); + const double a = avi_->val_; + const double b = bvi_->val_; + if (a >= digamma_diff_min_x || b >= digamma_diff_min_x) { + avi_->adj_ += adj_ * lbeta_partial(a, b); + bvi_->adj_ += adj_ * lbeta_partial(b, a); + } else { + // both arguments small: the plain differences share digamma(a + b) + const double digamma_ab = digamma(a + b); + avi_->adj_ += adj_ * (digamma(a) - digamma_ab); + bvi_->adj_ += adj_ * (digamma(b) - digamma_ab); + } } }; @@ -26,18 +48,14 @@ class lbeta_vd_vari : public op_vd_vari { public: lbeta_vd_vari(vari* avi, double b) : op_vd_vari(lbeta(avi->val_, b), avi, b) {} - void chain() { - avi_->adj_ += adj_ * (digamma(avi_->val_) - digamma(avi_->val_ + bd_)); - } + void chain() { avi_->adj_ += adj_ * lbeta_partial(avi_->val_, bd_); } }; class lbeta_dv_vari : public op_dv_vari { public: lbeta_dv_vari(double a, vari* bvi) : op_dv_vari(lbeta(a, bvi->val_), a, bvi) {} - void chain() { - bvi_->adj_ += adj_ * (digamma(bvi_->val_) - digamma(ad_ + bvi_->val_)); - } + void chain() { bvi_->adj_ += adj_ * lbeta_partial(bvi_->val_, ad_); } }; } // namespace internal From b05cbd11cfbcc8b13bbf5c455d67c5fe932b2a93 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:27 +0300 Subject: [PATCH 11/15] Avoid the lgamma cancellation in student_t, multi_student_t and yule_simon_lpmf --- .../prob/multi_student_t_cholesky_lpdf.hpp | 20 ++-- stan/math/prim/prob/multi_student_t_lpdf.hpp | 7 +- stan/math/prim/prob/student_t_lpdf.hpp | 96 +++++++++++++++++-- stan/math/prim/prob/yule_simon_lpmf.hpp | 30 +++--- 4 files changed, 120 insertions(+), 33 deletions(-) diff --git a/stan/math/prim/prob/multi_student_t_cholesky_lpdf.hpp b/stan/math/prim/prob/multi_student_t_cholesky_lpdf.hpp index d119460be62..f608cefc8b9 100644 --- a/stan/math/prim/prob/multi_student_t_cholesky_lpdf.hpp +++ b/stan/math/prim/prob/multi_student_t_cholesky_lpdf.hpp @@ -6,13 +6,14 @@ #include #include #include -#include +#include #include #include #include #include #include #include +#include #include #include #include @@ -147,12 +148,15 @@ inline return_type_t multi_student_t_cholesky_lpdf( matrix_partials_t L_deriv; const auto& half_nu = to_ref_if::value>(0.5 * nu_val); + // digamma(nu/2 + p/2) - digamma(nu/2) and lgamma(nu/2 + p/2) - + // lgamma(nu/2) without the cancellation for large nu; the second is + // lgamma(p/2) - lbeta(p/2, nu/2) const auto& digamma_vals = to_ref_if>( - digamma(half_nu + 0.5 * num_dims) - digamma(half_nu)); + digamma_diff(half_nu, 0.5 * num_dims)); if constexpr (include_summand::value) { - lp += lgamma(0.5 * nu_plus_dims) * size_vec; - lp += -lgamma(0.5 * nu_val) * size_vec; + lp += (lgamma(0.5 * num_dims) - lbeta(0.5 * num_dims, half_nu)) + * size_vec; lp += -(0.5 * num_dims) * log(nu_val) * size_vec; } @@ -315,8 +319,8 @@ inline return_type_t multi_student_t_cholesky_lpdf( if constexpr (is_autodiff_v) { T_partials_return half_nu = 0.5 * nu_val; - T_partials_return digamma_vals - = digamma(half_nu + 0.5 * size_y) - digamma(half_nu); + // digamma(nu/2 + p/2) - digamma(nu/2) without the cancellation + T_partials_return digamma_vals = digamma_diff(half_nu, 0.5 * size_y); T_partials_return G = dot_product(scaled_diff, y_val_minus_mu_val); partials<1>(ops_partials) @@ -325,8 +329,8 @@ inline return_type_t multi_student_t_cholesky_lpdf( } if constexpr (include_summand::value) { - lp += lgamma(0.5 * (nu_val + size_y)); - lp += -lgamma(0.5 * nu_val); + // lgamma(nu/2 + p/2) - lgamma(nu/2) as lgamma(p/2) - lbeta(p/2, nu/2) + lp += lgamma(0.5 * size_y) - lbeta(0.5 * size_y, 0.5 * nu_val); lp += -0.5 * size_y * log(nu_val); } diff --git a/stan/math/prim/prob/multi_student_t_lpdf.hpp b/stan/math/prim/prob/multi_student_t_lpdf.hpp index f4e0798d9c2..ea2155f7ba3 100644 --- a/stan/math/prim/prob/multi_student_t_lpdf.hpp +++ b/stan/math/prim/prob/multi_student_t_lpdf.hpp @@ -8,6 +8,7 @@ #include #include #include +#include #include #include #include @@ -103,8 +104,10 @@ inline return_type_t multi_student_t_lpdf( lp_type lp(0); if constexpr (include_summand::value) { - lp += lgamma(0.5 * (nu + num_dims)) * size_vec; - lp -= lgamma(0.5 * nu) * size_vec; + // lgamma((nu + p) / 2) - lgamma(nu / 2) is lgamma(p / 2) - lbeta(p / 2, + // nu / 2). The direct difference loses eps * nu; + // lbeta keeps the digits, and its derivative uses digamma_diff. + lp += (lgamma(0.5 * num_dims) - lbeta(0.5 * num_dims, 0.5 * nu)) * size_vec; lp -= (0.5 * num_dims) * log(nu) * size_vec; } diff --git a/stan/math/prim/prob/student_t_lpdf.hpp b/stan/math/prim/prob/student_t_lpdf.hpp index b2c625840b9..5ec585a3bd9 100644 --- a/stan/math/prim/prob/student_t_lpdf.hpp +++ b/stan/math/prim/prob/student_t_lpdf.hpp @@ -17,12 +17,90 @@ #include #include #include +#include +#include #include #include namespace stan { namespace math { +namespace internal { + +/** + * From this value of nu on, student_t_lpdf forms its nu terms with the + * asymptotic series below. Below it, the differences of two lgamma and of + * two digamma values lose about eps * lgamma(nu / 2), at most about 1e-14. + * For large nu that loss grows like eps * nu log(nu), and the result has no + * correct digits from about nu = 1e16. At nu = 40 the + * first omitted series term is below 1e-17 of the result. + */ +constexpr double student_t_series_min_nu = 40.0; + +/** + * Return lgamma((nu + 1) / 2) - lgamma(nu / 2) - log(nu) / 2, the part of + * the Student-t normalizing constant that depends on nu. For nu >= 40, + * with h = nu / 2, the asymptotic series + * + * lgamma(h + 1/2) - lgamma(h) = log(h) / 2 - 1 / (8 h) + 1 / (192 h^3) + * - 1 / (640 h^5) + 17 / (14336 h^7) - 31 / (18432 h^9) + O(h^-11) + * + * is used, in which log(h) / 2 - log(nu) / 2 = -log(2) / 2. + */ +struct student_t_nu_constant_fun { + template + static inline T fun(const T& nu) { + if (value_of_rec(nu) < student_t_series_min_nu) { + const T half_nu = 0.5 * nu; + return lgamma(half_nu + 0.5) - lgamma(half_nu) - 0.5 * log(nu); + } + const T z = 2.0 / nu; + const T z2 = square(z); + return -0.5 * LOG_TWO + + z + * (-0.125 + + z2 + * (1.0 / 192.0 + + z2 + * (-1.0 / 640.0 + + z2 + * (17.0 / 14336.0 + - z2 * (31.0 / 18432.0))))); + } +}; + +/** + * Return digamma((nu + 1) / 2) - digamma(nu / 2). For nu >= 40, with + * h = nu / 2, the derivative in h of the series above is used: + * + * 1 / (2 h) + 1 / (8 h^2) - 1 / (64 h^4) + 1 / (128 h^6) + * - 119 / (14336 h^8) + 279 / (18432 h^10) + O(h^-12). + */ +struct student_t_nu_digamma_fun { + template + static inline T fun(const T& nu) { + if (value_of_rec(nu) < student_t_series_min_nu) { + const T half_nu = 0.5 * nu; + return digamma(half_nu + 0.5) - digamma(half_nu); + } + const T z = 2.0 / nu; + const T z2 = square(z); + return z + * (0.5 + + z + * (0.125 + + z2 + * (-1.0 / 64.0 + + z2 + * (1.0 / 128.0 + + z2 + * (-119.0 / 14336.0 + + z2 * (279.0 / 18432.0)))))); + } +}; + +} // namespace internal + /** \ingroup prob_dists * The log of the Student-t density for the given y, nu, mean, and * scale parameter. The scale parameter must be greater @@ -107,8 +185,11 @@ inline return_type_t student_t_lpdf( logp -= LOG_SQRT_PI * N; } if constexpr (include_summand::value) { - logp += (sum(lgamma(half_nu + 0.5)) - sum(lgamma(half_nu)) - - 0.5 * sum(log(nu_val))) + // lgamma(nu/2 + 1/2) - lgamma(nu/2) - log(nu)/2, by a series for large + // nu; see internal::student_t_nu_constant_fun + logp += sum(apply_scalar_unary< + internal::student_t_nu_constant_fun, + std::decay_t>::apply(nu_val)) * N / math::size(nu); } if constexpr (include_summand::value) { @@ -134,12 +215,13 @@ inline return_type_t student_t_lpdf( / (1 + square_y_scaled_over_nu) - 1); if constexpr (is_autodiff_v) { - const auto& digamma_half_nu_plus_half = digamma(half_nu + 0.5); - const auto& digamma_half_nu = digamma(half_nu); + // digamma(nu/2 + 1/2) - digamma(nu/2), by a series for large nu; + // see internal::student_t_nu_digamma_fun + const auto& digamma_nu_term + = apply_scalar_unary>::apply(nu_val); edge<1>(ops_partials).partials_ - = 0.5 - * (digamma_half_nu_plus_half - digamma_half_nu - log1p_val - + rep_deriv / nu_val); + = 0.5 * (digamma_nu_term - log1p_val + rep_deriv / nu_val); } if constexpr (is_autodiff_v) { partials<3>(ops_partials) = rep_deriv / sigma_val; diff --git a/stan/math/prim/prob/yule_simon_lpmf.hpp b/stan/math/prim/prob/yule_simon_lpmf.hpp index 7ac83b3bf83..7ce7eeb4124 100644 --- a/stan/math/prim/prob/yule_simon_lpmf.hpp +++ b/stan/math/prim/prob/yule_simon_lpmf.hpp @@ -4,9 +4,10 @@ #include #include #include -#include +#include #include #include +#include #include #include #include @@ -61,24 +62,21 @@ inline return_type_t yule_simon_lpmf(const T_n &n, scalar_seq_view alpha_vec(alpha_ref); const size_t max_size_seq_view = max_size(n_ref, alpha_ref); T_partials_return logp(0.0); - if constexpr (include_summand::value) { - if constexpr (is_stan_scalar_v) { - logp += lgamma(n_ref) * max_size_seq_view; - } - } for (size_t i = 0; i < max_size_seq_view; i++) { - if constexpr (include_summand::value) { - if constexpr (!is_stan_scalar_v) { - logp += lgamma(n_vec.val(i)); - } + const T_partials_return alpha_plus_one = alpha_vec.val(i) + 1.0; + // lgamma(n) + lgamma(alpha + 1) - lgamma(n + alpha + 1) is + // lbeta(n, alpha + 1). Formed from the three lgamma values, the + // difference is of the size of alpha and loses eps * alpha. + logp += log(alpha_vec.val(i)) + lbeta(n_vec.val(i), alpha_plus_one); + if constexpr (!include_summand::value) { + // lbeta(n, alpha + 1) contains the constant lgamma(n) + logp -= lgamma(n_vec.val(i)); } - T_partials_return alpha_plus_one = alpha_vec.val(i) + 1.0; - logp += log(alpha_vec.val(i)) + lgamma(alpha_plus_one) - - lgamma(n_vec.val(i) + alpha_plus_one); if constexpr (is_autodiff_v) { - partials<0>(ops_partials)[i] += 1.0 / alpha_vec.val(i) - + digamma(alpha_plus_one) - - digamma(n_vec.val(i) + alpha_plus_one); + // digamma(alpha + 1) - digamma(n + alpha + 1) without the cancellation + partials<0>(ops_partials)[i] + += 1.0 / alpha_vec.val(i) + - digamma_diff(alpha_plus_one, n_vec.val(i)); } } return ops_partials.build(logp); From f24ee3a717f7fcb6e1a1d7a5dfc9279e769215ec Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:33 +0300 Subject: [PATCH 12/15] Test the large-shape accuracy of the repaired distributions --- test/unit/math/rev/prob/large_shapes_test.cpp | 436 ++++++++++++++++++ 1 file changed, 436 insertions(+) create mode 100644 test/unit/math/rev/prob/large_shapes_test.cpp diff --git a/test/unit/math/rev/prob/large_shapes_test.cpp b/test/unit/math/rev/prob/large_shapes_test.cpp new file mode 100644 index 00000000000..fd5b286e4d1 --- /dev/null +++ b/test/unit/math/rev/prob/large_shapes_test.cpp @@ -0,0 +1,436 @@ +#include +#include +#include +#include +#include +#include +#include +#include + +// Accuracy at large shape parameters of the distributions whose value or +// partials contain a difference f(x + k) - f(x), f in {lgamma, lbeta, +// digamma}, of a large shape x and a count or offset k. The plain +// differences keep no correct digits once x +// is near 1e15. beta_binomial_lpmf has its own tests in +// test/unit/math/{prim,rev}/prob/beta_binomial*_test.cpp. +// +// References: mpmath at 90 digits (value from mp.loggamma, partials from +// mp.digamma), checked against mp.diff at 140 digits; the cdfs as sums of +// the pmf. The arguments are written in hex so that they are exact. The +// partials of positive parameters are compared as x * d/dx, the gradient in +// log(x) that a sampler sees; the others are compared as they are. +// +// For each function the first rows are points where develop was wrong by +// more than 100 times the tolerance, and the last row is a point at a small +// shape where develop was correct. The cdfs of beta_binomial and +// beta_neg_binomial are limited by the tolerance of their 3F2 series, so +// their tolerances are larger. + +namespace large_shapes_test_internal { + +struct TestCase { + const char* tag; + std::vector args; + double value; + std::vector grads; +}; + +// (value, gradient) tolerance factors, applied as tol * max(1, |reference|) +std::pair tolerances(const std::string& tag) { + if (tag == "BBC" || tag == "BBLC" || tag == "BBLCC" || tag == "BNBC" + || tag == "BNBLCC") { + return {1e-6, 1e-4}; + } + return {1e-12, 1e-11}; +} + +// for each gradient, the index of the argument it is taken with respect +// to if that argument is a positive parameter (compared on the log scale), +// or -1 +std::vector log_scale_args(const std::string& tag) { + if (tag == "NB") { + return {1, 2}; + } else if (tag == "NB2") { + return {-1, 2}; + } else if (tag == "NB2L") { + return {-1, 2}; + } else if (tag == "GLM1") { + return {-1, -1, 4}; + } else if (tag == "GLM5") { + return {-1, -1, 12}; + } else if (tag == "BNB" || tag == "BNBC" || tag == "BNBLCC") { + return {1, 2, 3}; + } else if (tag == "DM") { + return {3, 4, 5}; + } else if (tag == "LKJ") { + return {3}; + } else if (tag == "LKJC") { + return {5}; + } else if (tag == "BBC" || tag == "BBLC" || tag == "BBLCC") { + return {2, 3}; + } else if (tag == "ST") { + return {-1, 1, -1, 3}; + } else if (tag == "MST" || tag == "MSTF") { + return {-1, -1, 2}; + } else if (tag == "LBETA") { + return {0, 1}; + } + return {1}; // YS, YSC, YSLC, YSLCC +} + +void eval(const std::string& tag, const std::vector& a, double& value, + std::vector& grads) { + using Eigen::Dynamic; + using Eigen::Matrix; + using stan::math::var; + var lp; + std::vector wrt; + auto to_int = [](double x) { return static_cast(x); }; + if (tag == "NB" || tag == "NB2" || tag == "NB2L") { + var p1 = a[1], p2 = a[2]; + if (tag == "NB") { + lp = stan::math::neg_binomial_lpmf(to_int(a[0]), p1, p2); + } else if (tag == "NB2") { + lp = stan::math::neg_binomial_2_lpmf(to_int(a[0]), p1, p2); + } else { + lp = stan::math::neg_binomial_2_log_lpmf(to_int(a[0]), p1, p2); + } + wrt = {p1, p2}; + } else if (tag == "GLM1" || tag == "GLM5") { + const int n_obs = (tag == "GLM1") ? 1 : 5; + std::vector y(n_obs); + Eigen::MatrixXd x(n_obs, 1); + for (int i = 0; i < n_obs; ++i) { + y[i] = to_int(a[i]); + x(i, 0) = a[n_obs + i]; + } + var alpha = a[2 * n_obs]; + Matrix beta(1); + beta(0) = a[2 * n_obs + 1]; + var phi = a[2 * n_obs + 2]; + lp = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta, phi); + wrt = {alpha, beta(0), phi}; + } else if (tag == "BNB" || tag == "BNBC" || tag == "BNBLCC") { + var r = a[1], alpha = a[2], beta = a[3]; + const int n = to_int(a[0]); + if (tag == "BNB") { + lp = stan::math::beta_neg_binomial_lpmf(n, r, alpha, beta); + } else if (tag == "BNBC") { + lp = stan::math::beta_neg_binomial_cdf(n, r, alpha, beta); + } else { + lp = stan::math::beta_neg_binomial_lccdf(n, r, alpha, beta); + } + wrt = {r, alpha, beta}; + } else if (tag == "DM") { + std::vector ns{to_int(a[0]), to_int(a[1]), to_int(a[2])}; + Matrix alpha(3); + alpha << a[3], a[4], a[5]; + lp = stan::math::dirichlet_multinomial_lpmf(ns, alpha); + wrt = {alpha(0), alpha(1), alpha(2)}; + } else if (tag == "LKJ") { + Eigen::MatrixXd omega(3, 3); + omega << 1.0, a[0], a[1], a[0], 1.0, a[2], a[1], a[2], 1.0; + var eta = a[3]; + lp = stan::math::lkj_corr_lpdf(omega, eta); + wrt = {eta}; + } else if (tag == "LKJC") { + Eigen::MatrixXd L = Eigen::MatrixXd::Zero(3, 3); + L(0, 0) = 1.0; + L(1, 0) = a[0]; + L(1, 1) = a[1]; + L(2, 0) = a[2]; + L(2, 1) = a[3]; + L(2, 2) = a[4]; + var eta = a[5]; + lp = stan::math::lkj_corr_cholesky_lpdf(L, eta); + wrt = {eta}; + } else if (tag == "BBC" || tag == "BBLC" || tag == "BBLCC") { + var alpha = a[2], beta = a[3]; + const int n = to_int(a[0]), N = to_int(a[1]); + if (tag == "BBC") { + lp = stan::math::beta_binomial_cdf(n, N, alpha, beta); + } else if (tag == "BBLC") { + lp = stan::math::beta_binomial_lcdf(n, N, alpha, beta); + } else { + lp = stan::math::beta_binomial_lccdf(n, N, alpha, beta); + } + wrt = {alpha, beta}; + } else if (tag == "ST") { + var y = a[0], nu = a[1], mu = a[2], sigma = a[3]; + lp = stan::math::student_t_lpdf(y, nu, mu, sigma); + wrt = {y, nu, mu, sigma}; + } else if (tag == "MST" || tag == "MSTF") { + // dimension 2, mu = 0, identity scale: y1, y2, nu + Matrix y(2); + y << a[0], a[1]; + Eigen::VectorXd mu = Eigen::VectorXd::Zero(2); + Eigen::MatrixXd scale = Eigen::MatrixXd::Identity(2, 2); + var nu = a[2]; + if (tag == "MST") { + lp = stan::math::multi_student_t_cholesky_lpdf(y, nu, mu, scale); + } else { + lp = stan::math::multi_student_t_lpdf(y, nu, mu, scale); + } + wrt = {y(0), y(1), nu}; + } else if (tag == "LBETA") { + var d = a[0], x = a[1]; + lp = stan::math::lbeta(d, x); + wrt = {d, x}; + } else if (tag == "YS") { + var alpha = a[1]; + lp = stan::math::yule_simon_lpmf(to_int(a[0]), alpha); + wrt = {alpha}; + } else { + var alpha = a[1]; + const int n = to_int(a[0]); + if (tag == "YSC") { + lp = stan::math::yule_simon_cdf(n, alpha); + } else if (tag == "YSLC") { + lp = stan::math::yule_simon_lcdf(n, alpha); + } else { + lp = stan::math::yule_simon_lccdf(n, alpha); + } + wrt = {alpha}; + } + lp.grad(); + value = lp.val(); + grads.clear(); + for (const var& v : wrt) { + grads.push_back(v.adj()); + } + stan::math::recover_memory(); +} + +// clang-format off +const std::vector test_cases = { + // neg_binomial_lpmf(n | alpha, beta): n, alpha, beta + {"NB", {0x1.8000000000000p+1, 0x1.c6bf526340000p+49, 0x1.2f2a36ecd5555p+48}, -1.4959226032237274, {1.3124999999999966e-30, 5.6249999999999838e-31}}, + {"NB", {0x1.0000000000000p+0, 0x1.c6bf526340000p+49, 0x1.2f2a36ecd5555p+48}, -1.9013877113318889, {-1.9999999999999957e-15, 5.9999999999999829e-15}}, + {"NB", {0x1.9000000000000p+5, 0x1.c6bf526340000p+49, 0x1.2309ce5400000p+44}, -2.8766166803657541, {2.4999999999998758e-29, 0.0}}, + {"NB", {0x0.0p+0, 0x1.9000000000000p+6, 0x1.0aaaaaaaaaaabp+5}, -2.9558802241544401, {-2.9558802241544401e-2, 8.7378640776699017e-2}}, + // neg_binomial_2_lpmf(n | mu, phi): n, mu, phi + {"NB2", {0x1.c800000000000p+5, 0x1.8000000000000p+1, 0x1.2a05f20000000p+33}, -1.1677494780996510e+2, {1.7999999994600000e+1, -1.4294999940379000e-17}}, + {"NB2", {0x1.4000000000000p+2, 0x1.9000000000000p+5, 0x1.2a05f20000000p+33}, -3.5227376614641316e+1, {-8.9999999550000002e-1, -1.0099999929136667e-17}}, + {"NB2", {0x1.0000000000000p+0, 0x1.8000000000000p+1, 0x1.2a05f20000000p+33}, -1.9013877111818903, {-6.6666666646666667e-1, -1.4999999991000000e-20}}, + {"NB2", {0x0.0p+0, 0x1.8000000000000p+1, 0x1.9000000000000p+6}, -2.9558802241544403, {-9.7087378640776699e-1, -4.3258864931139302e-4}}, + // neg_binomial_2_log_lpmf(n | eta, phi): n, eta, phi + {"NB2L", {0x1.8000000000000p+1, 0x1.193ea7aad030bp+0, 0x1.c6bf526340000p+49}, -1.4959226032237274, {-2.7213891705004510e-16, 1.4999999999999960e-30}}, + {"NB2L", {0x1.4000000000000p+2, 0x1.f4bd2b7ac1bafp+1, 0x1.c6bf526340000p+49}, -3.5227376715640301e+1, {-4.4999999999997745e+1, -1.0099999999999289e-27}}, + {"NB2L", {0x1.0000000000000p+1, -0x1.6d3c324e13f50p-2, 0x1.37807ed5e8000p+50}, -2.1064970684374103, {1.2999999999999994, 8.2582982577654707e-32}}, + {"NB2L", {0x0.0p+0, 0x1.193ea7aad030bp+0, 0x1.9000000000000p+6}, -2.9558802241544405, {-2.9126213592233012, -4.3258864931139310e-4}}, + // neg_binomial_2_log_glm_lpmf, one observation: y, x, alpha, beta, phi + {"GLM1", {0x0.0p+0, 0x1.0000000000000p-1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.c6bf526340000p+49}, -3.9550767229205693, {-3.9550767229205615, -1.9775383614602807, -7.8213159420940446e-30}}, + {"GLM1", {0x1.4000000000000p+2, 0x1.0000000000000p-1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.c6bf526340000p+49}, -1.8675684657026251, {1.0449232770794187, 5.2246163853970937e-1, 1.9540676725087929e-30}}, + {"GLM1", {0x0.0p+0, 0x1.0000000000000p-1, -0x1.999999999999ap-3, 0x1.8000000000000p-1, 0x1.37807ed5e8000p+50}, -1.1912462166123576, {-1.1912462166123571, -5.9562310830617854e-1, -3.7803493755481261e-31}}, + {"GLM1", {0x0.0p+0, 0x1.0000000000000p-1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.9000000000000p+6}, -3.8788665246722319, {-3.8046018026251338, -1.9023009013125669, -7.4264722047098149e-4}}, + // neg_binomial_2_log_glm_lpmf, five observations: y1..y5, x1..x5, alpha, beta, phi + {"GLM5", {0x1.0000000000000p+1, 0x1.0000000000000p+2, 0x1.8000000000000p+1, 0x1.4000000000000p+2, 0x1.8000000000000p+2, 0x1.0000000000000p-1, -0x1.0000000000000p-2, 0x1.4000000000000p+0, 0x0.0p+0, 0x1.0000000000000p+1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.c6bf526340000p+49}, -1.3267966555421410e+1, {-8.0507631204932393, -1.8705862362560038e+1, -2.2918189242185528e-29}}, + {"GLM5", {0x0.0p+0, 0x1.0000000000000p+0, 0x1.4000000000000p+2, 0x1.8000000000000p+1, 0x1.c800000000000p+5, 0x1.0000000000000p-1, -0x1.0000000000000p-2, 0x1.4000000000000p+0, 0x0.0p+0, 0x1.0000000000000p+1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.c6bf526340000p+49}, -5.5025862739499810e+1, {3.7949236879506146e+1, 8.5544137637438705e+1, -9.8183556706826074e-28}}, + {"GLM5", {0x1.0000000000000p+1, 0x1.0000000000000p+2, 0x1.8000000000000p+1, 0x1.4000000000000p+2, 0x1.8000000000000p+2, 0x1.0000000000000p-1, -0x1.0000000000000p-2, 0x1.4000000000000p+0, 0x0.0p+0, 0x1.0000000000000p+1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.d1a94a2000000p+39}, -1.3267966555398514e+1, {-8.0507631203930684, -1.8705862362370543e+1, -2.2918189241713975e-23}}, + {"GLM5", {0x0.0p+0, 0x1.0000000000000p+0, 0x1.4000000000000p+2, 0x1.8000000000000p+1, 0x1.c800000000000p+5, 0x1.0000000000000p-1, -0x1.0000000000000p-2, 0x1.4000000000000p+0, 0x0.0p+0, 0x1.0000000000000p+1, 0x1.0000000000000p+0, 0x1.8000000000000p-1, 0x1.9000000000000p+6}, -4.7240731087544224e+1, {3.3378922648644250e+1, 7.6036039699167587e+1, -6.2120571346832957e-2}}, + // beta_neg_binomial_lpmf(n | r, alpha, beta): n, r, alpha, beta + {"BNB", {0x1.4000000000000p+2, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -2.6841050769298924, {-3.5297736803318772e-1, 1.6666666666666617e-15, -8.3333333333333083e-16}}, + {"BNB", {0x1.4000000000000p+3, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -2.6389577451069742, {6.9616704560883204e-2, 1.6666666666666591e-30, 4.1666666666666470e-31}}, + {"BNB", {0x1.0000000000000p+0, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -4.2890886390146075, {-8.9861228866810702e-1, 2.9999999999999917e-15, -1.4999999999999983e-15}}, + {"BNB", {0x0.0p+0, 0x1.9000000000000p+6, 0x1.9000000000000p+6, 0x1.9000000000000p+6}, -5.2468933188272313e+1, {-4.0629959884472565e-1, 2.8935383163709856e-1, -4.0629959884472565e-1}}, + // dirichlet_multinomial_lpmf(ns | alpha): n1, n2, n3, alpha1, alpha2, alpha3 + {"DM", {0x0.0p+0, 0x1.0000000000000p+0, 0x0.0p+0, 0x1.6bcc41e900000p+47, 0x1.10d9316ec0000p+48, 0x1.c6bf526340000p+48}, -1.2039728043259360, {-1.0000000000000000e-15, 2.3333333333333333e-15, -1.0000000000000000e-15}}, + {"DM", {0x1.0000000000000p+0, 0x1.0000000000000p+1, 0x1.0000000000000p+1, 0x1.6bcc41e900000p+47, 0x1.10d9316ec0000p+48, 0x1.c6bf526340000p+48}, -2.0024805005437123, {9.9999999999999700e-30, 1.6666666666666656e-15, -9.9999999999999400e-16}}, + {"DM", {0x1.4000000000000p+2, 0x0.0p+0, 0x0.0p+0, 0x1.6bcc41e900000p+47, 0x1.10d9316ec0000p+48, 0x1.c6bf526340000p+48}, -8.0471895621704619, {1.9999999999999760e-14, -4.9999999999999900e-15, -4.9999999999999900e-15}}, + {"DM", {0x0.0p+0, 0x1.0000000000000p+0, 0x0.0p+0, 0x1.4000000000000p+4, 0x1.e000000000000p+4, 0x1.9000000000000p+5}, -1.2039728043259360, {-1.0000000000000000e-2, 2.3333333333333333e-2, -1.0000000000000000e-2}}, + // lkj_corr_lpdf(Omega | eta): Omega(1, 0), Omega(2, 0), Omega(2, 1), eta + {"LKJ", {0x0.0p+0, 0x0.0p+0, 0x0.0p+0, 0x1.c6bf526340000p+49}, 5.0091069763591928e+1, {1.4999999999999999e-15}}, + {"LKJ", {0x0.0p+0, 0x0.0p+0, 0x0.0p+0, 0x1.772aa3f848000p+51}, 5.1881953466300579e+1, {4.5454545454545453e-16}}, + {"LKJ", {0x0.0p+0, 0x0.0p+0, 0x0.0p+0, 0x1.802ba9f400000p+41}, 4.1520320547827412e+1, {4.5454545454544307e-13}}, + {"LKJ", {0x1.3333333333333p-2, -0x1.999999999999ap-3, 0x1.999999999999ap-4, 0x1.9000000000000p+6}, -1.1130679230833304e+1, {-1.4988714303399179e-1}}, + // lkj_corr_cholesky_lpdf(L | eta): L(1, 0), L(1, 1), L(2, 0), L(2, 1), L(2, 2), eta + {"LKJC", {0x0.0p+0, 0x1.0000000000000p+0, 0x0.0p+0, 0x0.0p+0, 0x1.0000000000000p+0, 0x1.c6bf526340000p+49}, 5.0091069763591928e+1, {1.4999999999999999e-15}}, + {"LKJC", {0x0.0p+0, 0x1.0000000000000p+0, 0x0.0p+0, 0x0.0p+0, 0x1.0000000000000p+0, 0x1.772aa3f848000p+51}, 5.1881953466300579e+1, {4.5454545454545453e-16}}, + {"LKJC", {0x0.0p+0, 0x1.0000000000000p+0, 0x0.0p+0, 0x0.0p+0, 0x1.0000000000000p+0, 0x1.802ba9f400000p+41}, 4.1520320547827412e+1, {4.5454545454544307e-13}}, + {"LKJC", {0x1.3333333333333p-2, 0x1.e86ab810ea912p-1, -0x1.999999999999ap-3, 0x1.57808173fc2ddp-3, 0x1.ee40264218109p-1, 0x1.9000000000000p+6}, -1.1177834570568931e+1, {-1.4988714303399185e-1}}, + // yule_simon_{cdf,lcdf,lccdf}(n | alpha): n, alpha + {"YSC", {0x1.0000000000000p+0, 0x1.9000000000000p+6}, 9.9009900990099010e-1, {9.8029604940692089e-5}}, + {"YSLC", {0x1.0000000000000p+0, 0x1.9000000000000p+6}, -9.9503308531680828e-3, {9.9009900990099010e-5}}, + {"YSLCC", {0x1.0000000000000p+1, 0x1.c6bf526340000p+49}, -6.8384405609261428e+1, {-1.9999999999999970e-15}}, + {"YSLCC", {0x1.0000000000000p+0, 0x1.c6bf526340000p+49}, -3.4538776394910686e+1, {-9.9999999999999900e-16}}, + {"YSLCC", {0x1.4000000000000p+2, 0x1.c6bf526340000p+49}, -1.6790639023177140e+2, {-4.9999999999999850e-15}}, + {"YSLCC", {0x1.0000000000000p+0, 0x1.9000000000000p+6}, -4.6151205168412595, {-9.9009900990099010e-3}}, + // beta_binomial_{cdf,lcdf,lccdf}(n | N, alpha, beta): n, N, alpha, beta + {"BBC", {0x1.d000000000000p+4, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, 5.2836214512816302e-1, {-1.8707495486824231e-15, 6.2358318289414104e-16}}, + {"BBC", {0x1.4000000000000p+2, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, 1.9080827003311692e-9, {-4.6543958285535218e-23, 1.5514652761844824e-23}}, + {"BBC", {0x1.0000000000000p+0, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, 9.6433473051139671e-14, {-2.7266564505209334e-27, 9.0888548350696083e-28}}, + {"BBC", {0x0.0p+0, 0x1.d400000000000p+6, 0x1.9000000000000p+6, 0x1.9000000000000p+6}, 2.0588830493356332e-26, {-9.5019176900109660e-27, 6.5044482281184328e-27}}, + {"BBLC", {0x0.0p+0, 0x1.d400000000000p+6, 0x1.3880000000000p+13, 0x1.3880000000000p+13}, -8.0760883216610174e+1, {-5.8331005942734647e-3, 5.7995618892517663e-3}}, + {"BBLC", {0x1.4000000000000p+2, 0x1.d400000000000p+6, 0x1.3880000000000p+13, 0x1.3880000000000p+13}, -6.1834457741779739e+1, {-5.3377886925755501e-3, 5.3097367250871483e-3}}, + {"BBLC", {0x1.4000000000000p+2, 0x1.d400000000000p+6, 0x1.e848000000000p+19, 0x1.e848000000000p+19}, -6.2113769736395483e+1, {-5.3543714610347851e-5, 5.3540877034666121e-5}}, + {"BBLC", {0x1.c800000000000p+5, 0x1.d400000000000p+6, 0x1.9000000000000p+6, 0x1.9000000000000p+6}, -8.1708867859107758e-1, {-3.8676835548171366e-2, 3.8433660832149159e-2}}, + {"BBLCC", {0x1.d000000000000p+4, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, -7.5154384451605549e-1, {3.9664957538041156e-15, -1.3221652512680385e-15}}, + {"BBLCC", {0x1.4000000000000p+2, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, -1.9080827021515590e-9, {4.6543958374344940e-23, -1.5514652791448065e-23}}, + {"BBLCC", {0x1.0000000000000p+0, 0x1.d400000000000p+6, 0x1.c6bf526340000p+49, 0x1.550f7dca70000p+51}, -9.6433473051144321e-14, {2.7266564505211963e-27, -9.0888548350704848e-28}}, + {"BBLCC", {0x0.0p+0, 0x1.d400000000000p+6, 0x1.9000000000000p+6, 0x1.9000000000000p+6}, -2.0588830493356332e-26, {9.5019176900109660e-27, -6.5044482281184328e-27}}, + // beta_neg_binomial_{cdf,lccdf}(n | r, alpha, beta): n, r, alpha, beta + {"BNBC", {0x0.0p+0, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, 4.1152263374485871e-3, {-4.5210382249716626e-3, 1.3717421124828587e-17, -6.8587105624143073e-18}}, + {"BNBC", {0x1.4000000000000p+2, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, 2.1312808006909550e-1, {-1.1257829350011265e-1, 4.5521516029060565e-16, -2.2760758014530300e-16}}, + {"BNBC", {0x1.0000000000000p+0, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, 1.7832647462277188e-2, {-1.6847681416578131e-2, 5.4869684499314275e-17, -2.7434842249657186e-17}}, + {"BNBC", {0x0.0p+0, 0x1.4000000000000p+3, 0x1.4000000000000p+3, 0x1.4000000000000p+3}, 4.6119797244235025e-3, {-1.9089636233770314e-3, 1.4059955145634730e-3, -1.9089636233770314e-3}}, + {"BNBLCC", {0x0.0p+0, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -4.1237171838620870e-3, {4.5397202011079093e-3, -1.3774104683195648e-17, 6.8870523415978377e-18}}, + {"BNBLCC", {0x1.4000000000000p+2, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -2.3968978849662940e-1, {1.4307067090410113e-1, -5.7851239669421455e-16, 2.8925619834710749e-16}}, + {"BNBLCC", {0x1.4000000000000p+3, 0x1.4000000000000p+2, 0x1.c6bf526340000p+49, 0x1.c6bf526340000p+50}, -9.0618005875034648e-1, {3.6151768966033925e-1, -1.7679265277287130e-15, 8.8396326386435604e-16}}, + {"BNBLCC", {0x0.0p+0, 0x1.4000000000000p+3, 0x1.4000000000000p+3, 0x1.4000000000000p+3}, -4.6226477159237375e-3, {1.9178085173744892e-3, -1.4125099819608222e-3, 1.9178085173744892e-3}}, + // student_t_lpdf(y | nu, mu, sigma): y, nu, mu, sigma + {"ST", {0x1.0000000000000p+0, 0x1.6bcc41e900000p+46, 0x0.0p+0, 0x1.0000000000000p+0}, -1.4189385332046777, {-1.0000000000000000, 4.9999999999999833e-29, 1.0000000000000000, 2.6727647100921956e-51}}, + {"ST", {-0x1.cd1e504efb30cp+1, 0x1.7cc4b890abebfp+45, 0x1.64703afcf3380p-4, 0x1.c43477d4ae376p+3}, -3.6014210338578538, {1.8475570774651872e-2, 1.0330563009181428e-28, -1.8475570774651872e-2, -6.5940664388226970e-2}}, + {"ST", {0x1.8000000000000p+1, 0x1.6bcc41e900000p+46, 0x0.0p+0, 0x1.0000000000000p+1}, -2.7370857137646191, {-7.4999999999999062e-1, 1.0937500000001266e-29, 7.4999999999999062e-1, 6.2499999999998594e-1}}, + {"ST", {-0x1.c8dd60359f470p-1, 0x1.da533b967d6d3p-4, -0x1.572a29ad6231cp-1, 0x1.9d9944276e11dp-3}, -1.6062891011707614, {4.5853946438927804, 6.8870497923078714, -4.5853946438927804, 9.0518700772802596e-2}}, + // multi_student_t_cholesky_lpdf (MST) and multi_student_t_lpdf (MSTF), + // dimension 2, mu = 0, identity scale: y1, y2, nu + {"MST", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.d1a94a2000000p+39}, -2.8378770664103455, {-1.0000000000000000, -1.0000000000000000, 9.9999999999866667e-25}}, + {"MST", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.74876e8000000p+36}, -2.8378770664193455, {-1.0000000000000000, -1.0000000000000000, 9.9999999998666667e-23}}, + {"MST", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.dcd6500000000p+29}, -2.8378770674093455, {-1.0000000000000000, -1.0000000000000000, 9.9999999866666667e-19}}, + {"MST", {-0x1.03f559f1746d0p-1, 0x1.22539ac62abe0p-3, 0x1.681d9f093b62ap-6}, -4.4798146609821203, {3.4235928669135459, -9.5588370652474091e-1, 4.1318417102814693e+1}}, + {"MSTF", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.d1a94a2000000p+39}, -2.8378770664103455, {-1.0000000000000000, -1.0000000000000000, 9.9999999999866667e-25}}, + {"MSTF", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.74876e8000000p+36}, -2.8378770664193455, {-1.0000000000000000, -1.0000000000000000, 9.9999999998666667e-23}}, + {"MSTF", {0x1.0000000000000p+0, 0x1.0000000000000p+0, 0x1.dcd6500000000p+29}, -2.8378770674093455, {-1.0000000000000000, -1.0000000000000000, 9.9999999866666667e-19}}, + {"MSTF", {-0x1.03f559f1746d0p-1, 0x1.22539ac62abe0p-3, 0x1.681d9f093b62ap-6}, -4.4798146609821203, {3.4235928669135459, -9.5588370652474091e-1, 4.1318417102814693e+1}}, + // yule_simon_lpmf(n | alpha): n, alpha + {"YS", {0x1.0000000000000p+0, 0x1.d1a94a2000000p+39}, -9.9999999999950000e-13, {9.9999999999900000e-25}}, + {"YS", {0x1.0000000000000p+0, 0x1.74876e8000000p+36}, -9.9999999999500000e-12, {9.9999999999000000e-23}}, + {"YS", {0x1.0000000000000p+0, 0x1.2a05f20000000p+33}, -9.9999999995000000e-11, {9.9999999990000000e-21}}, + {"YS", {0x1.0000000000000p+0, 0x1.cc271540e654dp-2}, -1.1710409729900522, {1.5353924179257093}}, + // lbeta(d, x) with var arguments: the derivative digamma(x) - digamma(d + x) + {"LBETA", {0x1.8000000000000p+0, 0x1.6bcc41e900000p+46}, -4.8475069190510208e+1, {-3.2199701327938073e+1, -1.4999999999999962e-14}}, + {"LBETA", {0x1.0000000000000p+0, 0x1.6bcc41e900000p+46}, -3.2236191301916640e+1, {-3.2813406966818177e+1, -1.0000000000000000e-14}}, + {"LBETA", {0x1.0000000000000p-1, 0x1.6bcc41e900000p+46}, -1.5545730708033618e+1, {-3.4199701327938063e+1, -5.0000000000000125e-15}}, + {"LBETA", {0x1.ba1a7a7909c76p-7, 0x1.f4f68d902030cp-6}, 4.6705195755596318, {-5.1474613897767773e+1, -1.0033990131067132e+1}}, +}; +// clang-format on +} // namespace large_shapes_test_internal + +TEST(ProbDistributions, large_shapes_value_and_log_scale_gradients) { + using large_shapes_test_internal::eval; + using large_shapes_test_internal::log_scale_args; + using large_shapes_test_internal::test_cases; + using large_shapes_test_internal::tolerances; + for (const auto& t : test_cases) { + const std::string tag(t.tag); + double value; + std::vector grads; + eval(tag, t.args, value, grads); + const auto tol = tolerances(tag); + EXPECT_NEAR(value, t.value, tol.first * std::max(1.0, std::fabs(t.value))) + << tag << " value, last argument " << t.args.back(); + const std::vector log_args = log_scale_args(tag); + ASSERT_EQ(grads.size(), t.grads.size()) << tag; + for (size_t j = 0; j < grads.size(); ++j) { + const double scale = log_args[j] >= 0 ? t.args[log_args[j]] : 1.0; + const double got = scale * grads[j]; + const double expected = scale * t.grads[j]; + EXPECT_NEAR(got, expected, + tol.second * std::max(1.0, std::fabs(expected))) + << tag << " gradient " << j << ", last argument " << t.args.back(); + } + } +} + +TEST(ProbDistributions, lkj_corr_eta_gradient_at_one) { + // At eta == 1.0 exactly, develop returned the constant without its + // derivative, so d/deta lost sum_k psi(eta + (K - 1) / 2) - psi(eta + + // (K - 1 - k) / 2). eta = exp(0) = 1 is the value at the default + // initialization of an unconstrained parameter. Reference from mpmath. + Eigen::MatrixXd omega(3, 3); + omega << 1.0, 0.3, -0.2, 0.3, 1.0, 0.1, -0.2, 0.1, 1.0; + stan::math::var eta = 1.0; + stan::math::var lp = stan::math::lkj_corr_lpdf(omega, eta); + lp.grad(); + EXPECT_NEAR(lp.val(), -1.596312591138855, 1e-14); + EXPECT_NEAR(eta.adj(), 1.2214197179296566, 1e-14); + stan::math::recover_memory(); +} + +TEST(ProbDistributions, large_shapes_propto_dropped_terms) { + // full - propto must be the sum of the terms that propto drops, computed + // here from their formulas. Some of these terms are inside a function + // call: lgamma(n) in lbeta(n, alpha + 1), lgamma(y + 1) in lchoose. + using stan::math::var; + using std::lgamma; + using std::log; + + // student_t_lpdf, parameter nu (both forms of the nu terms: nu < 40 and + // nu >= 40), data y, mu, and sigma = 1: drops log(sqrt(pi)) per term + { + Eigen::VectorXd y(3); + y << 0.3, -1.2, 2.0; + Eigen::VectorXd nu_d(3); + nu_d << 3.5, 55.0, 500.0; + Eigen::Matrix nu = nu_d.cast(); + var full = stan::math::student_t_lpdf(y, nu, 0.1, 1.0); + var prop = stan::math::student_t_lpdf(y, nu, 0.1, 1.0); + EXPECT_NEAR(full.val() - prop.val(), -3 * stan::math::LOG_SQRT_PI, 1e-13); + stan::math::recover_memory(); + } + // yule_simon_lpmf, parameter alpha: drops lgamma(n) + { + std::vector n{1, 3, 7}; + Eigen::VectorXd alpha_d(3); + alpha_d << 2.5, 0.7, 30.0; + Eigen::Matrix alpha = alpha_d.cast(); + var full = stan::math::yule_simon_lpmf(n, alpha); + var prop = stan::math::yule_simon_lpmf(n, alpha); + EXPECT_NEAR(full.val() - prop.val(), + lgamma(1.0) + lgamma(3.0) + lgamma(7.0), 1e-13); + stan::math::recover_memory(); + } + // beta_neg_binomial_lpmf, parameters r, alpha, beta: drops -lgamma(n + 1) + { + std::vector n{3, 5}; + var r = 2.5; + var a = 3.5; + var b = 1.5; + var full = stan::math::beta_neg_binomial_lpmf(n, r, a, b); + var prop = stan::math::beta_neg_binomial_lpmf(n, r, a, b); + EXPECT_NEAR(full.val() - prop.val(), -(lgamma(4.0) + lgamma(6.0)), 1e-13); + stan::math::recover_memory(); + } + // neg_binomial_2_log_glm_lpmf: drops -lgamma(y + 1); y theta if x, alpha + // and beta are data; phi log(phi) - lgamma(phi) + lgamma(y + phi) if phi + // is data + { + std::vector y{3, 5, 2}; + Eigen::MatrixXd x(3, 2); + x << 0.5, -0.2, 1.1, 0.3, -0.4, 0.8; + Eigen::VectorXd beta_d(2); + beta_d << 0.4, -0.3; + const double alpha = 0.3; + const double phi_d = 3.5; + Eigen::VectorXd theta = (x * beta_d).array() + alpha; + double lgamma_y1 = 0; + double y_theta = 0; + double phi_terms = 0; + for (size_t i = 0; i < y.size(); ++i) { + lgamma_y1 += lgamma(y[i] + 1.0); + y_theta += y[i] * theta(i); + phi_terms += phi_d * log(phi_d) - lgamma(phi_d) + lgamma(y[i] + phi_d); + } + Eigen::Matrix beta = beta_d.cast(); + var phi = phi_d; + // parameters beta and phi + var full = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta, + phi); + var prop + = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta, phi); + EXPECT_NEAR(full.val() - prop.val(), -lgamma_y1, 1e-13); + // parameter phi; x, alpha and beta are data + full = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta_d, + phi); + prop = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta_d, + phi); + EXPECT_NEAR(full.val() - prop.val(), -lgamma_y1 + y_theta, 1e-13); + // parameter beta; phi is data + full = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta, + phi_d); + prop = stan::math::neg_binomial_2_log_glm_lpmf(y, x, alpha, beta, + phi_d); + EXPECT_NEAR(full.val() - prop.val(), -lgamma_y1 + phi_terms, 1e-13); + stan::math::recover_memory(); + } +} From 9b9cc5192f5eeec53ff2998d757d32dd2dbb0e25 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 15:11:47 +0300 Subject: [PATCH 13/15] Port the large-shape repairs to OpenCL --- .../kernel_generator/elt_function_cl.hpp | 104 ++++++++++ .../kernels/device_functions/digamma_diff.hpp | 115 +++++++++++ .../device_functions/log_beta_ratio.hpp | 168 ++++++++++++++++ .../kernels/neg_binomial_2_log_glm_lpmf.hpp | 82 ++++---- stan/math/opencl/prim/beta_binomial_lpmf.hpp | 17 +- .../prim/neg_binomial_2_log_glm_lpmf.hpp | 41 ++-- .../opencl/prim/neg_binomial_2_log_lpmf.hpp | 6 +- stan/math/opencl/prim/neg_binomial_2_lpmf.hpp | 5 +- stan/math/opencl/prim/neg_binomial_lpmf.hpp | 9 +- stan/math/opencl/prim/student_t_lpdf.hpp | 46 ++++- .../device_functions/digamma_diff_test.cpp | 142 +++++++++++++ .../kernel_generator/elt_function_cl_test.cpp | 47 +++++ .../opencl/rev/beta_binomial_lpmf_test.cpp | 91 +++++++++ .../rev/neg_binomial_2_log_glm_lpmf_test.cpp | 186 ++++++++++++++++++ .../rev/neg_binomial_2_log_lpmf_test.cpp | 73 +++++++ .../opencl/rev/neg_binomial_2_lpmf_test.cpp | 64 ++++++ .../opencl/rev/neg_binomial_lpmf_test.cpp | 66 +++++++ .../math/opencl/rev/student_t_lpdf_test.cpp | 63 ++++++ 18 files changed, 1258 insertions(+), 67 deletions(-) create mode 100644 stan/math/opencl/kernels/device_functions/digamma_diff.hpp create mode 100644 stan/math/opencl/kernels/device_functions/log_beta_ratio.hpp create mode 100644 test/unit/math/opencl/device_functions/digamma_diff_test.cpp diff --git a/stan/math/opencl/kernel_generator/elt_function_cl.hpp b/stan/math/opencl/kernel_generator/elt_function_cl.hpp index df349a7b483..d5d42c38b2b 100644 --- a/stan/math/opencl/kernel_generator/elt_function_cl.hpp +++ b/stan/math/opencl/kernel_generator/elt_function_cl.hpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include #include @@ -14,6 +15,7 @@ #include #include #include +#include #include #include #include @@ -410,6 +412,108 @@ const std::vector lbeta_::includes{ stan::math::opencl_kernels::lgamma_stirling_device_function, stan::math::opencl_kernels::lgamma_stirling_diff_device_function, stan::math::opencl_kernels::lbeta_device_function}; + +ADD_BINARY_FUNCTION_WITH_INCLUDES( + digamma_diff, stan::math::opencl_kernels::digamma_device_function, + stan::math::opencl_kernels::digamma_diff_device_function) + +/** + * Represents lbeta(alpha + n, beta + m) - lbeta(alpha, beta) in kernel + * generator expressions. See the device function stan_log_beta_ratio. + * @tparam T1 type of the first shape + * @tparam T2 type of the second shape + * @tparam T3 type of the first count + * @tparam T4 type of the second count + */ +template +class log_beta_ratio_ : public elt_function_cl, + double, T1, T2, T3, T4> { + using base = elt_function_cl, double, T1, T2, + T3, T4>; + using base::arguments_; + + public: + using base::cols; + using base::rows; + static const std::vector includes; + explicit log_beta_ratio_(T1&& alpha, T2&& beta, T3&& n, T4&& m) + : base("stan_log_beta_ratio", std::forward(alpha), + std::forward(beta), std::forward(n), std::forward(m)) { + const std::array arg_rows{{this->template get_arg<0>().rows(), + this->template get_arg<1>().rows(), + this->template get_arg<2>().rows(), + this->template get_arg<3>().rows()}}; + const std::array arg_cols{{this->template get_arg<0>().cols(), + this->template get_arg<1>().cols(), + this->template get_arg<2>().cols(), + this->template get_arg<3>().cols()}}; + for (int i = 0; i < 4; i++) { + for (int j = i + 1; j < 4; j++) { + if (arg_rows[i] != base::dynamic && arg_rows[j] != base::dynamic) { + check_size_match("log_beta_ratio", "Rows of ", "an argument", + arg_rows[i], "rows of ", "another argument", + arg_rows[j]); + } + if (arg_cols[i] != base::dynamic && arg_cols[j] != base::dynamic) { + check_size_match("log_beta_ratio", "Columns of ", "an argument", + arg_cols[i], "columns of ", "another argument", + arg_cols[j]); + } + } + } + } + inline auto deep_copy() const { + auto&& arg1_copy = this->template get_arg<0>().deep_copy(); + auto&& arg2_copy = this->template get_arg<1>().deep_copy(); + auto&& arg3_copy = this->template get_arg<2>().deep_copy(); + auto&& arg4_copy = this->template get_arg<3>().deep_copy(); + return log_beta_ratio_, + std::remove_reference_t, + std::remove_reference_t, + std::remove_reference_t>{ + std::move(arg1_copy), std::move(arg2_copy), std::move(arg3_copy), + std::move(arg4_copy)}; + } + inline std::pair extreme_diagonals() const { + return {-rows() + 1, cols() - 1}; + } +}; + +/** + * Returns lbeta(alpha + n, beta + m) - lbeta(alpha, beta) for shapes + * alpha, beta > 0 and integer counts n, m >= 0, without the cancellation of + * the two lbeta values for large shapes. This is the kernel generator + * version of stan::math::internal::log_beta_ratio(). + * @tparam T1 type of the first shape + * @tparam T2 type of the second shape + * @tparam T3 type of the first count + * @tparam T4 type of the second count + * @param alpha first shape + * @param beta second shape + * @param n first count + * @param m second count + * @return expression for the difference of the two lbeta values + */ +template * = nullptr, + require_any_not_stan_scalar_t* = nullptr> +inline log_beta_ratio_, as_operation_cl_t, + as_operation_cl_t, as_operation_cl_t> +log_beta_ratio(T1&& alpha, T2&& beta, T3&& n, T4&& m) { + return log_beta_ratio_, as_operation_cl_t, + as_operation_cl_t, as_operation_cl_t>( + as_operation_cl(std::forward(alpha)), + as_operation_cl(std::forward(beta)), + as_operation_cl(std::forward(n)), + as_operation_cl(std::forward(m))); +} + +template +const std::vector log_beta_ratio_::includes{ + stan::math::opencl_kernels::lgamma_stirling_device_function, + stan::math::opencl_kernels::lgamma_stirling_diff_device_function, + stan::math::opencl_kernels::lbeta_device_function, + stan::math::opencl_kernels::log_beta_ratio_device_function}; ADD_BINARY_FUNCTION_WITH_INCLUDES( log_inv_logit_diff, opencl_kernels::log1p_exp_device_function, opencl_kernels::log1m_exp_device_function, diff --git a/stan/math/opencl/kernels/device_functions/digamma_diff.hpp b/stan/math/opencl/kernels/device_functions/digamma_diff.hpp new file mode 100644 index 00000000000..4211b91297a --- /dev/null +++ b/stan/math/opencl/kernels/device_functions/digamma_diff.hpp @@ -0,0 +1,115 @@ +#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_DIGAMMA_DIFF_HPP +#define STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_DIGAMMA_DIFF_HPP +#ifdef STAN_OPENCL + +#include +#include + +namespace stan { +namespace math { +namespace opencl_kernels { + +// \cond +static constexpr const char* digamma_diff_device_function + = "\n" + "#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_DIGAMMA_DIFF\n" + "#define " + "STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_DIGAMMA_DIFF\n" STRINGIFY( + // \endcond + /** \ingroup opencl_device_functions + * Return the difference of the digamma function at two arguments + * that differ by a nonnegative offset, psi(x + d) - psi(x). + * + * The plain difference of two digamma calls loses all accuracy + * when x is large. This function has a relative error of a few + * ulp for all x > 0 and d >= 0. It uses the same method as + * stan::math::digamma_diff(): for an integer d from 0 to 8 the + * sum of 1 / (x + j), j < d; for x < 10 and d >= 10 the plain + * difference, which does not cancel much there; otherwise shift x + * up to y >= 10 with the recurrence psi(y + 1) = psi(y) + 1 / y, + * then use the asymptotic expansion of psi(y + d) - psi(y), with + * the differences of the powers of 1 / y^2 and 1 / (y + d)^2 + * formed without cancellation. Needs the digamma device function. + * + * @param x first argument, positive + * @param d offset, nonnegative + * @return psi(x + d) - psi(x), or NaN if x is not positive or d + * is negative + */ + double digamma_diff(double x, double d) { + if (isnan(x) || isnan(d) || !(x > 0) || d < 0) { + return NAN; + } + if (isinf(d)) { + return INFINITY; + } + if (isinf(x)) { + return 0.0; + } + // a count d from 0 to 8: sum of the positive terms 1 / (x + j), + // from the smallest (exactly 1 / x for d = 1) + if (d <= 8.0 && d == floor(d)) { + double sum = 0.0; + for (int j = (int)d - 1; j >= 0; --j) { + sum += 1.0 / (x + j); + } + return sum; + } + // x < 10 and d >= 10: the plain difference does not cancel much + if (x < 10.0 && d >= 10.0) { + return digamma(x + d) - digamma(x); + } + // B_{2i} / (2i), i = 1..8 + const double coeffs[8] = { + 1.0 / 12.0, -1.0 / 120.0, 1.0 / 252.0, -1.0 / 240.0, + 1.0 / 132.0, -691.0 / 32760.0, 1.0 / 12.0, -3617.0 / 8160.0}; + // each shift term is d / (y (y + d)) = 1 / y - 1 / (y + d); for + // d < y it is formed as (d / (y + d)) / y, which neither cancels + // nor overflows; for d >= y the parts 1 / y are summed separately + // and added last, so that for small x the dominant 1 / x keeps its + // correct rounding (for d = 1 the result is exactly 1 / x) + double inv_sum = 0.0; + double shift_sum = 0.0; + double y = x; + while (y < 10.0) { + if (d >= y) { + inv_sum += 1.0 / y; + shift_sum -= 1.0 / (y + d); + } else { + shift_sum += (d / (y + d)) / y; + } + y += 1.0; + } + const double y_plus_d = y + d; + const double d_frac = d / y_plus_d; + // the square of the inverse underflows to 0 where the inverse of + // the square would overflow + const double inv_y = 1.0 / y; + const double inv_y_plus_d = 1.0 / y_plus_d; + const double u = inv_y * inv_y; + const double v = inv_y_plus_d * inv_y_plus_d; + // u - v = u (d / (y + d)) (1 + y / (y + d)), without cancellation + const double u_minus_v = u * d_frac * (1.0 + y / y_plus_d); + // u^i - v^i = (u - v) h_i with h_1 = 1, h_(i+1) = u h_i + v^i + double h = 1.0; + double v_pow = v; + double series = coeffs[0]; + for (int i = 1; i < 8; ++i) { + h = u * h + v_pow; + v_pow *= v; + series += coeffs[i] * h; + } + return inv_sum + + (shift_sum + log1p(d / y) + 0.5 * d_frac / y + + u_minus_v * series); + } + // \cond + ) "\n#endif\n"; // NOLINT +// \endcond + +} // namespace opencl_kernels +} // namespace math +} // namespace stan + +#endif +#endif diff --git a/stan/math/opencl/kernels/device_functions/log_beta_ratio.hpp b/stan/math/opencl/kernels/device_functions/log_beta_ratio.hpp new file mode 100644 index 00000000000..f42a0e86167 --- /dev/null +++ b/stan/math/opencl/kernels/device_functions/log_beta_ratio.hpp @@ -0,0 +1,168 @@ +#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_LOG_BETA_RATIO_HPP +#define STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_LOG_BETA_RATIO_HPP +#ifdef STAN_OPENCL + +#include +#include + +namespace stan { +namespace math { +namespace opencl_kernels { + +// \cond +static constexpr const char* log_beta_ratio_device_function + = "\n" + "#ifndef STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_LOG_BETA_RATIO\n" + "#define " + "STAN_MATH_OPENCL_KERNELS_DEVICE_FUNCTIONS_LOG_BETA_RATIO\n" STRINGIFY( + // \endcond + /** \ingroup opencl_device_functions + * Return log1p(t) - t for 0 <= t < 1 without cancellation. See + * stan::math::internal::log1pmx(). + * + * @param t argument in [0, 1) + * @return log1p(t) - t + */ + double stan_log1pmx(double t) { + if (t > 0.25) { + return log1p(t) - t; + } + const double r = t / (2.0 + t); + const double r2 = r * r; + double s = 1.0 / 21.0; + s = 1.0 / 19.0 + r2 * s; + s = 1.0 / 17.0 + r2 * s; + s = 1.0 / 15.0 + r2 * s; + s = 1.0 / 13.0 + r2 * s; + s = 1.0 / 11.0 + r2 * s; + s = 1.0 / 9.0 + r2 * s; + s = 1.0 / 7.0 + r2 * s; + s = 1.0 / 5.0 + r2 * s; + s = 1.0 / 3.0 + r2 * s; + return r * (2.0 * r2 * s - t); + } + + /** \ingroup opencl_device_functions + * Return the count that stan_log_beta_ratio_term() removes from + * its term: k if 0 < k < x, else 0. + * + * @param x shape + * @param k count, a nonnegative integer + * @return the removed count + */ + double stan_log_beta_ratio_removed(double x, double k) { + if (k == 0 || !(k / x < 1.0)) { + return 0.0; + } + return k; + } + + /** \ingroup opencl_device_functions + * Return (x - 1/2) log1p(k / x) minus the count that + * stan_log_beta_ratio_removed() returns. For k < x the term is + * about k and is formed as x (log1p(k / x) - k / x) + * - log1p(k / x) / 2. See + * stan::math::internal::stirling_log1p_term(). + * + * @param x shape, at least LGAMMA_STIRLING_DIFF_USEFUL + * @param k count, a nonnegative integer + * @return the term without the removed count + */ + double stan_log_beta_ratio_term(double x, double k) { + if (k == 0) { + return 0.0; + } + const double t = k / x; + if (t < 1.0) { + return x * stan_log1pmx(t) - 0.5 * log1p(t); + } + return (x - 0.5) * log1p(t); + } + + /** \ingroup opencl_device_functions + * Return an estimate of the size of the terms that lbeta(a, b) + * adds up: min(a, b) (1 + log1p(max / min)). Only used to choose + * between two forms. + * + * @param a first argument + * @param b second argument + * @return the size estimate + */ + double stan_lbeta_terms_size(double a, double b) { + const double small = fmin(a, b); + const double large = fmax(a, b); + return small * (1.0 + log1p(large / small)); + } + + /** \ingroup opencl_device_functions + * Return lbeta(alpha + n, beta + m) - lbeta(alpha, beta) for + * shapes alpha, beta > 0 and integer counts n, m >= 0, without + * the cancellation of the two lbeta values for large shapes. + * This is + * stan::math::internal::log_beta_ratio() with its shape-only + * part (internal::log_beta_ratio_denominator()) computed inside. + * + * For shapes of at least LGAMMA_STIRLING_DIFF_USEFUL, every lgamma + * is written as its Stirling form plus lgamma_stirling_diff(). + * The terms linear in the arguments cancel exactly. The function + * takes this form or the lbeta form, whichever has the smaller + * terms. + * + * @param alpha first shape, positive + * @param beta second shape, positive + * @param n first count, a nonnegative integer + * @param m second count, a nonnegative integer + * @return lbeta(alpha + n, beta + m) - lbeta(alpha, beta) + */ + double stan_log_beta_ratio(double alpha, double beta, double n, + double m) { + if (!(alpha >= LGAMMA_STIRLING_DIFF_USEFUL + && beta >= LGAMMA_STIRLING_DIFF_USEFUL)) { + return stan_lbeta(alpha + n, beta + m) - stan_lbeta(alpha, beta); + } + const double total_count = n + m; + const double alpha_plus_beta = alpha + beta; + const double total = alpha_plus_beta + total_count; + // the removed counts are integers, so their sums are exact + const double removed_plus = stan_log_beta_ratio_removed(alpha, n) + + stan_log_beta_ratio_removed(beta, m); + const double removed_minus + = stan_log_beta_ratio_removed(alpha_plus_beta, total_count); + const double term_alpha = stan_log_beta_ratio_term(alpha, n); + const double term_beta = stan_log_beta_ratio_term(beta, m); + const double term_total + = stan_log_beta_ratio_term(alpha_plus_beta, total_count); + // n log(p) + m log(q) with p + q = 1: take the log1m form for + // the larger of p and q, so that a ratio near 1 keeps its digits + const double p = (alpha + n) / total; + const double q = (beta + m) / total; + const double log_n = (p < q) ? n * log(p) : n * log1p(-q); + const double log_m = (p < q) ? m * log1p(-p) : m * log(q); + const double size_stirling = fabs(term_alpha) + fabs(term_beta) + + fabs(term_total) + fabs(log_n) + + fabs(log_m); + const double size_lbeta + = stan_lbeta_terms_size(alpha, beta) + + stan_lbeta_terms_size(alpha + n, beta + m); + if (size_stirling <= size_lbeta) { + const double denominator + = lgamma_stirling_diff(alpha) + lgamma_stirling_diff(beta) + - lgamma_stirling_diff(alpha_plus_beta); + return term_alpha + term_beta - term_total + + (removed_plus - removed_minus) + log_n + log_m + + lgamma_stirling_diff(alpha + n) + + lgamma_stirling_diff(beta + m) + - lgamma_stirling_diff(total) - denominator; + } + return stan_lbeta(alpha + n, beta + m) - stan_lbeta(alpha, beta); + } + // \cond + ) "\n#endif\n"; // NOLINT +// \endcond + +} // namespace opencl_kernels +} // namespace math +} // namespace stan + +#endif +#endif diff --git a/stan/math/opencl/kernels/neg_binomial_2_log_glm_lpmf.hpp b/stan/math/opencl/kernels/neg_binomial_2_log_glm_lpmf.hpp index 3a31d5dfc26..9b76dd12a01 100644 --- a/stan/math/opencl/kernels/neg_binomial_2_log_glm_lpmf.hpp +++ b/stan/math/opencl/kernels/neg_binomial_2_log_glm_lpmf.hpp @@ -3,7 +3,12 @@ #ifdef STAN_OPENCL #include +#include #include +#include +#include +#include +#include #include namespace stan { @@ -46,14 +51,18 @@ static constexpr const char* neg_binomial_2_log_glm_kernel_code = STRINGIFY( * @param need_phi_derivative whether phi_derivative needs to be computed * @param need_phi_derivative_sum whether phi_derivative_sum needs to be * computed - * @param need_logp1 interpreted as boolean - whether first part logp_global - * needs to be computed - * @param need_logp2 interpreted as boolean - whether second part - * logp_global needs to be computed - * @param need_logp3 interpreted as boolean - whether third part logp_global - * needs to be computed - * @param need_logp4 interpreted as boolean - whether fourth part - * logp_global needs to be computed + * @param need_logp_phi interpreted as boolean - whether the term + * binomial_coefficient_log(y + phi - 1, y) of logp_global needs to be + * computed + * @param need_add_lgamma_y1 interpreted as boolean - whether to add + * lgamma(y + 1), which propto drops but binomial_coefficient_log + * contains + * @param need_sub_phi_log_phi interpreted as boolean - whether to + * subtract phi log(phi), which propto drops for data phi but the + * log1p_exp terms contain + * @param need_sub_y_theta interpreted as boolean - whether to subtract + * y theta, which propto drops for data x, alpha and beta but the + * log1p_exp terms contain */ __kernel void neg_binomial_2_log_glm( __global double* logp_global, __global double* theta_derivative_global, @@ -65,8 +74,8 @@ static constexpr const char* neg_binomial_2_log_glm_kernel_code = STRINGIFY( const int is_alpha_vector, const int is_phi_vector, const int need_theta_derivative, const int need_theta_derivative_sum, const int need_phi_derivative, const int need_phi_derivative_sum, - const int need_logp1, const int need_logp2, const int need_logp3, - const int need_logp4) { + const int need_logp_phi, const int need_add_lgamma_y1, + const int need_sub_phi_log_phi, const int need_sub_y_theta) { const int gid = get_global_id(0); const int lid = get_local_id(0); const int lsize = get_local_size(0); @@ -91,28 +100,25 @@ static constexpr const char* neg_binomial_2_log_glm_kernel_code = STRINGIFY( } theta += alpha[gid * is_alpha_vector]; double log_phi = log(phi); - double logsumexp_theta_logphi; - if (theta > log_phi) { - logsumexp_theta_logphi = theta + log1p_exp(log_phi - theta); - } else { - logsumexp_theta_logphi = log_phi + log1p_exp(theta - log_phi); - } + // log1p(mu / phi) with mu = exp(theta) + double log1p_exp_theta_m_log_phi = log1p_exp(theta - log_phi); double y_plus_phi = y + phi; - if (need_logp1) { - logp -= lgamma(y + 1); + // The log pmf is lchoose(y + phi - 1, y) - y log1p(phi / mu) + // - phi log1p(mu / phi). Formed from lgamma(phi), lgamma(y + phi) + // and phi log(phi), the terms cancel for large phi + logp + -= y * log1p_exp(log_phi - theta) + phi * log1p_exp_theta_m_log_phi; + if (need_logp_phi) { + logp += binomial_coefficient_log(y_plus_phi - 1, y); } - if (need_logp2) { - logp -= lgamma(phi); - if (phi != 0) { - logp += phi * log(phi); - } + if (need_add_lgamma_y1) { + logp += lgamma(y + 1); } - logp -= y_plus_phi * logsumexp_theta_logphi; - if (need_logp3) { - logp += y * theta; + if (need_sub_phi_log_phi) { + logp -= phi * log_phi; } - if (need_logp4) { - logp += lgamma(y_plus_phi); + if (need_sub_y_theta) { + logp -= y * theta; } double theta_exp = exp(theta); theta_derivative = y - theta_exp * y_plus_phi / (theta_exp + phi); @@ -120,9 +126,11 @@ static constexpr const char* neg_binomial_2_log_glm_kernel_code = STRINGIFY( theta_derivative_global[gid] = theta_derivative; } if (need_phi_derivative) { - phi_derivative = 1 - y_plus_phi / (theta_exp + phi) + log_phi - - logsumexp_theta_logphi + digamma(y_plus_phi) - - digamma(phi); + // 1 - (y + phi) / (mu + phi) and log(phi) - log(mu + phi) are + // formed without cancellation, and digamma_diff replaces + // digamma(y + phi) - digamma(phi) + phi_derivative = (theta_exp - y) / (theta_exp + phi) + - log1p_exp_theta_m_log_phi + digamma_diff(phi, y); if (!need_phi_derivative_sum) { phi_derivative_global[gid] = phi_derivative; } @@ -196,10 +204,14 @@ static constexpr const char* neg_binomial_2_log_glm_kernel_code = STRINGIFY( const kernel_cl - neg_binomial_2_log_glm("neg_binomial_2_log_glm", - {digamma_device_function, log1p_exp_device_function, - neg_binomial_2_log_glm_kernel_code}, - {{"REDUCTION_STEP_SIZE", 4}, {"LOCAL_SIZE_", 64}}); + neg_binomial_2_log_glm( + "neg_binomial_2_log_glm", + {digamma_device_function, digamma_diff_device_function, + log1p_exp_device_function, lgamma_stirling_device_function, + lgamma_stirling_diff_device_function, lbeta_device_function, + binomial_coefficient_log_device_function, + neg_binomial_2_log_glm_kernel_code}, + {{"REDUCTION_STEP_SIZE", 4}, {"LOCAL_SIZE_", 64}}); } // namespace opencl_kernels } // namespace math diff --git a/stan/math/opencl/prim/beta_binomial_lpmf.hpp b/stan/math/opencl/prim/beta_binomial_lpmf.hpp index 39406d1b38e..7563d779483 100644 --- a/stan/math/opencl/prim/beta_binomial_lpmf.hpp +++ b/stan/math/opencl/prim/beta_binomial_lpmf.hpp @@ -7,7 +7,7 @@ #include #include #include -#include +#include #include #include #include @@ -77,16 +77,17 @@ inline return_type_t beta_binomial_lpmf( auto beta_pos_finite = beta_val > 0.0 && isfinite(beta_val); auto return_neg_inf = (n < 0 || n > N) + constant(0, N_size, 1); - auto lbeta_diff - = lbeta(n + alpha_val, N - n + beta_val) - lbeta(alpha_val, beta_val); - auto digamma_diff - = digamma(alpha_val + beta_val) - digamma(N + alpha_val + beta_val); + // lbeta(n + alpha, N - n + beta) - lbeta(alpha, beta) without the + // cancellation of the two lbeta values for large shapes + auto lbeta_diff = log_beta_ratio(alpha_val, beta_val, n, N - n); auto logp_expr = colwise_sum(static_select::value>( binomial_coefficient_log(N, n) + lbeta_diff, lbeta_diff)); - auto alpha_deriv = digamma(n + alpha_val) + digamma_diff - digamma(alpha_val); - auto beta_deriv - = digamma(N - n + beta_val) + digamma_diff - digamma(beta_val); + // Each partial is a sum of two digamma differences psi(x + k) - psi(x); + // digamma_diff forms them without cancellation + auto digamma_diff_total = digamma_diff(alpha_val + beta_val, N); + auto alpha_deriv = digamma_diff(alpha_val, n) - digamma_diff_total; + auto beta_deriv = digamma_diff(beta_val, N - n) - digamma_diff_total; matrix_cl logp_cl; matrix_cl return_neg_inf_cl; diff --git a/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp b/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp index 7ee8051b17e..2e40a7e2bcd 100644 --- a/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp +++ b/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp @@ -13,7 +13,7 @@ #include #include #include -#include +#include #include #include #include @@ -129,13 +129,21 @@ neg_binomial_2_log_glm_lpmf(const T_y_cl& y, const T_x_cl& x, = is_autodiff_v || need_phi_derivative_sum; matrix_cl phi_derivative_cl( need_phi_derivative ? (need_phi_derivative_sum ? wgs : N) : 0, 1); - const bool need_logp1 = include_summand::value; - const bool need_logp2 - = include_summand::value && is_phi_vector; - const bool need_logp3 - = include_summand::value; - const bool need_logp4 = include_summand::value - && (is_y_vector || is_phi_vector); + // binomial_coefficient_log(y + phi - 1, y) is computed in the kernel if it + // differs between instances, and below otherwise + const bool need_logp_phi = include_summand::value + && (is_y_vector || is_phi_vector); + // Under propto, the kernel (or the code below, for a scalar y or phi) + // removes -lgamma(y + 1) (in binomial_coefficient_log), phi log(phi) for + // data phi, and y theta for data x, alpha and beta. + const bool need_add_lgamma_y1 = !include_summand::value + && include_summand::value + && is_y_vector; + const bool need_sub_phi_log_phi = !include_summand::value + && !include_summand::value + && is_phi_vector; + const bool need_sub_y_theta + = !include_summand::value; matrix_cl logp_cl(wgs, 1); try { @@ -145,7 +153,8 @@ neg_binomial_2_log_glm_lpmf(const T_y_cl& y, const T_x_cl& x, y_val_cl, x_val, alpha_val_cl, beta_val, phi_val_cl, N, M, is_y_vector, is_alpha_vector, is_phi_vector, need_theta_derivative, need_theta_derivative_sum, need_phi_derivative, need_phi_derivative_sum, - need_logp1, need_logp2, need_logp3, need_logp4); + need_logp_phi, need_add_lgamma_y1, need_sub_phi_log_phi, + need_sub_y_theta); } catch (const cl::Error& e) { check_opencl_error(function, e); } @@ -167,12 +176,18 @@ neg_binomial_2_log_glm_lpmf(const T_y_cl& y, const T_x_cl& x, = isfinite(phi_val) && phi_val > 0; } - if constexpr (include_summand::value && !is_phi_vector) { - logp += N * (multiply_log(phi_val, phi_val) - lgamma(phi_val)); - } if constexpr (include_summand::value && !is_y_vector && !is_phi_vector) { - logp += lgamma(y_val + phi_val) * N; + logp += N * binomial_coefficient_log(y_val + phi_val - 1, y_val); + } + if constexpr (!include_summand::value + && include_summand::value && !is_y_vector) { + logp += N * lgamma(y_val + 1.0); + } + if constexpr (!include_summand::value + && !include_summand::value + && !is_phi_vector) { + logp -= N * multiply_log(phi_val, phi_val); } auto ops_partials = make_partials_propagator(x, alpha, beta, phi); diff --git a/stan/math/opencl/prim/neg_binomial_2_log_lpmf.hpp b/stan/math/opencl/prim/neg_binomial_2_log_lpmf.hpp index 2e5554f3d7f..79bf7b0da1c 100644 --- a/stan/math/opencl/prim/neg_binomial_2_log_lpmf.hpp +++ b/stan/math/opencl/prim/neg_binomial_2_log_lpmf.hpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -92,9 +93,10 @@ neg_binomial_2_log_lpmf(const T_n_cl& n, const T_log_location_cl& eta, logp2 + elt_multiply(n, eta_val), logp2)); auto eta_deriv = n - elt_multiply(n_plus_phi, exp_eta_over_exp_eta_phi); + // digamma_diff replaces digamma(n + phi) - digamma(phi), which loses all + // accuracy for large phi auto phi_deriv = exp_eta_over_exp_eta_phi - elt_divide(n, exp_eta + phi_val) - - log1p_exp_eta_m_logphi - digamma(phi_val) - + digamma(n_plus_phi); + - log1p_exp_eta_m_logphi + digamma_diff(phi_val, n); matrix_cl logp_cl; matrix_cl eta_deriv_cl; diff --git a/stan/math/opencl/prim/neg_binomial_2_lpmf.hpp b/stan/math/opencl/prim/neg_binomial_2_lpmf.hpp index 9dc4eb9a774..ee6692f01f7 100644 --- a/stan/math/opencl/prim/neg_binomial_2_lpmf.hpp +++ b/stan/math/opencl/prim/neg_binomial_2_lpmf.hpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -91,8 +92,10 @@ inline return_type_t neg_binomial_2_lpmf( auto log_term = select(mu_val < phi_val, log1p(-elt_divide(mu_val, mu_plus_phi)), log_phi - log_mu_plus_phi); + // digamma_diff replaces digamma(n + phi) - digamma(phi), which loses all + // accuracy for large phi auto phi_deriv = elt_divide(mu_val - n, mu_plus_phi) + log_term - - digamma(phi_val) + digamma(n_plus_phi); + + digamma_diff(phi_val, n); matrix_cl logp_cl; matrix_cl mu_deriv_cl; diff --git a/stan/math/opencl/prim/neg_binomial_lpmf.hpp b/stan/math/opencl/prim/neg_binomial_lpmf.hpp index 6df44dfb881..4177d6c71fa 100644 --- a/stan/math/opencl/prim/neg_binomial_lpmf.hpp +++ b/stan/math/opencl/prim/neg_binomial_lpmf.hpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -72,11 +73,11 @@ inline return_type_t neg_binomial_lpmf( function, "Inverse scale parameter", beta_val, "positive finite"); auto beta_positive_finite = 0 < beta_val && isfinite(beta_val); - auto digamma_alpha = digamma(alpha_val); auto log1p_inv_beta = log1p(elt_divide(1.0, beta_val)); auto log1p_beta = log1p(beta_val); + // alpha / beta - alpha / (1 + beta), without the cancellation auto lambda_m_alpha_over_1p_beta - = elt_divide(alpha_val, beta_val) - elt_divide(alpha_val, 1.0 + beta_val); + = elt_divide(elt_divide(alpha_val, beta_val), 1.0 + beta_val); auto logp1 = -elt_multiply(alpha_val, log1p_inv_beta) - elt_multiply(n, log1p_beta); @@ -86,7 +87,9 @@ inline return_type_t neg_binomial_lpmf( + binomial_coefficient_log(n + alpha_val - 1.0, alpha_val - 1.0), logp1)); - auto alpha_deriv = digamma(alpha_val + n) - digamma_alpha - log1p_inv_beta; + // digamma_diff replaces digamma(alpha + n) - digamma(alpha), which loses + // all accuracy for large alpha + auto alpha_deriv = digamma_diff(alpha_val, n) - log1p_inv_beta; auto beta_deriv = lambda_m_alpha_over_1p_beta - elt_divide(n, beta_val + 1.0); matrix_cl logp_cl; diff --git a/stan/math/opencl/prim/student_t_lpdf.hpp b/stan/math/opencl/prim/student_t_lpdf.hpp index f358f00bce6..b79b3d34fd4 100644 --- a/stan/math/opencl/prim/student_t_lpdf.hpp +++ b/stan/math/opencl/prim/student_t_lpdf.hpp @@ -7,6 +7,8 @@ #include #include #include +#include +#include #include #include @@ -99,9 +101,30 @@ inline return_type_t student_t_lpdf( auto log1p_val = log1p(square_y_scaled_over_nu); auto logp1 = -elt_multiply((half_nu + 0.5), log1p_val); + // The nu terms as on the CPU (internal::student_t_nu_constant_fun and + // internal::student_t_nu_digamma_fun): differences of lgamma and digamma + // values below nu = 40, asymptotic series in z = 2 / nu from nu = 40 on, + // without the cancellation for large nu. + auto nu_small = nu_val < internal::student_t_series_min_nu; + auto z = elt_divide(2.0, nu_val); + auto z2 = elt_multiply(z, z); + auto nu_constant_series + = -0.5 * LOG_TWO + + elt_multiply( + z, + -0.125 + + elt_multiply( + z2, 1.0 / 192.0 + + elt_multiply( + z2, -1.0 / 640.0 + + elt_multiply( + z2, 17.0 / 14336.0 + - (31.0 / 18432.0) * z2)))); + auto nu_constant = select( + nu_small, lgamma(half_nu + 0.5) - lgamma(half_nu) - 0.5 * log(nu_val), + nu_constant_series); auto logp2 = static_select::value>( - logp1 + lgamma(half_nu + 0.5) - lgamma(half_nu) - 0.5 * log(nu_val), - logp1); + logp1 + nu_constant, logp1); auto logp_expr = colwise_sum(static_select::value>( logp2 - log(sigma_val), logp2)); @@ -114,9 +137,22 @@ inline return_type_t student_t_lpdf( auto rep_deriv = elt_divide(elt_multiply(nu_val + 1, square_y_scaled_over_nu), 1 + square_y_scaled_over_nu) - 1; - auto nu_deriv = 0.5 - * (digamma(half_nu + 0.5) - digamma(half_nu) - log1p_val - + elt_divide(rep_deriv, nu_val)); + auto nu_digamma_series = elt_multiply( + z, 0.5 + + elt_multiply( + z, 0.125 + + elt_multiply( + z2, -1.0 / 64.0 + + elt_multiply( + z2, 1.0 / 128.0 + + elt_multiply( + z2, -119.0 / 14336.0 + + (279.0 / 18432.0) + * z2))))); + auto nu_digamma = select(nu_small, digamma(half_nu + 0.5) - digamma(half_nu), + nu_digamma_series); + auto nu_deriv + = 0.5 * (nu_digamma - log1p_val + elt_divide(rep_deriv, nu_val)); auto sigma_deriv = elt_divide(rep_deriv, sigma_val); matrix_cl logp_cl; diff --git a/test/unit/math/opencl/device_functions/digamma_diff_test.cpp b/test/unit/math/opencl/device_functions/digamma_diff_test.cpp new file mode 100644 index 00000000000..1520148cde8 --- /dev/null +++ b/test/unit/math/opencl/device_functions/digamma_diff_test.cpp @@ -0,0 +1,142 @@ +#ifdef STAN_OPENCL +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +static const std::string test_digamma_diff_kernel_code + = STRINGIFY(__kernel void test(__global double *C, __global double *A, + __global double *B) { + const int i = get_global_id(0); + C[i] = digamma_diff(A[i], B[i]); + }); + +const stan::math::opencl_kernels::kernel_cl< + stan::math::opencl_kernels::out_buffer, + stan::math::opencl_kernels::in_buffer, + stan::math::opencl_kernels::in_buffer> + digamma_diff_kernel( + "test", {stan::math::opencl_kernels::digamma_device_function, + stan::math::opencl_kernels::digamma_diff_device_function, + test_digamma_diff_kernel_code}); + +namespace { +Eigen::VectorXd digamma_diff_cl(const Eigen::VectorXd &x, + const Eigen::VectorXd &d) { + stan::math::matrix_cl x_cl(x); + stan::math::matrix_cl d_cl(d); + stan::math::matrix_cl res_cl(x.size(), 1); + digamma_diff_kernel(cl::NDRange(x.size()), res_cl, x_cl, d_cl); + return stan::math::from_matrix_cl(res_cl); +} +} // namespace + +TEST(MathMatrixCL, digamma_diff) { + Eigen::VectorXd x = Eigen::VectorXd::Random(1000).array() * 15 + 15.01; + Eigen::VectorXd d = Eigen::VectorXd::Random(1000).array() * 50 + 50; + Eigen::VectorXd res = digamma_diff_cl(x, d); + stan::test::expect_near_rel("digamma_diff (OpenCL)", res, + stan::math::digamma_diff(x, d)); +} + +TEST(MathMatrixCL, digamma_diff_large_shapes_match_cpu) { + // log-uniform x in [1e-3, 1e16] and d in [1e-3, 1e8]; the CPU function + // has a relative error of a few ulp + const int n = 4096; + Eigen::VectorXd x + = (Eigen::VectorXd::Random(n).array() * 9.5 + 6.5) * std::log(10.0); + Eigen::VectorXd d + = (Eigen::VectorXd::Random(n).array() * 5.5 + 2.5) * std::log(10.0); + x = x.array().exp(); + d = d.array().exp(); + Eigen::VectorXd res = digamma_diff_cl(x, d); + Eigen::VectorXd cpu = stan::math::digamma_diff(x, d); + double max_rel = 0; + for (int i = 0; i < n; ++i) { + const double rel = std::fabs(res(i) - cpu(i)) / std::fabs(cpu(i)); + max_rel = std::max(max_rel, rel); + EXPECT_NEAR(res(i), cpu(i), 1e-14 * std::fabs(cpu(i))) + << "x = " << x(i) << ", d = " << d(i); + } + std::cout << "digamma_diff OpenCL vs CPU, max relative difference " << max_rel + << " (" << max_rel / std::numeric_limits::epsilon() + << " eps)" << std::endl; +} + +TEST(MathMatrixCL, digamma_diff_reference_values) { + // psi(x + d) - psi(x) from mp.digamma at 120 digits for the exact double + // arguments, checked against the same computation at 80 digits + struct TestValue { + double x; + double d; + double value; + }; + const std::vector values = { + {0x1.56e1fc2f8f359p-997, 0x1.0000000000000p+0, 9.9999999999999997e+299}, + {0x1.0624dd2f1a9fcp-10, 0x1.0000000000000p+0, 9.9999999999999998e+2}, + {0x1.0000000000000p-1, 0x0.0p+0, 0.0}, + {0x1.0000000000000p-1, 0x1.b7cdfd9d7bdbbp-34, 4.9348021997032397e-10}, + {0x1.0000000000000p-1, 0x1.8000000000000p+1, 3.0666666666666667}, + {0x1.4000000000000p+1, 0x1.0000000000000p-2, 1.1574438433018941e-1}, + {0x1.d333333333333p+2, 0x1.c800000000000p+5, 2.2379430906763218}, + {0x1.3ff7ced916873p+3, 0x1.0000000000000p+0, 1.000100010001e-1}, + {0x1.4000000000000p+3, 0x1.0000000000000p+0, 1.0e-1}, + {0x1.9000000000000p+3, 0x1.e848000000000p+19, 1.1330326906617404e+1}, + {0x1.f400000000000p+9, 0x1.0000000000000p+0, 1.0e-3}, + {0x1.7d78400000000p+26, 0x1.8000000000000p+1, 2.9999999700000005e-8}, + {0x1.d1a94a2000000p+39, 0x1.c800000000000p+5, 5.6999999998404e-11}, + {0x1.c6bf526340000p+49, 0x1.8000000000000p+1, 2.999999999999997e-15}, + {0x1.c6bf526340000p+49, 0x1.c6bf526340000p+49, 6.9314718055994556e-1}, + {0x1.550f7dca70000p+51, 0x1.d400000000000p+6, 3.8999999999999246e-14}, + {0x1.5af1d78b58c40p+66, 0x1.0000000000000p+0, 1.0e-20}, + {0x1.0000000000000p+0, 0x1.7e43c8800759cp+996, 6.9135274356311524e+2}, + {0x1.7e43c8800759cp+996, 0x1.7e43c8800759cp+996, 6.9314718055994531e-1}, + }; + const int n = values.size(); + Eigen::VectorXd x(n); + Eigen::VectorXd d(n); + for (int i = 0; i < n; ++i) { + x(i) = values[i].x; + d(i) = values[i].d; + } + Eigen::VectorXd res = digamma_diff_cl(x, d); + double max_rel = 0; + for (int i = 0; i < n; ++i) { + const double ref = values[i].value; + if (ref == 0) { + EXPECT_EQ(res(i), 0.0) << "x = " << x(i) << ", d = " << d(i); + continue; + } + const double rel = std::fabs(res(i) - ref) / std::fabs(ref); + max_rel = std::max(max_rel, rel); + EXPECT_NEAR(res(i), ref, 1e-14 * std::fabs(ref)) + << "x = " << x(i) << ", d = " << d(i); + } + std::cout << "digamma_diff OpenCL vs mpmath, max relative error " << max_rel + << " (" << max_rel / std::numeric_limits::epsilon() + << " eps)" << std::endl; +} + +TEST(MathMatrixCL, digamma_diff_edge_cases) { + const double inf = std::numeric_limits::infinity(); + Eigen::VectorXd x(7); + x << NAN, 1.5, 0.0, -1.0, 1.5, inf, 1.5; + Eigen::VectorXd d(7); + d << 1.0, NAN, 1.0, 1.0, -1.0, 3.0, inf; + Eigen::VectorXd res = digamma_diff_cl(x, d); + for (int i = 0; i < 5; ++i) { + EXPECT_TRUE(std::isnan(res(i))) << "x = " << x(i) << ", d = " << d(i); + } + EXPECT_EQ(res(5), 0.0); + EXPECT_EQ(res(6), inf); +} + +#endif diff --git a/test/unit/math/opencl/kernel_generator/elt_function_cl_test.cpp b/test/unit/math/opencl/kernel_generator/elt_function_cl_test.cpp index f23bce2e8d7..0dc569dd2fa 100644 --- a/test/unit/math/opencl/kernel_generator/elt_function_cl_test.cpp +++ b/test/unit/math/opencl/kernel_generator/elt_function_cl_test.cpp @@ -260,6 +260,7 @@ TEST(KernelGenerator, multiple_operations_with_includes_test) { TEST_BINARY_FUNCTION(beta) TEST_BINARY_FUNCTION(binomial_coefficient_log) +TEST_BINARY_FUNCTION(digamma_diff) TEST_BINARY_FUNCTION(fdim) TEST_BINARY_FUNCTION(fmax) TEST_BINARY_FUNCTION(fmin) @@ -269,4 +270,50 @@ TEST_BINARY_FUNCTION(lmultiply) TEST_BINARY_FUNCTION(multiply_log) TEST_BINARY_FUNCTION(pow) +TEST(KernelGenerator, log_beta_ratio_test) { + // shapes below and above 10, and above 1e15, with integer counts + MatrixXd alpha(3, 3); + alpha << 0.3, 2.5, 9.9, 10.0, 11.2, 1e6, 3e13, 1e15, 5e17; + MatrixXd beta(3, 3); + beta << 4.1, 0.7, 12.0, 25.0, 1e8, 15.0, 7e12, 4.84, 2e17; + MatrixXi n(3, 3); + n << 0, 3, 1, 5, 0, 1000, 400, 5, 57; + MatrixXi m(3, 3); + m << 2, 0, 7, 15, 1, 3, 600, 112, 60; + MatrixXi m_size(3, 2); + m_size << 2, 0, 7, 15, 1, 3; + + matrix_cl alpha_cl(alpha); + matrix_cl beta_cl(beta); + matrix_cl n_cl(n); + matrix_cl m_cl(m); + matrix_cl m_size_cl(m_size); + + EXPECT_THROW(stan::math::log_beta_ratio(alpha_cl, beta_cl, n_cl, m_size_cl), + std::invalid_argument); + + auto cpu = [](double a, double b, double k, double l) { + return stan::math::internal::log_beta_ratio( + a, b, k, l, stan::math::internal::log_beta_ratio_denominator(a, b)); + }; + matrix_cl res1_cl + = stan::math::log_beta_ratio(alpha_cl, beta_cl, n_cl, m_cl); + matrix_cl res2_cl + = stan::math::log_beta_ratio(alpha_cl, 3.5, n_cl, 4); + MatrixXd res1 = stan::math::from_matrix_cl(res1_cl); + MatrixXd res2 = stan::math::from_matrix_cl(res2_cl); + for (int i = 0; i < 3; ++i) { + for (int j = 0; j < 3; ++j) { + const double correct1 = cpu(alpha(i, j), beta(i, j), n(i, j), m(i, j)); + EXPECT_NEAR(res1(i, j), correct1, + 1e-13 * std::max(1.0, std::fabs(correct1))) + << "i = " << i << ", j = " << j; + const double correct2 = cpu(alpha(i, j), 3.5, n(i, j), 4); + EXPECT_NEAR(res2(i, j), correct2, + 1e-13 * std::max(1.0, std::fabs(correct2))) + << "i = " << i << ", j = " << j; + } + } +} + #endif diff --git a/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp b/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp index cdac00afb20..9960d9eaff3 100644 --- a/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp @@ -3,6 +3,9 @@ #include #include #include +#include +#include +#include #include TEST(ProbDistributionsBetaBinomial, error_checking) { @@ -216,4 +219,92 @@ TEST(ProbDistributionsBetaBinomial, opencl_matches_cpu_big) { beta.transpose().eval()); } +namespace beta_binomial_opencl_test_internal { +struct TestValue { + int n; + int N; + double alpha; + double beta; + double value; + double grad_log_alpha; // alpha * d/dalpha + double grad_log_beta; // beta * d/dbeta +}; + +// The tests above compare OpenCL against the CPU and cannot see an error +// that both share. These are absolute references from mpmath at 80 digits: +// the value from mp.loggamma, the gradients from mp.digamma differences, +// both checked at 130 digits (the gradients with mp.diff). The first four +// rows are as in +// test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp. The shapes are in +// hex so that they are exact. The plain differences of lbeta and digamma +// values keep no correct digits for shapes near 1e15. +// The last row is a small shape where the plain differences are correct. +const std::vector test_values = { + {57, 117, 0x1.1f43fcc4b662cp+45, 0x1.1f43fcc4b662cp+45, -2.6471538352642870, + -1.4999999999974545, 1.4999999999981384}, + {400, 1000, 0x1.b48eb57e00000p+44, 0x1.977420dc00000p+42, + -4.1354072540579752e+2, -4.1081081080252486e+2, 4.1081081078769344e+2}, + {0, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -3.3658802476858363e+1, -2.9249999999999996e+1, 2.9249999999999990e+1}, + {117, 117, 0x1.6345785d8a000p+56, 0x1.0a741a4627800p+58, + -1.6219644025102715e+2, 8.7749999999999936e+1, -8.7749999999999987e+1}, + // one shape below 10 + {5, 117, 0x1.c6bf526340000p+49, 0x1.35c28f5c28f5cp+2, + -3.4143943462813216e+3, -1.1199999999999266e+2, 1.5906435196801570e+1}, + {30, 117, 0x1.4000000000000p+1, 0x1.6bcc41e900000p+46, + -8.2341924150797106e+2, 6.9065498632679532, -2.9999999999966625e+1}, + {5, 20, 0x1.4000000000000p+3, 0x1.9000000000000p+4, -1.8540068216786033, + -3.4626330704764243e-1, 5.0911144710080701e-1}, +}; + +template +void expect_reference(const TestValue& t, const stan::math::var& lp, + const T_alpha& alpha_adj, const T_beta& beta_adj, + const std::string& signature) { + EXPECT_NEAR(lp.val(), t.value, 1e-12 * std::max(1.0, std::fabs(t.value))) + << signature << ": n = " << t.n << ", N = " << t.N + << ", alpha = " << t.alpha << ", beta = " << t.beta; + EXPECT_NEAR(t.alpha * alpha_adj, t.grad_log_alpha, + 1e-11 * std::max(1.0, std::fabs(t.grad_log_alpha))) + << signature << ": n = " << t.n << ", N = " << t.N + << ", alpha = " << t.alpha << ", beta = " << t.beta; + EXPECT_NEAR(t.beta * beta_adj, t.grad_log_beta, + 1e-11 * std::max(1.0, std::fabs(t.grad_log_beta))) + << signature << ": n = " << t.n << ", N = " << t.N + << ", alpha = " << t.alpha << ", beta = " << t.beta; +} +} // namespace beta_binomial_opencl_test_internal + +TEST(ProbDistributionsBetaBinomial, opencl_large_shapes_reference) { + using beta_binomial_opencl_test_internal::expect_reference; + using beta_binomial_opencl_test_internal::test_values; + using stan::math::var; + for (const auto& t : test_values) { + const std::vector n{t.n}; + const std::vector N{t.N}; + stan::math::matrix_cl n_cl(n); + stan::math::matrix_cl N_cl(N); + + // shapes as vectors: the kernel generator functions run on the device + Eigen::Matrix alpha(1); + alpha << t.alpha; + Eigen::Matrix beta(1); + beta << t.beta; + auto alpha_cl = stan::math::to_matrix_cl(alpha); + auto beta_cl = stan::math::to_matrix_cl(beta); + var lp = stan::math::beta_binomial_lpmf(n_cl, N_cl, alpha_cl, beta_cl); + lp.grad(); + expect_reference(t, lp, alpha(0).adj(), beta(0).adj(), "vector shapes"); + stan::math::recover_memory(); + + // shapes as scalars: alpha + beta is formed on the host + var alpha_s = t.alpha; + var beta_s = t.beta; + var lp_s = stan::math::beta_binomial_lpmf(n_cl, N_cl, alpha_s, beta_s); + lp_s.grad(); + expect_reference(t, lp_s, alpha_s.adj(), beta_s.adj(), "scalar shapes"); + stan::math::recover_memory(); + } +} + #endif diff --git a/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp b/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp index 3f7c08a1684..53058a466f1 100644 --- a/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp @@ -3,6 +3,9 @@ #include #include #include +#include +#include +#include #include using Eigen::Array; @@ -237,4 +240,187 @@ TEST(ProbDistributionsNegBinomial2LogGLM, opencl_matches_cpu_big) { stan::math::test::compare_cpu_opencl_prim_rev( neg_binomial_2_log_glm_lpmf_functor_propto, y, x, alpha, beta, phi); } + +TEST(ProbDistributionsNegBinomial2LogGLM, opencl_large_shapes_reference) { + // The tests above compare OpenCL against the CPU and cannot see an error + // that both share. These are absolute references from mpmath at 90 digits, + // checked against 140 digits, as in + // test/unit/math/rev/prob/large_shapes_test.cpp. The + // value formed from lgamma(phi), lgamma(y + phi) and phi log(phi), and + // the phi partial formed from digamma(y + phi) - digamma(phi), keep no + // correct digits for phi near 1e15. One attribute; the phi partial is + // compared as phi * d/dphi. The last row is a small phi where the plain + // differences are correct. + struct TestValue { + vector y; + vector x; + double alpha; + double beta; + double phi; + double value; + double d_alpha; + double d_beta; + double d_phi; + }; + const vector x5{0.5, -0.25, 1.25, 0.0, 2.0}; + const vector test_values = { + {{0}, + {0.5}, + 1.0, + 0.75, + 0x1.c6bf526340000p+49, + -3.9550767229205693, + -3.9550767229205615, + -1.9775383614602807, + -7.8213159420940446e-30}, + {{5}, + {0.5}, + 1.0, + 0.75, + 0x1.c6bf526340000p+49, + -1.8675684657026251, + 1.0449232770794187, + 5.2246163853970937e-1, + 1.9540676725087929e-30}, + {{0}, + {0.5}, + -0x1.999999999999ap-3, + 0.75, + 0x1.37807ed5e8000p+50, + -1.1912462166123576, + -1.1912462166123571, + -5.9562310830617854e-1, + -3.7803493755481261e-31}, + {{2, 4, 3, 5, 6}, + x5, + 1.0, + 0.75, + 0x1.c6bf526340000p+49, + -1.3267966555421410e+1, + -8.0507631204932393, + -1.8705862362560038e+1, + -2.2918189242185528e-29}, + {{0, 1, 5, 3, 57}, + x5, + 1.0, + 0.75, + 0x1.c6bf526340000p+49, + -5.5025862739499810e+1, + 3.7949236879506146e+1, + 8.5544137637438705e+1, + -9.8183556706826074e-28}, + {{2, 4, 3, 5, 6}, + x5, + 1.0, + 0.75, + 0x1.d1a94a2000000p+39, + -1.3267966555398514e+1, + -8.0507631203930684, + -1.8705862362370543e+1, + -2.2918189241713975e-23}, + {{0}, + {0.5}, + 1.0, + 0.75, + 0x1.9000000000000p+6, + -3.8788665246722319, + -3.8046018026251338, + -1.9023009013125669, + -7.4264722047098149e-4}, + {{0, 1, 5, 3, 57}, + x5, + 1.0, + 0.75, + 0x1.9000000000000p+6, + -4.7240731087544224e+1, + 3.3378922648644250e+1, + 7.6036039699167587e+1, + -6.2120571346832957e-2}, + }; + auto expect_reference = [](const TestValue& t, const var& lp, + double alpha_adj, double beta_adj, double phi_adj, + const char* signature) { + const std::string where = std::string(signature) + + ": y[0] = " + std::to_string(t.y[0]) + + ", N = " + std::to_string(t.y.size()) + + ", phi = " + std::to_string(t.phi); + EXPECT_NEAR(lp.val(), t.value, 1e-12 * std::max(1.0, std::fabs(t.value))) + << where; + EXPECT_NEAR(alpha_adj, t.d_alpha, + 1e-11 * std::max(1.0, std::fabs(t.d_alpha))) + << where; + EXPECT_NEAR(beta_adj, t.d_beta, 1e-11 * std::max(1.0, std::fabs(t.d_beta))) + << where; + const double gp = t.phi * t.d_phi; + EXPECT_NEAR(t.phi * phi_adj, gp, 1e-11 * std::max(1.0, std::fabs(gp))) + << where; + }; + for (const auto& t : test_values) { + const int N = t.y.size(); + Matrix x(N, 1); + for (int i = 0; i < N; ++i) { + x(i, 0) = t.x[i]; + } + matrix_cl y_cl(t.y); + matrix_cl x_cl(x); + + // y vector, alpha and phi scalars: the kernel sums the phi partial + { + var alpha = t.alpha; + Matrix beta(1); + beta << t.beta; + var phi = t.phi; + auto beta_cl = stan::math::to_matrix_cl(beta); + var lp = stan::math::neg_binomial_2_log_glm_lpmf(y_cl, x_cl, alpha, + beta_cl, phi); + lp.grad(); + expect_reference(t, lp, alpha.adj(), beta(0).adj(), phi.adj(), + "scalar alpha and phi"); + stan::math::recover_memory(); + } + + // alpha and phi vectors: the kernel computes the value term and the phi + // partial of each instance + { + Matrix alpha(N); + Matrix beta(1); + beta << t.beta; + Matrix phi(N); + for (int i = 0; i < N; ++i) { + alpha(i) = t.alpha; + phi(i) = t.phi; + } + auto alpha_cl = stan::math::to_matrix_cl(alpha); + auto beta_cl = stan::math::to_matrix_cl(beta); + auto phi_cl = stan::math::to_matrix_cl(phi); + var lp = stan::math::neg_binomial_2_log_glm_lpmf(y_cl, x_cl, alpha_cl, + beta_cl, phi_cl); + lp.grad(); + double alpha_adj = 0; + double phi_adj = 0; + for (int i = 0; i < N; ++i) { + alpha_adj += alpha(i).adj(); + phi_adj += phi(i).adj(); + } + expect_reference(t, lp, alpha_adj, beta(0).adj(), phi_adj, + "vector alpha and phi"); + stan::math::recover_memory(); + } + + // y, alpha and phi scalars: the value term is computed on the host + if (N == 1) { + var alpha = t.alpha; + Matrix beta(1); + beta << t.beta; + var phi = t.phi; + auto beta_cl = stan::math::to_matrix_cl(beta); + var lp = stan::math::neg_binomial_2_log_glm_lpmf(t.y[0], x_cl, alpha, + beta_cl, phi); + lp.grad(); + expect_reference(t, lp, alpha.adj(), beta(0).adj(), phi.adj(), + "scalar y, alpha and phi"); + stan::math::recover_memory(); + } + } +} #endif diff --git a/test/unit/math/opencl/rev/neg_binomial_2_log_lpmf_test.cpp b/test/unit/math/opencl/rev/neg_binomial_2_log_lpmf_test.cpp index d9e2e89b852..e1e972cd00a 100644 --- a/test/unit/math/opencl/rev/neg_binomial_2_log_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/neg_binomial_2_log_lpmf_test.cpp @@ -182,4 +182,77 @@ TEST(ProbDistributionsNegBinomial2Log, opencl_matches_cpu_eta_phi_scalar) { neg_binomial_2_log_lpmf_functor_propto, n, eta, phi); } +TEST(ProbDistributionsNegBinomial2Log, opencl_large_shapes_reference) { + // The tests above compare OpenCL against the CPU and cannot see an error + // that both share. These are absolute references from mpmath at 90 digits, + // checked against 140 digits, as in + // test/unit/math/rev/prob/large_shapes_test.cpp. The difference + // digamma(n + phi) - digamma(phi) keeps no correct digits for large phi. + // The phi partial is compared as phi * d/dphi. At phi = 1e15 the rounding + // errors of the develop OpenCL formula cancel to 0 and it passes by + // chance; the rows at phi = 1e10 and 1e12 show its error (3.6e-5 and + // 1.4e-9 in phi * d/dphi). The last row is a small phi where the plain + // difference is correct. + using stan::math::var; + struct TestValue { + int n; + double eta; + double phi; + double value; + double d_eta; + double d_phi; + }; + const std::vector test_values = { + {3, 0x1.193ea7aad030bp+0, 0x1.c6bf526340000p+49, -1.4959226032237274, + -2.7213891705004510e-16, 1.4999999999999960e-30}, + {5, 0x1.f4bd2b7ac1bafp+1, 0x1.c6bf526340000p+49, -3.5227376715640301e+1, + -4.4999999999997745e+1, -1.0099999999999289e-27}, + {2, -0x1.6d3c324e13f50p-2, 0x1.37807ed5e8000p+50, -2.1064970684374103, + 1.2999999999999994, 8.2582982577654707e-32}, + {3, 0x1.193ea7aad030bp+0, 0x1.2a05f20000000p+33, -1.4959226033737259, + -2.7213891696840424e-16, 1.4999999996e-20}, + {5, 0x1.f4bd2b7ac1bafp+1, 0x1.2a05f20000000p+33, -3.5227376614641312e+1, + -4.4999999774999996e+1, -1.0099999929136665e-17}, + {57, 0x1.193ea7aad030bp+0, 0x1.d1a94a2000000p+39, -1.1677494795148559e+2, + 5.3999999999838e+1, -1.429499999940379e-21}, + {0, 0x1.193ea7aad030bp+0, 0x1.9000000000000p+6, -2.9558802241544405, + -2.9126213592233012, -4.3258864931139310e-4}, + }; + auto expect_reference = [](const TestValue& t, const var& lp, double eta_adj, + double phi_adj, const char* signature) { + EXPECT_NEAR(lp.val(), t.value, 1e-12 * std::max(1.0, std::fabs(t.value))) + << signature << ": n = " << t.n << ", eta = " << t.eta + << ", phi = " << t.phi; + EXPECT_NEAR(eta_adj, t.d_eta, 1e-11 * std::max(1.0, std::fabs(t.d_eta))) + << signature << ": n = " << t.n << ", eta = " << t.eta + << ", phi = " << t.phi; + const double gp = t.phi * t.d_phi; + EXPECT_NEAR(t.phi * phi_adj, gp, 1e-11 * std::max(1.0, std::fabs(gp))) + << signature << ": n = " << t.n << ", eta = " << t.eta + << ", phi = " << t.phi; + }; + for (const auto& t : test_values) { + const std::vector n{t.n}; + stan::math::matrix_cl n_cl(n); + + Eigen::Matrix eta(1); + eta << t.eta; + Eigen::Matrix phi(1); + phi << t.phi; + auto eta_cl = stan::math::to_matrix_cl(eta); + auto phi_cl = stan::math::to_matrix_cl(phi); + var lp = stan::math::neg_binomial_2_log_lpmf(n_cl, eta_cl, phi_cl); + lp.grad(); + expect_reference(t, lp, eta(0).adj(), phi(0).adj(), "vector"); + stan::math::recover_memory(); + + var eta_s = t.eta; + var phi_s = t.phi; + var lp_s = stan::math::neg_binomial_2_log_lpmf(n_cl, eta_s, phi_s); + lp_s.grad(); + expect_reference(t, lp_s, eta_s.adj(), phi_s.adj(), "scalar"); + stan::math::recover_memory(); + } +} + #endif diff --git a/test/unit/math/opencl/rev/neg_binomial_2_lpmf_test.cpp b/test/unit/math/opencl/rev/neg_binomial_2_lpmf_test.cpp index e7e863610d9..ab9e479db09 100644 --- a/test/unit/math/opencl/rev/neg_binomial_2_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/neg_binomial_2_lpmf_test.cpp @@ -188,4 +188,68 @@ TEST(ProbDistributionsNegBinomial2, opencl_scalar_n_mu) { neg_binomial_2_lpmf_functor_propto, n, mu, phi); } +TEST(muProbDistributionsNegBinomial2, opencl_large_shapes_reference) { + // The tests above compare OpenCL against the CPU and cannot see an error + // that both share. These are absolute references from mpmath at 90 digits, + // checked against 140 digits, as in + // test/unit/math/rev/prob/large_shapes_test.cpp. The difference + // digamma(n + phi) - digamma(phi) keeps no correct digits for large phi. + // The phi partial is compared as phi * d/dphi. The last row is a small + // phi where the plain difference is correct. + using stan::math::var; + struct TestValue { + int n; + double mu; + double phi; + double value; + double d_mu; + double d_phi; + }; + const std::vector test_values = { + {57, 0x1.8000000000000p+1, 0x1.2a05f20000000p+33, -1.1677494780996510e+2, + 1.7999999994600000e+1, -1.4294999940379000e-17}, + {5, 0x1.9000000000000p+5, 0x1.2a05f20000000p+33, -3.5227376614641316e+1, + -8.9999999550000002e-1, -1.0099999929136667e-17}, + {1, 0x1.8000000000000p+1, 0x1.2a05f20000000p+33, -1.9013877111818903, + -6.6666666646666667e-1, -1.4999999991000000e-20}, + {0, 0x1.8000000000000p+1, 0x1.9000000000000p+6, -2.9558802241544403, + -9.7087378640776699e-1, -4.3258864931139302e-4}, + }; + auto expect_reference = [](const TestValue& t, const var& lp, double mu_adj, + double phi_adj, const char* signature) { + EXPECT_NEAR(lp.val(), t.value, 1e-12 * std::max(1.0, std::fabs(t.value))) + << signature << ": n = " << t.n << ", mu = " << t.mu + << ", phi = " << t.phi; + EXPECT_NEAR(mu_adj, t.d_mu, 1e-11 * std::max(1.0, std::fabs(t.d_mu))) + << signature << ": n = " << t.n << ", mu = " << t.mu + << ", phi = " << t.phi; + const double gp = t.phi * t.d_phi; + EXPECT_NEAR(t.phi * phi_adj, gp, 1e-11 * std::max(1.0, std::fabs(gp))) + << signature << ": n = " << t.n << ", mu = " << t.mu + << ", phi = " << t.phi; + }; + for (const auto& t : test_values) { + const std::vector n{t.n}; + stan::math::matrix_cl n_cl(n); + + Eigen::Matrix mu(1); + mu << t.mu; + Eigen::Matrix phi(1); + phi << t.phi; + auto mu_cl = stan::math::to_matrix_cl(mu); + auto phi_cl = stan::math::to_matrix_cl(phi); + var lp = stan::math::neg_binomial_2_lpmf(n_cl, mu_cl, phi_cl); + lp.grad(); + expect_reference(t, lp, mu(0).adj(), phi(0).adj(), "vector"); + stan::math::recover_memory(); + + var mu_s = t.mu; + var phi_s = t.phi; + var lp_s = stan::math::neg_binomial_2_lpmf(n_cl, mu_s, phi_s); + lp_s.grad(); + expect_reference(t, lp_s, mu_s.adj(), phi_s.adj(), "scalar"); + stan::math::recover_memory(); + } +} + #endif diff --git a/test/unit/math/opencl/rev/neg_binomial_lpmf_test.cpp b/test/unit/math/opencl/rev/neg_binomial_lpmf_test.cpp index 28014935e8e..bebe1c8c77e 100644 --- a/test/unit/math/opencl/rev/neg_binomial_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/neg_binomial_lpmf_test.cpp @@ -173,4 +173,70 @@ TEST(muProbDistributionsNegBinomial, opencl_matches_cpu_big) { beta.transpose().eval()); } +TEST(muProbDistributionsNegBinomial, opencl_large_shapes_reference) { + // The tests above compare OpenCL against the CPU and cannot see an error + // that both share. These are absolute references from mpmath at 90 digits, + // checked against 140 digits, as in + // test/unit/math/rev/prob/large_shapes_test.cpp. The difference + // digamma(alpha + n) - digamma(alpha) keeps no correct digits for alpha + // near 1e15. The partials are compared as x * d/dx. The last row is a + // small shape where the plain difference is correct. + using stan::math::var; + struct TestValue { + int n; + double alpha; + double beta; + double value; + double d_alpha; + double d_beta; + }; + const std::vector test_values = { + {3, 0x1.c6bf526340000p+49, 0x1.2f2a36ecd5555p+48, -1.4959226032237274, + 1.3124999999999966e-30, 5.6249999999999838e-31}, + {1, 0x1.c6bf526340000p+49, 0x1.2f2a36ecd5555p+48, -1.9013877113318889, + -1.9999999999999957e-15, 5.9999999999999829e-15}, + {50, 0x1.c6bf526340000p+49, 0x1.2309ce5400000p+44, -2.8766166803657541, + 2.4999999999998758e-29, 0.0}, + {0, 0x1.9000000000000p+6, 0x1.0aaaaaaaaaaabp+5, -2.9558802241544401, + -2.9558802241544401e-2, 8.7378640776699017e-2}, + }; + auto expect_reference = [](const TestValue& t, const var& lp, + double alpha_adj, double beta_adj, + const char* signature) { + EXPECT_NEAR(lp.val(), t.value, 1e-12 * std::max(1.0, std::fabs(t.value))) + << signature << ": n = " << t.n << ", alpha = " << t.alpha + << ", beta = " << t.beta; + const double ga = t.alpha * t.d_alpha; + EXPECT_NEAR(t.alpha * alpha_adj, ga, 1e-11 * std::max(1.0, std::fabs(ga))) + << signature << ": n = " << t.n << ", alpha = " << t.alpha + << ", beta = " << t.beta; + const double gb = t.beta * t.d_beta; + EXPECT_NEAR(t.beta * beta_adj, gb, 1e-11 * std::max(1.0, std::fabs(gb))) + << signature << ": n = " << t.n << ", alpha = " << t.alpha + << ", beta = " << t.beta; + }; + for (const auto& t : test_values) { + const std::vector n{t.n}; + stan::math::matrix_cl n_cl(n); + + Eigen::Matrix alpha(1); + alpha << t.alpha; + Eigen::Matrix beta(1); + beta << t.beta; + auto alpha_cl = stan::math::to_matrix_cl(alpha); + auto beta_cl = stan::math::to_matrix_cl(beta); + var lp = stan::math::neg_binomial_lpmf(n_cl, alpha_cl, beta_cl); + lp.grad(); + expect_reference(t, lp, alpha(0).adj(), beta(0).adj(), "vector"); + stan::math::recover_memory(); + + var alpha_s = t.alpha; + var beta_s = t.beta; + var lp_s = stan::math::neg_binomial_lpmf(n_cl, alpha_s, beta_s); + lp_s.grad(); + expect_reference(t, lp_s, alpha_s.adj(), beta_s.adj(), "scalar"); + stan::math::recover_memory(); + } +} + #endif diff --git a/test/unit/math/opencl/rev/student_t_lpdf_test.cpp b/test/unit/math/opencl/rev/student_t_lpdf_test.cpp index 886f9687735..705648002bf 100644 --- a/test/unit/math/opencl/rev/student_t_lpdf_test.cpp +++ b/test/unit/math/opencl/rev/student_t_lpdf_test.cpp @@ -3,6 +3,8 @@ #include #include #include +#include +#include #include TEST(ProbDistributionsStudentT, error_checking) { @@ -160,4 +162,65 @@ TEST(ProbDistributionsStudentT, opencl_matches_cpu_big) { nu.transpose().eval(), mu.transpose().eval(), sigma.transpose().eval()); } +TEST(ProbDistributionsStudentT, opencl_large_shapes_reference) { + // The tests above compare OpenCL against the CPU and cannot see an error + // that both share. These are absolute references from mpmath, as in + // test/unit/math/rev/prob/large_shapes_test.cpp. The differences + // lgamma(nu/2 + 1/2) - lgamma(nu/2) and digamma(nu/2 + 1/2) - + // digamma(nu/2) keep no correct digits for large nu. The nu and sigma + // partials are compared on the log scale. The last row is a small nu + // where the plain differences are correct. + using stan::math::var; + struct TestValue { + double y; + double nu; + double mu; + double sigma; + double value; + double d_y; + double d_nu; + double d_mu; + double d_sigma; + }; + const std::vector test_values = { + {0x1.0000000000000p+0, 0x1.6bcc41e900000p+46, 0x0.0p+0, + 0x1.0000000000000p+0, -1.4189385332046777, -1.0000000000000000, + 4.9999999999999833e-29, 1.0000000000000000, 2.6727647100921956e-51}, + {-0x1.cd1e504efb30cp+1, 0x1.7cc4b890abebfp+45, 0x1.64703afcf3380p-4, + 0x1.c43477d4ae376p+3, -3.6014210338578538, 1.8475570774651872e-2, + 1.0330563009181428e-28, -1.8475570774651872e-2, -6.5940664388226970e-2}, + {0x1.8000000000000p+1, 0x1.6bcc41e900000p+46, 0x0.0p+0, + 0x1.0000000000000p+1, -2.7370857137646191, -7.4999999999999062e-1, + 1.0937500000001266e-29, 7.4999999999999062e-1, 6.2499999999998594e-1}, + {-0x1.c8dd60359f470p-1, 0x1.da533b967d6d3p-4, -0x1.572a29ad6231cp-1, + 0x1.9d9944276e11dp-3, -1.6062891011707614, 4.5853946438927804, + 6.8870497923078714, -4.5853946438927804, 9.0518700772802596e-2}, + }; + auto near = [](double got, double expected, double tol) { + return std::fabs(got - expected) + <= tol * std::max(1.0, std::fabs(expected)); + }; + for (const auto& t : test_values) { + Eigen::Matrix y(1), nu(1), mu(1), sigma(1); + y << t.y; + nu << t.nu; + mu << t.mu; + sigma << t.sigma; + auto y_cl = stan::math::to_matrix_cl(y); + auto nu_cl = stan::math::to_matrix_cl(nu); + auto mu_cl = stan::math::to_matrix_cl(mu); + auto sigma_cl = stan::math::to_matrix_cl(sigma); + var lp = stan::math::student_t_lpdf(y_cl, nu_cl, mu_cl, sigma_cl); + lp.grad(); + EXPECT_TRUE(near(lp.val(), t.value, 1e-12)) << "value, nu = " << t.nu; + EXPECT_TRUE(near(y(0).adj(), t.d_y, 1e-11)) << "d_y, nu = " << t.nu; + EXPECT_TRUE(near(t.nu * nu(0).adj(), t.nu * t.d_nu, 1e-11)) + << "nu * d_nu, nu = " << t.nu; + EXPECT_TRUE(near(mu(0).adj(), t.d_mu, 1e-11)) << "d_mu, nu = " << t.nu; + EXPECT_TRUE(near(t.sigma * sigma(0).adj(), t.sigma * t.d_sigma, 1e-11)) + << "sigma * d_sigma, nu = " << t.nu; + stan::math::recover_memory(); + } +} + #endif From f950044e363e99303b3e9064519d034d324a70af Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Thu, 8 Oct 2026 19:45:27 +0300 Subject: [PATCH 14/15] Fix the n range check and the phi derivative sum in two OpenCL lpmfs --- stan/math/opencl/prim/beta_binomial_lpmf.hpp | 3 +- .../prim/neg_binomial_2_log_glm_lpmf.hpp | 5 ++- .../opencl/rev/beta_binomial_lpmf_test.cpp | 33 ++++++++++++++++++ .../rev/neg_binomial_2_log_glm_lpmf_test.cpp | 34 +++++++++++++++++++ 4 files changed, 71 insertions(+), 4 deletions(-) diff --git a/stan/math/opencl/prim/beta_binomial_lpmf.hpp b/stan/math/opencl/prim/beta_binomial_lpmf.hpp index 7563d779483..c81d5baef67 100644 --- a/stan/math/opencl/prim/beta_binomial_lpmf.hpp +++ b/stan/math/opencl/prim/beta_binomial_lpmf.hpp @@ -95,8 +95,9 @@ inline return_type_t beta_binomial_lpmf( matrix_cl beta_deriv_cl; results(check_N_nonnegative, check_alpha_pos_finite, check_beta_pos_finite, - logp_cl, alpha_deriv_cl, beta_deriv_cl) + logp_cl, return_neg_inf_cl, alpha_deriv_cl, beta_deriv_cl) = expressions(N_nonnegative, alpha_pos_finite, beta_pos_finite, logp_expr, + return_neg_inf, calc_if>(alpha_deriv), calc_if>(beta_deriv)); diff --git a/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp b/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp index 2e40a7e2bcd..ec9622752d1 100644 --- a/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp +++ b/stan/math/opencl/prim/neg_binomial_2_log_glm_lpmf.hpp @@ -124,9 +124,8 @@ neg_binomial_2_log_glm_lpmf(const T_y_cl& y, const T_x_cl& x, const bool need_theta_derivative_sum = need_theta_derivative && !is_alpha_vector; matrix_cl theta_derivative_sum_cl(wgs, 1); - const bool need_phi_derivative_sum = !is_alpha_vector; - const bool need_phi_derivative - = is_autodiff_v || need_phi_derivative_sum; + const bool need_phi_derivative = is_autodiff_v; + const bool need_phi_derivative_sum = need_phi_derivative && !is_phi_vector; matrix_cl phi_derivative_cl( need_phi_derivative ? (need_phi_derivative_sum ? wgs : N) : 0, 1); // binomial_coefficient_log(y + phi - 1, y) is computed in the kernel if it diff --git a/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp b/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp index 9960d9eaff3..bf5e12ffd4b 100644 --- a/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/beta_binomial_lpmf_test.cpp @@ -80,6 +80,39 @@ TEST(ProbDistributionsBetaBinomial, error_checking) { std::domain_error); } +TEST(ProbDistributionsBetaBinomial, opencl_n_outside_support) { + // n < 0 or n > N has probability 0: LOG_ZERO, also under propto, and no + // gradient, as on the CPU + std::vector N{5, 5, 5}; + Eigen::VectorXd alpha(3); + alpha << 0.3, 1.8, 1.3; + Eigen::VectorXd beta(3); + beta << 0.3, 1.8, 1.2; + stan::math::matrix_cl N_cl(N); + stan::math::matrix_cl alpha_cl(alpha); + stan::math::matrix_cl beta_cl(beta); + for (const std::vector& n : + {std::vector{2, 6, 1}, std::vector{2, -1, 1}}) { + stan::math::matrix_cl n_cl(n); + EXPECT_EQ(stan::math::beta_binomial_lpmf(n, N, alpha, beta), + stan::math::LOG_ZERO); + EXPECT_EQ(stan::math::beta_binomial_lpmf(n_cl, N_cl, alpha_cl, beta_cl), + stan::math::LOG_ZERO); + stan::math::var_value> alpha_v + = stan::math::to_matrix_cl(alpha); + stan::math::var lp + = stan::math::beta_binomial_lpmf(n_cl, N_cl, alpha_v, beta_cl); + stan::math::var lp_propto + = stan::math::beta_binomial_lpmf(n_cl, N_cl, alpha_v, beta_cl); + EXPECT_EQ(lp.val(), stan::math::LOG_ZERO); + EXPECT_EQ(lp_propto.val(), stan::math::LOG_ZERO); + (lp + lp_propto).grad(); + Eigen::VectorXd alpha_adj = stan::math::from_matrix_cl(alpha_v.adj()); + EXPECT_TRUE((alpha_adj.array() == 0).all()) << alpha_adj.transpose(); + stan::math::recover_memory(); + } +} + auto beta_binomial_lpmf_functor = [](const auto& n, const auto& N, const auto& alpha, const auto& beta) { return stan::math::beta_binomial_lpmf(n, N, alpha, beta); diff --git a/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp b/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp index 53058a466f1..577903144ba 100644 --- a/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp +++ b/test/unit/math/opencl/rev/neg_binomial_2_log_glm_lpmf_test.cpp @@ -220,6 +220,40 @@ TEST(ProbDistributionsNegBinomial2LogGLM, neg_binomial_2_log_glm_lpmf_functor_propto, y, x, alpha, beta, phi); } +TEST(ProbDistributionsNegBinomial2LogGLM, opencl_matches_cpu_mixed_alpha_phi) { + // A scalar alpha with a vector phi, and a vector alpha with a scalar phi. + // N = 153 needs more than one work group. + for (int N : {3, 153}) { + int M = 2; + // deterministic values, so that the random inputs of the later tests do + // not change + vector y(N); + Matrix x(N, M); + Matrix alpha_vec(N, 1); + Matrix phi_vec(N, 1); + for (int i = 0; i < N; i++) { + y[i] = (i * 7) % 13; + x(i, 0) = std::sin(0.7 * i); + x(i, 1) = std::cos(1.3 * i); + alpha_vec(i) = 0.5 * std::sin(0.3 * i); + phi_vec(i) = 0.1 + 0.05 * (i % 29); + } + Matrix beta(M, 1); + beta << 0.3, -0.2; + double alpha = 0.3; + double phi = 13.2; + + stan::math::test::compare_cpu_opencl_prim_rev( + neg_binomial_2_log_glm_lpmf_functor, y, x, alpha, beta, phi_vec); + stan::math::test::compare_cpu_opencl_prim_rev( + neg_binomial_2_log_glm_lpmf_functor_propto, y, x, alpha, beta, phi_vec); + stan::math::test::compare_cpu_opencl_prim_rev( + neg_binomial_2_log_glm_lpmf_functor, y, x, alpha_vec, beta, phi); + stan::math::test::compare_cpu_opencl_prim_rev( + neg_binomial_2_log_glm_lpmf_functor_propto, y, x, alpha_vec, beta, phi); + } +} + TEST(ProbDistributionsNegBinomial2LogGLM, opencl_matches_cpu_big) { int N = 153; int M = 71; From 468ff26dbbe32fa2c662e03ae0474cce54431270 Mon Sep 17 00:00:00 2001 From: Aki Vehtari Date: Fri, 9 Oct 2026 12:48:16 +0300 Subject: [PATCH 15/15] Use the AgradRev fixture in the new rev tests and remove a C-style cast --- .../math/opencl/kernels/device_functions/digamma_diff.hpp | 2 +- test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp | 4 ++-- test/unit/math/rev/prob/large_shapes_test.cpp | 8 +++++--- 3 files changed, 8 insertions(+), 6 deletions(-) diff --git a/stan/math/opencl/kernels/device_functions/digamma_diff.hpp b/stan/math/opencl/kernels/device_functions/digamma_diff.hpp index 4211b91297a..97253a3ce87 100644 --- a/stan/math/opencl/kernels/device_functions/digamma_diff.hpp +++ b/stan/math/opencl/kernels/device_functions/digamma_diff.hpp @@ -50,7 +50,7 @@ static constexpr const char* digamma_diff_device_function // from the smallest (exactly 1 / x for d = 1) if (d <= 8.0 && d == floor(d)) { double sum = 0.0; - for (int j = (int)d - 1; j >= 0; --j) { + for (int j = convert_int(d) - 1; j >= 0; --j) { sum += 1.0 / (x + j); } return sum; diff --git a/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp b/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp index 1ff5e18e742..8db6b0bed92 100644 --- a/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp +++ b/test/unit/math/rev/prob/beta_binomial_lpmf_test.cpp @@ -43,7 +43,7 @@ std::vector testValues = { }; } // namespace beta_binomial_lpmf_rev_test_internal -TEST(ProbDistributionsBetaBinomial, log_shape_gradients) { +TEST_F(AgradRev, ProbDistributionsBetaBinomial_log_shape_gradients) { using beta_binomial_lpmf_rev_test_internal::TestValue; using beta_binomial_lpmf_rev_test_internal::testValues; using stan::math::var; @@ -72,7 +72,7 @@ TEST(ProbDistributionsBetaBinomial, log_shape_gradients) { } } -TEST(ProbDistributionsBetaBinomial, log_concentration_gradient) { +TEST_F(AgradRev, ProbDistributionsBetaBinomial_log_concentration_gradient) { // alpha = beta = exp(lc) / 2, n = 57, N = 117. The // gradient in lc goes to 0 like 27 / exp(lc). develop returned 0 or // noise (-0.14 at lc = 32) from lc = 20 on, which made the log density a diff --git a/test/unit/math/rev/prob/large_shapes_test.cpp b/test/unit/math/rev/prob/large_shapes_test.cpp index fd5b286e4d1..834e942cf11 100644 --- a/test/unit/math/rev/prob/large_shapes_test.cpp +++ b/test/unit/math/rev/prob/large_shapes_test.cpp @@ -202,6 +202,7 @@ void eval(const std::string& tag, const std::vector& a, double& value, } // clang-format off +// NOLINTBEGIN(whitespace/line_length) const std::vector test_cases = { // neg_binomial_lpmf(n | alpha, beta): n, alpha, beta {"NB", {0x1.8000000000000p+1, 0x1.c6bf526340000p+49, 0x1.2f2a36ecd5555p+48}, -1.4959226032237274, {1.3124999999999966e-30, 5.6249999999999838e-31}}, @@ -303,10 +304,11 @@ const std::vector test_cases = { {"LBETA", {0x1.0000000000000p-1, 0x1.6bcc41e900000p+46}, -1.5545730708033618e+1, {-3.4199701327938063e+1, -5.0000000000000125e-15}}, {"LBETA", {0x1.ba1a7a7909c76p-7, 0x1.f4f68d902030cp-6}, 4.6705195755596318, {-5.1474613897767773e+1, -1.0033990131067132e+1}}, }; +// NOLINTEND // clang-format on } // namespace large_shapes_test_internal -TEST(ProbDistributions, large_shapes_value_and_log_scale_gradients) { +TEST_F(AgradRev, ProbDistributions_large_shapes_value_and_log_scale_gradients) { using large_shapes_test_internal::eval; using large_shapes_test_internal::log_scale_args; using large_shapes_test_internal::test_cases; @@ -332,7 +334,7 @@ TEST(ProbDistributions, large_shapes_value_and_log_scale_gradients) { } } -TEST(ProbDistributions, lkj_corr_eta_gradient_at_one) { +TEST_F(AgradRev, ProbDistributions_lkj_corr_eta_gradient_at_one) { // At eta == 1.0 exactly, develop returned the constant without its // derivative, so d/deta lost sum_k psi(eta + (K - 1) / 2) - psi(eta + // (K - 1 - k) / 2). eta = exp(0) = 1 is the value at the default @@ -347,7 +349,7 @@ TEST(ProbDistributions, lkj_corr_eta_gradient_at_one) { stan::math::recover_memory(); } -TEST(ProbDistributions, large_shapes_propto_dropped_terms) { +TEST_F(AgradRev, ProbDistributions_large_shapes_propto_dropped_terms) { // full - propto must be the sum of the terms that propto drops, computed // here from their formulas. Some of these terms are inside a function // call: lgamma(n) in lbeta(n, alpha + 1), lgamma(y + 1) in lchoose.