// Copyright Maksym Zhelyeznyakov 2025-2026. // Distributed under the Boost Software License, Version 1.0. // (See accompanying file LICENSE_1_0.txt or copy at // https://www.boost.org/LICENSE_1_0.txt) #ifndef REVERSE_MODE_AUTODIFF_EXPRESSION_TEMPLATE_BASE_HPP #define REVERSE_MODE_AUTODIFF_EXPRESSION_TEMPLATE_BASE_HPP #include #include namespace boost { namespace math { namespace differentiation { namespace reverse_mode { /* forward declarations for utitlity functions */ struct expression_base {}; template struct expression; template class rvar; template struct abstract_binary_expression; template struct abstract_unary_expression; template class gradient_node; // forward declaration for tape namespace detail { template using void_t = void; // Check if T has a 'value_type' alias template struct has_value_type : std::false_type {}; template struct has_value_type> : std::true_type {}; template struct has_binary_sub_types : std::false_type {}; template struct has_binary_sub_types> : std::true_type {}; template struct has_unary_sub_type : std::false_type {}; template struct has_unary_sub_type> : std::true_type {}; template struct count_rvar_impl { static constexpr std::size_t value = 0; }; template struct count_rvar_impl, DerivativeOrder> { static constexpr std::size_t value = 1; }; template struct count_rvar_impl< RealType, DerivativeOrder, std::enable_if_t::value && !std::is_same>::value && !has_unary_sub_type::value>> { static constexpr std::size_t value = count_rvar_impl::value + count_rvar_impl::value; }; template struct count_rvar_impl< RealType, DerivativeOrder, typename std::enable_if_t< has_unary_sub_type::value && !std::is_same>::value && !has_binary_sub_types::value>> { static constexpr std::size_t value = count_rvar_impl::value; }; template constexpr std::size_t count_rvars = detail::count_rvar_impl::value; template struct is_expression : std::is_base_of::type> {}; template struct rvar_type_impl { using type = rvar; }; template struct rvar_type_impl { using type = RealType; }; } // namespace detail template using rvar_t = typename detail::rvar_type_impl::type; template struct expression : expression_base { /* @brief * base expression class * */ using value_type = RealType; static constexpr size_t order_v = DerivativeOrder; using derived_type = DerivedExpression; static constexpr size_t num_literals = 0; using inner_t = rvar_t; inner_t evaluate() const { return static_cast(this)->evaluate(); } template void propagatex(gradient_node *node, inner_t adj) const { return static_cast(this)->template propagatex(node, adj); } }; template struct abstract_binary_expression : public expression< RealType, DerivativeOrder, abstract_binary_expression> { using lhs_type = LHS; using rhs_type = RHS; using value_type = RealType; using inner_t = rvar_t; const lhs_type lhs; const rhs_type rhs; explicit abstract_binary_expression( const expression &left_hand_expr, const expression &right_hand_expr) : lhs(static_cast(left_hand_expr)) , rhs(static_cast(right_hand_expr)){}; inner_t evaluate() const { return static_cast(this)->evaluate(); }; template void propagatex(gradient_node *node, inner_t adj) const { const inner_t lv = lhs.evaluate(); const inner_t rv = rhs.evaluate(); const inner_t v = evaluate(); const inner_t partial_l = ConcreteBinaryOperation::left_derivative(lv, rv, v); const inner_t partial_r = ConcreteBinaryOperation::right_derivative(lv, rv, v); constexpr size_t num_lhs_args = detail::count_rvars; constexpr size_t num_rhs_args = detail::count_rvars; propagate_lhs(node, adj * partial_l); propagate_rhs(node, adj * partial_r); } private: /* everything here just emulates c++17 if constexpr */ template 0), int>::type = 0> void propagate_lhs(gradient_node *node, inner_t adj) const { lhs.template propagatex(node, adj); } template::type = 0> void propagate_lhs(gradient_node *, inner_t) const {} template 0), int>::type = 0> void propagate_rhs(gradient_node *node, inner_t adj) const { rhs.template propagatex(node, adj); } template::type = 0> void propagate_rhs(gradient_node *, inner_t) const {} }; template struct abstract_unary_expression : public expression< RealType, DerivativeOrder, abstract_unary_expression> { using arg_type = ARG; using value_type = RealType; using inner_t = rvar_t; const arg_type arg; const RealType constant; explicit abstract_unary_expression(const expression &arg_expr, const RealType &constant) : arg(static_cast(arg_expr)) , constant(constant){}; inner_t evaluate() const { return static_cast(this)->evaluate(); }; template void propagatex(gradient_node *node, inner_t adj) const { inner_t argv = arg.evaluate(); inner_t v = evaluate(); inner_t partial_arg = ConcreteUnaryOperation::derivative(argv, v, constant); arg.template propagatex(node, adj * partial_arg); } }; } // namespace reverse_mode } // namespace differentiation } // namespace math } // namespace boost #endif // REVERSE_MODE_AUTODIFF_EXPRESSION_TEMPLATE_BASE_HPP