Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 23 additions & 16 deletions stan/math/prim/err/check_3F2_converges.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -50,24 +50,31 @@ inline void check_3F2_converges(const char* function, const T_a1& a1,
check_not_nan("check_3F2_converges", "b2", b2);
check_not_nan("check_3F2_converges", "z", z);

int num_terms = 0;
// A numerator parameter a = -m (m a nonnegative integer) makes the terms
// zero from term m + 1 on, so the series is a polynomial that ends at the
// smallest such m.
bool is_polynomial = false;
double num_terms = 0;
auto add_end = [&](const auto& a) {
if (is_nonpositive_integer(a)) {
const double m = fabs(value_of_rec(a));
num_terms = is_polynomial ? std::fmin(num_terms, m) : m;
is_polynomial = true;
}
};
add_end(a1);
add_end(a2);
add_end(a3);

if (is_nonpositive_integer(a1) && fabs(a1) >= num_terms) {
is_polynomial = true;
num_terms = floor(fabs(value_of_rec(a1)));
}
if (is_nonpositive_integer(a2) && fabs(a2) >= num_terms) {
is_polynomial = true;
num_terms = floor(fabs(value_of_rec(a2)));
}
if (is_nonpositive_integer(a3) && fabs(a3) >= num_terms) {
is_polynomial = true;
num_terms = floor(fabs(value_of_rec(a3)));
}

bool is_undefined = (is_nonpositive_integer(b1) && fabs(b1) <= num_terms)
|| (is_nonpositive_integer(b2) && fabs(b2) <= num_terms);
// A denominator parameter b = -m (m a nonnegative integer) is a pole from
// term m + 1 on, because (b)_k is zero for k > m. A polynomial ends at
// term num_terms, so there b is a pole only if m < num_terms; in an
// infinite series b is always a pole.
auto is_pole = [&](const auto& b) {
return is_nonpositive_integer(b)
&& (!is_polynomial || fabs(value_of_rec(b)) < num_terms);
};
bool is_undefined = is_pole(b1) || is_pole(b2);

if (is_polynomial && !is_undefined) {
return;
Expand Down
1 change: 1 addition & 0 deletions stan/math/prim/fun.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,7 @@
#include <stan/math/prim/fun/hypergeometric_2F1.hpp>
#include <stan/math/prim/fun/hypergeometric_2F2.hpp>
#include <stan/math/prim/fun/hypergeometric_3F2.hpp>
#include <stan/math/prim/fun/hypergeometric_3F2_tail_bound.hpp>
#include <stan/math/prim/fun/hypergeometric_pFq.hpp>
#include <stan/math/prim/fun/hypot.hpp>
#include <stan/math/prim/fun/identity_matrix.hpp>
Expand Down
123 changes: 115 additions & 8 deletions stan/math/prim/fun/grad_F32.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,12 @@
#include <stan/math/prim/fun/constants.hpp>
#include <stan/math/prim/fun/exp.hpp>
#include <stan/math/prim/fun/fabs.hpp>
#include <stan/math/prim/fun/hypergeometric_3F2_tail_bound.hpp>
#include <stan/math/prim/fun/inv.hpp>
#include <stan/math/prim/fun/is_nonpositive_integer.hpp>
#include <stan/math/prim/fun/log.hpp>
#include <stan/math/prim/fun/value_of_rec.hpp>
#include <array>
#include <cmath>

namespace stan {
Expand All @@ -21,6 +25,12 @@ namespace math {
* to within <code>precision</code> or throwing when the
* function takes <code>max_steps</code> steps.
*
* A terminating series (a1, a2 or a3 a non-positive integer) is a
* polynomial. Its sum stops at the first zero term, or earlier when bounds
* of the sums of all remaining terms of the series and of every requested
* gradient are below 1e-17 times the absolute values of their partial sums;
* <code>precision</code> does not apply to it.
*
* This power-series representation converges for all gradients
* under the same conditions as the 3F2 function itself.
*
Expand All @@ -45,7 +55,9 @@ namespace math {
* @param[in] b1 b1 see generalized hypergeometric function definition.
* @param[in] b2 b2 see generalized hypergeometric function definition.
* @param[in] z z see generalized hypergeometric function definition.
* @param[in] precision precision of the infinite sum
* @param[in] precision precision of the infinite sum: the sum stops after a
* term below <code>precision</code> in absolute value. Not used for a
* terminating series.
* @param[in] max_steps number of steps to take
*/
template <bool grad_a1 = true, bool grad_a2 = true, bool grad_a3 = true,
Expand Down Expand Up @@ -77,22 +89,59 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
for (int i = 0; i < 6; ++i) {
log_g_old_sign[i] = 1.0;
}
// A terminating series stops when bounds of the sums of all remaining
// terms of the series and of every requested gradient are negligible
// against the partial sums.
const bool is_polynomial = is_nonpositive_integer(a1)
|| is_nonpositive_integer(a2)
|| is_nonpositive_integer(a3);
constexpr double polynomial_precision = 1e-17;
const std::array<double, 3> a_val{value_of_rec(a1), value_of_rec(a2),
value_of_rec(a3)};
const std::array<double, 2> b_val{value_of_rec(b1), value_of_rec(b2)};
const double abs_z = std::fabs(value_of_rec(z));
// index of the last term that can be nonzero
double last = max_steps + 1.0;
for (double a_i : a_val) {
if (a_i <= 0.0 && a_i == std::floor(a_i)) {
last = std::fmin(last, -a_i);
}
}
auto is_negligible = [](double increment, const auto& sum) {
return increment <= polynomial_precision * std::fabs(value_of_rec(sum));
};
double f_sum = 1.0;
std::array<T1, 6> term{0};
for (int k = 0; k <= max_steps; ++k) {
// A numerator parameter that has reached zero ends the series (a
// polynomial). Stop before the ratio, which is 0 / 0 when a denominator
// parameter reaches zero at the same k.
if (a1 + k == 0 || a2 + k == 0 || a3 + k == 0) {
return;
}
T1 p = (a1 + k) * (a2 + k) * (a3 + k) / ((b1 + k) * (b2 + k) * (1 + k));
if (p == 0) {
return;
}

log_t_new += log(fabs(p)) + log_z;
log_t_new_sign = p >= 0.0 ? log_t_new_sign : -log_t_new_sign;
bool negligible = false;
double t_new = 0.0;
if (is_polynomial) {
t_new = std::exp(value_of_rec(log_t_new));
f_sum += value_of_rec(log_t_new_sign) * t_new;
negligible = is_negligible(t_new, f_sum);
}
if constexpr (grad_a1) {
term[0]
= log_g_old_sign[0] * log_t_old_sign * exp(log_g_old[0] - log_t_old)
+ inv(a1 + k);
log_g_old[0] = log_t_new + log(fabs(term[0]));
log_g_old_sign[0] = term[0] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[0] += log_g_old_sign[0] * exp(log_g_old[0]);
const T1 increment = exp(log_g_old[0]);
g[0] += log_g_old_sign[0] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[0]);
}

if constexpr (grad_a2) {
Expand All @@ -101,7 +150,9 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
+ inv(a2 + k);
log_g_old[1] = log_t_new + log(fabs(term[1]));
log_g_old_sign[1] = term[1] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[1] += log_g_old_sign[1] * exp(log_g_old[1]);
const T1 increment = exp(log_g_old[1]);
g[1] += log_g_old_sign[1] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[1]);
}

if constexpr (grad_a3) {
Expand All @@ -110,7 +161,9 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
+ inv(a3 + k);
log_g_old[2] = log_t_new + log(fabs(term[2]));
log_g_old_sign[2] = term[2] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[2] += log_g_old_sign[2] * exp(log_g_old[2]);
const T1 increment = exp(log_g_old[2]);
g[2] += log_g_old_sign[2] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[2]);
}

if constexpr (grad_b1) {
Expand All @@ -119,7 +172,9 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
- inv(b1 + k);
log_g_old[3] = log_t_new + log(fabs(term[3]));
log_g_old_sign[3] = term[3] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[3] += log_g_old_sign[3] * exp(log_g_old[3]);
const T1 increment = exp(log_g_old[3]);
g[3] += log_g_old_sign[3] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[3]);
}

if constexpr (grad_b2) {
Expand All @@ -128,7 +183,9 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
- inv(b2 + k);
log_g_old[4] = log_t_new + log(fabs(term[4]));
log_g_old_sign[4] = term[4] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[4] += log_g_old_sign[4] * exp(log_g_old[4]);
const T1 increment = exp(log_g_old[4]);
g[4] += log_g_old_sign[4] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[4]);
}

if constexpr (grad_z) {
Expand All @@ -137,10 +194,60 @@ inline void grad_F32(T1* g, const T2& a1, const T3& a2, const T4& a3,
+ inv(z);
log_g_old[5] = log_t_new + log(fabs(term[5]));
log_g_old_sign[5] = term[5] >= 0.0 ? log_t_new_sign : -log_t_new_sign;
g[5] += log_g_old_sign[5] * exp(log_g_old[5]);
const T1 increment = exp(log_g_old[5]);
g[5] += log_g_old_sign[5] * increment;
negligible = negligible && is_negligible(value_of_rec(increment), g[5]);
}

if (log_t_new <= log(precision)) {
if (is_polynomial) {
if (negligible) {
// The remaining terms t_{k + 2}, ..., t_last have the ratios r_j for
// j = k + 1, ..., last - 1. A gradient term is t_j s_j, where s_j
// (term[i] for j = k + 1) grows by at most c per step.
const double n_left = last - k - 1;
if (n_left <= 0) {
return;
}
const double lo = k + 1;
const double hi = last - 1;
const auto sums = internal::hypergeometric_3F2_tail_sums(
internal::hypergeometric_3F2_ratio_bound(a_val, b_val, abs_z, lo,
hi),
n_left);
auto tail_is_negligible = [&](const auto& s, double c,
const auto& sum) {
return is_negligible(
t_new
* (std::fabs(value_of_rec(s)) * sums.first + c * sums.second),
sum);
};
auto growth = [&](double p) {
return 1.0 / internal::hypergeometric_3F2_min_abs(p, lo, hi);
};
bool done = is_negligible(t_new * sums.first, f_sum);
if constexpr (grad_a1) {
done = done && tail_is_negligible(term[0], growth(a_val[0]), g[0]);
}
if constexpr (grad_a2) {
done = done && tail_is_negligible(term[1], growth(a_val[1]), g[1]);
}
if constexpr (grad_a3) {
done = done && tail_is_negligible(term[2], growth(a_val[2]), g[2]);
}
if constexpr (grad_b1) {
done = done && tail_is_negligible(term[3], growth(b_val[0]), g[3]);
}
if constexpr (grad_b2) {
done = done && tail_is_negligible(term[4], growth(b_val[1]), g[4]);
}
if constexpr (grad_z) {
done = done && tail_is_negligible(term[5], 1.0 / abs_z, g[5]);
}
if (done) {
return;
}
}
} else if (log_t_new <= log(precision)) {
return; // implicit abs
}

Expand Down
82 changes: 74 additions & 8 deletions stan/math/prim/fun/hypergeometric_3F2.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,19 +8,44 @@
#include <stan/math/prim/fun/to_vector.hpp>
#include <stan/math/prim/fun/constants.hpp>
#include <stan/math/prim/fun/fabs.hpp>
#include <stan/math/prim/fun/hypergeometric_3F2_tail_bound.hpp>
#include <stan/math/prim/fun/hypergeometric_pFq.hpp>
#include <stan/math/prim/fun/sum.hpp>
#include <stan/math/prim/fun/sign.hpp>
#include <stan/math/prim/fun/value_of_rec.hpp>
#include <array>
#include <cmath>

namespace stan {
namespace math {
namespace internal {
/**
* Sum the power series of the hypergeometric function 3F2 term by term.
*
* `hypergeometric_3F2` calls this function at z = 1 with sum(b) <= sum(a).
* There `check_3F2_converges` accepts only a terminating series (a
* polynomial), so the sum stops at the first zero term. It stops earlier
* when a bound of the sum of all remaining terms is below `precision` times
* the absolute value of the partial sum.
*
* @tparam Ta type of Eigen/Std vector 'a' arguments
* @tparam Tb type of Eigen/Std vector 'b' arguments
* @tparam Tz type of z argument
* @param[in] a numerator parameters
* @param[in] b denominator parameters
* @param[in] z argument
* @param[in] precision relative precision of the sum. The default 1e-17 is
* below half an ulp of the sum.
* @param[in] max_steps number of steps to take
* @return the sum of the series
* @throw std::domain_error if the sum overflows or needs more than
* max_steps steps
*/
template <typename Ta, typename Tb, typename Tz,
require_all_vector_t<Ta, Tb>* = nullptr,
require_stan_scalar_t<Tz>* = nullptr>
inline return_type_t<Ta, Tb, Tz> hypergeometric_3F2_infsum(
const Ta& a, const Tb& b, const Tz& z, double precision = 1e-6,
const Ta& a, const Tb& b, const Tz& z, double precision = 1e-17,
int max_steps = 1e5) {
using T_return = return_type_t<Ta, Tb, Tz>;
Eigen::Array<scalar_type_t<Ta>, 3, 1> a_array = as_array_or_scalar(a);
Expand All @@ -37,9 +62,46 @@ inline return_type_t<Ta, Tb, Tz> hypergeometric_3F2_infsum(
int z_sign = sign(value_of_rec(z));
int t_sign = z_sign * a_signs.prod() * b_signs.prod();

// For the bound of the remaining terms: the parameters, and the index of
// the last term that can be nonzero (the first zero numerator ends the
// series, otherwise the steps end the sum)
const std::array<double, 3> a_val{value_of_rec(a_array[0]),
value_of_rec(a_array[1]),
value_of_rec(a_array[2])};
const std::array<double, 2> b_val{value_of_rec(b_array[0]),
value_of_rec(b_array[1])};
const double abs_z = std::fabs(value_of_rec(z));
double last = max_steps + 1.0;
for (double a_i : a_val) {
if (a_i <= 0.0 && a_i == std::floor(a_i)) {
last = std::fmin(last, -a_i);
}
}

int k = 0;
const double log_precision = log(precision);
while (k <= max_steps && log_t >= log_precision) {
double abs_term = 1.0;
while (k <= max_steps) {
// A numerator parameter that has reached zero makes this term and every
// later term zero: the series is a polynomial and has ended. Without
// this stop the sign below is 0 while the magnitude keeps growing, and
// 0 * inf gives NaN.
if ((value_of_rec(a_array) == 0.0).any()) {
return t_acc;
}
// Stop when the last term t_k and a bound of the sum of the remaining
// terms t_{k + 1}, ..., t_last are negligible against the partial sum.
// The ratios of the remaining terms are r_j for j = k, ..., last - 1.
const double abs_sum = std::fabs(value_of_rec(t_acc));
if (abs_term <= precision * abs_sum) {
const double ratio_bound
= hypergeometric_3F2_ratio_bound(a_val, b_val, abs_z, k, last - 1);
const double tail
= abs_term
* hypergeometric_3F2_tail_sums(ratio_bound, last - k).first;
if (tail <= precision * abs_sum) {
return t_acc;
}
}
// Replace zero values with 1 prior to taking the log so that we accumulate
// 0.0 rather than -inf
const auto& abs_apk = math::fabs((a_array == 0).select(1.0, a_array));
Expand All @@ -50,7 +112,9 @@ inline return_type_t<Ta, Tb, Tz> hypergeometric_3F2_infsum(
}

log_t += p + log_z;
t_acc += t_sign * exp(log_t);
const auto term = exp(log_t);
t_acc += t_sign * term;
abs_term = value_of_rec(term);

if (is_inf(t_acc)) {
throw_domain_error("hypergeometric_3F2", "sum (output)", t_acc,
Expand All @@ -61,9 +125,10 @@ inline return_type_t<Ta, Tb, Tz> hypergeometric_3F2_infsum(
b_array += 1.0;
a_signs = sign(value_of_rec(a_array));
b_signs = sign(value_of_rec(b_array));
t_sign = a_signs.prod() * b_signs.prod() * t_sign;
t_sign = z_sign * a_signs.prod() * b_signs.prod() * t_sign;
}
if (k == max_steps) {
// The loop ends with k = max_steps + 1 when the steps run out
if (k > max_steps) {
throw_domain_error("hypergeometric_3F2", "k (internal counter)", max_steps,
"exceeded iterations, hypergeometric function did not ",
"converge.");
Expand Down Expand Up @@ -123,8 +188,9 @@ inline auto hypergeometric_3F2(const Ta& a, const Tb& b, const Tz& z) {
check_3F2_converges("hypergeometric_3F2", a_ref[0], a_ref[1], a_ref[2],
b_ref[0], b_ref[1], z);
// Boost's pFq throws convergence errors in some cases, fallback to naive
// infinite-sum approach (tests pass for these)
if (z == 1.0 && (sum(b_ref) - sum(a_ref)) < 0.0) {
// infinite-sum approach (tests pass for these). At z = 1 Boost also throws
// when sum(b) == sum(a), also for a terminating series.
if (z == 1.0 && (sum(b_ref) - sum(a_ref)) <= 0.0) {
return internal::hypergeometric_3F2_infsum(a_ref, b_ref, z);
}
return hypergeometric_pFq(a_ref, b_ref, z);
Expand Down
Loading
Loading