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
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
9 changes: 4 additions & 5 deletions stan/math/rev/fun/mdivide_left_tri.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -360,16 +360,15 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) {
auto arena_A_val = to_arena(arena_A.val());

arena_t<ret_type> res
= arena_A_val.template triangularView<TriView>().solve(
arena_B.val_op());
= arena_A_val.template triangularView<TriView>().solve(arena_B.val());

reverse_pass_callback([arena_A, arena_B, arena_A_val, res]() mutable {
promote_scalar_t<double, T2> adjB
= arena_A_val.template triangularView<TriView>().transpose().solve(
res.adj());

arena_B.adj() += adjB;
arena_A.adj() -= (adjB * res.val_op().transpose().eval())
arena_A.adj() -= (adjB * res.val().transpose().eval())
.template triangularView<TriView>();
});

Expand All @@ -386,7 +385,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) {
= arena_A_val.template triangularView<TriView>().transpose().solve(
res.adj());

arena_A.adj() -= (adjB * res.val_op().transpose().eval())
arena_A.adj() -= (adjB * res.val().transpose().eval())
.template triangularView<TriView>();
});

Expand All @@ -396,7 +395,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) {
arena_t<promote_scalar_t<var, T2>> arena_B = B;

arena_t<ret_type> res
= arena_A.template triangularView<TriView>().solve(arena_B.val_op());
= arena_A.template triangularView<TriView>().solve(arena_B.val());

reverse_pass_callback([arena_A, arena_B, res]() mutable {
promote_scalar_t<double, T2> adjB
Expand Down
4 changes: 2 additions & 2 deletions stan/math/rev/fun/multiply.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ inline auto multiply(T1&& A, T2&& B) {
arena_t<promote_scalar_t<var, T2>> arena_B(std::forward<T2>(B));
using return_t
= return_var_matrix_t<decltype(arena_A * value_of(B).eval()), T1, T2>;
arena_t<return_t> res = arena_A * arena_B.val_op();
arena_t<return_t> res = arena_A * arena_B.val();
reverse_pass_callback([arena_B, arena_A, res]() mutable {
arena_B.adj() += arena_A.transpose() * res.adj_op();
});
Expand All @@ -65,7 +65,7 @@ inline auto multiply(T1&& A, T2&& B) {
using return_t
= return_var_matrix_t<decltype(value_of(arena_A).eval() * arena_B), T1,
T2>;
arena_t<return_t> res = arena_A.val_op() * arena_B;
arena_t<return_t> res = arena_A.val() * arena_B;
reverse_pass_callback([arena_A, arena_B, res]() mutable {
arena_A.adj() += res.adj_op() * arena_B.transpose();
});
Expand Down
Loading