Skip to content

Commit 5b59600

Browse files
authored
Merge pull request #3351 from stan-dev/val-op
Alway use CwiseUnaryOp for `.val()` of `var` types
2 parents 20bf734 + c8fc616 commit 5b59600

24 files changed

Lines changed: 232 additions & 216 deletions

doxygen/contributor_help_pages/getting_started.md

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -660,9 +660,15 @@ The values of `x` have the same shape as `x`.
660660
661661
### Values and adjoint extensions to Eigen
662662
663-
The matrix and vector autodiff types come with an extra `.val()` and `.adj()`, member functions called `.val_op()` and `.adj_op()`.
664-
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.
665-
See [here](https://github.com/stan-dev/math/issues/2653) for the details and other information for when this workaround is needed.
663+
By default, `.val()` and `.adj()` return a `CwiseUnaryView` when the matrix/vector is non-`const`,
664+
and returns a `CwiseUnaryOp` when the matrix/vector is `const`. The exception is that
665+
calling `.val()` on a matrix/vector of `var` types will always return a `CwiseUnaryOp`, as the
666+
underlying value is always `const`.
667+
668+
However, using a `CwiseUnaryView` in some Eigen operations (e.g., multiplication, transposition)
669+
can result in a compilation error. To workaround this, the `.val_op()` and `.adj_op()` member
670+
functions have been added to explicitly request a `CwiseUnaryOp` regardless of whether the
671+
matrix/vector is `const` or not.
666672
667673
The member functions `.val()` and `.val_op()` return expressions that evaluate to the values
668674
of the autodiff matrix.

stan/math/prim/eigen_plugins.h

Lines changed: 22 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -82,6 +82,7 @@ struct val_Op{
8282
double& operator()(double& v) const { return v; }
8383
};
8484

85+
8586
/**
8687
* Coefficient-wise function applying val_Op struct to a matrix of const var
8788
* or vari* and returning a view to the const matrix of doubles containing
@@ -94,16 +95,31 @@ val() const { return CwiseUnaryOp<val_Op, const Derived>(derived());
9495
/**
9596
* Coefficient-wise function applying val_Op struct to a matrix of var
9697
* or vari* and returning a view to the values
97-
*/
98+
*/
99+
template <
100+
typename T = Scalar,
101+
std::enable_if_t<
102+
!std::disjunction_v<
103+
std::is_arithmetic<std::decay_t<T>>,
104+
is_fvar<std::decay_t<T>>
105+
>
106+
>* = nullptr>
107+
inline CwiseUnaryOp<val_Op, Derived>
108+
val() { return CwiseUnaryOp<val_Op, Derived>(derived());
109+
}
110+
111+
template <
112+
typename T = Scalar,
113+
std::enable_if_t<
114+
std::disjunction_v<
115+
std::is_arithmetic<std::decay_t<T>>,
116+
is_fvar<std::decay_t<T>>
117+
>
118+
>* = nullptr>
98119
inline CwiseUnaryView<val_Op, Derived>
99120
val() { return CwiseUnaryView<val_Op, Derived>(derived());
100121
}
101122

102-
/**
103-
* Coefficient-wise function applying val_Op struct to a matrix of var
104-
* or vari* and returning a view to the matrix of doubles containing
105-
* the values
106-
*/
107123
inline CwiseUnaryOp<val_Op, Derived>
108124
val_op() { return CwiseUnaryOp<val_Op, Derived>(derived());
109125
}

stan/math/rev/constraint/stochastic_column_constrain.hpp

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ inline plain_type_t<T> stochastic_column_constrain(const T& y) {
3030
const auto M = y.cols();
3131
arena_t<T> arena_y = y;
3232

33-
arena_t<ret_type> arena_x = stochastic_column_constrain(arena_y.val_op());
33+
arena_t<ret_type> arena_x = stochastic_column_constrain(arena_y.val());
3434

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

42-
auto&& x_val = arena_x.val_op();
42+
auto&& x_val = arena_x.val();
4343
auto&& x_adj = arena_x.adj_op();
4444

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

8383
double lp_val = 0;
8484
arena_t<ret_type> arena_x
85-
= stochastic_column_constrain(arena_y.val_op(), lp_val);
85+
= stochastic_column_constrain(arena_y.val(), lp_val);
8686
lp += lp_val;
8787

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

95-
auto&& x_val = arena_x.val_op();
95+
auto&& x_val = arena_x.val();
9696
auto&& x_adj = arena_x.adj_op();
9797

9898
const auto x_val_rows = x_val.rows();

stan/math/rev/constraint/stochastic_row_constrain.hpp

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ inline auto stochastic_row_constrain(const T& y) {
2828
const auto M = y.cols();
2929
arena_t<T> arena_y = y;
3030

31-
arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val_op());
31+
arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val());
3232

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

40-
auto&& x_val = arena_x.val_op();
40+
auto&& x_val = arena_x.val();
4141
auto&& x_adj = arena_x.adj_op();
4242

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

8181
double lp_val = 0;
82-
arena_t<ret_type> arena_x
83-
= stochastic_row_constrain(arena_y.val_op(), lp_val);
82+
arena_t<ret_type> arena_x = stochastic_row_constrain(arena_y.val(), lp_val);
8483
lp += lp_val;
8584

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

93-
auto&& x_val = arena_x.val_op();
92+
auto&& x_val = arena_x.val();
9493
auto&& x_adj = arena_x.adj_op();
9594

9695
const auto x_val_cols = x_val.cols();

stan/math/rev/fun/eigendecompose_sym.hpp

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -41,20 +41,19 @@ inline auto eigendecompose_sym(const T& m) {
4141

4242
reverse_pass_callback([eigenvals, arena_m, eigenvecs]() mutable {
4343
// eigenvalue reverse calculation
44-
auto value_adj = eigenvecs.val_op() * eigenvals.adj().asDiagonal()
45-
* eigenvecs.val_op().transpose();
44+
auto value_adj = eigenvecs.val() * eigenvals.adj().asDiagonal()
45+
* eigenvecs.val().transpose();
4646
// eigenvector reverse calculation
4747
const auto p = arena_m.val().cols();
48-
Eigen::MatrixXd f
49-
= (1
50-
/ (eigenvals.val_op().rowwise().replicate(p).transpose()
51-
- eigenvals.val_op().rowwise().replicate(p))
52-
.array());
48+
Eigen::MatrixXd f = (1
49+
/ (eigenvals.val().rowwise().replicate(p).transpose()
50+
- eigenvals.val().rowwise().replicate(p))
51+
.array());
5352
f.diagonal().setZero();
5453
auto vector_adj
55-
= eigenvecs.val_op()
56-
* f.cwiseProduct(eigenvecs.val_op().transpose() * eigenvecs.adj_op())
57-
* eigenvecs.val_op().transpose();
54+
= eigenvecs.val()
55+
* f.cwiseProduct(eigenvecs.val().transpose() * eigenvecs.adj_op())
56+
* eigenvecs.val().transpose();
5857

5958
arena_m.adj() += value_adj + vector_adj;
6059
});

stan/math/rev/fun/eigenvectors_sym.hpp

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -43,9 +43,9 @@ inline auto eigenvectors_sym(const T& m) {
4343
.array());
4444
f.diagonal().setZero();
4545
arena_m.adj()
46-
+= eigenvecs.val_op()
47-
* f.cwiseProduct(eigenvecs.val_op().transpose() * eigenvecs.adj_op())
48-
* eigenvecs.val_op().transpose();
46+
+= eigenvecs.val()
47+
* f.cwiseProduct(eigenvecs.val().transpose() * eigenvecs.adj_op())
48+
* eigenvecs.val().transpose();
4949
});
5050

5151
return return_t(eigenvecs);

stan/math/rev/fun/generalized_inverse.hpp

Lines changed: 10 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -24,15 +24,13 @@ template <typename T1, typename T2>
2424
inline auto generalized_inverse_lambda(T1& G_arena, T2& inv_G) {
2525
return [G_arena, inv_G]() mutable {
2626
G_arena.adj()
27-
+= -(inv_G.val_op().transpose() * inv_G.adj_op()
28-
* inv_G.val_op().transpose())
29-
+ (-G_arena.val_op() * inv_G.val_op()
27+
+= -(inv_G.val().transpose() * inv_G.adj_op() * inv_G.val().transpose())
28+
+ (-G_arena.val() * inv_G.val()
3029
+ Eigen::MatrixXd::Identity(G_arena.rows(), inv_G.cols()))
31-
* inv_G.adj_op().transpose() * inv_G.val_op()
32-
* inv_G.val_op().transpose()
33-
+ inv_G.val_op().transpose() * inv_G.val_op()
34-
* inv_G.adj_op().transpose()
35-
* (-inv_G.val_op() * G_arena.val_op()
30+
* inv_G.adj_op().transpose() * inv_G.val()
31+
* inv_G.val().transpose()
32+
+ inv_G.val().transpose() * inv_G.val() * inv_G.adj_op().transpose()
33+
* (-inv_G.val() * G_arena.val()
3634
+ Eigen::MatrixXd::Identity(inv_G.rows(), G_arena.cols()));
3735
};
3836
}
@@ -83,17 +81,17 @@ inline auto generalized_inverse(const VarMat& G) {
8381
}
8482
} else if (G.rows() < G.cols()) {
8583
arena_t<VarMat> G_arena(G);
86-
arena_t<ret_type> inv_G((G_arena.val_op() * G_arena.val_op().transpose())
84+
arena_t<ret_type> inv_G((G_arena.val() * G_arena.val().transpose())
8785
.ldlt()
88-
.solve(G_arena.val_op())
86+
.solve(G_arena.val())
8987
.transpose());
9088
reverse_pass_callback(internal::generalized_inverse_lambda(G_arena, inv_G));
9189
return ret_type(inv_G);
9290
} else {
9391
arena_t<VarMat> G_arena(G);
94-
arena_t<ret_type> inv_G((G_arena.val_op().transpose() * G_arena.val_op())
92+
arena_t<ret_type> inv_G((G_arena.val().transpose() * G_arena.val())
9593
.ldlt()
96-
.solve(G_arena.val_op().transpose()));
94+
.solve(G_arena.val().transpose()));
9795
reverse_pass_callback(internal::generalized_inverse_lambda(G_arena, inv_G));
9896
return ret_type(inv_G);
9997
}

stan/math/rev/fun/inverse.hpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ inline auto inverse(const T& m) {
3030
}
3131

3232
arena_t<T> arena_m = m;
33-
arena_t<promote_scalar_t<double, T>> res_val = arena_m.val_op().inverse();
33+
arena_t<promote_scalar_t<double, T>> res_val = arena_m.val().inverse();
3434
arena_t<ret_type> res = res_val;
3535

3636
reverse_pass_callback([res, res_val, arena_m]() mutable {

stan/math/rev/fun/mdivide_left.hpp

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -42,24 +42,24 @@ inline auto mdivide_left(T1&& A, T2&& B) {
4242
if constexpr (is_autodiff_v<T1> && is_autodiff_v<T2>) {
4343
arena_t<T1> arena_A(std::forward<T1>(A));
4444
arena_t<T2> arena_B(std::forward<T2>(B));
45-
auto hqr_A_ptr = make_chainable_ptr(arena_A.val_op().householderQr());
46-
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val_op());
45+
auto hqr_A_ptr = make_chainable_ptr(arena_A.val().householderQr());
46+
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val());
4747
reverse_pass_callback([arena_A, arena_B, hqr_A_ptr, res]() mutable {
4848
promote_scalar_t<double, T2> adjB
4949
= hqr_A_ptr->householderQ()
5050
* hqr_A_ptr->matrixQR()
5151
.template triangularView<Eigen::Upper>()
5252
.transpose()
5353
.solve(res.adj());
54-
arena_A.adj() -= adjB * res.val_op().transpose();
54+
arena_A.adj() -= adjB * res.val().transpose();
5555
arena_B.adj() += adjB;
5656
});
5757

5858
return ret_type(res);
5959
} else if constexpr (is_autodiff_v<T2>) {
6060
arena_t<T2> arena_B(std::forward<T2>(B));
6161
auto hqr_A_ptr = make_chainable_ptr(value_of(A).householderQr());
62-
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val_op());
62+
arena_t<ret_type> res = hqr_A_ptr->solve(arena_B.val());
6363
reverse_pass_callback([arena_B, hqr_A_ptr, res]() mutable {
6464
arena_B.adj() += hqr_A_ptr->householderQ()
6565
* hqr_A_ptr->matrixQR()
@@ -70,15 +70,15 @@ inline auto mdivide_left(T1&& A, T2&& B) {
7070
return ret_type(res);
7171
} else {
7272
arena_t<T1> arena_A(std::forward<T1>(A));
73-
auto hqr_A_ptr = make_chainable_ptr(arena_A.val_op().householderQr());
73+
auto hqr_A_ptr = make_chainable_ptr(arena_A.val().householderQr());
7474
arena_t<ret_type> res = hqr_A_ptr->solve(value_of(B));
7575
reverse_pass_callback([arena_A, hqr_A_ptr, res]() mutable {
7676
arena_A.adj() -= hqr_A_ptr->householderQ()
7777
* hqr_A_ptr->matrixQR()
7878
.template triangularView<Eigen::Upper>()
7979
.transpose()
8080
.solve(res.adj())
81-
* res.val_op().transpose();
81+
* res.val().transpose();
8282
});
8383
return ret_type(res);
8484
}

stan/math/rev/fun/mdivide_left_ldlt.hpp

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -39,13 +39,13 @@ inline auto mdivide_left_ldlt(LDLT_factor<T1>& A, const T2& B) {
3939
if constexpr (is_autodiff_v<T1> && is_autodiff_v<T2>) {
4040
arena_t<promote_scalar_t<var, T2>> arena_B = B;
4141
arena_t<promote_scalar_t<var, T1>> arena_A = A.matrix();
42-
arena_t<ret_type> res = A.ldlt().solve(arena_B.val_op());
42+
arena_t<ret_type> res = A.ldlt().solve(arena_B.val());
4343
const auto* ldlt_ptr = make_chainable_ptr(A.ldlt());
4444

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

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

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

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

6262
return ret_type(res);
6363
} else {
6464
arena_t<promote_scalar_t<var, T2>> arena_B = B;
65-
arena_t<ret_type> res = A.ldlt().solve(arena_B.val_op());
65+
arena_t<ret_type> res = A.ldlt().solve(arena_B.val());
6666
const auto* ldlt_ptr = make_chainable_ptr(A.ldlt());
6767

6868
reverse_pass_callback([arena_B, ldlt_ptr, res]() mutable {

0 commit comments

Comments
 (0)