File indexing completed on 2026-08-17 08:48:11
0001
0002
0003
0004
0005 #ifndef BOOST_MATH_DIFFERENTIATION_AUTODIFF_HPP
0006 #define BOOST_MATH_DIFFERENTIATION_AUTODIFF_HPP
0007
0008 #include <boost/math/constants/constants.hpp>
0009
0010 #if defined(BOOST_MATH_REVERSE_MODE_ET_OFF) && defined(BOOST_MATH_REVERSE_MODE_ET_ON)
0011 #error "Cannot define both BOOST_MATH_REVERSE_MODE_ET_OFF and BOOST_MATH_REVERSE_MODE_ET_ON"
0012 #endif
0013
0014 #if !defined(BOOST_MATH_REVERSE_MODE_ET_OFF) && !defined(BOOST_MATH_REVERSE_MODE_ET_ON)
0015 #define BOOST_MATH_REVERSE_MODE_ET_ON
0016 #endif
0017
0018 #ifdef BOOST_MATH_REVERSE_MODE_ET_ON
0019 #include <boost/math/differentiation/detail/reverse_mode_autodiff_basic_ops_et.hpp>
0020 #include <boost/math/differentiation/detail/reverse_mode_autodiff_stl_et.hpp>
0021 #else
0022 #include <boost/math/differentiation/detail/reverse_mode_autodiff_basic_ops_no_et.hpp>
0023 #include <boost/math/differentiation/detail/reverse_mode_autodiff_stl_no_et.hpp>
0024 #endif
0025
0026 #include <boost/math/differentiation/detail/reverse_mode_autodiff_comparison_operator_overloads.hpp>
0027 #include <boost/math/differentiation/detail/reverse_mode_autodiff_erf_overloads.hpp>
0028 #include <boost/math/differentiation/detail/reverse_mode_autodiff_expression_template_base.hpp>
0029 #include <boost/math/differentiation/detail/reverse_mode_autodiff_memory_management.hpp>
0030 #include <boost/math/special_functions/acosh.hpp>
0031 #include <boost/math/special_functions/asinh.hpp>
0032 #include <boost/math/special_functions/atanh.hpp>
0033 #include <boost/math/special_functions/digamma.hpp>
0034 #include <boost/math/special_functions/erf.hpp>
0035 #include <boost/math/special_functions/lambert_w.hpp>
0036 #include <boost/math/special_functions/polygamma.hpp>
0037 #include <boost/math/special_functions/round.hpp>
0038 #include <boost/math/special_functions/trunc.hpp>
0039 #include <boost/math/tools/config.hpp>
0040 #include <boost/math/tools/promotion.hpp>
0041 #include <cstddef>
0042 #include <iostream>
0043 #include <type_traits>
0044 #include <vector>
0045 #define BOOST_MATH_BUFFER_SIZE 65536
0046
0047 namespace boost {
0048 namespace math {
0049 namespace differentiation {
0050 namespace reverse_mode {
0051
0052
0053 template<typename RealType, size_t DerivativeOrder, class DerivedExpression>
0054 struct expression;
0055
0056 template<typename RealType, size_t DerivativeOrder>
0057 class rvar;
0058
0059 template<typename RealType,
0060 size_t DerivativeOrder,
0061 typename LHS,
0062 typename RHS,
0063 typename ConcreteBinaryOperation>
0064 struct abstract_binary_expression;
0065
0066 template<typename RealType, size_t DerivativeOrder, typename ARG, typename ConcreteBinaryOperation>
0067 struct abstract_unary_expression;
0068
0069 template<typename RealType, size_t DerivativeOrder>
0070 class gradient_node;
0071
0072 template<typename RealType, size_t DerivativeOrder, size_t buffer_size = BOOST_MATH_BUFFER_SIZE>
0073 class gradient_tape
0074 {
0075
0076
0077 private:
0078
0079 using inner_t = rvar_t<RealType, DerivativeOrder - 1>;
0080
0081
0082 detail::flat_linear_allocator<inner_t, buffer_size> adjoints_;
0083 detail::flat_linear_allocator<inner_t, buffer_size> derivatives_;
0084 detail::flat_linear_allocator<gradient_node<RealType, DerivativeOrder>, buffer_size>
0085 gradient_nodes_;
0086 detail::flat_linear_allocator<gradient_node<RealType, DerivativeOrder> *, buffer_size>
0087 argument_nodes_;
0088
0089
0090 template<size_t n>
0091 gradient_node<RealType, DerivativeOrder> *fill_node_at_compile_time(
0092 std::true_type, gradient_node<RealType, DerivativeOrder> *node_ptr)
0093 {
0094 node_ptr->derivatives_ = derivatives_.template emplace_back_n<n>();
0095 node_ptr->argument_nodes_ = argument_nodes_.template emplace_back_n<n>();
0096 return node_ptr;
0097 }
0098
0099 template<size_t n>
0100 gradient_node<RealType, DerivativeOrder> *fill_node_at_compile_time(
0101 std::false_type, gradient_node<RealType, DerivativeOrder> *node_ptr)
0102 {
0103 node_ptr->derivatives_ = nullptr;
0104 node_ptr->argument_adjoints_ = nullptr;
0105 node_ptr->argument_nodes_ = nullptr;
0106 return node_ptr;
0107 }
0108
0109 public:
0110
0111
0112
0113
0114 using adjoint_iterator = typename detail::flat_linear_allocator<inner_t, buffer_size>::iterator;
0115 using derivatives_iterator =
0116 typename detail::flat_linear_allocator<inner_t, buffer_size>::iterator;
0117 using gradient_nodes_iterator =
0118 typename detail::flat_linear_allocator<gradient_node<RealType, DerivativeOrder>,
0119 buffer_size>::iterator;
0120 using argument_nodes_iterator =
0121 typename detail::flat_linear_allocator<gradient_node<RealType, DerivativeOrder> *,
0122 buffer_size>::iterator;
0123
0124 gradient_tape() { clear(); };
0125
0126 gradient_tape(const gradient_tape &) = delete;
0127 gradient_tape &operator=(const gradient_tape &) = delete;
0128 gradient_tape(gradient_tape &&other) = delete;
0129 gradient_tape operator=(gradient_tape &&other) = delete;
0130 ~gradient_tape() noexcept { clear(); }
0131 void clear()
0132 {
0133 adjoints_.clear();
0134 derivatives_.clear();
0135 gradient_nodes_.clear();
0136 argument_nodes_.clear();
0137 }
0138
0139
0140 gradient_node<RealType, DerivativeOrder> *emplace_leaf_node()
0141 {
0142 gradient_node<RealType, DerivativeOrder> *node = &*gradient_nodes_.emplace_back();
0143 node->adjoint_ = adjoints_.emplace_back();
0144 node->derivatives_ = derivatives_iterator();
0145 node->argument_nodes_ = argument_nodes_iterator();
0146
0147 return node;
0148 };
0149
0150
0151 gradient_node<RealType, DerivativeOrder> *emplace_active_unary_node()
0152 {
0153 gradient_node<RealType, DerivativeOrder> *node = &*gradient_nodes_.emplace_back();
0154 node->n_ = 1;
0155 node->adjoint_ = adjoints_.emplace_back();
0156 node->derivatives_ = derivatives_.emplace_back();
0157
0158 return node;
0159 };
0160
0161
0162 template<size_t n>
0163 gradient_node<RealType, DerivativeOrder> *emplace_active_multi_node()
0164 {
0165 gradient_node<RealType, DerivativeOrder> *node = &*gradient_nodes_.emplace_back();
0166 node->n_ = n;
0167 node->adjoint_ = adjoints_.emplace_back();
0168
0169 return fill_node_at_compile_time<n>(std::integral_constant<bool, (n > 0)>{}, node);
0170 }
0171
0172
0173 gradient_node<RealType, DerivativeOrder> *emplace_active_multi_node(size_t n)
0174 {
0175 gradient_node<RealType, DerivativeOrder> *node = &*gradient_nodes_.emplace_back();
0176 node->n_ = n;
0177 node->adjoint_ = adjoints_.emplace_back();
0178 if (n > 0) {
0179 node->derivatives_ = derivatives_.emplace_back_n(n);
0180 node->argument_nodes_ = argument_nodes_.emplace_back_n(n);
0181 }
0182 return node;
0183 };
0184
0185 void zero_grad()
0186 {
0187 const RealType zero = RealType(0.0);
0188 adjoints_.fill(zero);
0189 }
0190
0191
0192 auto begin() { return gradient_nodes_.begin(); }
0193 auto end() { return gradient_nodes_.end(); }
0194 auto find(gradient_node<RealType, DerivativeOrder> *node)
0195 {
0196 return gradient_nodes_.find(node);
0197 };
0198 void add_checkpoint()
0199 {
0200 gradient_nodes_.add_checkpoint();
0201 adjoints_.add_checkpoint();
0202 derivatives_.add_checkpoint();
0203 argument_nodes_.add_checkpoint();
0204 };
0205
0206 auto last_checkpoint() { return gradient_nodes_.last_checkpoint(); };
0207 auto first_checkpoint() { return gradient_nodes_.last_checkpoint(); };
0208 auto checkpoint_at(size_t index) { return gradient_nodes_.get_checkpoint_at(index); };
0209 void rewind_to_last_checkpoint()
0210 {
0211 gradient_nodes_.rewind_to_last_checkpoint();
0212 adjoints_.rewind_to_last_checkpoint();
0213 derivatives_.rewind_to_last_checkpoint();
0214 argument_nodes_.rewind_to_last_checkpoint();
0215 };
0216 void rewind_to_checkpoint_at(size_t index)
0217
0218 {
0219 gradient_nodes_.rewind_to_checkpoint_at(index);
0220 adjoints_.rewind_to_checkpoint_at(index);
0221 derivatives_.rewind_to_checkpoint_at(index);
0222 argument_nodes_.rewind_to_checkpoint_at(index);
0223 }
0224
0225
0226 void rewind()
0227 {
0228 gradient_nodes_.rewind();
0229 adjoints_.rewind();
0230 derivatives_.rewind();
0231 argument_nodes_.rewind();
0232 }
0233
0234
0235 gradient_node<RealType, DerivativeOrder> &operator[](size_t i) { return gradient_nodes_[i]; }
0236 const gradient_node<RealType, DerivativeOrder> &operator[](size_t i) const
0237 {
0238 return gradient_nodes_[i];
0239 }
0240 };
0241
0242 template<typename RealType, size_t DerivativeOrder>
0243 class gradient_node
0244 {
0245
0246
0247
0248
0249 public:
0250 using adjoint_iterator = typename gradient_tape<RealType, DerivativeOrder>::adjoint_iterator;
0251 using derivatives_iterator =
0252 typename gradient_tape<RealType, DerivativeOrder>::derivatives_iterator;
0253 using argument_nodes_iterator =
0254 typename gradient_tape<RealType, DerivativeOrder>::argument_nodes_iterator;
0255
0256 private:
0257 size_t n_;
0258 using inner_t = rvar_t<RealType, DerivativeOrder - 1>;
0259
0260
0261
0262 adjoint_iterator adjoint_;
0263 derivatives_iterator derivatives_;
0264 argument_nodes_iterator argument_nodes_;
0265
0266 public:
0267 friend class gradient_tape<RealType, DerivativeOrder>;
0268 friend class rvar<RealType, DerivativeOrder>;
0269
0270 gradient_node() = default;
0271 explicit gradient_node(const size_t n)
0272 : n_(n)
0273 , adjoint_(nullptr)
0274 , derivatives_(nullptr)
0275 {}
0276 explicit gradient_node(const size_t n,
0277 RealType *adjoint,
0278 RealType *derivatives,
0279 rvar<RealType, DerivativeOrder> **arguments)
0280 : n_(n)
0281 , adjoint_(adjoint)
0282 , derivatives_(derivatives)
0283 { static_cast<void>(arguments); }
0284
0285 inner_t get_adjoint_v() const { return *adjoint_; }
0286 inner_t get_derivative_v(size_t arg_id) const { return derivatives_[static_cast<ptrdiff_t>(arg_id)]; };
0287 inner_t get_argument_adjoint_v(size_t arg_id) const
0288 {
0289 return *argument_nodes_[static_cast<ptrdiff_t>(arg_id)]->adjoint_;
0290 }
0291
0292 adjoint_iterator get_adjoint_ptr() { return adjoint_; }
0293 adjoint_iterator get_adjoint_ptr() const { return adjoint_; };
0294 void update_adjoint_v(inner_t value) { *adjoint_ = value; };
0295 void update_derivative_v(size_t arg_id, inner_t value) { derivatives_[static_cast<ptrdiff_t>(arg_id)] = value; };
0296 void update_argument_adj_v(size_t arg_id, inner_t value)
0297 {
0298 argument_nodes_[static_cast<ptrdiff_t>(arg_id)]->update_adjoint_v(value);
0299 };
0300 void update_argument_ptr_at(size_t arg_id, gradient_node<RealType, DerivativeOrder> *node_ptr)
0301 {
0302 argument_nodes_[static_cast<ptrdiff_t>(arg_id)] = node_ptr;
0303 }
0304
0305 void backward()
0306 {
0307 if (!n_)
0308 return;
0309
0310 using boost::math::differentiation::reverse_mode::fabs;
0311 using std::fabs;
0312 if (!adjoint_ || fabs(*adjoint_) < 2 * std::numeric_limits<RealType>::epsilon())
0313 return;
0314
0315 if (!argument_nodes_)
0316 return;
0317
0318 if (!derivatives_)
0319 return;
0320
0321 for (size_t i = 0; i < n_; ++i) {
0322 auto adjoint = get_adjoint_v();
0323 auto derivative = get_derivative_v(i);
0324 auto argument_adjoint = get_argument_adjoint_v(i);
0325 update_argument_adj_v(i, argument_adjoint + derivative * adjoint);
0326 }
0327 }
0328 };
0329
0330
0331 template<typename RealType, size_t DerivativeOrder>
0332 inline gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &get_active_tape()
0333 {
0334 static BOOST_MATH_THREAD_LOCAL gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE>
0335 tape;
0336 return tape;
0337 }
0338
0339 template<typename RealType, size_t DerivativeOrder = 1>
0340 class rvar : public expression<RealType, DerivativeOrder, rvar<RealType, DerivativeOrder>>
0341 {
0342 private:
0343 using inner_t = rvar_t<RealType, DerivativeOrder - 1>;
0344 friend class gradient_node<RealType, DerivativeOrder>;
0345 inner_t value_;
0346 gradient_node<RealType, DerivativeOrder> *node_ = nullptr;
0347 template<typename, size_t>
0348 friend class rvar;
0349
0350
0351
0352
0353 template<size_t target_order, size_t current_order>
0354 struct get_value_at_impl
0355 {
0356 static_assert(target_order <= current_order, "Requested depth exceeds variable order.");
0357
0358
0359
0360 static auto &get(rvar<RealType, current_order> &v)
0361 {
0362 return get_value_at_impl<target_order, current_order - 1>::get(v.get_value());
0363 }
0364
0365
0366 static const auto &get(const rvar<RealType, current_order> &v)
0367 {
0368 return get_value_at_impl<target_order, current_order - 1>::get(v.get_value());
0369 }
0370 };
0371
0372
0373
0374 template<size_t target_order>
0375 struct get_value_at_impl<target_order, target_order>
0376 {
0377
0378
0379 static auto &get(rvar<RealType, target_order> &v) { return v; }
0380
0381
0382 static const auto &get(const rvar<RealType, target_order> &v) { return v; }
0383 };
0384
0385 void make_leaf_node()
0386 {
0387 gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &tape
0388 = get_active_tape<RealType, DerivativeOrder>();
0389 node_ = tape.emplace_leaf_node();
0390 }
0391
0392 void make_unary_node()
0393 {
0394 gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &tape
0395 = get_active_tape<RealType, DerivativeOrder>();
0396 node_ = tape.emplace_active_unary_node();
0397 }
0398
0399 void make_multi_node(size_t n)
0400 {
0401 gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &tape
0402 = get_active_tape<RealType, DerivativeOrder>();
0403 node_ = tape.emplace_active_multi_node(n);
0404 }
0405
0406 template<size_t n>
0407 void make_multi_node()
0408 {
0409 gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &tape
0410 = get_active_tape<RealType, DerivativeOrder>();
0411 node_ = tape.template emplace_active_multi_node<n>();
0412 }
0413
0414 template<typename E>
0415 void make_rvar_from_expr(const expression<RealType, DerivativeOrder, E> &expr)
0416 {
0417 make_multi_node<detail::count_rvars<E, DerivativeOrder>>();
0418 expr.template propagatex<0>(node_, inner_t(static_cast<RealType>(1.0)));
0419 }
0420 RealType get_item_impl(std::true_type) const
0421 {
0422 return value_.get_item_impl(std::integral_constant<bool, (DerivativeOrder - 1 > 1)>{});
0423 }
0424
0425 RealType get_item_impl(std::false_type) const { return value_; }
0426
0427 public:
0428 using value_type = RealType;
0429 static constexpr size_t DerivativeOrder_v = DerivativeOrder;
0430 rvar()
0431 : value_()
0432 {
0433 make_leaf_node();
0434 }
0435 rvar(const RealType value)
0436 : value_(inner_t{static_cast<RealType>(value)})
0437 {
0438 make_leaf_node();
0439 }
0440
0441 rvar &operator=(RealType v)
0442 {
0443 value_ = inner_t(v);
0444 if (node_ == nullptr) {
0445 make_leaf_node();
0446 }
0447 return *this;
0448 }
0449 rvar(const rvar<RealType, DerivativeOrder> &other) = default;
0450 rvar &operator=(const rvar<RealType, DerivativeOrder> &other) = default;
0451
0452 template<size_t arg_index>
0453 void propagatex(gradient_node<RealType, DerivativeOrder> *node, inner_t adj) const
0454 {
0455 node->update_derivative_v(arg_index, adj);
0456 node->update_argument_ptr_at(arg_index, node_);
0457 }
0458
0459 template<class E>
0460 rvar(const expression<RealType, DerivativeOrder, E> &expr)
0461 {
0462 value_ = expr.evaluate();
0463 make_rvar_from_expr(expr);
0464 }
0465
0466 template<typename T,
0467 typename = std::enable_if_t<is_floating_point_v<T> && !is_same_v<T, RealType>>>
0468 rvar(T v)
0469 : value_(inner_t{static_cast<RealType>(v)})
0470 {
0471 make_leaf_node();
0472 }
0473
0474 template<class E>
0475 rvar &operator=(const expression<RealType, DerivativeOrder, E> &expr)
0476 {
0477 value_ = expr.evaluate();
0478 make_rvar_from_expr(expr);
0479 return *this;
0480 }
0481
0482 template<class E>
0483 rvar<RealType, DerivativeOrder> &operator+=(const expression<RealType, DerivativeOrder, E> &expr)
0484 {
0485 *this = *this + expr;
0486 return *this;
0487 }
0488
0489 template<class E>
0490 rvar<RealType, DerivativeOrder> &operator*=(const expression<RealType, DerivativeOrder, E> &expr)
0491 {
0492 *this = *this * expr;
0493 return *this;
0494 }
0495
0496 template<class E>
0497 rvar<RealType, DerivativeOrder> &operator-=(const expression<RealType, DerivativeOrder, E> &expr)
0498 {
0499 *this = *this - expr;
0500 return *this;
0501 }
0502
0503 template<class E>
0504 rvar<RealType, DerivativeOrder> &operator/=(const expression<RealType, DerivativeOrder, E> &expr)
0505 {
0506 *this = *this / expr;
0507 return *this;
0508 }
0509
0510 rvar<RealType, DerivativeOrder> &operator+=(const RealType &v)
0511 {
0512 *this = *this + v;
0513 return *this;
0514 }
0515
0516 rvar<RealType, DerivativeOrder> &operator*=(const RealType &v)
0517 {
0518 *this = *this * v;
0519 return *this;
0520 }
0521
0522 rvar<RealType, DerivativeOrder> &operator-=(const RealType &v)
0523 {
0524 *this = *this - v;
0525 return *this;
0526 }
0527
0528 rvar<RealType, DerivativeOrder> &operator/=(const RealType &v)
0529 {
0530 *this = *this / v;
0531 return *this;
0532 }
0533
0534
0535 const inner_t &adjoint() const { return *node_->get_adjoint_ptr(); }
0536 inner_t &adjoint() { return *node_->get_adjoint_ptr(); }
0537
0538 const inner_t &evaluate() const { return value_; };
0539 inner_t &get_value() { return value_; };
0540
0541 explicit operator RealType() const { return item(); }
0542
0543 explicit operator int() const { return static_cast<int>(item()); }
0544 explicit operator long() const { return static_cast<long>(item()); }
0545 explicit operator long long() const { return static_cast<long long>(item()); }
0546
0547
0548
0549
0550
0551 template<size_t N>
0552 auto &get_value_at()
0553 {
0554 static_assert(N <= DerivativeOrder, "Requested depth exceeds variable order.");
0555 return get_value_at_impl<N, DerivativeOrder>::get(*this);
0556 }
0557
0558
0559 template<size_t N>
0560 const auto &get_value_at() const
0561 {
0562 static_assert(N <= DerivativeOrder, "Requested depth exceeds variable order.");
0563 return get_value_at_impl<N, DerivativeOrder>::get(*this);
0564 }
0565
0566 RealType item() const
0567 {
0568 return get_item_impl(std::integral_constant<bool, (DerivativeOrder > 1)>{});
0569 }
0570
0571 void backward()
0572 {
0573 gradient_tape<RealType, DerivativeOrder, BOOST_MATH_BUFFER_SIZE> &tape
0574 = get_active_tape<RealType, DerivativeOrder>();
0575 auto it = tape.find(node_);
0576 it->update_adjoint_v(inner_t(static_cast<RealType>(1.0)));
0577 while (it != tape.begin()) {
0578 it->backward();
0579 --it;
0580 }
0581 it->backward();
0582 }
0583 };
0584
0585 template<typename RealType, size_t DerivativeOrder>
0586 std::ostream &operator<<(std::ostream &os, const rvar<RealType, DerivativeOrder> var)
0587 {
0588 os << "rvar<" << DerivativeOrder << ">(" << var.item() << "," << var.adjoint() << ")";
0589 return os;
0590 }
0591
0592 template<typename RealType, size_t DerivativeOrder, typename E>
0593 std::ostream &operator<<(std::ostream &os, const expression<RealType, DerivativeOrder, E> &expr)
0594 {
0595 rvar<RealType, DerivativeOrder> tmp = expr;
0596 os << "rvar<" << DerivativeOrder << ">(" << tmp.item() << "," << tmp.adjoint() << ")";
0597 return os;
0598 }
0599
0600 template<typename RealType, size_t DerivativeOrder>
0601 rvar<RealType, DerivativeOrder> make_rvar(const RealType v)
0602 {
0603 static_assert(DerivativeOrder > 0, "rvar order must be >= 1");
0604 return rvar<RealType, DerivativeOrder>(v);
0605 }
0606 template<typename RealType, size_t DerivativeOrder, typename E>
0607 rvar<RealType, DerivativeOrder> make_rvar(const expression<RealType, DerivativeOrder, E> &expr)
0608 {
0609 static_assert(DerivativeOrder > 0, "rvar order must be >= 1");
0610 return rvar<RealType, DerivativeOrder>(expr);
0611 }
0612
0613 namespace detail {
0614
0615
0616
0617
0618
0619
0620 template<typename RealType, size_t DerivativeOrder>
0621 struct grad_op_impl
0622 {
0623 std::vector<rvar<RealType, DerivativeOrder - 1>> operator()(
0624 rvar<RealType, DerivativeOrder> &f, std::vector<rvar<RealType, DerivativeOrder> *> &x)
0625 {
0626 auto &tape = get_active_tape<RealType, DerivativeOrder>();
0627 tape.zero_grad();
0628 f.backward();
0629
0630 std::vector<rvar<RealType, DerivativeOrder - 1>> gradient_vector;
0631 gradient_vector.reserve(x.size());
0632
0633 for (auto &xi : x) {
0634 gradient_vector.emplace_back(xi->adjoint());
0635 }
0636 return gradient_vector;
0637 }
0638 };
0639
0640
0641
0642
0643 template<typename T>
0644 struct grad_op_impl<T, 1>
0645 {
0646 std::vector<T> operator()(rvar<T, 1> &f, std::vector<rvar<T, 1> *> &x)
0647 {
0648 gradient_tape<T, 1, BOOST_MATH_BUFFER_SIZE> &tape = get_active_tape<T, 1>();
0649 tape.zero_grad();
0650 f.backward();
0651 std::vector<T> gradient_vector;
0652 gradient_vector.reserve(x.size());
0653 for (auto &xi : x) {
0654 gradient_vector.push_back(xi->adjoint());
0655 }
0656 return gradient_vector;
0657 }
0658 };
0659
0660
0661
0662
0663
0664 template<size_t N,
0665 typename RealType,
0666 size_t DerivativeOrder_1,
0667 size_t DerivativeOrder_2,
0668 typename Enable = void>
0669 struct grad_nd_impl
0670 {
0671 auto operator()(rvar<RealType, DerivativeOrder_1> &f,
0672 std::vector<rvar<RealType, DerivativeOrder_2> *> &x)
0673 {
0674 static_assert(N > 1, "N must be greater than 1 for this template");
0675
0676 auto current_grad = grad(f, x);
0677
0678 std::vector<decltype(grad_nd_impl<N - 1, RealType, DerivativeOrder_1 - 1, DerivativeOrder_2>()(
0679 current_grad[0], x))>
0680 result;
0681 result.reserve(current_grad.size());
0682
0683 for (auto &g_i : current_grad) {
0684 result.push_back(
0685 grad_nd_impl<N - 1, RealType, DerivativeOrder_1 - 1, DerivativeOrder_2>()(g_i, x));
0686 }
0687 return result;
0688 }
0689 };
0690
0691
0692 template<typename RealType, size_t DerivativeOrder_1, size_t DerivativeOrder_2>
0693 struct grad_nd_impl<1, RealType, DerivativeOrder_1, DerivativeOrder_2>
0694 {
0695 auto operator()(rvar<RealType, DerivativeOrder_1> &f,
0696 std::vector<rvar<RealType, DerivativeOrder_2> *> &x)
0697 {
0698 return grad(f, x);
0699 }
0700 };
0701
0702 template<typename ptr>
0703 struct rvar_order;
0704
0705 template<typename RealType, size_t DerivativeOrder>
0706 struct rvar_order<rvar<RealType, DerivativeOrder> *>
0707 {
0708 static constexpr size_t value = DerivativeOrder;
0709 };
0710
0711 }
0712
0713
0714
0715
0716
0717
0718
0719
0720
0721
0722
0723 template<typename RealType, size_t DerivativeOrder_1, size_t DerivativeOrder_2>
0724 auto grad(rvar<RealType, DerivativeOrder_1> &f, std::vector<rvar<RealType, DerivativeOrder_2> *> &x)
0725 {
0726 static_assert(DerivativeOrder_1 <= DerivativeOrder_2,
0727 "variable differentiating w.r.t. must have order >= function order");
0728 std::vector<rvar<RealType, DerivativeOrder_1> *> xx;
0729 xx.reserve(x.size());
0730 for (auto &xi : x)
0731 xx.push_back(&(xi->template get_value_at<DerivativeOrder_1>()));
0732 return detail::grad_op_impl<RealType, DerivativeOrder_1>{}(f, xx);
0733 }
0734
0735
0736 template<typename RealType, size_t DerivativeOrder_1, typename First, typename... Other>
0737 auto grad(rvar<RealType, DerivativeOrder_1> &f, First first, Other... other)
0738 {
0739 constexpr size_t DerivativeOrder_2 = detail::rvar_order<First>::value;
0740 static_assert(DerivativeOrder_1 <= DerivativeOrder_2,
0741 "variable differentiating w.r.t. must have order >= function order");
0742 std::vector<rvar<RealType, DerivativeOrder_2> *> x_vec = {first, other...};
0743 return grad(f, x_vec);
0744 }
0745
0746
0747
0748
0749
0750
0751
0752
0753 template<typename RealType, size_t DerivativeOrder_1, size_t DerivativeOrder_2>
0754 auto hess(rvar<RealType, DerivativeOrder_1> &f, std::vector<rvar<RealType, DerivativeOrder_2> *> &x)
0755 {
0756 return detail::grad_nd_impl<2, RealType, DerivativeOrder_1, DerivativeOrder_2>{}(f, x);
0757 }
0758
0759
0760 template<typename RealType, size_t DerivativeOrder_1, typename First, typename... Other>
0761 auto hess(rvar<RealType, DerivativeOrder_1> &f, First first, Other... other)
0762 {
0763 constexpr size_t DerivativeOrder_2 = detail::rvar_order<First>::value;
0764 std::vector<rvar<RealType, DerivativeOrder_2> *> x_vec = {first, other...};
0765 return hess(f, x_vec);
0766 }
0767
0768
0769
0770
0771
0772
0773
0774 template<size_t N, typename RealType, size_t DerivativeOrder_1, size_t DerivativeOrder_2>
0775 auto grad_nd(rvar<RealType, DerivativeOrder_1> &f,
0776 std::vector<rvar<RealType, DerivativeOrder_2> *> &x)
0777 {
0778 static_assert(DerivativeOrder_1 >= N, "Function order must be at least N");
0779 static_assert(DerivativeOrder_2 >= DerivativeOrder_1,
0780 "Variable order must be at least function order");
0781
0782 return detail::grad_nd_impl<N, RealType, DerivativeOrder_1, DerivativeOrder_2>()(f, x);
0783 }
0784
0785
0786
0787 template<size_t N, typename ftype, typename First, typename... Other>
0788 auto grad_nd(ftype &f, First first, Other... other)
0789 {
0790 using RealType = typename ftype::value_type;
0791 constexpr size_t DerivativeOrder_1 = detail::rvar_order<ftype *>::value;
0792 constexpr size_t DerivativeOrder_2 = detail::rvar_order<First>::value;
0793 std::vector<rvar<RealType, DerivativeOrder_2> *> x_vec = {first, other...};
0794 return detail::grad_nd_impl<N, RealType, DerivativeOrder_1, DerivativeOrder_1>{}(f, x_vec);
0795 }
0796 }
0797 }
0798 }
0799 }
0800 namespace std {
0801
0802
0803 template<typename RealType, size_t DerivativeOrder>
0804 class numeric_limits<boost::math::differentiation::reverse_mode::rvar<RealType, DerivativeOrder>>
0805 : public numeric_limits<typename boost::math::differentiation::reverse_mode::
0806 rvar<RealType, DerivativeOrder>::value_type>
0807 {};
0808 }
0809 #endif