Skip to content
Merged
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
148 changes: 95 additions & 53 deletions stan/math/mix/functor/laplace_marginal_density_estimator.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,8 +10,9 @@
#include <stan/math/prim/fun/quad_form_diag.hpp>
#include <stan/math/prim/fun/value_of.hpp>
#include <stan/math/prim/functor/iter_tuple_nested.hpp>
#include <unsupported/Eigen/MatrixFunctions>
#include <algorithm>
#include <cmath>
#include <limits>
#include <mutex>
#include <iomanip>

Expand Down Expand Up @@ -219,57 +220,76 @@ struct laplace_density_estimates {
};

/**
* Returns the principal square root of a block diagonal matrix.
* Returns the principal square root of a symmetric positive semi-definite
* block diagonal matrix.
*
* Each block is symmetrised and decomposed with a symmetric eigensolver.
* Eigenvalues that are negative only at rounding level, as the zero
* eigenvalues of a rank-deficient block are (for example the negative
* Hessian of a likelihood with more latent variables than observations),
* are clamped to zero. An eigenvalue below
* `-block_size * epsilon * max(|eigenvalues|)` means the block is not
* positive semi-definite.
*
* @tparam WRootMat A type inheriting from `Eigen::EigenBase`.
* @param W_root The output matrix to store the square root.
* @param W The input block diagonal matrix.
* @param block_size The size of each block in the block diagonal matrix.
* @throw std::domain_error if a block has non-finite entries or is not
* positive semi-definite.
*/
template <typename WRootMat>
inline void block_matrix_sqrt(WRootMat& W_root,
const Eigen::SparseMatrix<double>& W,
const Eigen::Index block_size) {
int n_block = W.cols() / block_size;
const Eigen::Index n_block = W.cols() / block_size;
Eigen::MatrixXd local_block(block_size, block_size);
Eigen::MatrixXd local_block_sqrt(block_size, block_size);
Eigen::MatrixXd sqrt_t_mat = Eigen::MatrixXd::Zero(block_size, block_size);
Eigen::SelfAdjointEigenSolver<Eigen::MatrixXd> eigensolver;
// No block operation available for sparse matrices, so we have to loop
// See https://eigen.tuxfamily.org/dox/group__TutorialSparse.html#title7
for (int i = 0; i < n_block; i++) {
sqrt_t_mat.setZero();
for (Eigen::Index i = 0; i < n_block; i++) {
local_block
= W.block(i * block_size, i * block_size, block_size, block_size);
if (!local_block.array().isFinite().any()) {
throw std::domain_error(
std::string("Error in block_matrix_sqrt: "
"NaNs detected in block diagonal starting at (")
+ std::to_string(i) + ", " + std::to_string(i) + ")");
if (unlikely(!local_block.array().isFinite().all())) {
[](auto i) STAN_COLD_PATH {
throw std::domain_error(
std::string("Error in block_matrix_sqrt: "
"non-finite values detected in block diagonal "
"starting at (")
+ std::to_string(i) + ", " + std::to_string(i) + ")");
}(i);
}
// Issue here, sqrt is done over T of the complex schur
Eigen::RealSchur<Eigen::MatrixXd> schurOfA(local_block);
// Compute Schur decomposition of arg
const auto& t_mat = schurOfA.matrixT();
const auto& u_mat = schurOfA.matrixU();
// Check if diagonal of schur is not positive
if ((t_mat.diagonal().array() < 0).any()) {
throw std::domain_error(
std::string("Error in block_matrix_sqrt: "
"values less than 0 detected in block diagonal's schur "
"decomposition starting at (")
+ std::to_string(i) + ", " + std::to_string(i) + ")");
local_block_sqrt = 0.5 * (local_block + local_block.transpose());
eigensolver.compute(local_block_sqrt);
if (unlikely(eigensolver.info() != Eigen::Success)) {
[](auto i) STAN_COLD_PATH {
throw std::domain_error(
std::string("Error in block_matrix_sqrt: "
"eigendecomposition failed for block diagonal "
"starting at (")
+ std::to_string(i) + ", " + std::to_string(i) + ")");
}(i);
}
try {
// Compute square root of T
Eigen::matrix_sqrt_quasi_triangular(t_mat, sqrt_t_mat);
// Compute square root of arg
local_block_sqrt = u_mat * sqrt_t_mat * u_mat.adjoint();
} catch (const std::exception& e) {
throw std::domain_error(
"Error in block_matrix_sqrt: "
"The matrix is not positive definite");
const Eigen::VectorXd eigenvalues = eigensolver.eigenvalues();
const double tolerance = block_size * std::numeric_limits<double>::epsilon()
* eigenvalues.cwiseAbs().maxCoeff();
if (unlikely(eigenvalues.minCoeff() < -tolerance)) {
[](auto&& i, auto&& eigenvalues) {
throw std::domain_error(
std::string("Error in block_matrix_sqrt: block diagonal starting "
"at (")
+ std::to_string(i) + ", " + std::to_string(i)
+ ") is not positive semi-definite (smallest eigenvalue "
+ std::to_string(eigenvalues.minCoeff()) + ")");
}(i, eigenvalues);
}
for (int k = 0; k < block_size; k++) {
for (int j = 0; j < block_size; j++) {
local_block_sqrt.noalias()
= eigensolver.eigenvectors()
* eigenvalues.cwiseMax(0.0).cwiseSqrt().asDiagonal()
* eigensolver.eigenvectors().transpose();
for (Eigen::Index k = 0; k < block_size; k++) {
for (Eigen::Index j = 0; j < block_size; j++) {
W_root.coeffRef(i * block_size + j, i * block_size + k)
= local_block_sqrt(j, k);
}
Expand Down Expand Up @@ -1021,37 +1041,51 @@ inline auto run_newton_loop(SolverPolicy& solver, NewtonStateT& state,
}

/**
* @brief Log a solver fallback event to the provided stream.
* @param[in] allow_fallthrough If false, throw instead of logging
* @brief Throw for a solver failure when falling through to the next solver
* is not allowed.
* @param[in] context Context string for the message
* @param[in] iter Current iteration number
* @param[in] failed_solver Name of the solver that failed
* @param[in] e Exception that caused the failure
*/
[[noreturn]] inline void throw_solver_failure(std::string_view context,
Eigen::Index iter,
std::string_view failed_solver,
const std::exception& e) {
std::ostringstream os;
os << context << ": " << failed_solver << " failed at iteration " << iter
<< " and allow_fallthrough is false. Reason: " << e.what();
throw std::domain_error(os.str());
}

/**
* @brief Log a solver fallback event to the provided stream, if any.
* @param[in,out] msgs Output stream (may be nullptr)
* @param[in] context Context string for the log
* @param[in] iter Current iteration number
* @param[in] failed_solver Name of the solver that failed
* @param[in] next_solver Name of the solver being attempted next
* @param[in] e Exception that caused the fallback
*/
inline void log_solver_fallback(const bool allow_fallthrough,
std::ostream* msgs, std::string_view context,
inline void log_solver_fallback(std::ostream* msgs, std::string_view context,
Eigen::Index iter,
std::string_view failed_solver,
std::string_view next_solver,
const std::exception& e) {
if (!msgs) {
return;
}
// Build once so we don't interleave with other logs.
std::ostringstream os;
std::string msg_type = allow_fallthrough ? "WARNING" : "ERROR";
os << "[" << context << "] " << msg_type << ": solver fallback\n"
os << "[" << context << "] WARNING: solver fallback\n"
<< " " << std::left << std::setw(12) << "iteration:" << iter << "\n"
<< " " << std::left << std::setw(12) << "failed:" << failed_solver << "\n"
<< " " << std::left << std::setw(12) << "reason:" << e.what() << "\n"
<< " " << std::left << std::setw(12) << "action:"
<< "trying " << next_solver << "\n"
<< "note: this warning message will only be displayed once."
<< "\n";
if (allow_fallthrough && msgs) {
(*msgs) << os.str();
} else {
throw std::domain_error(std::string("[") + std::string(context) + "]");
}
(*msgs) << os.str();
}

template <bool InitTheta, typename Opts>
Expand Down Expand Up @@ -1114,7 +1148,8 @@ inline auto create_update_fun(ObjFun&& obj_fun, ThetaGradFun&& theta_grad_f,
};
}

static STAN_THREADS_DEF std::once_flag fallback_warning;
static STAN_THREADS_DEF std::once_flag fallback_warning_1_2;
static STAN_THREADS_DEF std::once_flag fallback_warning_2_3;
/**
* For a latent Gaussian model with hyperparameters phi and
* latent variables theta, and observations y, this function computes
Expand Down Expand Up @@ -1213,13 +1248,16 @@ inline auto laplace_marginal_density_est(
const std::string solver_type
= (options.hessian_block_size == 1) ? "Diagonal" : "Block";
std::string failed = "solver 1 (" + solver_type + " Hessian-root Cholesky)";
if (!options.allow_fallthrough) {
throw_solver_failure("laplace_marginal_density", step_iter, failed, e);
}
std::call_once(
fallback_warning,
[](auto&&... args) {
fallback_warning_1_2,
[](auto&&... args) STAN_COLD_PATH {
log_solver_fallback(std::forward<decltype(args)>(args)...);
},
options.allow_fallthrough, msgs, "laplace_marginal_density", step_iter,
std::move(failed), "solver 2 (Covariance-root Cholesky)", e);
msgs, "laplace_marginal_density", step_iter, std::move(failed),
"solver 2 (Covariance-root Cholesky)", e);
}
try {
if (options.solver == 2 || options.allow_fallthrough) {
Expand All @@ -1228,12 +1266,16 @@ inline auto laplace_marginal_density_est(
covariance, update_fun, msgs);
}
} catch (const std::exception& e) {
if (!options.allow_fallthrough) {
throw_solver_failure("laplace_marginal_density", step_iter,
"solver 2 (Covariance-root Cholesky)", e);
}
std::call_once(
fallback_warning,
[](auto&&... args) {
fallback_warning_2_3,
[](auto&&... args) STAN_COLD_PATH {
log_solver_fallback(std::forward<decltype(args)>(args)...);
},
options.allow_fallthrough, msgs, "laplace_marginal_density", step_iter,
msgs, "laplace_marginal_density", step_iter,
"solver 2 (Covariance-root Cholesky)", "solver 3 (General LU solver)",
e);
}
Expand Down
93 changes: 93 additions & 0 deletions test/unit/math/laplace/block_matrix_sqrt_test.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
#include <stan/math/mix.hpp>
#include <gtest/gtest.h>
#include <cmath>
#include <limits>

namespace {

// sparse block-diagonal matrix with the block pattern the Laplace solver
// reserves for W_r (see CholeskyWSolverBlock)
Eigen::SparseMatrix<double> block_pattern(int n_blocks, int block_size) {
const int n = n_blocks * block_size;
Eigen::SparseMatrix<double> m(n, n);
m.reserve(Eigen::VectorXi::Constant(n, block_size));
for (int b = 0; b < n_blocks; ++b) {
for (int k = 0; k < block_size; ++k) {
for (int j = 0; j < block_size; ++j) {
m.insert(b * block_size + j, b * block_size + k) = 1.0;
}
}
}
m.makeCompressed();
return m;
}

Eigen::SparseMatrix<double> block_diag(
const std::vector<Eigen::MatrixXd>& blocks) {
const int block_size = blocks[0].rows();
Eigen::SparseMatrix<double> w = block_pattern(blocks.size(), block_size);
for (std::size_t b = 0; b < blocks.size(); ++b) {
for (int k = 0; k < block_size; ++k) {
for (int j = 0; j < block_size; ++j) {
w.coeffRef(b * block_size + j, b * block_size + k) = blocks[b](j, k);
}
}
}
return w;
}

void expect_principal_sqrt(const Eigen::SparseMatrix<double>& w,
int block_size) {
Eigen::SparseMatrix<double> w_root
= block_pattern(w.rows() / block_size, block_size);
EXPECT_NO_THROW(
stan::math::internal::block_matrix_sqrt(w_root, w, block_size));
const Eigen::MatrixXd root = w_root;
EXPECT_TRUE(root.isApprox(root.transpose(), 1e-12));
Eigen::SelfAdjointEigenSolver<Eigen::MatrixXd> eig(root);
EXPECT_GE(eig.eigenvalues().minCoeff(), -1e-12);
EXPECT_TRUE((root * root).isApprox(Eigen::MatrixXd(w), 1e-10));
}

} // namespace

TEST(LaplaceBlockMatrixSqrt, PositiveDefiniteBlocks) {
Eigen::MatrixXd m(3, 3);
m << 1.0, 0.3, -0.2, 0.5, 2.0, 0.1, -0.4, 0.2, 1.5;
Eigen::MatrixXd a = m * m.transpose() + Eigen::MatrixXd::Identity(3, 3);
Eigen::MatrixXd b = 2.0 * Eigen::MatrixXd::Identity(3, 3);
expect_principal_sqrt(block_diag({a, b}), 3);
}

// A negative Hessian with more latent variables than observations is only
// positive semi-definite: its zero eigenvalues come out of floating point
// as tiny values of either sign and must not be rejected.
TEST(LaplaceBlockMatrixSqrt, RankDeficientBlockIsAccepted) {
Eigen::MatrixXd z(3, 6);
z << 1.0, 0.5, -0.3, 0.8, 0.1, -0.6, 0.2, -1.1, 0.4, 0.3, 0.9, 0.7, -0.5, 0.6,
1.2, -0.2, 0.4, 0.3;
Eigen::MatrixXd w = z.transpose() * z; // rank 3 of 6
expect_principal_sqrt(block_diag({w}), 6);

Eigen::MatrixXd ones = Eigen::MatrixXd::Ones(2, 2); // rank 1 of 2
Eigen::MatrixXd spd = Eigen::MatrixXd::Identity(2, 2);
expect_principal_sqrt(block_diag({spd, ones}), 2);
}

TEST(LaplaceBlockMatrixSqrt, IndefiniteBlockThrows) {
Eigen::MatrixXd indefinite(2, 2);
indefinite << 1.0, 0.0, 0.0, -1.0;
Eigen::SparseMatrix<double> w = block_diag({indefinite});
Eigen::SparseMatrix<double> w_root = block_pattern(1, 2);
EXPECT_THROW(stan::math::internal::block_matrix_sqrt(w_root, w, 2),
std::domain_error);
}

TEST(LaplaceBlockMatrixSqrt, NonFiniteBlockThrows) {
Eigen::MatrixXd nan_block = Eigen::MatrixXd::Identity(2, 2);
nan_block(0, 1) = std::numeric_limits<double>::quiet_NaN();
Eigen::SparseMatrix<double> w = block_diag({nan_block});
Eigen::SparseMatrix<double> w_root = block_pattern(1, 2);
EXPECT_THROW(stan::math::internal::block_matrix_sqrt(w_root, w, 2),
std::domain_error);
}
Loading
Loading