|
2 | 2 | #include <stan/math/rev/core.hpp> |
3 | 3 | #include <stan/math.hpp> |
4 | 4 | #include <test/unit/math/rev/util.hpp> |
| 5 | +#include <tuple> |
| 6 | +#include <type_traits> |
5 | 7 | #include <vector> |
6 | 8 |
|
7 | 9 | using stan::math::var; |
@@ -306,3 +308,45 @@ TEST_F(AgradRev, Rev_deep_copy_vars_std_vector_eigen_matrix_var_arg) { |
306 | 308 | EXPECT_NE(out[i](j).vi_, arg[i](j).vi_); |
307 | 309 | } |
308 | 310 | } |
| 311 | + |
| 312 | +TEST_F(AgradRev, Rev_deep_copy_vars_tuple_data_arg) { |
| 313 | + const std::tuple<> empty; |
| 314 | + const auto arg = std::make_tuple(5, Eigen::VectorXd::Ones(2).eval()); |
| 315 | + |
| 316 | + auto empty_out = stan::math::deep_copy_vars(empty); |
| 317 | + auto out = stan::math::deep_copy_vars(arg); |
| 318 | + |
| 319 | + static_assert(std::is_same_v<decltype(empty_out), std::tuple<>>); |
| 320 | + static_assert(std::is_same_v<decltype(out), |
| 321 | + std::tuple<const int&, const Eigen::VectorXd&>>); |
| 322 | + EXPECT_EQ(&std::get<0>(out), &std::get<0>(arg)); |
| 323 | + EXPECT_EQ(&std::get<1>(out), &std::get<1>(arg)); |
| 324 | +} |
| 325 | + |
| 326 | +TEST_F(AgradRev, Rev_deep_copy_vars_nested_tuple_var_arg) { |
| 327 | + Eigen::Matrix<var, Eigen::Dynamic, 1> vars(2); |
| 328 | + vars << 2.0, 3.0; |
| 329 | + auto arg = std::make_tuple(1.0, var(4.0), vars, std::make_tuple(var(5.0), 6)); |
| 330 | + |
| 331 | + auto out = stan::math::deep_copy_vars(arg); |
| 332 | + |
| 333 | + static_assert(std::is_reference_v<std::tuple_element_t<0, decltype(out)>>); |
| 334 | + EXPECT_EQ(&std::get<0>(out), &std::get<0>(arg)); |
| 335 | + EXPECT_FLOAT_EQ(std::get<1>(out).val(), std::get<1>(arg).val()); |
| 336 | + EXPECT_NE(std::get<1>(out).vi_, std::get<1>(arg).vi_); |
| 337 | + for (int i = 0; i < vars.size(); ++i) { |
| 338 | + EXPECT_FLOAT_EQ(std::get<2>(out)(i).val(), std::get<2>(arg)(i).val()); |
| 339 | + EXPECT_NE(std::get<2>(out)(i).vi_, std::get<2>(arg)(i).vi_); |
| 340 | + } |
| 341 | + EXPECT_FLOAT_EQ(std::get<0>(std::get<3>(out)).val(), |
| 342 | + std::get<0>(std::get<3>(arg)).val()); |
| 343 | + EXPECT_NE(std::get<0>(std::get<3>(out)).vi_, |
| 344 | + std::get<0>(std::get<3>(arg)).vi_); |
| 345 | + EXPECT_EQ(&std::get<1>(std::get<3>(out)), &std::get<1>(std::get<3>(arg))); |
| 346 | + |
| 347 | + auto rvalue_out = stan::math::deep_copy_vars( |
| 348 | + std::make_tuple(7.0, var(8.0), std::make_tuple(9))); |
| 349 | + static_assert(std::is_same_v<decltype(rvalue_out), |
| 350 | + std::tuple<double, var, std::tuple<int>>>); |
| 351 | + EXPECT_FLOAT_EQ(8.0, std::get<1>(rvalue_out).val()); |
| 352 | +} |
0 commit comments