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
83 changes: 2 additions & 81 deletions stan/math/prim/prob/normal_lccdf.hpp
Original file line number Diff line number Diff line change
@@ -1,19 +1,7 @@
#ifndef STAN_MATH_PRIM_PROB_NORMAL_LCCDF_HPP
#define STAN_MATH_PRIM_PROB_NORMAL_LCCDF_HPP

#include <stan/math/prim/meta.hpp>
#include <stan/math/prim/err.hpp>
#include <stan/math/prim/fun/constants.hpp>
#include <stan/math/prim/fun/erf.hpp>
#include <stan/math/prim/fun/erfc.hpp>
#include <stan/math/prim/fun/exp.hpp>
#include <stan/math/prim/fun/log.hpp>
#include <stan/math/prim/fun/max_size.hpp>
#include <stan/math/prim/fun/scalar_seq_view.hpp>
#include <stan/math/prim/fun/size_zero.hpp>
#include <stan/math/prim/fun/value_of.hpp>
#include <stan/math/prim/functor/partials_propagator.hpp>
#include <cmath>
#include <stan/math/prim/prob/normal_lcdf.hpp>

namespace stan {
namespace math {
Expand All @@ -24,74 +12,7 @@ template <typename T_y, typename T_loc, typename T_scale,
inline return_type_t<T_y, T_loc, T_scale> normal_lccdf(const T_y& y,
const T_loc& mu,
const T_scale& sigma) {
using T_partials_return = partials_return_t<T_y, T_loc, T_scale>;
using std::exp;
using std::log;
using T_y_ref = ref_type_t<T_y>;
using T_mu_ref = ref_type_t<T_loc>;
using T_sigma_ref = ref_type_t<T_scale>;
static constexpr const char* function = "normal_lccdf";
check_consistent_sizes(function, "Random variable", y, "Location parameter",
mu, "Scale parameter", sigma);
T_y_ref y_ref = y;
T_mu_ref mu_ref = mu;
T_sigma_ref sigma_ref = sigma;
check_not_nan(function, "Random variable", y_ref);
check_finite(function, "Location parameter", mu_ref);
check_positive(function, "Scale parameter", sigma_ref);

if (size_zero(y, mu, sigma)) {
return 0;
}

T_partials_return ccdf_log(0.0);
auto ops_partials = make_partials_propagator(y_ref, mu_ref, sigma_ref);

scalar_seq_view<T_y_ref> y_vec(y_ref);
scalar_seq_view<T_mu_ref> mu_vec(mu_ref);
scalar_seq_view<T_sigma_ref> sigma_vec(sigma_ref);
size_t N = max_size(y, mu, sigma);

for (size_t n = 0; n < N; n++) {
const T_partials_return y_dbl = y_vec.val(n);
const T_partials_return mu_dbl = mu_vec.val(n);
const T_partials_return sigma_dbl = sigma_vec.val(n);

const T_partials_return scaled_diff
= (y_dbl - mu_dbl) / (sigma_dbl * SQRT_TWO);

T_partials_return one_m_erf;
if (scaled_diff < -37.5 * INV_SQRT_TWO) {
one_m_erf = 2.0;
} else if (scaled_diff < -5.0 * INV_SQRT_TWO) {
one_m_erf = 2.0 - erfc(-scaled_diff);
} else if (scaled_diff > 8.25 * INV_SQRT_TWO) {
one_m_erf = 0.0;
} else {
one_m_erf = 1.0 - erf(scaled_diff);
}

ccdf_log += LOG_HALF + log(one_m_erf);

if constexpr (is_any_autodiff_v<T_y, T_loc, T_scale>) {
const T_partials_return rep_deriv_div_sigma
= scaled_diff > 8.25 * INV_SQRT_TWO
? INFTY
: SQRT_TWO_OVER_SQRT_PI * exp(-scaled_diff * scaled_diff)
/ one_m_erf / sigma_dbl;
if constexpr (is_autodiff_v<T_y>) {
partials<0>(ops_partials)[n] -= rep_deriv_div_sigma;
}
if constexpr (is_autodiff_v<T_loc>) {
partials<1>(ops_partials)[n] += rep_deriv_div_sigma;
}
if constexpr (is_autodiff_v<T_scale>) {
partials<2>(ops_partials)[n]
+= rep_deriv_div_sigma * scaled_diff * SQRT_TWO;
}
}
}
return ops_partials.build(ccdf_log);
return normal_lcdf(-as_array_or_scalar(y), -as_array_or_scalar(mu), sigma);
}

} // namespace math
Expand Down
60 changes: 2 additions & 58 deletions stan/math/prim/prob/std_normal_lccdf.hpp
Original file line number Diff line number Diff line change
@@ -1,19 +1,7 @@
#ifndef STAN_MATH_PRIM_PROB_STD_NORMAL_LCCDF_HPP
#define STAN_MATH_PRIM_PROB_STD_NORMAL_LCCDF_HPP

#include <stan/math/prim/meta.hpp>
#include <stan/math/prim/err.hpp>
#include <stan/math/prim/fun/constants.hpp>
#include <stan/math/prim/fun/erf.hpp>
#include <stan/math/prim/fun/erfc.hpp>
#include <stan/math/prim/fun/exp.hpp>
#include <stan/math/prim/fun/log.hpp>
#include <stan/math/prim/fun/scalar_seq_view.hpp>
#include <stan/math/prim/fun/size.hpp>
#include <stan/math/prim/fun/size_zero.hpp>
#include <stan/math/prim/fun/value_of.hpp>
#include <stan/math/prim/functor/partials_propagator.hpp>
#include <cmath>
#include <stan/math/prim/prob/std_normal_lcdf.hpp>

namespace stan {
namespace math {
Expand All @@ -22,51 +10,7 @@ template <
typename T_y,
require_all_not_nonscalar_prim_or_rev_kernel_expression_t<T_y>* = nullptr>
inline return_type_t<T_y> std_normal_lccdf(const T_y& y) {
using T_partials_return = partials_return_t<T_y>;
using std::exp;
using std::log;
using T_y_ref = ref_type_t<T_y>;
static constexpr const char* function = "std_normal_lccdf";
T_y_ref y_ref = y;
check_not_nan(function, "Random variable", y_ref);

if (size_zero(y)) {
return 0;
}

T_partials_return lccdf(0.0);
auto ops_partials = make_partials_propagator(y_ref);

scalar_seq_view<T_y_ref> y_vec(y_ref);
size_t N = stan::math::size(y);

for (size_t n = 0; n < N; n++) {
const T_partials_return y_dbl = y_vec.val(n);
const T_partials_return scaled_y = y_dbl * INV_SQRT_TWO;

T_partials_return one_m_erf;
if (y_dbl < -37.5) {
one_m_erf = 2.0;
} else if (y_dbl < -5.0) {
one_m_erf = 2.0 - erfc(-scaled_y);
} else if (y_dbl > 8.25) {
one_m_erf = 0.0;
} else {
one_m_erf = 1.0 - erf(scaled_y);
}

lccdf += LOG_HALF + log(one_m_erf);

if constexpr (is_autodiff_v<T_y>) {
const T_partials_return rep_deriv
= y_dbl > 8.25
? INFTY
: SQRT_TWO_OVER_SQRT_PI * exp(-scaled_y * scaled_y) / one_m_erf;
partials<0>(ops_partials)[n] -= rep_deriv;
}
}

return ops_partials.build(lccdf);
return std_normal_lcdf(-as_array_or_scalar(y));
}

} // namespace math
Expand Down
137 changes: 108 additions & 29 deletions test/unit/math/prim/prob/normal_ccdf_log_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,34 +16,113 @@ TEST(ProbNormal, ccdf_log_matches_lccdf) {
TEST(ProbNormal, lccdf_tail) {
using stan::math::normal_lccdf;

EXPECT_FLOAT_EQ(-6.661338147750941214694e-16, normal_lccdf(-8.0, 0, 1));
EXPECT_FLOAT_EQ(-3.186340080674249758114e-14, normal_lccdf(-7.5, 0, 1));
EXPECT_FLOAT_EQ(-1.279865102788699562477e-12, normal_lccdf(-7.0, 0, 1));
EXPECT_FLOAT_EQ(-4.015998644826973564545e-11, normal_lccdf(-6.5, 0, 1));
EXPECT_FLOAT_EQ(-9.865877009111571184118e-10, normal_lccdf(-6.0, 0, 1));
EXPECT_FLOAT_EQ(-1.898956265833514866414e-08, normal_lccdf(-5.5, 0, 1));
EXPECT_FLOAT_EQ(-2.866516130081049047962e-07, normal_lccdf(-5.0, 0, 1));
EXPECT_FLOAT_EQ(-3.397678896843115195074e-06, normal_lccdf(-4.5, 0, 1));
EXPECT_FLOAT_EQ(-3.167174337748932124543e-05, normal_lccdf(-4.0, 0, 1));
EXPECT_FLOAT_EQ(-0.0002326561413768195969113, normal_lccdf(-3.5, 0, 1));
EXPECT_FLOAT_EQ(-0.001350809964748202673598, normal_lccdf(-3.0, 0, 1));
EXPECT_FLOAT_EQ(-0.0062290254858600267035, normal_lccdf(-2.5, 0, 1));
EXPECT_FLOAT_EQ(-0.02301290932896348992442, normal_lccdf(-2.0, 0, 1));
EXPECT_FLOAT_EQ(-0.06914345561223400604689, normal_lccdf(-1.5, 0, 1));
EXPECT_FLOAT_EQ(-0.1727537790234499048836, normal_lccdf(-1.0, 0, 1));
// The test values come from 4.6.1
//
// q <- seq(-10, 37.5, by = 0.5)
// for (i in 1:length(q)) {
// cat(
// sprintf(
// "EXPECT_FLOAT_EQ(%.22g, normal_lccdf(%.1f, 0, 1));\n",
// pnorm(q[i], lower.tail = FALSE, log.p = TRUE),
// q[i]
// )
// )
// }

EXPECT_FLOAT_EQ(-7.619853024160526919908e-24, normal_lccdf(-10.0, 0, 1));
EXPECT_FLOAT_EQ(-1.049451507536260815732e-21, normal_lccdf(-9.5, 0, 1));
EXPECT_FLOAT_EQ(-1.128588405953840782741e-19, normal_lccdf(-9.0, 0, 1));
EXPECT_FLOAT_EQ(-9.479534822203319190782e-18, normal_lccdf(-8.5, 0, 1));
EXPECT_FLOAT_EQ(-6.220960574271786832586e-16, normal_lccdf(-8.0, 0, 1));
EXPECT_FLOAT_EQ(-3.19089167291094746711e-14, normal_lccdf(-7.5, 0, 1));
EXPECT_FLOAT_EQ(-1.279812543886654064677e-12, normal_lccdf(-7.0, 0, 1));
EXPECT_FLOAT_EQ(-4.016000583939758853991e-11, normal_lccdf(-6.5, 0, 1));
EXPECT_FLOAT_EQ(-9.865876455243755941787e-10, normal_lccdf(-6.0, 0, 1));
EXPECT_FLOAT_EQ(-1.898956264618946382412e-08, normal_lccdf(-5.5, 0, 1));
EXPECT_FLOAT_EQ(-2.866516129637635770404e-07, normal_lccdf(-5.0, 0, 1));
EXPECT_FLOAT_EQ(-3.397678896834465718134e-06, normal_lccdf(-4.5, 0, 1));
EXPECT_FLOAT_EQ(-3.167174337748926703532e-05, normal_lccdf(-4.0, 0, 1));
EXPECT_FLOAT_EQ(-0.0002326561413768044451859, normal_lccdf(-3.5, 0, 1));
EXPECT_FLOAT_EQ(-0.001350809964748193783141, normal_lccdf(-3.0, 0, 1));
EXPECT_FLOAT_EQ(-0.006229025485860002417371, normal_lccdf(-2.5, 0, 1));
EXPECT_FLOAT_EQ(-0.02301290932896349339387, normal_lccdf(-2.0, 0, 1));
EXPECT_FLOAT_EQ(-0.0691434556122339921691, normal_lccdf(-1.5, 0, 1));
EXPECT_FLOAT_EQ(-0.172753779023449877128, normal_lccdf(-1.0, 0, 1));
EXPECT_FLOAT_EQ(-0.3689464152886565151412, normal_lccdf(-0.5, 0, 1));
EXPECT_FLOAT_EQ(-0.6931471805599452862268, normal_lccdf(0, 0, 1));
EXPECT_FLOAT_EQ(-1.175911761593618320987, normal_lccdf(0.5, 0, 1));
EXPECT_FLOAT_EQ(-1.841021645009263352222, normal_lccdf(1.0, 0, 1));
EXPECT_FLOAT_EQ(-2.705944400823889317564, normal_lccdf(1.5, 0, 1));
EXPECT_FLOAT_EQ(-3.78318433368203210776, normal_lccdf(2.0, 0, 1));
EXPECT_FLOAT_EQ(-5.081648277278686620662, normal_lccdf(2.5, 0, 1));
EXPECT_FLOAT_EQ(-6.607726221510342945464, normal_lccdf(3.0, 0, 1));
EXPECT_FLOAT_EQ(-8.366065308344028395027, normal_lccdf(3.5, 0, 1));
EXPECT_FLOAT_EQ(-10.36010148652728979357, normal_lccdf(4.0, 0, 1));
EXPECT_FLOAT_EQ(-12.59241973571053385683, normal_lccdf(4.5, 0, 1));
EXPECT_FLOAT_EQ(-15.06499839383403838156, normal_lccdf(5.0, 0, 1));
EXPECT_FLOAT_EQ(-17.77937635198566113104, normal_lccdf(5.5, 0, 1));
EXPECT_FLOAT_EQ(-20.73676889383495947072, normal_lccdf(6.0, 0, 1));
EXPECT_FLOAT_EQ(-23.93814997800869548428, normal_lccdf(6.5, 0, 1));
EXPECT_FLOAT_EQ(-0.6931471805599452862268, normal_lccdf(0.0, 0, 1));
EXPECT_FLOAT_EQ(-1.175911761593618543031, normal_lccdf(0.5, 0, 1));
EXPECT_FLOAT_EQ(-1.841021645009263574266, normal_lccdf(1.0, 0, 1));
EXPECT_FLOAT_EQ(-2.705944400823889761654, normal_lccdf(1.5, 0, 1));
EXPECT_FLOAT_EQ(-3.78318433368203166367, normal_lccdf(2.0, 0, 1));
EXPECT_FLOAT_EQ(-5.081648277278690173375, normal_lccdf(2.5, 0, 1));
EXPECT_FLOAT_EQ(-6.607726221510349162713, normal_lccdf(3.0, 0, 1));
EXPECT_FLOAT_EQ(-8.36606530834409412023, normal_lccdf(3.5, 0, 1));
EXPECT_FLOAT_EQ(-10.36010148652729156993, normal_lccdf(4.0, 0, 1));
EXPECT_FLOAT_EQ(-12.59241973571307937618, normal_lccdf(4.5, 0, 1));
EXPECT_FLOAT_EQ(-15.06499839398872531149, normal_lccdf(5.0, 0, 1));
EXPECT_FLOAT_EQ(-17.77937635262525972735, normal_lccdf(5.5, 0, 1));
EXPECT_FLOAT_EQ(-20.73676894997470654403, normal_lccdf(6.0, 0, 1));
EXPECT_FLOAT_EQ(-23.93814949516183787637, normal_lccdf(6.5, 0, 1));
EXPECT_FLOAT_EQ(-27.38430749881107573174, normal_lccdf(7.0, 0, 1));
EXPECT_FLOAT_EQ(-31.07589090289000210987, normal_lccdf(7.5, 0, 1));
EXPECT_FLOAT_EQ(-35.0134371599145524101, normal_lccdf(8.0, 0, 1));
EXPECT_FLOAT_EQ(-39.19739642821767233727, normal_lccdf(8.5, 0, 1));
EXPECT_FLOAT_EQ(-43.62814911333211398414, normal_lccdf(9.0, 0, 1));
EXPECT_FLOAT_EQ(-48.30601929896523216712, normal_lccdf(9.5, 0, 1));
EXPECT_FLOAT_EQ(-53.23128515051246978373, normal_lccdf(10.0, 0, 1));
EXPECT_FLOAT_EQ(-58.40418706107324453569, normal_lccdf(10.5, 0, 1));
EXPECT_FLOAT_EQ(-63.82493409442371756768, normal_lccdf(11.0, 0, 1));
EXPECT_FLOAT_EQ(-69.49370912909535036306, normal_lccdf(11.5, 0, 1));
EXPECT_FLOAT_EQ(-75.41067300156879582573, normal_lccdf(12.0, 0, 1));
EXPECT_FLOAT_EQ(-81.57596787074388089422, normal_lccdf(12.5, 0, 1));
EXPECT_FLOAT_EQ(-87.98971997102252373679, normal_lccdf(13.0, 0, 1));
EXPECT_FLOAT_EQ(-94.65204188128289786164, normal_lccdf(13.5, 0, 1));
EXPECT_FLOAT_EQ(-101.5630344074499618046, normal_lccdf(14.0, 0, 1));
EXPECT_FLOAT_EQ(-108.7227881543204688342, normal_lccdf(14.5, 0, 1));
EXPECT_FLOAT_EQ(-116.1313848457116932877, normal_lccdf(15.0, 0, 1));
EXPECT_FLOAT_EQ(-123.788898439410374408, normal_lccdf(15.5, 0, 1));
EXPECT_FLOAT_EQ(-131.6953960737596958097, normal_lccdf(16.0, 0, 1));
EXPECT_FLOAT_EQ(-139.850938875285208951, normal_lccdf(16.5, 0, 1));
EXPECT_FLOAT_EQ(-148.255582650980386461, normal_lccdf(17.0, 0, 1));
EXPECT_FLOAT_EQ(-156.9093784843464050027, normal_lccdf(17.5, 0, 1));
EXPECT_FLOAT_EQ(-165.8123732507141880888, normal_lccdf(18.0, 0, 1));
EXPECT_FLOAT_EQ(-174.9646100645466049173, normal_lccdf(18.5, 0, 1));
EXPECT_FLOAT_EQ(-184.3661286691609575428, normal_lccdf(19.0, 0, 1));
EXPECT_FLOAT_EQ(-194.0169657774974893982, normal_lccdf(19.5, 0, 1));
EXPECT_FLOAT_EQ(-203.9171553710972659701, normal_lccdf(20.0, 0, 1));
EXPECT_FLOAT_EQ(-214.066728963263813057, normal_lccdf(20.5, 0, 1));
EXPECT_FLOAT_EQ(-224.4657158314144851374, normal_lccdf(21.0, 0, 1));
EXPECT_FLOAT_EQ(-235.114143222833007485, normal_lccdf(21.5, 0, 1));
EXPECT_FLOAT_EQ(-246.0120365373809079301, normal_lccdf(22.0, 0, 1));
EXPECT_FLOAT_EQ(-257.1594194901841774481, normal_lccdf(22.5, 0, 1));
EXPECT_FLOAT_EQ(-268.5563142568631178619, normal_lccdf(23.0, 0, 1));
EXPECT_FLOAT_EQ(-280.2027416034976567971, normal_lccdf(23.5, 0, 1));
EXPECT_FLOAT_EQ(-292.0987210032077996402, normal_lccdf(24.0, 0, 1));
EXPECT_FLOAT_EQ(-304.2442707409637137062, normal_lccdf(24.5, 0, 1));
EXPECT_FLOAT_EQ(-316.6394080080202684258, normal_lccdf(25.0, 0, 1));
EXPECT_FLOAT_EQ(-329.2841489871795488398, normal_lccdf(25.5, 0, 1));
EXPECT_FLOAT_EQ(-342.1785089299278297403, normal_lccdf(26.0, 0, 1));
EXPECT_FLOAT_EQ(-355.3225022263559935709, normal_lccdf(26.5, 0, 1));
EXPECT_FLOAT_EQ(-368.7161424686563577779, normal_lccdf(27.0, 0, 1));
EXPECT_FLOAT_EQ(-382.3594425088898560716, normal_lccdf(27.5, 0, 1));
EXPECT_FLOAT_EQ(-396.252414511631059213, normal_lccdf(28.0, 0, 1));
EXPECT_FLOAT_EQ(-410.3950700020255908385, normal_lccdf(28.5, 0, 1));
EXPECT_FLOAT_EQ(-424.7874199097301470829, normal_lccdf(29.0, 0, 1));
EXPECT_FLOAT_EQ(-439.4294746091502474883, normal_lccdf(29.5, 0, 1));
EXPECT_FLOAT_EQ(-454.3212439563432099021, normal_lccdf(30.0, 0, 1));
EXPECT_FLOAT_EQ(-469.4627373229121189979, normal_lccdf(30.5, 0, 1));
EXPECT_FLOAT_EQ(-484.8539636271792687694, normal_lccdf(31.0, 0, 1));
EXPECT_FLOAT_EQ(-500.494931362897091276, normal_lccdf(31.5, 0, 1));
EXPECT_FLOAT_EQ(-516.3856486257253664007, normal_lccdf(32.0, 0, 1));
EXPECT_FLOAT_EQ(-532.5261231376803152671, normal_lccdf(32.5, 0, 1));
EXPECT_FLOAT_EQ(-548.9163622697381015314, normal_lccdf(33.0, 0, 1));
EXPECT_FLOAT_EQ(-565.5563730627579843713, normal_lccdf(33.5, 0, 1));
EXPECT_FLOAT_EQ(-582.4461622468717223455, normal_lccdf(34.0, 0, 1));
EXPECT_FLOAT_EQ(-599.5857362594723554139, normal_lccdf(34.5, 0, 1));
EXPECT_FLOAT_EQ(-616.9751012619225321032, normal_lccdf(35.0, 0, 1));
EXPECT_FLOAT_EQ(-634.6142631550883379532, normal_lccdf(35.5, 0, 1));
EXPECT_FLOAT_EQ(-652.5032275937984422853, normal_lccdf(36.0, 0, 1));
EXPECT_FLOAT_EQ(-670.642000000313714736, normal_lccdf(36.5, 0, 1));
EXPECT_FLOAT_EQ(-689.0305855768906440062, normal_lccdf(37.0, 0, 1));
EXPECT_FLOAT_EQ(-707.6689893175072256781, normal_lccdf(37.5, 0, 1));
}
Loading
Loading