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
60 changes: 59 additions & 1 deletion include/boost/math/distributions/poisson.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
#include <boost/math/special_functions/factorials.hpp> // factorials.
#include <boost/math/tools/roots.hpp> // for root finding.
#include <boost/math/distributions/detail/inv_discrete_quantile.hpp>
#include <boost/math/constants/constants.hpp>

namespace boost
{
Expand Down Expand Up @@ -141,6 +142,62 @@ namespace boost
return true;
} // bool check_dist_and_prob

template <class RealType, class Policy>
BOOST_MATH_GPU_ENABLED inline RealType stirlerr(const RealType& n, const Policy& pol) {
BOOST_MATH_STD_USING // for ADL of std functions.
using boost::math::lgamma;

// Use the direct formula for small n
if (n < boost::math::detail::minimum_argument_for_bernoulli_recursion<RealType>()) {
return lgamma(n + 1) - (n * log(n) - n + RealType(0.5) * log(RealType(2) * boost::math::constants::pi<RealType>() * n));
}

// Use the Stirling series approximation for large n
return boost::math::detail::bernoulli_stirling_series<RealType>(n, pol);

}

template <class RealType, class Policy>
BOOST_MATH_GPU_ENABLED inline RealType bd0(const RealType& mean, const RealType& k, const Policy& pol) {
BOOST_MATH_STD_USING // for ADL of std functions.

// Calculate v = (k - mean) / (k + mean) from Loader (2000) approximation
bool is_close = abs(k - mean) < RealType(0.1) * (k + mean);

if (is_close) { // Use the series approximation if |v| < 0.1
RealType v = (k - mean) / (k + mean);
RealType v2 = v * v;
RealType series_term = ((k - mean) * (k - mean)) / (k + mean);

RealType term = 2 * k * v;

// Series is trivially zero when k = mean, so skip the loop
if (term != RealType(0)) {
RealType target_epsilon = abs(term) * boost::math::tools::epsilon<RealType>();
const boost::math::size_t max_iterations = policies::get_max_series_iterations<Policy>();

for (boost::math::size_t i = 1U;; ++i) {
term *= v2;
RealType next = term / (2 * i + 1);
series_term += next;
// Break if the next term is less than the target epsilon
if (abs(next) < target_epsilon) {
break;
}
if (i > max_iterations) {
return policies::raise_evaluation_error<RealType>("bd0<%1%>()", "Series did not converge in the allotted iterations, best approximation was %1%", series_term, pol);
}
}
}
return series_term;
} else {
// Use the direct formula if |v| >= 0.1
RealType direct = (k == 0) ? RealType(0) : k * log(k / mean) + mean - k;
return direct;
}

}

} // namespace poisson_detail

BOOST_MATH_EXPORT template <class RealType = double, class Policy = policies::policy<> >
Expand Down Expand Up @@ -304,7 +361,8 @@ namespace boost
// Special case where k and lambda are both positive
if(k > 0 && mean > 0)
{
return -lgamma(k+1) + k*log(mean) - mean;
// Use the Loader (2000) saddle-point approximation for logpdf calculation
return -poisson_detail::stirlerr(k, Policy()) - poisson_detail::bd0(mean, k, Policy()) - RealType(0.5) * log(RealType(2) * boost::math::constants::pi<RealType>() * k);
}

result = log(pdf(dist, k));
Expand Down
29 changes: 20 additions & 9 deletions include/boost/math/special_functions/gamma.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -487,14 +487,9 @@ int minimum_argument_for_bernoulli_recursion()
}

template <class T, class Policy>
T scaled_tgamma_no_lanczos(const T& z, const Policy& pol, bool islog = false)
{
T bernoulli_stirling_series(const T& z, const Policy& pol) {
BOOST_MATH_STD_USING
//
// Calculates tgamma(z) / (z/e)^z
// Requires that our argument is large enough for Sterling's approximation to hold.
// Used internally when combining gamma's of similar magnitude without logarithms.
//

BOOST_MATH_ASSERT(minimum_argument_for_bernoulli_recursion<T>() <= z);

// Perform the Bernoulli series expansion of Stirling's approximation.
Expand Down Expand Up @@ -527,17 +522,33 @@ T scaled_tgamma_no_lanczos(const T& z, const Policy& pol, bool islog = false)
}
if (n > number_of_bernoullis_b2n)
// Safety net, we hope to never get here:
return policies::raise_evaluation_error("scaled_tgamma_no_lanczos<%1%>()", "Exceeded maximum series iterations without reaching convergence, best approximation was %1%", T(exp(sum) * half_ln_two_pi_over_z), pol); // LCOV_EXCL_LINE
return policies::raise_evaluation_error("bernoulli_stirling_series<%1%>()", "Exceeded maximum series iterations without reaching convergence, best approximation was %1%", sum, pol); // LCOV_EXCL_LINE

sum += term;

// Sanity check for divergence:
T fterm = fabs(term);
if(fterm > last_term)
// Safety net, we hope to never get here:
return policies::raise_evaluation_error("scaled_tgamma_no_lanczos<%1%>()", "Series became divergent without reaching convergence, best approximation was %1%", T(exp(sum) * half_ln_two_pi_over_z), pol); // LCOV_EXCL_LINE
return policies::raise_evaluation_error("bernoulli_stirling_series<%1%>()", "Series became divergent without reaching convergence, best approximation was %1%", sum, pol); // LCOV_EXCL_LINE
last_term = fterm;
}
return sum;
}

template <class T, class Policy>
T scaled_tgamma_no_lanczos(const T& z, const Policy& pol, bool islog = false)
{
BOOST_MATH_STD_USING
//
// Calculates tgamma(z) / (z/e)^z
// Requires that our argument is large enough for Sterling's approximation to hold.
// Used internally when combining gamma's of similar magnitude without logarithms.
//
BOOST_MATH_ASSERT(minimum_argument_for_bernoulli_recursion<T>() <= z);

T sum = bernoulli_stirling_series(z, pol);
const T half_ln_two_pi_over_z = sqrt(boost::math::constants::two_pi<T>() / z);

// Complete Stirling's approximation.
T scaled_gamma_value = islog ? T(sum + log(half_ln_two_pi_over_z)) : T(exp(sum) * half_ln_two_pi_over_z);
Expand Down
63 changes: 62 additions & 1 deletion test/test_poisson.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -244,7 +244,68 @@ void test_spots(RealType)
static_cast<RealType>(20)), // K>> mean
log(static_cast<RealType>(8.277463646553730E-009)), // probability.
tolerance);


// New test cases for Loader (2000) saddle-point approximation. Probs already
// in log space. Values calculated using mpmath (1000-digit precision).
BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(14)), // mean 14.
static_cast<RealType>(14)),
static_cast<RealType>(-2.244418568125061), // probability (already in log space).
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(20)), // mean 20.
static_cast<RealType>(18)),
static_cast<RealType>(-2.472264284061216), // probability (already in log space).
tolerance);

// Cases below require around 15+ significant decimal digits to represent
// k / mean meaningfully, so skip for float.
if (std::numeric_limits<RealType>::digits10 > 15)
{
BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1000000)), // mean 1000000.
static_cast<RealType>(1300000)),
static_cast<RealType>(-41081.501683746894),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1e8)), // mean 1e8.
static_cast<RealType>(8e7)),
static_cast<RealType>(-2148525.91257035),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1e10)), // mean 1e10.
static_cast<RealType>(105e9)),
static_cast<RealType>(-151894402015.7727),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1e15)), // mean 1e15.
static_cast<RealType>(8e14)), // |v| > 0.1 boundary
static_cast<RealType>(-21485158948650.273),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1e15)), // mean 1e15.
static_cast<RealType>(12e14)), // |v| < 0.1 boundary
static_cast<RealType>(-18785868152763.832),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(1e16)), // mean 1e16.
static_cast<RealType>(1e16)), // old formula returns 0.0 here
static_cast<RealType>(-19.339619277157038),
tolerance);

BOOST_CHECK_CLOSE(
logpdf(poisson_distribution<RealType>(static_cast<RealType>(5e15)), // mean 5e15.
static_cast<RealType>(5e15)),
static_cast<RealType>(-18.993045686877064),
tolerance);
}

// CDF
BOOST_CHECK_CLOSE(
cdf(poisson_distribution<RealType>(static_cast<RealType>(1)), // mean unity.
Expand Down