diff --git a/stan/math/prim/eigen_plugins.h b/stan/math/prim/eigen_plugins.h index d21b5917322..da2d5c6b38b 100644 --- a/stan/math/prim/eigen_plugins.h +++ b/stan/math/prim/eigen_plugins.h @@ -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 @@ -94,16 +95,31 @@ val() const { return CwiseUnaryOp(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>, + is_fvar> + > + >* = nullptr> +inline CwiseUnaryOp +val() { return CwiseUnaryOp(derived()); +} + +template < + typename T = Scalar, + std::enable_if_t< + std::disjunction_v< + std::is_arithmetic>, + is_fvar> + > + >* = nullptr> inline CwiseUnaryView val() { return CwiseUnaryView(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() { return CwiseUnaryOp(derived()); } diff --git a/stan/math/rev/constraint/stochastic_column_constrain.hpp b/stan/math/rev/constraint/stochastic_column_constrain.hpp index 725b5483541..bbc0e046b38 100644 --- a/stan/math/rev/constraint/stochastic_column_constrain.hpp +++ b/stan/math/rev/constraint/stochastic_column_constrain.hpp @@ -30,7 +30,7 @@ inline plain_type_t stochastic_column_constrain(const T& y) { const auto M = y.cols(); arena_t arena_y = y; - arena_t arena_x = stochastic_column_constrain(arena_y.val_op()); + arena_t arena_x = stochastic_column_constrain(arena_y.val()); if (unlikely(N == 0 || M == 0)) { return arena_x; @@ -39,7 +39,7 @@ inline plain_type_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()); @@ -82,7 +82,7 @@ inline plain_type_t stochastic_column_constrain(const T& y, double lp_val = 0; arena_t 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)) { @@ -92,7 +92,7 @@ inline plain_type_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(); diff --git a/stan/math/rev/constraint/stochastic_row_constrain.hpp b/stan/math/rev/constraint/stochastic_row_constrain.hpp index bf9acd724d5..80fa461c3a6 100644 --- a/stan/math/rev/constraint/stochastic_row_constrain.hpp +++ b/stan/math/rev/constraint/stochastic_row_constrain.hpp @@ -28,7 +28,7 @@ inline auto stochastic_row_constrain(const T& y) { const auto M = y.cols(); arena_t arena_y = y; - arena_t arena_x = stochastic_row_constrain(arena_y.val_op()); + arena_t arena_x = stochastic_row_constrain(arena_y.val()); if (unlikely(N == 0 || M == 0)) { return arena_x; @@ -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()); @@ -79,8 +79,7 @@ inline plain_type_t stochastic_row_constrain(const T& y, arena_t arena_y = y; double lp_val = 0; - arena_t arena_x - = stochastic_row_constrain(arena_y.val_op(), lp_val); + arena_t arena_x = stochastic_row_constrain(arena_y.val(), lp_val); lp += lp_val; if (unlikely(N == 0 || M == 0)) { @@ -90,7 +89,7 @@ inline plain_type_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(); diff --git a/stan/math/rev/fun/eigendecompose_sym.hpp b/stan/math/rev/fun/eigendecompose_sym.hpp index 17663bd7ce5..873da174049 100644 --- a/stan/math/rev/fun/eigendecompose_sym.hpp +++ b/stan/math/rev/fun/eigendecompose_sym.hpp @@ -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; }); diff --git a/stan/math/rev/fun/eigenvectors_sym.hpp b/stan/math/rev/fun/eigenvectors_sym.hpp index 6f2d2bfca4d..707acd38ac5 100644 --- a/stan/math/rev/fun/eigenvectors_sym.hpp +++ b/stan/math/rev/fun/eigenvectors_sym.hpp @@ -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); diff --git a/stan/math/rev/fun/generalized_inverse.hpp b/stan/math/rev/fun/generalized_inverse.hpp index 11c0eb18ada..c8297d29264 100644 --- a/stan/math/rev/fun/generalized_inverse.hpp +++ b/stan/math/rev/fun/generalized_inverse.hpp @@ -24,15 +24,13 @@ template 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())); }; } @@ -83,17 +81,17 @@ inline auto generalized_inverse(const VarMat& G) { } } else if (G.rows() < G.cols()) { arena_t G_arena(G); - arena_t inv_G((G_arena.val_op() * G_arena.val_op().transpose()) + arena_t 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 G_arena(G); - arena_t inv_G((G_arena.val_op().transpose() * G_arena.val_op()) + arena_t 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); } diff --git a/stan/math/rev/fun/inverse.hpp b/stan/math/rev/fun/inverse.hpp index 7d44cb5ef68..4d34df0175c 100644 --- a/stan/math/rev/fun/inverse.hpp +++ b/stan/math/rev/fun/inverse.hpp @@ -30,7 +30,7 @@ inline auto inverse(const T& m) { } arena_t arena_m = m; - arena_t> res_val = arena_m.val_op().inverse(); + arena_t> res_val = arena_m.val().inverse(); arena_t res = res_val; reverse_pass_callback([res, res_val, arena_m]() mutable { diff --git a/stan/math/rev/fun/mdivide_left.hpp b/stan/math/rev/fun/mdivide_left.hpp index 4af112e0064..a440103fc42 100644 --- a/stan/math/rev/fun/mdivide_left.hpp +++ b/stan/math/rev/fun/mdivide_left.hpp @@ -42,8 +42,8 @@ inline auto mdivide_left(T1&& A, T2&& B) { if constexpr (is_autodiff_v && is_autodiff_v) { arena_t arena_A(std::forward(A)); arena_t arena_B(std::forward(B)); - auto hqr_A_ptr = make_chainable_ptr(arena_A.val_op().householderQr()); - arena_t res = hqr_A_ptr->solve(arena_B.val_op()); + auto hqr_A_ptr = make_chainable_ptr(arena_A.val().householderQr()); + arena_t res = hqr_A_ptr->solve(arena_B.val()); reverse_pass_callback([arena_A, arena_B, hqr_A_ptr, res]() mutable { promote_scalar_t adjB = hqr_A_ptr->householderQ() @@ -51,7 +51,7 @@ inline auto mdivide_left(T1&& A, T2&& B) { .template triangularView() .transpose() .solve(res.adj()); - arena_A.adj() -= adjB * res.val_op().transpose(); + arena_A.adj() -= adjB * res.val().transpose(); arena_B.adj() += adjB; }); @@ -59,7 +59,7 @@ inline auto mdivide_left(T1&& A, T2&& B) { } else if constexpr (is_autodiff_v) { arena_t arena_B(std::forward(B)); auto hqr_A_ptr = make_chainable_ptr(value_of(A).householderQr()); - arena_t res = hqr_A_ptr->solve(arena_B.val_op()); + arena_t 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() @@ -70,7 +70,7 @@ inline auto mdivide_left(T1&& A, T2&& B) { return ret_type(res); } else { arena_t arena_A(std::forward(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 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() @@ -78,7 +78,7 @@ inline auto mdivide_left(T1&& A, T2&& B) { .template triangularView() .transpose() .solve(res.adj()) - * res.val_op().transpose(); + * res.val().transpose(); }); return ret_type(res); } diff --git a/stan/math/rev/fun/mdivide_left_ldlt.hpp b/stan/math/rev/fun/mdivide_left_ldlt.hpp index d140dbce053..aa423aee6cc 100644 --- a/stan/math/rev/fun/mdivide_left_ldlt.hpp +++ b/stan/math/rev/fun/mdivide_left_ldlt.hpp @@ -39,13 +39,13 @@ inline auto mdivide_left_ldlt(LDLT_factor& A, const T2& B) { if constexpr (is_autodiff_v && is_autodiff_v) { arena_t> arena_B = B; arena_t> arena_A = A.matrix(); - arena_t res = A.ldlt().solve(arena_B.val_op()); + arena_t 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 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; }); @@ -56,13 +56,13 @@ inline auto mdivide_left_ldlt(LDLT_factor& 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> arena_B = B; - arena_t res = A.ldlt().solve(arena_B.val_op()); + arena_t 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 { diff --git a/stan/math/rev/fun/mdivide_left_spd.hpp b/stan/math/rev/fun/mdivide_left_spd.hpp index 654ce14e25f..26472797c28 100644 --- a/stan/math/rev/fun/mdivide_left_spd.hpp +++ b/stan/math/rev/fun/mdivide_left_spd.hpp @@ -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 arena_A_llt = A_llt.matrixL(); - arena_t res = A_llt.solve(arena_B.val_op()); + arena_t res = A_llt.solve(arena_B.val()); reverse_pass_callback([arena_A, arena_B, arena_A_llt, res]() mutable { promote_scalar_t adjB = res.adj(); @@ -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; }); @@ -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); @@ -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); @@ -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 arena_A_llt = A_llt.matrixL(); - arena_t res = A_llt.solve(arena_B.val_op()); + arena_t res = A_llt.solve(arena_B.val()); reverse_pass_callback([arena_B, arena_A_llt, res]() mutable { promote_scalar_t adjB = res.adj(); diff --git a/stan/math/rev/fun/mdivide_left_tri.hpp b/stan/math/rev/fun/mdivide_left_tri.hpp index f0e0f7f81b0..bd95a9ca289 100644 --- a/stan/math/rev/fun/mdivide_left_tri.hpp +++ b/stan/math/rev/fun/mdivide_left_tri.hpp @@ -360,8 +360,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) { auto arena_A_val = to_arena(arena_A.val()); arena_t res - = arena_A_val.template triangularView().solve( - arena_B.val_op()); + = arena_A_val.template triangularView().solve(arena_B.val()); reverse_pass_callback([arena_A, arena_B, arena_A_val, res]() mutable { promote_scalar_t adjB @@ -369,7 +368,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) { 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(); }); @@ -386,7 +385,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) { = arena_A_val.template triangularView().transpose().solve( res.adj()); - arena_A.adj() -= (adjB * res.val_op().transpose().eval()) + arena_A.adj() -= (adjB * res.val().transpose().eval()) .template triangularView(); }); @@ -396,7 +395,7 @@ inline auto mdivide_left_tri(T1 &&A, T2 &&B) { arena_t> arena_B = B; arena_t res - = arena_A.template triangularView().solve(arena_B.val_op()); + = arena_A.template triangularView().solve(arena_B.val()); reverse_pass_callback([arena_A, arena_B, res]() mutable { promote_scalar_t adjB diff --git a/stan/math/rev/fun/multiply.hpp b/stan/math/rev/fun/multiply.hpp index 597c964ef9b..c2b55433e8d 100644 --- a/stan/math/rev/fun/multiply.hpp +++ b/stan/math/rev/fun/multiply.hpp @@ -54,7 +54,7 @@ inline auto multiply(T1&& A, T2&& B) { arena_t> arena_B(std::forward(B)); using return_t = return_var_matrix_t; - arena_t res = arena_A * arena_B.val_op(); + arena_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(); }); @@ -65,7 +65,7 @@ inline auto multiply(T1&& A, T2&& B) { using return_t = return_var_matrix_t; - arena_t res = arena_A.val_op() * arena_B; + arena_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(); }); diff --git a/stan/math/rev/fun/singular_values.hpp b/stan/math/rev/fun/singular_values.hpp index cc5da35f692..dd286aea421 100644 --- a/stan/math/rev/fun/singular_values.hpp +++ b/stan/math/rev/fun/singular_values.hpp @@ -29,7 +29,7 @@ inline auto singular_values(const EigMat& m) { auto arena_m = to_arena(m); Eigen::JacobiSVD svd( - arena_m.val_op(), Eigen::ComputeThinU | Eigen::ComputeThinV); + arena_m.val(), Eigen::ComputeThinU | Eigen::ComputeThinV); arena_t singular_values = svd.singularValues(); diff --git a/stan/math/rev/fun/sqrt.hpp b/stan/math/rev/fun/sqrt.hpp index c4d40696cfc..7b6cd6b7144 100644 --- a/stan/math/rev/fun/sqrt.hpp +++ b/stan/math/rev/fun/sqrt.hpp @@ -62,8 +62,8 @@ inline auto sqrt(const T& a) { return make_callback_var( a.val().array().sqrt().matrix(), [a](auto& vi) mutable { a.adj().array() - += (vi.val_op().array() == 0.0) - .select(0.0, vi.adj().array() / (2.0 * vi.val_op().array())); + += (vi.val().array() == 0.0) + .select(0.0, vi.adj().array() / (2.0 * vi.val().array())); }); } diff --git a/stan/math/rev/fun/svd.hpp b/stan/math/rev/fun/svd.hpp index 3f6b6ccbd1c..903ac88cefd 100644 --- a/stan/math/rev/fun/svd.hpp +++ b/stan/math/rev/fun/svd.hpp @@ -64,31 +64,31 @@ inline auto svd(const EigMat& m) { reverse_pass_callback([arena_m, arena_U, singular_values, arena_V, arena_Fp, arena_Fm]() mutable { // SVD-U reverse mode - Eigen::MatrixXd UUadjT = arena_U.val_op().transpose() * arena_U.adj_op(); + Eigen::MatrixXd UUadjT = arena_U.val().transpose() * arena_U.adj_op(); auto u_adj - = .5 * arena_U.val_op() + = .5 * arena_U.val() * (arena_Fp.array() * (UUadjT - UUadjT.transpose()).array()) .matrix() - * arena_V.val_op().transpose() + * arena_V.val().transpose() + (Eigen::MatrixXd::Identity(arena_m.rows(), arena_m.rows()) - - arena_U.val_op() * arena_U.val_op().transpose()) + - arena_U.val() * arena_U.val().transpose()) * arena_U.adj_op() - * singular_values.val_op().asDiagonal().inverse() - * arena_V.val_op().transpose(); + * singular_values.val().asDiagonal().inverse() + * arena_V.val().transpose(); // Singular values reverse mode - auto d_adj = arena_U.val_op() * singular_values.adj().asDiagonal() - * arena_V.val_op().transpose(); + auto d_adj = arena_U.val() * singular_values.adj().asDiagonal() + * arena_V.val().transpose(); // SVD-V reverse mode - Eigen::MatrixXd VTVadj = arena_V.val_op().transpose() * arena_V.adj_op(); + Eigen::MatrixXd VTVadj = arena_V.val().transpose() * arena_V.adj_op(); auto v_adj - = 0.5 * arena_U.val_op() + = 0.5 * arena_U.val() * (arena_Fm.array() * (VTVadj - VTVadj.transpose()).array()) .matrix() - * arena_V.val_op().transpose() - + arena_U.val_op() * singular_values.val_op().asDiagonal().inverse() + * arena_V.val().transpose() + + arena_U.val() * singular_values.val().asDiagonal().inverse() * arena_V.adj_op().transpose() * (Eigen::MatrixXd::Identity(arena_m.cols(), arena_m.cols()) - - arena_V.val_op() * arena_V.val_op().transpose()); + - arena_V.val() * arena_V.val().transpose()); arena_m.adj() += u_adj + d_adj + v_adj; }); diff --git a/stan/math/rev/fun/svd_U.hpp b/stan/math/rev/fun/svd_U.hpp index 5a695c7e653..25d3c3be7fb 100644 --- a/stan/math/rev/fun/svd_U.hpp +++ b/stan/math/rev/fun/svd_U.hpp @@ -33,7 +33,7 @@ inline auto svd_U(const EigMat& m) { auto arena_m = to_arena(m); Eigen::JacobiSVD> svd( - arena_m.val_op().eval(), Eigen::ComputeThinU | Eigen::ComputeThinV); + arena_m.val().eval(), Eigen::ComputeThinU | Eigen::ComputeThinV); auto arena_D = to_arena(svd.singularValues()); @@ -53,19 +53,19 @@ inline auto svd_U(const EigMat& m) { arena_t arena_U = svd.matrixU(); auto arena_V = to_arena(svd.matrixV()); - reverse_pass_callback([arena_m, arena_U, arena_D, arena_V, - arena_Fp]() mutable { - Eigen::MatrixXd UUadjT = arena_U.val_op().transpose() * arena_U.adj_op(); - arena_m.adj() - += .5 * arena_U.val_op() - * (arena_Fp.array() * (UUadjT - UUadjT.transpose()).array()) - .matrix() - * arena_V.transpose() - + (Eigen::MatrixXd::Identity(arena_m.rows(), arena_m.rows()) - - arena_U.val_op() * arena_U.val_op().transpose()) - * arena_U.adj_op() * arena_D.asDiagonal().inverse() - * arena_V.transpose(); - }); + reverse_pass_callback( + [arena_m, arena_U, arena_D, arena_V, arena_Fp]() mutable { + Eigen::MatrixXd UUadjT = arena_U.val().transpose() * arena_U.adj_op(); + arena_m.adj() + += .5 * arena_U.val() + * (arena_Fp.array() * (UUadjT - UUadjT.transpose()).array()) + .matrix() + * arena_V.transpose() + + (Eigen::MatrixXd::Identity(arena_m.rows(), arena_m.rows()) + - arena_U.val() * arena_U.val().transpose()) + * arena_U.adj_op() * arena_D.asDiagonal().inverse() + * arena_V.transpose(); + }); return ret_type(arena_U); } diff --git a/stan/math/rev/fun/svd_V.hpp b/stan/math/rev/fun/svd_V.hpp index afeeb4d9a4a..38b2d1f116b 100644 --- a/stan/math/rev/fun/svd_V.hpp +++ b/stan/math/rev/fun/svd_V.hpp @@ -33,7 +33,7 @@ inline auto svd_V(const EigMat& m) { auto arena_m = to_arena(m); Eigen::JacobiSVD>> svd( - arena_m.val_op().eval(), Eigen::ComputeThinU | Eigen::ComputeThinV); + arena_m.val().eval(), Eigen::ComputeThinU | Eigen::ComputeThinV); auto arena_D = to_arena(svd.singularValues()); @@ -55,16 +55,16 @@ inline auto svd_V(const EigMat& m) { reverse_pass_callback([arena_m, arena_U, arena_D, arena_V, arena_Fm]() mutable { - Eigen::MatrixXd VTVadj = arena_V.val_op().transpose() * arena_V.adj_op(); + Eigen::MatrixXd VTVadj = arena_V.val().transpose() * arena_V.adj_op(); arena_m.adj() += 0.5 * arena_U * (arena_Fm.array() * (VTVadj - VTVadj.transpose()).array()) .matrix() - * arena_V.val_op().transpose() + * arena_V.val().transpose() + arena_U * arena_D.asDiagonal().inverse() * arena_V.adj_op().transpose() * (Eigen::MatrixXd::Identity(arena_m.cols(), arena_m.cols()) - - arena_V.val_op() * arena_V.val_op().transpose()); + - arena_V.val() * arena_V.val().transpose()); }); return ret_type(arena_V); diff --git a/stan/math/rev/fun/tcrossprod.hpp b/stan/math/rev/fun/tcrossprod.hpp index e5beb438c44..afab2bf9f05 100644 --- a/stan/math/rev/fun/tcrossprod.hpp +++ b/stan/math/rev/fun/tcrossprod.hpp @@ -26,12 +26,12 @@ inline auto tcrossprod(const T& M) { using ret_type = return_var_matrix_t< Eigen::Matrix, T>; arena_t arena_M = M; - arena_t res = arena_M.val_op() * arena_M.val_op().transpose(); + arena_t res = arena_M.val() * arena_M.val().transpose(); if (likely(M.size() > 0)) { reverse_pass_callback([res, arena_M]() mutable { arena_M.adj() - += (res.adj_op() + res.adj_op().transpose()) * arena_M.val_op(); + += (res.adj_op() + res.adj_op().transpose()) * arena_M.val(); }); } diff --git a/stan/math/rev/fun/trace.hpp b/stan/math/rev/fun/trace.hpp index e7091b8f240..c0e627dad7b 100644 --- a/stan/math/rev/fun/trace.hpp +++ b/stan/math/rev/fun/trace.hpp @@ -24,7 +24,7 @@ template * = nullptr> inline auto trace(T&& m) { arena_t arena_m(std::forward(m)); - return make_callback_var(arena_m.val_op().trace(), + return make_callback_var(arena_m.val().trace(), [arena_m](const auto& vi) mutable { arena_m.adj().diagonal().array() += vi.adj(); }); diff --git a/stan/math/rev/fun/trace_dot.hpp b/stan/math/rev/fun/trace_dot.hpp index 636303be6c4..b9cb99cdffd 100644 --- a/stan/math/rev/fun/trace_dot.hpp +++ b/stan/math/rev/fun/trace_dot.hpp @@ -42,14 +42,14 @@ inline var trace_dot(Mat1&& A, Mat2&& B) { auto res_val = arena_A.val().cwiseProduct(arena_B.val().transpose()).sum(); return make_callback_var(res_val, [arena_A, arena_B](auto&& res) mutable { if constexpr (is_var_matrix::value) { - arena_A.adj().noalias() += res.adj() * arena_B.val_op().transpose(); + arena_A.adj().noalias() += res.adj() * arena_B.val().transpose(); } else { - arena_A.adj() += res.adj() * arena_B.val_op().transpose(); + arena_A.adj() += res.adj() * arena_B.val().transpose(); } if constexpr (is_var_matrix::value) { - arena_B.adj().noalias() += res.adj() * arena_A.val_op().transpose(); + arena_B.adj().noalias() += res.adj() * arena_A.val().transpose(); } else { - arena_B.adj() += res.adj() * arena_A.val_op().transpose(); + arena_B.adj() += res.adj() * arena_A.val().transpose(); } }); } else if constexpr (is_autodiff_v) { diff --git a/stan/math/rev/fun/trace_gen_inv_quad_form_ldlt.hpp b/stan/math/rev/fun/trace_gen_inv_quad_form_ldlt.hpp index b5d0f0b36c0..fd2d64c8ea3 100644 --- a/stan/math/rev/fun/trace_gen_inv_quad_form_ldlt.hpp +++ b/stan/math/rev/fun/trace_gen_inv_quad_form_ldlt.hpp @@ -45,30 +45,30 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, LDLT_factor& A, arena_t> arena_A = A.matrix(); arena_t> arena_B = B; arena_t> arena_D = D; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); - auto BTAsolveB = to_arena(arena_B.val_op().transpose() * AsolveB); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); + auto BTAsolveB = to_arena(arena_B.val().transpose() * AsolveB); - var res = (arena_D.val_op() * BTAsolveB).trace(); + var res = (arena_D.val() * BTAsolveB).trace(); - reverse_pass_callback( - [arena_A, BTAsolveB, AsolveB, arena_B, arena_D, res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback([arena_A, BTAsolveB, AsolveB, arena_B, arena_D, + res]() mutable { + double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().transpose() - * AsolveB.transpose(); - arena_B.adj() += C_adj * AsolveB - * (arena_D.val_op() + arena_D.val_op().transpose()); - arena_D.adj() += C_adj * BTAsolveB; - }); + arena_A.adj() + -= C_adj * AsolveB * arena_D.val().transpose() * AsolveB.transpose(); + arena_B.adj() + += C_adj * AsolveB * (arena_D.val() + arena_D.val().transpose()); + arena_D.adj() += C_adj * BTAsolveB; + }); return res; } else if constexpr (is_all_autodiff_v && is_constant_v) { arena_t> arena_A = A.matrix(); arena_t> arena_B = B; arena_t> arena_D = value_of(D); - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_D * arena_B.val_op().transpose() * AsolveB).trace(); + var res = (arena_D * arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([arena_A, AsolveB, arena_B, arena_D, res]() mutable { double C_adj = res.adj(); @@ -86,16 +86,16 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, LDLT_factor& A, auto AsolveB = to_arena(A.ldlt().solve(value_of(B_ref))); auto BTAsolveB = to_arena(value_of(B_ref).transpose() * AsolveB); - var res = (arena_D.val_op() * BTAsolveB).trace(); + var res = (arena_D.val() * BTAsolveB).trace(); - reverse_pass_callback( - [arena_A, BTAsolveB, AsolveB, arena_D, res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback([arena_A, BTAsolveB, AsolveB, arena_D, + res]() mutable { + double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().transpose() - * AsolveB.transpose(); - arena_D.adj() += C_adj * BTAsolveB; - }); + arena_A.adj() + -= C_adj * AsolveB * arena_D.val().transpose() * AsolveB.transpose(); + arena_D.adj() += C_adj * BTAsolveB; + }); return res; } else if constexpr (is_autodiff_v && is_constant_all_v) { @@ -109,25 +109,25 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, LDLT_factor& A, reverse_pass_callback([arena_A, AsolveB, arena_D, res]() mutable { double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().transpose() - * AsolveB.transpose(); + arena_A.adj() + -= C_adj * AsolveB * arena_D.transpose() * AsolveB.transpose(); }); return res; } else if constexpr (is_constant_v && is_all_autodiff_v) { arena_t> arena_B = B; arena_t> arena_D = D; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); - auto BTAsolveB = to_arena(arena_B.val_op().transpose() * AsolveB); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); + auto BTAsolveB = to_arena(arena_B.val().transpose() * AsolveB); - var res = (arena_D.val_op() * BTAsolveB).trace(); + var res = (arena_D.val() * BTAsolveB).trace(); reverse_pass_callback( [BTAsolveB, AsolveB, arena_B, arena_D, res]() mutable { double C_adj = res.adj(); - arena_B.adj() += C_adj * AsolveB - * (arena_D.val_op() + arena_D.val_op().transpose()); + arena_B.adj() + += C_adj * AsolveB * (arena_D.val() + arena_D.val().transpose()); arena_D.adj() += C_adj * BTAsolveB; }); @@ -135,9 +135,9 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, LDLT_factor& A, } else if constexpr (is_constant_all_v && is_autodiff_v) { arena_t> arena_B = B; arena_t> arena_D = value_of(D); - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_D * arena_B.val_op().transpose() * AsolveB).trace(); + var res = (arena_D * arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([AsolveB, arena_B, arena_D, res]() mutable { arena_B.adj() += res.adj() * AsolveB * (arena_D + arena_D.transpose()); @@ -150,7 +150,7 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, LDLT_factor& A, auto BTAsolveB = to_arena(value_of(B_ref).transpose() * A.ldlt().solve(value_of(B_ref))); - var res = (arena_D.val_op() * BTAsolveB).trace(); + var res = (arena_D.val() * BTAsolveB).trace(); reverse_pass_callback([BTAsolveB, arena_D, res]() mutable { arena_D.adj() += res.adj() * BTAsolveB; @@ -194,30 +194,30 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, const LDLT_factor& A, arena_t> arena_A = A.matrix(); arena_t> arena_B = B; arena_t> arena_D = D; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); - auto BTAsolveB = to_arena(arena_B.val_op().transpose() * AsolveB); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); + auto BTAsolveB = to_arena(arena_B.val().transpose() * AsolveB); - var res = (arena_D.val_op().asDiagonal() * BTAsolveB).trace(); + var res = (arena_D.val().asDiagonal() * BTAsolveB).trace(); - reverse_pass_callback( - [arena_A, BTAsolveB, AsolveB, arena_B, arena_D, res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback([arena_A, BTAsolveB, AsolveB, arena_B, arena_D, + res]() mutable { + double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().asDiagonal() - * AsolveB.transpose(); - arena_B.adj() += C_adj * AsolveB * 2 * arena_D.val_op().asDiagonal(); - arena_D.adj() += C_adj * BTAsolveB.diagonal(); - }); + arena_A.adj() + -= C_adj * AsolveB * arena_D.val().asDiagonal() * AsolveB.transpose(); + arena_B.adj() += C_adj * AsolveB * 2 * arena_D.val().asDiagonal(); + arena_D.adj() += C_adj * BTAsolveB.diagonal(); + }); return res; } else if constexpr (is_all_autodiff_v && is_constant_v) { arena_t> arena_A = A.matrix(); arena_t> arena_B = B; arena_t> arena_D = value_of(D); - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_D.asDiagonal() * arena_B.val_op().transpose() * AsolveB) - .trace(); + var res + = (arena_D.asDiagonal() * arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([arena_A, AsolveB, arena_B, arena_D, res]() mutable { double C_adj = res.adj(); @@ -235,16 +235,16 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, const LDLT_factor& A, auto AsolveB = to_arena(A.ldlt().solve(value_of(B_ref))); auto BTAsolveB = to_arena(value_of(B_ref).transpose() * AsolveB); - var res = (arena_D.val_op().asDiagonal() * BTAsolveB).trace(); + var res = (arena_D.val().asDiagonal() * BTAsolveB).trace(); - reverse_pass_callback( - [arena_A, BTAsolveB, AsolveB, arena_D, res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback([arena_A, BTAsolveB, AsolveB, arena_D, + res]() mutable { + double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().asDiagonal() - * AsolveB.transpose(); - arena_D.adj() += C_adj * BTAsolveB.diagonal(); - }); + arena_A.adj() + -= C_adj * AsolveB * arena_D.val().asDiagonal() * AsolveB.transpose(); + arena_D.adj() += C_adj * BTAsolveB.diagonal(); + }); return res; } else if constexpr (is_autodiff_v && is_constant_all_v) { @@ -259,24 +259,24 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, const LDLT_factor& A, reverse_pass_callback([arena_A, AsolveB, arena_D, res]() mutable { double C_adj = res.adj(); - arena_A.adj() -= C_adj * AsolveB * arena_D.val_op().asDiagonal() - * AsolveB.transpose(); + arena_A.adj() + -= C_adj * AsolveB * arena_D.val().asDiagonal() * AsolveB.transpose(); }); return res; } else if constexpr (is_constant_v && is_all_autodiff_v) { arena_t> arena_B = B; arena_t> arena_D = D; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); - auto BTAsolveB = to_arena(arena_B.val_op().transpose() * AsolveB); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); + auto BTAsolveB = to_arena(arena_B.val().transpose() * AsolveB); - var res = (arena_D.val_op().asDiagonal() * BTAsolveB).trace(); + var res = (arena_D.val().asDiagonal() * BTAsolveB).trace(); reverse_pass_callback( [BTAsolveB, AsolveB, arena_B, arena_D, res]() mutable { double C_adj = res.adj(); - arena_B.adj() += C_adj * AsolveB * 2 * arena_D.val_op().asDiagonal(); + arena_B.adj() += C_adj * AsolveB * 2 * arena_D.val().asDiagonal(); arena_D.adj() += C_adj * BTAsolveB.diagonal(); }); @@ -284,10 +284,10 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, const LDLT_factor& A, } else if constexpr (is_constant_all_v && is_autodiff_v) { arena_t> arena_B = B; arena_t> arena_D = value_of(D); - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_D.asDiagonal() * arena_B.val_op().transpose() * AsolveB) - .trace(); + var res + = (arena_D.asDiagonal() * arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([AsolveB, arena_B, arena_D, res]() mutable { arena_B.adj() += res.adj() * AsolveB * 2 * arena_D.asDiagonal(); @@ -300,7 +300,7 @@ inline var trace_gen_inv_quad_form_ldlt(const Td& D, const LDLT_factor& A, auto BTAsolveB = to_arena(value_of(B_ref).transpose() * A.ldlt().solve(value_of(B_ref))); - var res = (arena_D.val_op().asDiagonal() * BTAsolveB).trace(); + var res = (arena_D.val().asDiagonal() * BTAsolveB).trace(); reverse_pass_callback([BTAsolveB, arena_D, res]() mutable { arena_D.adj() += res.adj() * BTAsolveB.diagonal(); diff --git a/stan/math/rev/fun/trace_gen_quad_form.hpp b/stan/math/rev/fun/trace_gen_quad_form.hpp index 98b6a2a7b8c..4466a126024 100644 --- a/stan/math/rev/fun/trace_gen_quad_form.hpp +++ b/stan/math/rev/fun/trace_gen_quad_form.hpp @@ -146,8 +146,8 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { arena_t> arena_A = A; arena_t> arena_B = B; - auto arena_BDT = to_arena(arena_B.val_op() * arena_D.val_op().transpose()); - auto arena_AB = to_arena(arena_A.val_op() * arena_B.val_op()); + auto arena_BDT = to_arena(arena_B.val() * arena_D.val().transpose()); + auto arena_AB = to_arena(arena_A.val() * arena_B.val()); var res = (arena_BDT.transpose() * arena_AB).trace(); @@ -155,13 +155,13 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { [arena_A, arena_B, arena_D, arena_BDT, arena_AB, res]() mutable { double C_adj = res.adj(); - arena_A.adj() += C_adj * arena_BDT * arena_B.val_op().transpose(); + arena_A.adj() += C_adj * arena_BDT * arena_B.val().transpose(); arena_B.adj() += C_adj - * (arena_AB * arena_D.val_op() - + arena_A.val_op().transpose() * arena_BDT); + * (arena_AB * arena_D.val() + + arena_A.val().transpose() * arena_BDT); - arena_D.adj() += C_adj * (arena_AB.transpose() * arena_B.val_op()); + arena_D.adj() += C_adj * (arena_AB.transpose() * arena_B.val()); }); return res; @@ -170,20 +170,20 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { arena_t> arena_A = A; arena_t> arena_B = B; - auto arena_BDT = to_arena(arena_B.val_op() * arena_D.transpose()); - auto arena_AB = to_arena(arena_A.val_op() * arena_B.val_op()); + auto arena_BDT = to_arena(arena_B.val() * arena_D.transpose()); + auto arena_AB = to_arena(arena_A.val() * arena_B.val()); var res = (arena_BDT.transpose() * arena_AB).trace(); - reverse_pass_callback([arena_A, arena_B, arena_D, arena_BDT, arena_AB, - res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback( + [arena_A, arena_B, arena_D, arena_BDT, arena_AB, res]() mutable { + double C_adj = res.adj(); - arena_A.adj() += C_adj * arena_BDT * arena_B.val_op().transpose(); - arena_B.adj() - += C_adj - * (arena_AB * arena_D + arena_A.val_op().transpose() * arena_BDT); - }); + arena_A.adj() += C_adj * arena_BDT * arena_B.val().transpose(); + arena_B.adj() + += C_adj + * (arena_AB * arena_D + arena_A.val().transpose() * arena_BDT); + }); return res; } else if constexpr (is_all_autodiff_v && is_constant_v) { @@ -191,10 +191,10 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { arena_t> arena_A = A; arena_t> arena_B = value_of(B); - auto arena_BDT = to_arena(arena_B.val_op() * arena_D.val_op().transpose()); - auto arena_AB = to_arena(arena_A.val_op() * arena_B.val_op()); + auto arena_BDT = to_arena(arena_B * arena_D.val().transpose()); + auto arena_AB = to_arena(arena_A.val() * arena_B); - var res = (arena_BDT.transpose() * arena_A.val_op() * arena_B).trace(); + var res = (arena_BDT.transpose() * arena_A.val() * arena_B).trace(); reverse_pass_callback( [arena_A, arena_B, arena_D, arena_BDT, arena_AB, res]() mutable { @@ -212,10 +212,10 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { auto arena_BDT = to_arena(arena_B * arena_D); - var res = (arena_BDT.transpose() * arena_A.val_op() * arena_B).trace(); + var res = (arena_BDT.transpose() * arena_A.val() * arena_B).trace(); reverse_pass_callback([arena_A, arena_B, arena_BDT, res]() mutable { - arena_A.adj() += res.adj() * arena_BDT * arena_B.val_op().transpose(); + arena_A.adj() += res.adj() * arena_BDT * arena_B.transpose(); }); return res; @@ -224,21 +224,21 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { arena_t> arena_A = value_of(A); arena_t> arena_B = B; - auto arena_AB = to_arena(arena_A * arena_B.val_op()); - auto arena_BDT = to_arena(arena_B.val_op() * arena_D.val_op()); + auto arena_AB = to_arena(arena_A * arena_B.val()); + auto arena_BDT = to_arena(arena_B.val() * arena_D.val()); var res = (arena_BDT.transpose() * arena_AB).trace(); - reverse_pass_callback([arena_A, arena_B, arena_D, arena_AB, arena_BDT, - res]() mutable { - double C_adj = res.adj(); + reverse_pass_callback( + [arena_A, arena_B, arena_D, arena_AB, arena_BDT, res]() mutable { + double C_adj = res.adj(); - arena_B.adj() - += C_adj - * (arena_AB * arena_D.val_op() + arena_A.transpose() * arena_BDT); + arena_B.adj() + += C_adj + * (arena_AB * arena_D.val() + arena_A.transpose() * arena_BDT); - arena_D.adj() += C_adj * (arena_AB.transpose() * arena_B.val_op()); - }); + arena_D.adj() += C_adj * (arena_AB.transpose() * arena_B.val()); + }); return res; } else if constexpr (is_constant_all_v && is_autodiff_v) { @@ -246,17 +246,16 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { arena_t> arena_A = value_of(A); arena_t> arena_B = B; - auto arena_AB = to_arena(arena_A * arena_B.val_op()); - auto arena_BDT = to_arena(arena_B.val_op() * arena_D.val_op()); + auto arena_AB = to_arena(arena_A * arena_B.val()); + auto arena_BDT = to_arena(arena_B.val() * arena_D); var res = (arena_BDT.transpose() * arena_AB).trace(); - reverse_pass_callback( - [arena_A, arena_B, arena_D, arena_AB, arena_BDT, res]() mutable { - arena_B.adj() += res.adj() - * (arena_AB * arena_D.val_op() - + arena_A.val_op().transpose() * arena_BDT); - }); + reverse_pass_callback([arena_A, arena_B, arena_D, arena_AB, arena_BDT, + res]() mutable { + arena_B.adj() + += res.adj() * (arena_AB * arena_D + arena_A.transpose() * arena_BDT); + }); return res; } else if constexpr (is_constant_all_v && is_autodiff_v) { @@ -266,7 +265,7 @@ inline var trace_gen_quad_form(const Td& D, const Ta& A, const Tb& B) { auto arena_AB = to_arena(arena_A * arena_B); - var res = (arena_D.val_op() * arena_B.transpose() * arena_AB).trace(); + var res = (arena_D.val() * arena_B.transpose() * arena_AB).trace(); reverse_pass_callback([arena_AB, arena_B, arena_D, res]() mutable { arena_D.adj() += res.adj() * (arena_AB.transpose() * arena_B); diff --git a/stan/math/rev/fun/trace_inv_quad_form_ldlt.hpp b/stan/math/rev/fun/trace_inv_quad_form_ldlt.hpp index 24773ac64c2..8a03a4f46c1 100644 --- a/stan/math/rev/fun/trace_inv_quad_form_ldlt.hpp +++ b/stan/math/rev/fun/trace_inv_quad_form_ldlt.hpp @@ -39,9 +39,9 @@ inline var trace_inv_quad_form_ldlt(LDLT_factor& A, const T2& B) { if constexpr (is_autodiff_v && is_autodiff_v) { arena_t arena_A = A.matrix(); arena_t arena_B = B; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_B.val_op().transpose() * AsolveB).trace(); + var res = (arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([arena_A, AsolveB, arena_B, res]() mutable { arena_A.adj() += -res.adj() * AsolveB * AsolveB.transpose(); @@ -64,9 +64,9 @@ inline var trace_inv_quad_form_ldlt(LDLT_factor& A, const T2& B) { return res; } else { arena_t arena_B = B; - auto AsolveB = to_arena(A.ldlt().solve(arena_B.val_op())); + auto AsolveB = to_arena(A.ldlt().solve(arena_B.val())); - var res = (arena_B.val_op().transpose() * AsolveB).trace(); + var res = (arena_B.val().transpose() * AsolveB).trace(); reverse_pass_callback([AsolveB, arena_B, res]() mutable { arena_B.adj() += 2 * res.adj() * AsolveB;