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
12 changes: 9 additions & 3 deletions doxygen/contributor_help_pages/getting_started.md
Original file line number Diff line number Diff line change
Expand Up @@ -660,9 +660,15 @@ The values of `x` have the same shape as `x`.

### Values and adjoint extensions to Eigen

The matrix and vector autodiff types come with an extra `.val()` and `.adj()`, member functions called `.val_op()` and `.adj_op()`.
These `*_op()` member functions are used as a workaround for a bug in Eigen where transpose expressions will be inaccessible because of an incorrect const reference.
See [here](https://github.com/stan-dev/math/issues/2653) for the details and other information for when this workaround is needed.
By default, `.val()` and `.adj()` return a `CwiseUnaryView` when the matrix/vector is non-`const`,
and returns a `CwiseUnaryOp` when the matrix/vector is `const`. The exception is that
calling `.val()` on a matrix/vector of `var` types will always return a `CwiseUnaryOp`, as the
underlying value is always `const`.

However, using a `CwiseUnaryView` in some Eigen operations (e.g., multiplication, transposition)
can result in a compilation error. To workaround this, the `.val_op()` and `.adj_op()` member
functions have been added to explicitly request a `CwiseUnaryOp` regardless of whether the
matrix/vector is `const` or not.

The member functions `.val()` and `.val_op()` return expressions that evaluate to the values
of the autodiff matrix.
Expand Down
28 changes: 22 additions & 6 deletions stan/math/prim/eigen_plugins.h
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,7 @@ struct val_Op{
double& operator()(double& v) const { return v; }
};


/**
* Coefficient-wise function applying val_Op struct to a matrix of const var
* or vari* and returning a view to the const matrix of doubles containing
Expand All @@ -94,16 +95,31 @@ val() const { return CwiseUnaryOp<val_Op, const Derived>(derived());
/**
* Coefficient-wise function applying val_Op struct to a matrix of var
* or vari* and returning a view to the values
*/
*/
template <
typename T = Scalar,
std::enable_if_t<
!std::disjunction_v<
std::is_arithmetic<std::decay_t<T>>,
is_fvar<std::decay_t<T>>
>
>* = nullptr>
inline CwiseUnaryOp<val_Op, Derived>
val() { return CwiseUnaryOp<val_Op, Derived>(derived());
}

template <
typename T = Scalar,
std::enable_if_t<
std::disjunction_v<
std::is_arithmetic<std::decay_t<T>>,
is_fvar<std::decay_t<T>>
>
>* = nullptr>
inline CwiseUnaryView<val_Op, Derived>
val() { return CwiseUnaryView<val_Op, Derived>(derived());
}

/**
* Coefficient-wise function applying val_Op struct to a matrix of var
* or vari* and returning a view to the matrix of doubles containing
* the values
*/
inline CwiseUnaryOp<val_Op, Derived>
val_op() { return CwiseUnaryOp<val_Op, Derived>(derived());
}
Expand Down
8 changes: 4 additions & 4 deletions stan/math/rev/constraint/stochastic_column_constrain.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ inline plain_type_t<T> stochastic_column_constrain(const T& y) {
const auto M = y.cols();
arena_t<T> arena_y = y;

arena_t<ret_type> arena_x = stochastic_column_constrain(arena_y.val_op());
arena_t<ret_type> arena_x = stochastic_column_constrain(arena_y.val());

if (unlikely(N == 0 || M == 0)) {
return arena_x;
Expand All @@ -39,7 +39,7 @@ inline plain_type_t<T> stochastic_column_constrain(const T& y) {
reverse_pass_callback([arena_y, arena_x]() mutable {
const auto M = arena_y.cols();

auto&& x_val = arena_x.val_op();
auto&& x_val = arena_x.val();
auto&& x_adj = arena_x.adj_op();

Eigen::VectorXd x_pre_softmax_adj(x_val.rows());
Expand Down Expand Up @@ -82,7 +82,7 @@ inline plain_type_t<T> stochastic_column_constrain(const T& y,

double lp_val = 0;
arena_t<ret_type> arena_x
= stochastic_column_constrain(arena_y.val_op(), lp_val);
= stochastic_column_constrain(arena_y.val(), lp_val);
lp += lp_val;

if (unlikely(N == 0 || M == 0)) {
Expand All @@ -92,7 +92,7 @@ inline plain_type_t<T> stochastic_column_constrain(const T& y,
reverse_pass_callback([arena_y, arena_x, lp]() mutable {
const auto M = arena_y.cols();

auto&& x_val = arena_x.val_op();
auto&& x_val = arena_x.val();
auto&& x_adj = arena_x.adj_op();

const auto x_val_rows = x_val.rows();
Expand Down
9 changes: 4 additions & 5 deletions stan/math/rev/constraint/stochastic_row_constrain.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ inline auto stochastic_row_constrain(const T& y) {
const auto M = y.cols();
arena_t<T> arena_y = y;

arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val_op());
arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val());

if (unlikely(N == 0 || M == 0)) {
return arena_x;
Expand All @@ -37,7 +37,7 @@ inline auto stochastic_row_constrain(const T& y) {
reverse_pass_callback([arena_y, arena_x]() mutable {
const auto N = arena_y.rows();

auto&& x_val = arena_x.val_op();
auto&& x_val = arena_x.val();
auto&& x_adj = arena_x.adj_op();

Eigen::VectorXd x_pre_softmax_adj(x_val.cols());
Expand Down Expand Up @@ -79,8 +79,7 @@ inline plain_type_t<T> stochastic_row_constrain(const T& y,
arena_t<T> arena_y = y;

double lp_val = 0;
arena_t<ret_type> arena_x
= stochastic_row_constrain(arena_y.val_op(), lp_val);
arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val(), lp_val);
lp += lp_val;

if (unlikely(N == 0 || M == 0)) {
Expand All @@ -90,7 +89,7 @@ inline plain_type_t<T> stochastic_row_constrain(const T& y,
reverse_pass_callback([arena_y, arena_x, lp]() mutable {
const auto N = arena_y.rows();

auto&& x_val = arena_x.val_op();
auto&& x_val = arena_x.val();
auto&& x_adj = arena_x.adj_op();

const auto x_val_cols = x_val.cols();
Expand Down
19 changes: 9 additions & 10 deletions stan/math/rev/fun/eigendecompose_sym.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -41,20 +41,19 @@ inline auto eigendecompose_sym(const T& m) {

reverse_pass_callback([eigenvals, arena_m, eigenvecs]() mutable {
// eigenvalue reverse calculation
auto value_adj = eigenvecs.val_op() * eigenvals.adj().asDiagonal()
* eigenvecs.val_op().transpose();
auto value_adj = eigenvecs.val() * eigenvals.adj().asDiagonal()
* eigenvecs.val().transpose();
// eigenvector reverse calculation
const auto p = arena_m.val().cols();
Eigen::MatrixXd f
= (1
/ (eigenvals.val_op().rowwise().replicate(p).transpose()
- eigenvals.val_op().rowwise().replicate(p))
.array());
Eigen::MatrixXd f = (1
/ (eigenvals.val().rowwise().replicate(p).transpose()
- eigenvals.val().rowwise().replicate(p))
.array());
f.diagonal().setZero();
auto vector_adj
= eigenvecs.val_op()
* f.cwiseProduct(eigenvecs.val_op().transpose() * eigenvecs.adj_op())
* eigenvecs.val_op().transpose();
= eigenvecs.val()
* f.cwiseProduct(eigenvecs.val().transpose() * eigenvecs.adj_op())
* eigenvecs.val().transpose();

arena_m.adj() += value_adj + vector_adj;
});
Expand Down
6 changes: 3 additions & 3 deletions stan/math/rev/fun/eigenvectors_sym.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -43,9 +43,9 @@ inline auto eigenvectors_sym(const T& m) {
.array());
f.diagonal().setZero();
arena_m.adj()
+= eigenvecs.val_op()
* f.cwiseProduct(eigenvecs.val_op().transpose() * eigenvecs.adj_op())
* eigenvecs.val_op().transpose();
+= eigenvecs.val()
* f.cwiseProduct(eigenvecs.val().transpose() * eigenvecs.adj_op())
* eigenvecs.val().transpose();
});

return return_t(eigenvecs);
Expand Down
22 changes: 10 additions & 12 deletions stan/math/rev/fun/generalized_inverse.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -24,15 +24,13 @@ template <typename T1, typename T2>
inline auto generalized_inverse_lambda(T1& G_arena, T2& inv_G) {
return [G_arena, inv_G]() mutable {
G_arena.adj()
+= -(inv_G.val_op().transpose() * inv_G.adj_op()
* inv_G.val_op().transpose())
+ (-G_arena.val_op() * inv_G.val_op()
+= -(inv_G.val().transpose() * inv_G.adj_op() * inv_G.val().transpose())
+ (-G_arena.val() * inv_G.val()
+ Eigen::MatrixXd::Identity(G_arena.rows(), inv_G.cols()))
* inv_G.adj_op().transpose() * inv_G.val_op()
* inv_G.val_op().transpose()
+ inv_G.val_op().transpose() * inv_G.val_op()
* inv_G.adj_op().transpose()
* (-inv_G.val_op() * G_arena.val_op()
* inv_G.adj_op().transpose() * inv_G.val()
* inv_G.val().transpose()
+ inv_G.val().transpose() * inv_G.val() * inv_G.adj_op().transpose()
* (-inv_G.val() * G_arena.val()
+ Eigen::MatrixXd::Identity(inv_G.rows(), G_arena.cols()));
};
}
Expand Down Expand Up @@ -83,17 +81,17 @@ inline auto generalized_inverse(const VarMat& G) {
}
} else if (G.rows() < G.cols()) {
arena_t<VarMat> G_arena(G);
arena_t<ret_type> inv_G((G_arena.val_op() * G_arena.val_op().transpose())
arena_t<ret_type> inv_G((G_arena.val() * G_arena.val().transpose())
.ldlt()
.solve(G_arena.val_op())
.solve(G_arena.val())
.transpose());
reverse_pass_callback(internal::generalized_inverse_lambda(G_arena, inv_G));
return ret_type(inv_G);
} else {
arena_t<VarMat> G_arena(G);
arena_t<ret_type> inv_G((G_arena.val_op().transpose() * G_arena.val_op())
arena_t<ret_type> inv_G((G_arena.val().transpose() * G_arena.val())
.ldlt()
.solve(G_arena.val_op().transpose()));
.solve(G_arena.val().transpose()));
reverse_pass_callback(internal::generalized_inverse_lambda(G_arena, inv_G));
return ret_type(inv_G);
}
Expand Down
2 changes: 1 addition & 1 deletion stan/math/rev/fun/inverse.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ inline auto inverse(const T& m) {
}

arena_t<T> arena_m = m;
arena_t<promote_scalar_t<double, T>> res_val = arena_m.val_op().inverse();
arena_t<promote_scalar_t<double, T>> res_val = arena_m.val().inverse();
arena_t<ret_type> res = res_val;

reverse_pass_callback([res, res_val, arena_m]() mutable {
Expand Down
12 changes: 6 additions & 6 deletions stan/math/rev/fun/mdivide_left.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -42,24 +42,24 @@ inline auto mdivide_left(T1&& A, T2&& B) {
if constexpr (is_autodiff_v<T1> && is_autodiff_v<T2>) {
arena_t<T1> arena_A(std::forward<T1>(A));
arena_t<T2> arena_B(std::forward<T2>(B));
auto hqr_A_ptr = make_chainable_ptr(arena_A.val_op().householderQr());
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val_op());
auto hqr_A_ptr = make_chainable_ptr(arena_A.val().householderQr());
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val());
reverse_pass_callback([arena_A, arena_B, hqr_A_ptr, res]() mutable {
promote_scalar_t<double, T2> adjB
= hqr_A_ptr->householderQ()
* hqr_A_ptr->matrixQR()
.template triangularView<Eigen::Upper>()
.transpose()
.solve(res.adj());
arena_A.adj() -= adjB * res.val_op().transpose();
arena_A.adj() -= adjB * res.val().transpose();
arena_B.adj() += adjB;
});

return ret_type(res);
} else if constexpr (is_autodiff_v<T2>) {
arena_t<T2> arena_B(std::forward<T2>(B));
auto hqr_A_ptr = make_chainable_ptr(value_of(A).householderQr());
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val_op());
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val());
reverse_pass_callback([arena_B, hqr_A_ptr, res]() mutable {
arena_B.adj() += hqr_A_ptr->householderQ()
* hqr_A_ptr->matrixQR()
Expand All @@ -70,15 +70,15 @@ inline auto mdivide_left(T1&& A, T2&& B) {
return ret_type(res);
} else {
arena_t<T1> arena_A(std::forward<T1>(A));
auto hqr_A_ptr = make_chainable_ptr(arena_A.val_op().householderQr());
auto hqr_A_ptr = make_chainable_ptr(arena_A.val().householderQr());
arena_t<ret_type> res = hqr_A_ptr->solve(value_of(B));
reverse_pass_callback([arena_A, hqr_A_ptr, res]() mutable {
arena_A.adj() -= hqr_A_ptr->householderQ()
* hqr_A_ptr->matrixQR()
.template triangularView<Eigen::Upper>()
.transpose()
.solve(res.adj())
* res.val_op().transpose();
* res.val().transpose();
});
return ret_type(res);
}
Expand Down
8 changes: 4 additions & 4 deletions stan/math/rev/fun/mdivide_left_ldlt.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -39,13 +39,13 @@ inline auto mdivide_left_ldlt(LDLT_factor<T1>& A, const T2& B) {
if constexpr (is_autodiff_v<T1> && is_autodiff_v<T2>) {
arena_t<promote_scalar_t<var, T2>> arena_B = B;
arena_t<promote_scalar_t<var, T1>> arena_A = A.matrix();
arena_t<ret_type> res = A.ldlt().solve(arena_B.val_op());
arena_t<ret_type> res = A.ldlt().solve(arena_B.val());
const auto* ldlt_ptr = make_chainable_ptr(A.ldlt());

reverse_pass_callback([arena_A, arena_B, ldlt_ptr, res]() mutable {
promote_scalar_t<double, T2> adjB = ldlt_ptr->solve(res.adj());

arena_A.adj() -= adjB * res.val_op().transpose();
arena_A.adj() -= adjB * res.val().transpose();
arena_B.adj() += adjB;
});

Expand All @@ -56,13 +56,13 @@ inline auto mdivide_left_ldlt(LDLT_factor<T1>& A, const T2& B) {
const auto* ldlt_ptr = make_chainable_ptr(A.ldlt());

reverse_pass_callback([arena_A, ldlt_ptr, res]() mutable {
arena_A.adj() -= ldlt_ptr->solve(res.adj()) * res.val_op().transpose();
arena_A.adj() -= ldlt_ptr->solve(res.adj()) * res.val().transpose();
});

return ret_type(res);
} else {
arena_t<promote_scalar_t<var, T2>> arena_B = B;
arena_t<ret_type> res = A.ldlt().solve(arena_B.val_op());
arena_t<ret_type> res = A.ldlt().solve(arena_B.val());
const auto* ldlt_ptr = make_chainable_ptr(A.ldlt());

reverse_pass_callback([arena_B, ldlt_ptr, res]() mutable {
Expand Down
12 changes: 6 additions & 6 deletions stan/math/rev/fun/mdivide_left_spd.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -276,12 +276,12 @@ inline auto mdivide_left_spd(const T1 &A, const T2 &B) {
check_symmetric("mdivide_left_spd", "A", arena_A.val());
check_not_nan("mdivide_left_spd", "A", arena_A.val());

auto A_llt = arena_A.val_op().llt();
auto A_llt = arena_A.val().llt();

check_pos_definite("mdivide_left_spd", "A", A_llt);

arena_t<Eigen::MatrixXd> arena_A_llt = A_llt.matrixL();
arena_t<ret_type> res = A_llt.solve(arena_B.val_op());
arena_t<ret_type> res = A_llt.solve(arena_B.val());

reverse_pass_callback([arena_A, arena_B, arena_A_llt, res]() mutable {
promote_scalar_t<double, T2> adjB = res.adj();
Expand All @@ -291,7 +291,7 @@ inline auto mdivide_left_spd(const T1 &A, const T2 &B) {
.transpose()
.solveInPlace(adjB);

arena_A.adj() -= adjB * res.val_op().transpose();
arena_A.adj() -= adjB * res.val().transpose();
arena_B.adj() += adjB;
});

Expand All @@ -302,7 +302,7 @@ inline auto mdivide_left_spd(const T1 &A, const T2 &B) {
check_symmetric("mdivide_left_spd", "A", arena_A.val());
check_not_nan("mdivide_left_spd", "A", arena_A.val());

auto A_llt = arena_A.val_op().llt();
auto A_llt = arena_A.val().llt();

check_pos_definite("mdivide_left_spd", "A", A_llt);

Expand All @@ -317,7 +317,7 @@ inline auto mdivide_left_spd(const T1 &A, const T2 &B) {
.transpose()
.solveInPlace(adjB);

arena_A.adj() -= adjB * res.val_op().transpose().eval();
arena_A.adj() -= adjB * res.val().transpose().eval();
});

return ret_type(res);
Expand All @@ -333,7 +333,7 @@ inline auto mdivide_left_spd(const T1 &A, const T2 &B) {
check_pos_definite("mdivide_left_spd", "A", A_llt);

arena_t<Eigen::MatrixXd> arena_A_llt = A_llt.matrixL();
arena_t<ret_type> res = A_llt.solve(arena_B.val_op());
arena_t<ret_type> res = A_llt.solve(arena_B.val());

reverse_pass_callback([arena_B, arena_A_llt, res]() mutable {
promote_scalar_t<double, T2> adjB = res.adj();
Expand Down
Loading