diff --git a/stan/math/mix/prob/laplace_latent_neg_binomial_2_log_rng.hpp b/stan/math/mix/prob/laplace_latent_neg_binomial_2_log_rng.hpp index 3cd8df1fbfd..0f764b3ae64 100644 --- a/stan/math/mix/prob/laplace_latent_neg_binomial_2_log_rng.hpp +++ b/stan/math/mix/prob/laplace_latent_neg_binomial_2_log_rng.hpp @@ -28,7 +28,8 @@ namespace math { * @tparam RNG A valid boost rng type * @param[in] y Observed counts. * @param[in] y_index Index indicating which group each observation belongs to. - * @param[in] eta Overdisperison parameter. + * @param[in] eta the overdispersion parameter: a scalar shared by all + * groups, or a vector with one entry per group. * @param[in] mean The mean of the latent normal variable. * \laplace_common_args * @param[in] hessian_block_size Block size for the Hessian approximation with @@ -72,7 +73,8 @@ inline Eigen::VectorXd laplace_latent_tol_neg_binomial_2_log_rng( * @tparam RNG A valid boost rng type * @param[in] y Observed counts. * @param[in] y_index Index indicating which group each observation belongs to. - * @param[in] eta Overdisperison parameter. + * @param[in] eta the overdispersion parameter: a scalar shared by all + * groups, or a vector with one entry per group. * @param[in] mean The mean of the latent normal variable. * \laplace_common_args * @param[in] hessian_block_size Block size for the Hessian approximation with diff --git a/stan/math/mix/prob/laplace_marginal_bernoulli_logit_lpmf.hpp b/stan/math/mix/prob/laplace_marginal_bernoulli_logit_lpmf.hpp index 4bbb6eeff20..79e1598486d 100644 --- a/stan/math/mix/prob/laplace_marginal_bernoulli_logit_lpmf.hpp +++ b/stan/math/mix/prob/laplace_marginal_bernoulli_logit_lpmf.hpp @@ -46,7 +46,7 @@ struct bernoulli_logit_likelihood { Eigen::VectorXd counts_per_group = Eigen::VectorXd::Zero(theta.size()); Eigen::VectorXd n_per_group = Eigen::VectorXd::Zero(theta.size()); - for (int i = 0; i < theta.size(); i++) { + for (size_t i = 0; i < y_index.size(); i++) { counts_per_group(y_index[i] - 1) += y[i]; n_per_group(y_index[i] - 1) += 1; } diff --git a/stan/math/mix/prob/laplace_marginal_neg_binomial_2_log_lpmf.hpp b/stan/math/mix/prob/laplace_marginal_neg_binomial_2_log_lpmf.hpp index 843f9d2e4dd..699a54b914a 100644 --- a/stan/math/mix/prob/laplace_marginal_neg_binomial_2_log_lpmf.hpp +++ b/stan/math/mix/prob/laplace_marginal_neg_binomial_2_log_lpmf.hpp @@ -20,12 +20,18 @@ #include #include #include +#include #include namespace stan { namespace math { struct neg_binomial_2_log_likelihood { + /** + * Returns the lpmf for a negative binomial (2nd parameterization, log + * link) across multiple groups. The dispersion `eta` is either a scalar + * shared by all groups or a vector with one entry per group. + */ template * = nullptr> inline auto operator()(const ThetaVec& theta, const Eta& eta, @@ -35,23 +41,49 @@ struct neg_binomial_2_log_likelihood { Eigen::VectorXi n_per_group = Eigen::VectorXi::Zero(theta.size()); Eigen::VectorXi counts_per_group = Eigen::VectorXi::Zero(theta.size()); - for (int i = 0; i < y.size(); i++) { + for (size_t i = 0; i < y.size(); i++) { n_per_group[y_index[i] - 1]++; counts_per_group[y_index[i] - 1] += y[i]; } Eigen::Map y_map(y.data(), y.size()); auto theta_offset = add(theta, mean); - auto log_eta = log(eta); - auto lse = to_ref(log_sum_exp(theta_offset, log_eta)); + if constexpr (is_stan_scalar::value) { + // one dispersion shared by all groups + auto log_eta = log(eta); + auto lse = to_ref(log_sum_exp(theta_offset, log_eta)); - return sum(binomial_coefficient_log(subtract(add(y_map, eta), 1.0), y_map)) - + sum(add( - // counts_per_group * (theta - log(eta + exp(theta))) - elt_multiply(counts_per_group, subtract(theta_offset, lse)), - // n_per_group * eta * (log(eta) - log(eta + exp(theta))) - elt_multiply(multiply(n_per_group, eta), - subtract(log_eta, lse)))); + return sum(binomial_coefficient_log(subtract(add(y_map, eta), 1.0), + y_map)) + + sum(add( + // counts_per_group * (theta - log(eta + exp(theta))) + elt_multiply(counts_per_group, subtract(theta_offset, lse)), + // n_per_group * eta * (log(eta) - log(eta + exp(theta))) + elt_multiply(multiply(n_per_group, eta), + subtract(log_eta, lse)))); + } else { + // one dispersion per group + check_size_match("neg_binomial_2_log_likelihood", "eta", eta.size(), + "theta", theta.size()); + const auto& eta_ref = to_ref(eta); + auto log_eta = to_ref(log(eta_ref)); + auto lse = to_ref(log_sum_exp(theta_offset, log_eta)); + // y + eta - 1 with the dispersion of each observation's group + Eigen::Matrix, Eigen::Dynamic, 1> y_plus_eta_m1( + y.size()); + for (size_t i = 0; i < y.size(); ++i) { + y_plus_eta_m1.coeffRef(i) + = eta_ref.coeff(y_index[i] - 1) + (y[i] - 1.0); + } + + return sum(binomial_coefficient_log(y_plus_eta_m1, y_map)) + + sum(add( + // counts_per_group * (theta - log(eta + exp(theta))) + elt_multiply(counts_per_group, subtract(theta_offset, lse)), + // n_per_group * eta * (log(eta) - log(eta + exp(theta))) + elt_multiply(elt_multiply(n_per_group, eta_ref), + subtract(log_eta, lse)))); + } } }; @@ -70,7 +102,8 @@ struct neg_binomial_2_log_likelihood { * @param[in] y observed counts. * @param[in] y_index group to which each observation belongs. Each group * is parameterized by one element of theta. - * @param[in] eta non-marginalized model parameters for the likelihood. + * @param[in] eta the overdispersion parameter: a scalar shared by all + * groups, or a vector with one entry per group. * @param[in] mean the mean of the latent normal variable * \laplace_common_args * @param[in] hessian_block_size Block size for the Hessian approximation with @@ -107,7 +140,8 @@ inline auto laplace_marginal_tol_neg_binomial_2_log_lpmf( * @param[in] y observed counts. * @param[in] y_index group to which each observation belongs. Each group * is parameterized by one element of theta. - * @param[in] eta Parameter argument for likelihood function. + * @param[in] eta the overdispersion parameter: a scalar shared by all + * groups, or a vector with one entry per group. * @param[in] mean the mean of the latent normal variable * \laplace_common_args * @param[in] hessian_block_size Block size for the Hessian approximation with @@ -129,8 +163,11 @@ inline auto laplace_marginal_neg_binomial_2_log_lpmf( } struct neg_binomial_2_log_likelihood_summary { + // Without a per-observation group index the dispersion cannot vary by + // group, so `eta` must be a scalar here. template * = nullptr> + require_eigen_vector_t* = nullptr, + require_stan_scalar_t* = nullptr> inline auto operator()(const ThetaVec& theta, const Eta& eta, const std::vector& y, const std::vector& n_per_group, diff --git a/stan/math/mix/prob/laplace_marginal_poisson_log_lpmf.hpp b/stan/math/mix/prob/laplace_marginal_poisson_log_lpmf.hpp index b8c715c6939..c9cb99fa629 100644 --- a/stan/math/mix/prob/laplace_marginal_poisson_log_lpmf.hpp +++ b/stan/math/mix/prob/laplace_marginal_poisson_log_lpmf.hpp @@ -39,15 +39,16 @@ struct poisson_log_likelihood { Eigen::VectorXd counts_per_group = Eigen::VectorXd::Zero(theta.size()); Eigen::VectorXd n_per_group = Eigen::VectorXd::Zero(theta.size()); - for (int i = 0; i < theta.size(); i++) { + double norm_constant = 0; + for (size_t i = 0; i < y_index.size(); i++) { counts_per_group(y_index[i] - 1) += y[i]; n_per_group(y_index[i] - 1) += 1; + norm_constant -= lgamma(y[i] + 1.0); } auto theta_offset = to_ref(add(theta, mean)); - return -sum(lgamma(add(counts_per_group, 1))) - + dot_product(theta_offset, counts_per_group) + return norm_constant + dot_product(theta_offset, counts_per_group) - dot_product(n_per_group, exp(theta_offset)); } }; diff --git a/test/unit/math/laplace/laplace_marginal_bernoulli_logit_lpmf_test.cpp b/test/unit/math/laplace/laplace_marginal_bernoulli_logit_lpmf_test.cpp index b0e503ec1cd..77dc178b1c5 100644 --- a/test/unit/math/laplace/laplace_marginal_bernoulli_logit_lpmf_test.cpp +++ b/test/unit/math/laplace/laplace_marginal_bernoulli_logit_lpmf_test.cpp @@ -78,4 +78,74 @@ TEST_P(laplace_marginal_bernoulli_logit_lpmf, phi_dim500) { LAPLACE_INSTANTIATE_TEST_SUITE_P(laplace_marginal_bernoulli_logit_lpmf); +// Reference likelihood computed per observation, without the grouped +// sufficient-statistics shortcut used by bernoulli_logit_likelihood. +struct bernoulli_logit_obs_likelihood { + template + auto operator()(const Theta& theta, const std::vector& y, + const std::vector& y_index, const Mean& mean, + std::ostream* /*pstream*/) const { + Eigen::Matrix, Eigen::Dynamic, 1> + theta_obs(y.size()); + for (size_t i = 0; i < y.size(); ++i) { + theta_obs(i) = theta(y_index[i] - 1) + mean(y_index[i] - 1); + } + return stan::math::bernoulli_logit_lpmf(y, theta_obs); + } +}; + +// More observations than latent variables: groups with several observations +// must aggregate all of them, not just the first theta.size() entries. +TEST(laplace_marginal_bernoulli_logit_lpmf_grouped, multiple_obs_per_group) { + using stan::math::laplace_marginal_tol; + using stan::math::laplace_marginal_tol_bernoulli_logit_lpmf; + + constexpr int dim_theta = 3; + const std::vector y{1, 0, 1, 1, 1, 0}; + const std::vector y_index{1, 2, 3, 1, 2, 1}; + const Eigen::VectorXd mean{{0.3, -0.2, 0.1}}; + const std::vector x{ + Eigen::VectorXd{{0.05100797, 0.16086164}}, + Eigen::VectorXd{{-0.59823393, 0.98701425}}, + Eigen::VectorXd{{0.31296868, -0.68926772}}}; + const Eigen::VectorXd theta_0 = Eigen::VectorXd::Zero(dim_theta); + constexpr double alpha = 1.6; + constexpr double rho = 0.45; + constexpr double tolerance = 1e-12; + constexpr int max_num_steps = 1000; + constexpr int hessian_block_size = 1; + constexpr int solver = 1; + constexpr int max_steps_line_search = 0; + + const double marginal = laplace_marginal_tol_bernoulli_logit_lpmf( + y, y_index, mean, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + const double reference = laplace_marginal_tol( + bernoulli_logit_obs_likelihood{}, std::forward_as_tuple(y, y_index, mean), + hessian_block_size, stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + EXPECT_NEAR(reference, marginal, 1e-6); + + // derivatives w.r.t. the mean with more observations than latents + constexpr stan::test::ad_tolerances tols{ + stan::test::ad_gradient_tols{1e-8, 1e-3}}; + auto f = [&](auto&& mean_arg) { + return laplace_marginal_tol_bernoulli_logit_lpmf( + y, y_index, mean_arg, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + }; + stan::test::expect_ad(tols, f, mean); +} + } // namespace diff --git a/test/unit/math/laplace/laplace_marginal_neg_binomial_log_lpmf_test.cpp b/test/unit/math/laplace/laplace_marginal_neg_binomial_log_lpmf_test.cpp index cf2c078a2f3..c7ab6eafd77 100644 --- a/test/unit/math/laplace/laplace_marginal_neg_binomial_log_lpmf_test.cpp +++ b/test/unit/math/laplace/laplace_marginal_neg_binomial_log_lpmf_test.cpp @@ -121,4 +121,107 @@ TEST_P(laplace_disease_map_test, laplace_marginal_neg_binomial_2_log_lpmf) { LAPLACE_INSTANTIATE_TEST_SUITE_P(laplace_disease_map_test); +// Reference likelihood computed per observation with the prim lpmf, +// indexing the dispersion by group when it is a vector. +struct neg_binomial_2_log_obs_likelihood { + template + auto operator()(const Theta& theta, const Eta& eta, const std::vector& y, + const std::vector& y_index, const Mean& mean, + std::ostream* /*pstream*/) const { + Eigen::Matrix, Eigen::Dynamic, 1> + theta_obs(y.size()); + Eigen::Matrix, Eigen::Dynamic, 1> eta_obs( + y.size()); + for (size_t i = 0; i < y.size(); ++i) { + theta_obs(i) = theta(y_index[i] - 1) + mean(y_index[i] - 1); + if constexpr (stan::is_stan_scalar::value) { + eta_obs(i) = eta; + } else { + eta_obs(i) = eta(y_index[i] - 1); + } + } + return stan::math::neg_binomial_2_log_lpmf(y, theta_obs, eta_obs); + } +}; + +class laplace_marginal_neg_binomial_log_lpmf_grouped : public ::testing::Test { + protected: + const std::vector y{1, 0, 5, 2, 7, 0}; + const std::vector y_index{1, 2, 3, 1, 2, 1}; + const Eigen::VectorXd mean{{0.3, -0.2, 0.1}}; + const std::vector x{ + Eigen::VectorXd{{0.05100797, 0.16086164}}, + Eigen::VectorXd{{-0.59823393, 0.98701425}}, + Eigen::VectorXd{{0.31296868, -0.68926772}}}; + const Eigen::VectorXd theta_0 = Eigen::VectorXd::Zero(3); + static constexpr double alpha = 1.6; + static constexpr double rho = 0.45; + static constexpr double tolerance = 1e-12; + static constexpr int max_num_steps = 1000; + static constexpr int hessian_block_size = 1; + static constexpr int solver = 1; + static constexpr int max_steps_line_search = 0; + + template + auto marginal(const Eta& eta) { + return stan::math::laplace_marginal_tol_neg_binomial_2_log_lpmf( + y, y_index, eta, mean, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + } + template + auto reference(const Eta& eta) { + return stan::math::laplace_marginal_tol( + neg_binomial_2_log_obs_likelihood{}, + std::forward_as_tuple(eta, y, y_index, mean), hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + } +}; + +// Scalar dispersion with more observations than latent variables. +TEST_F(laplace_marginal_neg_binomial_log_lpmf_grouped, scalar_eta) { + constexpr double eta = 1.5; + EXPECT_NEAR(reference(eta), marginal(eta), 1e-6); +} + +// A vector dispersion (one entry per group) must work and match the +// per-observation reference. This is what stanc generates for the Stan +// signature, which requires `vector eta`. +TEST_F(laplace_marginal_neg_binomial_log_lpmf_grouped, vector_eta) { + const Eigen::VectorXd eta{{1.5, 2.5, 0.75}}; + EXPECT_NEAR(reference(eta), marginal(eta), 1e-6); + + // a constant vector dispersion agrees with the scalar version + const Eigen::VectorXd eta_rep = Eigen::VectorXd::Constant(3, 1.5); + EXPECT_NEAR(marginal(1.5), marginal(eta_rep), 1e-8); + + // dispersion vector must have one entry per latent variable + const Eigen::VectorXd eta_bad = Eigen::VectorXd::Constant(2, 1.5); + EXPECT_THROW(marginal(eta_bad), std::invalid_argument); +} + +// eta and mean must remain differentiable (autodiff arguments). +TEST_F(laplace_marginal_neg_binomial_log_lpmf_grouped, vector_eta_ad) { + const Eigen::VectorXd eta{{1.5, 2.5, 0.75}}; + constexpr stan::test::ad_tolerances tols{ + stan::test::ad_gradient_tols{1e-8, 1e-3}}; + auto f = [&](auto&& eta_arg, auto&& mean_arg) { + return stan::math::laplace_marginal_tol_neg_binomial_2_log_lpmf( + y, y_index, eta_arg, mean_arg, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + }; + stan::test::expect_ad(tols, f, eta, mean); +} + } // namespace diff --git a/test/unit/math/laplace/laplace_marginal_poisson_log_lpmf_test.cpp b/test/unit/math/laplace/laplace_marginal_poisson_log_lpmf_test.cpp index 6b9d86db888..2e40c30b95a 100644 --- a/test/unit/math/laplace/laplace_marginal_poisson_log_lpmf_test.cpp +++ b/test/unit/math/laplace/laplace_marginal_poisson_log_lpmf_test.cpp @@ -235,4 +235,74 @@ TEST_P(laplace_disease_map_test, laplace_marginal_poisson_log_lpmf) { } LAPLACE_INSTANTIATE_TEST_SUITE_P(laplace_disease_map_test); +// Reference likelihood computed per observation, without the grouped +// sufficient-statistics shortcut used by poisson_log_likelihood. +struct poisson_log_obs_likelihood { + template + auto operator()(const Theta& theta, const std::vector& y, + const std::vector& y_index, const Mean& mean, + std::ostream* /*pstream*/) const { + Eigen::Matrix, Eigen::Dynamic, 1> + theta_obs(y.size()); + for (size_t i = 0; i < y.size(); ++i) { + theta_obs(i) = theta(y_index[i] - 1) + mean(y_index[i] - 1); + } + return stan::math::poisson_log_lpmf(y, theta_obs); + } +}; + +// More observations than latent variables: groups with several observations +// must aggregate all of them, not just the first theta.size() entries. +TEST(laplace_marginal_poisson_log_lpmf_grouped, multiple_obs_per_group) { + using stan::math::laplace_marginal_tol; + using stan::math::laplace_marginal_tol_poisson_log_lpmf; + + constexpr int dim_theta = 3; + const std::vector y{1, 0, 3, 2, 4, 1}; + const std::vector y_index{1, 2, 3, 1, 2, 1}; + const Eigen::VectorXd mean{{0.3, -0.2, 0.1}}; + const std::vector x{ + Eigen::VectorXd{{0.05100797, 0.16086164}}, + Eigen::VectorXd{{-0.59823393, 0.98701425}}, + Eigen::VectorXd{{0.31296868, -0.68926772}}}; + const Eigen::VectorXd theta_0 = Eigen::VectorXd::Zero(dim_theta); + constexpr double alpha = 1.6; + constexpr double rho = 0.45; + constexpr double tolerance = 1e-12; + constexpr int max_num_steps = 1000; + constexpr int hessian_block_size = 1; + constexpr int solver = 1; + constexpr int max_steps_line_search = 0; + + const double marginal = laplace_marginal_tol_poisson_log_lpmf( + y, y_index, mean, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + const double reference = laplace_marginal_tol( + poisson_log_obs_likelihood{}, std::forward_as_tuple(y, y_index, mean), + hessian_block_size, stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + EXPECT_NEAR(reference, marginal, 1e-6); + + // derivatives w.r.t. the mean with more observations than latents + constexpr stan::test::ad_tolerances tols{ + stan::test::ad_gradient_tols{1e-8, 1e-3}}; + auto f = [&](auto&& mean_arg) { + return laplace_marginal_tol_poisson_log_lpmf( + y, y_index, mean_arg, hessian_block_size, + stan::math::test::squared_kernel_functor{}, + std::forward_as_tuple(x, alpha, rho), + std::make_tuple(theta_0, tolerance, max_num_steps, solver, + max_steps_line_search, true), + nullptr); + }; + stan::test::expect_ad(tols, f, mean); +} + } // namespace