// 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_AUTODOFF_BASIC_OPERATOR_OVERLOADS_HPP #define REVERSE_MODE_AUTODOFF_BASIC_OPERATOR_OVERLOADS_HPP #include namespace boost { namespace math { namespace differentiation { namespace reverse_mode { /****************************************************************************************************************/ template struct add_expr : public abstract_binary_expression> { /* @brief addition * rvar+rvar * */ using inner_t = rvar_t; // Explicitly define constructor to forward to base class explicit add_expr(const expression &left_hand_expr, const expression &right_hand_expr) : abstract_binary_expression>(left_hand_expr, right_hand_expr) {} inner_t evaluate() const { return this->lhs.evaluate() + this->rhs.evaluate(); } static const inner_t left_derivative(const inner_t & /*l*/, const inner_t & /*r*/, const inner_t & /*v*/) { return inner_t(static_cast(1.0)); } static const inner_t right_derivative(const inner_t & /*l*/, const inner_t & /*r*/, const inner_t & /*v*/) { return inner_t(static_cast(1.0)); } }; template struct add_const_expr : public abstract_unary_expression> { /* @brief * rvar+float or float+rvar * */ using inner_t = rvar_t; explicit add_const_expr(const expression &arg_expr, const RealType v) : abstract_unary_expression>(arg_expr, v){}; inner_t evaluate() const { return this->arg.evaluate() + inner_t(this->constant); } static const inner_t derivative(const inner_t & /*argv*/, const inner_t & /*v*/, const RealType & /*constant*/) { return inner_t(static_cast(1.0)); } }; /****************************************************************************************************************/ template struct mult_expr : public abstract_binary_expression> { /* @brief multiplication * rvar * rvar * */ using inner_t = rvar_t; explicit mult_expr(const expression &left_hand_expr, const expression &right_hand_expr) : abstract_binary_expression>(left_hand_expr, right_hand_expr) {} inner_t evaluate() const { return this->lhs.evaluate() * this->rhs.evaluate(); }; static const inner_t left_derivative(const inner_t & /*l*/, const inner_t &r, const inner_t & /*v*/) noexcept { return r; }; static const inner_t right_derivative(const inner_t &l, const inner_t & /*r*/, const inner_t & /*v*/) noexcept { return l; }; }; template struct mult_const_expr : public abstract_unary_expression> { /* @brief * rvar+float or float+rvar * */ using inner_t = rvar_t; explicit mult_const_expr(const expression &arg_expr, const RealType v) : abstract_unary_expression>(arg_expr, v){}; inner_t evaluate() const { return this->arg.evaluate() * inner_t(this->constant); } static const inner_t derivative(const inner_t & /*argv*/, const inner_t & /*v*/, const RealType &constant) { return inner_t(constant); } }; /****************************************************************************************************************/ template struct sub_expr : public abstract_binary_expression> { /* @brief addition * rvar-rvar * */ using inner_t = rvar_t; // Explicitly define constructor to forward to base class explicit sub_expr(const expression &left_hand_expr, const expression &right_hand_expr) : abstract_binary_expression>(left_hand_expr, right_hand_expr) {} inner_t evaluate() const { return this->lhs.evaluate() - this->rhs.evaluate(); } static const inner_t left_derivative(const inner_t & /*l*/, const inner_t & /*r*/, const inner_t & /*v*/) { return inner_t(static_cast(1.0)); } static const inner_t right_derivative(const inner_t & /*l*/, const inner_t & /*r*/, const inner_t & /*v*/) { return inner_t(static_cast(-1.0)); } }; /****************************************************************************************************************/ template struct div_expr : public abstract_binary_expression> { /* @brief multiplication * rvar / rvar * */ using inner_t = rvar_t; // Explicitly define constructor to forward to base class explicit div_expr(const expression &left_hand_expr, const expression &right_hand_expr) : abstract_binary_expression>(left_hand_expr, right_hand_expr) {} inner_t evaluate() const { return this->lhs.evaluate() / this->rhs.evaluate(); }; static const inner_t left_derivative(const inner_t & /*l*/, const inner_t &r, const inner_t & /*v*/) { return static_cast(1.0) / r; }; static const inner_t right_derivative(const inner_t &l, const inner_t &r, const inner_t & /*v*/) { return -l / (r * r); }; }; template struct div_by_const_expr : public abstract_unary_expression> { /* @brief * rvar/float * */ using inner_t = rvar_t; explicit div_by_const_expr(const expression &arg_expr, const RealType v) : abstract_unary_expression>(arg_expr, v){}; inner_t evaluate() const { return this->arg.evaluate() / inner_t(this->constant); } static const inner_t derivative(const inner_t & /*argv*/, const inner_t & /*v*/, const RealType &constant) { return inner_t(1.0 / constant); } }; template struct const_div_by_expr : public abstract_unary_expression> { /** @brief * float/rvar * */ using inner_t = rvar_t; explicit const_div_by_expr(const expression &arg_expr, const RealType v) : abstract_unary_expression>(arg_expr, v){}; inner_t evaluate() const { return inner_t(this->constant) / this->arg.evaluate(); } static const inner_t derivative(const inner_t &argv, const inner_t & /*v*/, const RealType &constant) { return -inner_t{constant} / (argv * argv); } }; /****************************************************************************************************************/ } // namespace reverse_mode } // namespace differentiation } // namespace math } // namespace boost #endif