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
29 changes: 27 additions & 2 deletions stan/math/fwd/fun/lbeta.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,38 @@
#include <stan/math/fwd/core.hpp>

#include <stan/math/fwd/fun/digamma.hpp>
#include <stan/math/prim/fun/digamma_diff.hpp>
#include <stan/math/prim/fun/lbeta.hpp>
#include <stan/math/prim/fun/value_of_rec.hpp>

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 <typename T1, typename T2>
inline return_type_t<T1, T2> 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 <typename T>
inline fvar<T> lbeta(const fvar<T>& x1, const fvar<T>& x2) {
if (value_of_rec(x1.val_) >= internal::digamma_diff_min_x
|| value_of_rec(x2.val_) >= internal::digamma_diff_min_x) {
return fvar<T>(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<T>(lbeta(x1.val_, x2.val_),
x1.d_ * digamma(x1.val_) + x2.d_ * digamma(x2.val_)
- (x1.d_ + x2.d_) * digamma(x1.val_ + x2.val_));
Expand All @@ -20,13 +45,13 @@ inline fvar<T> lbeta(const fvar<T>& x1, const fvar<T>& x2) {
template <typename T>
inline fvar<T> lbeta(double x1, const fvar<T>& x2) {
return fvar<T>(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 <typename T>
inline fvar<T> lbeta(const fvar<T>& x1, double x2) {
return fvar<T>(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
Expand Down
104 changes: 104 additions & 0 deletions stan/math/opencl/kernel_generator/elt_function_cl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#include <stan/math/opencl/kernels/device_functions/binomial_coefficient_log.hpp>
#include <stan/math/opencl/kernels/device_functions/beta.hpp>
#include <stan/math/opencl/kernels/device_functions/digamma.hpp>
#include <stan/math/opencl/kernels/device_functions/digamma_diff.hpp>
#include <stan/math/opencl/kernels/device_functions/erfcx.hpp>
#include <stan/math/opencl/kernels/device_functions/inv_logit.hpp>
#include <stan/math/opencl/kernels/device_functions/inv_Phi.hpp>
Expand All @@ -14,6 +15,7 @@
#include <stan/math/opencl/kernels/device_functions/lgamma_stirling.hpp>
#include <stan/math/opencl/kernels/device_functions/lgamma_stirling_diff.hpp>
#include <stan/math/opencl/kernels/device_functions/lmultiply.hpp>
#include <stan/math/opencl/kernels/device_functions/log_beta_ratio.hpp>
#include <stan/math/opencl/kernels/device_functions/log_inv_logit.hpp>
#include <stan/math/opencl/kernels/device_functions/log_inv_logit_diff.hpp>
#include <stan/math/opencl/kernels/device_functions/log_diff_exp.hpp>
Expand Down Expand Up @@ -410,6 +412,108 @@ const std::vector<const char*> lbeta_<T1, T2>::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 <typename T1, typename T2, typename T3, typename T4>
class log_beta_ratio_ : public elt_function_cl<log_beta_ratio_<T1, T2, T3, T4>,
double, T1, T2, T3, T4> {
using base = elt_function_cl<log_beta_ratio_<T1, T2, T3, T4>, double, T1, T2,
T3, T4>;
using base::arguments_;

public:
using base::cols;
using base::rows;
static const std::vector<const char*> includes;
explicit log_beta_ratio_(T1&& alpha, T2&& beta, T3&& n, T4&& m)
: base("stan_log_beta_ratio", std::forward<T1>(alpha),
std::forward<T2>(beta), std::forward<T3>(n), std::forward<T4>(m)) {
const std::array<int, 4> 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<int, 4> 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<decltype(arg1_copy)>,
std::remove_reference_t<decltype(arg2_copy)>,
std::remove_reference_t<decltype(arg3_copy)>,
std::remove_reference_t<decltype(arg4_copy)>>{
std::move(arg1_copy), std::move(arg2_copy), std::move(arg3_copy),
std::move(arg4_copy)};
}
inline std::pair<int, int> 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 <typename T1, typename T2, typename T3, typename T4,
require_all_kernel_expressions_t<T1, T2, T3, T4>* = nullptr,
require_any_not_stan_scalar_t<T1, T2, T3, T4>* = nullptr>
inline log_beta_ratio_<as_operation_cl_t<T1>, as_operation_cl_t<T2>,
as_operation_cl_t<T3>, as_operation_cl_t<T4>>
log_beta_ratio(T1&& alpha, T2&& beta, T3&& n, T4&& m) {
return log_beta_ratio_<as_operation_cl_t<T1>, as_operation_cl_t<T2>,
as_operation_cl_t<T3>, as_operation_cl_t<T4>>(
as_operation_cl(std::forward<T1>(alpha)),
as_operation_cl(std::forward<T2>(beta)),
as_operation_cl(std::forward<T3>(n)),
as_operation_cl(std::forward<T4>(m)));
}

template <typename T1, typename T2, typename T3, typename T4>
const std::vector<const char*> log_beta_ratio_<T1, T2, T3, T4>::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,
Expand Down
115 changes: 115 additions & 0 deletions stan/math/opencl/kernels/device_functions/digamma_diff.hpp
Original file line number Diff line number Diff line change
@@ -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 <stan/math/opencl/stringify.hpp>
#include <string>

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 = convert_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
Loading
Loading