15#include <gch/small_vector.hpp>
17#include "sleipnir/autodiff/expression_type.hpp"
18#include "sleipnir/util/intrusive_shared_ptr.hpp"
19#include "sleipnir/util/pool.hpp"
21namespace slp::detail {
26inline constexpr bool USE_POOL_ALLOCATOR =
false;
28inline constexpr bool USE_POOL_ALLOCATOR =
true;
31template <
typename Scalar>
34template <
typename Scalar>
35constexpr void inc_ref_count(Expression<Scalar>* expr);
36template <
typename Scalar>
37constexpr void dec_ref_count(Expression<Scalar>* expr);
42template <
typename Scalar>
43using ExpressionPtr = IntrusiveSharedPtr<Expression<Scalar>>;
50template <
typename T,
typename... Args>
51static ExpressionPtr<typename T::Scalar> make_expression_ptr(Args&&... args) {
52 if constexpr (USE_POOL_ALLOCATOR) {
53 return allocate_intrusive_shared<T>(global_pool_allocator<T>(),
54 std::forward<Args>(args)...);
56 return make_intrusive_shared<T>(std::forward<Args>(args)...);
60template <
typename Scalar, ExpressionType T>
61struct BinaryMinusExpression;
63template <
typename Scalar, ExpressionType T>
64struct BinaryPlusExpression;
66template <
typename Scalar>
67struct ConstantExpression;
69template <
typename Scalar, ExpressionType T>
72template <
typename Scalar, ExpressionType T>
75template <
typename Scalar, ExpressionType T>
76struct UnaryMinusExpression;
82template <
typename Scalar>
83ExpressionPtr<Scalar> constant_ptr(Scalar value);
88template <
typename Scalar_>
104 std::array<ExpressionPtr<Scalar>, 2>
args{
nullptr,
nullptr};
148 return type() == ExpressionType::CONSTANT &&
val == constant;
157 using enum ExpressionType;
163 }
else if (
rhs->is_constant(
Scalar(0))) {
166 }
else if (
lhs->is_constant(
Scalar(1))) {
169 }
else if (
rhs->is_constant(
Scalar(1))) {
175 if (
lhs->type() == CONSTANT &&
rhs->type() == CONSTANT) {
176 return constant_ptr(
lhs->val *
rhs->val);
180 if (
lhs->type() == CONSTANT) {
181 if (
rhs->type() == LINEAR) {
183 }
else if (
rhs->type() == QUADRATIC) {
188 }
else if (
rhs->type() == CONSTANT) {
189 if (
lhs->type() == LINEAR) {
191 }
else if (
lhs->type() == QUADRATIC) {
196 }
else if (
lhs->type() == LINEAR &&
rhs->type() == LINEAR) {
209 using enum ExpressionType;
215 }
else if (
rhs->is_constant(
Scalar(1))) {
221 if (
lhs->type() == CONSTANT &&
rhs->type() == CONSTANT) {
222 return constant_ptr(
lhs->val /
rhs->val);
226 if (
rhs->type() == CONSTANT) {
227 if (
lhs->type() == LINEAR) {
229 }
else if (
lhs->type() == QUADRATIC) {
245 using enum ExpressionType;
252 }
else if (
rhs ==
nullptr ||
rhs->is_constant(
Scalar(0))) {
258 if (
lhs->type() == CONSTANT &&
rhs->type() == CONSTANT) {
259 return constant_ptr(
lhs->val +
rhs->val);
262 auto type = std::max(
lhs->type(),
rhs->type());
263 if (
type == LINEAR) {
266 }
else if (
type == QUADRATIC) {
290 using enum ExpressionType;
301 }
else if (
rhs->is_constant(
Scalar(0))) {
307 if (
lhs->type() == CONSTANT &&
rhs->type() == CONSTANT) {
308 return constant_ptr(
lhs->val -
rhs->val);
311 auto type = std::max(
lhs->type(),
rhs->type());
312 if (
type == LINEAR) {
315 }
else if (
type == QUADRATIC) {
328 using enum ExpressionType;
337 if (
lhs->type() == CONSTANT) {
338 return constant_ptr(-
lhs->val);
341 if (
lhs->type() == LINEAR) {
343 }
else if (
lhs->type() == QUADRATIC) {
371 virtual ExpressionType
type()
const = 0;
376 virtual std::string_view
name()
const = 0;
406 return constant_ptr(
Scalar(0));
417 return constant_ptr(
Scalar(0));
421template <
typename Scalar>
426template <
typename Scalar>
427ExpressionPtr<Scalar> cbrt(
const ExpressionPtr<Scalar>& x);
428template <
typename Scalar>
429ExpressionPtr<Scalar> exp(
const ExpressionPtr<Scalar>& x);
430template <
typename Scalar>
431ExpressionPtr<Scalar> sign(
const ExpressionPtr<Scalar>& x);
432template <
typename Scalar>
433ExpressionPtr<Scalar> sin(
const ExpressionPtr<Scalar>& x);
434template <
typename Scalar>
435ExpressionPtr<Scalar> sinh(
const ExpressionPtr<Scalar>& x);
436template <
typename Scalar>
437ExpressionPtr<Scalar> sqrt(
const ExpressionPtr<Scalar>& x);
443template <
typename Scalar, ExpressionType T>
455 ExpressionType
type()
const override {
return T; }
457 std::string_view
name()
const override {
return "binary minus"; }
480template <
typename Scalar, ExpressionType T>
492 ExpressionType
type()
const override {
return T; }
494 std::string_view
name()
const override {
return "binary plus"; }
516template <
typename Scalar>
529 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
531 std::string_view
name()
const override {
return "cbrt"; }
552template <
typename Scalar>
554 using enum ExpressionType;
558 if (x->type() == CONSTANT) {
559 if (x->val == Scalar(0)) {
562 }
else if (x->val == Scalar(-1) || x->val == Scalar(1)) {
565 return constant_ptr(cbrt(x->val));
569 return make_expression_ptr<CbrtExpression<Scalar>>(x);
575template <
typename Scalar>
585 ExpressionType
type()
const override {
return ExpressionType::CONSTANT; }
587 std::string_view
name()
const override {
return "constant"; }
593template <
typename Scalar>
606 ExpressionType
type()
const override {
return ExpressionType::LINEAR; }
608 std::string_view
name()
const override {
return "decision variable"; }
615template <
typename Scalar, ExpressionType T>
626 ExpressionType
type()
const override {
return T; }
628 std::string_view
name()
const override {
return "division"; }
655template <
typename Scalar, ExpressionType T>
666 ExpressionType
type()
const override {
return T; }
668 std::string_view
name()
const override {
return "multiplication"; }
695template <
typename Scalar, ExpressionType T>
705 ExpressionType
type()
const override {
return T; }
707 std::string_view
name()
const override {
return "unary minus"; }
722template <
typename Scalar>
731template <
typename Scalar>
732constexpr void dec_ref_count(Expression<Scalar>* expr) {
737 gch::small_vector<Expression<Scalar>*> stack;
738 stack.emplace_back(expr);
740 while (!stack.empty()) {
741 auto elem = stack.back();
746 if (--elem->ref_count == 0) {
747 if (elem->adjoint_expr !=
nullptr) {
748 stack.emplace_back(elem->adjoint_expr.get());
750 for (
auto& arg : elem->args) {
751 if (arg !=
nullptr) {
752 stack.emplace_back(arg.get());
758 if constexpr (USE_POOL_ALLOCATOR) {
759 auto alloc = global_pool_allocator<Expression<Scalar>>();
760 std::allocator_traits<
decltype(alloc)>::deallocate(
761 alloc, elem,
sizeof(Expression<Scalar>));
763 operator delete(elem);
772template <
typename Scalar>
785 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
787 std::string_view
name()
const override {
return "abs"; }
792 }
else if (x >
Scalar(0)) {
810template <
typename Scalar>
812 using enum ExpressionType;
816 if (x->is_constant(Scalar(0))) {
822 if (x->type() == CONSTANT) {
823 return constant_ptr(abs(x->val));
826 return make_expression_ptr<AbsExpression<Scalar>>(x);
832template <
typename Scalar>
845 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
847 std::string_view
name()
const override {
return "acos"; }
865template <
typename Scalar>
867 using enum ExpressionType;
871 if (x->is_constant(Scalar(0))) {
872 return constant_ptr(Scalar(std::numbers::pi) / Scalar(2));
876 if (x->type() == CONSTANT) {
877 return constant_ptr(acos(x->val));
880 return make_expression_ptr<AcosExpression<Scalar>>(x);
886template <
typename Scalar>
899 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
901 std::string_view
name()
const override {
return "asin"; }
919template <
typename Scalar>
921 using enum ExpressionType;
925 if (x->is_constant(Scalar(0))) {
931 if (x->type() == CONSTANT) {
932 return constant_ptr(asin(x->val));
935 return make_expression_ptr<AsinExpression<Scalar>>(x);
941template <
typename Scalar>
954 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
956 std::string_view
name()
const override {
return "atan"; }
973template <
typename Scalar>
975 using enum ExpressionType;
979 if (x->is_constant(Scalar(0))) {
985 if (x->type() == CONSTANT) {
986 return constant_ptr(atan(x->val));
989 return make_expression_ptr<AtanExpression<Scalar>>(x);
995template <
typename Scalar>
1010 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1012 std::string_view
name()
const override {
return "atan2"; }
1015 return this->
adjoint * x / (y * y + x * x);
1019 return this->
adjoint * -y / (y * y + x * x);
1040template <
typename Scalar>
1043 using enum ExpressionType;
1047 if (y->type() == CONSTANT && x->type() == CONSTANT) {
1048 return constant_ptr(atan2(y->val, x->val));
1051 return make_expression_ptr<Atan2Expression<Scalar>>(y, x);
1057template <
typename Scalar>
1070 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1072 std::string_view
name()
const override {
return "cos"; }
1076 return this->
adjoint * -sin(x);
1090template <
typename Scalar>
1092 using enum ExpressionType;
1096 if (x->is_constant(Scalar(0))) {
1097 return constant_ptr(Scalar(1));
1101 if (x->type() == CONSTANT) {
1102 return constant_ptr(cos(x->val));
1105 return make_expression_ptr<CosExpression<Scalar>>(x);
1111template <
typename Scalar>
1124 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1126 std::string_view
name()
const override {
return "cosh"; }
1130 return this->
adjoint * sinh(x);
1144template <
typename Scalar>
1146 using enum ExpressionType;
1150 if (x->is_constant(Scalar(0))) {
1151 return constant_ptr(Scalar(1));
1155 if (x->type() == CONSTANT) {
1156 return constant_ptr(cosh(x->val));
1159 return make_expression_ptr<CoshExpression<Scalar>>(x);
1165template <
typename Scalar>
1178 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1180 std::string_view
name()
const override {
return "erf"; }
1184 return this->
adjoint *
Scalar(2.0 * std::numbers::inv_sqrtpi) * exp(-x * x);
1191 constant_ptr(
Scalar(2.0 * std::numbers::inv_sqrtpi)) * exp(-x * x);
1199template <
typename Scalar>
1201 using enum ExpressionType;
1205 if (x->is_constant(Scalar(0))) {
1211 if (x->type() == CONSTANT) {
1212 return constant_ptr(erf(x->val));
1215 return make_expression_ptr<ErfExpression<Scalar>>(x);
1221template <
typename Scalar>
1234 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1236 std::string_view
name()
const override {
return "exp"; }
1240 return this->
adjoint * exp(x);
1254template <
typename Scalar>
1256 using enum ExpressionType;
1260 if (x->is_constant(Scalar(0))) {
1261 return constant_ptr(Scalar(1));
1265 if (x->type() == CONSTANT) {
1266 return constant_ptr(exp(x->val));
1269 return make_expression_ptr<ExpExpression<Scalar>>(x);
1272template <
typename Scalar>
1273ExpressionPtr<Scalar> hypot(
const ExpressionPtr<Scalar>& x,
1274 const ExpressionPtr<Scalar>& y);
1279template <
typename Scalar>
1294 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1296 std::string_view
name()
const override {
return "hypot"; }
1300 return this->
adjoint * x / hypot(x, y);
1305 return this->
adjoint * y / hypot(x, y);
1326template <
typename Scalar>
1329 using enum ExpressionType;
1333 if (x->is_constant(Scalar(0))) {
1335 }
else if (y->is_constant(Scalar(0))) {
1340 if (x->type() == CONSTANT && y->type() == CONSTANT) {
1341 return constant_ptr(hypot(x->val, y->val));
1344 return make_expression_ptr<HypotExpression<Scalar>>(x, y);
1350template <
typename Scalar>
1362 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1364 std::string_view
name()
const override {
return "is nonnegative"; }
1371template <
typename Scalar>
1373 if (x->type() == ExpressionType::CONSTANT) {
1374 return constant_ptr(x->val >= Scalar(0) ? Scalar(1) : Scalar(0));
1377 return make_expression_ptr<IsNonnegativeExpression<Scalar>>(x);
1383template <
typename Scalar>
1395 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1397 std::string_view
name()
const override {
return "is positive"; }
1404template <
typename Scalar>
1406 if (x->type() == ExpressionType::CONSTANT) {
1407 return constant_ptr(x->val > Scalar(0) ? Scalar(1) : Scalar(0));
1410 return make_expression_ptr<IsPositiveExpression<Scalar>>(x);
1416template <
typename Scalar>
1429 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1431 std::string_view
name()
const override {
return "log"; }
1446template <
typename Scalar>
1448 using enum ExpressionType;
1452 if (x->is_constant(Scalar(0))) {
1458 if (x->type() == CONSTANT) {
1459 return constant_ptr(log(x->val));
1462 return make_expression_ptr<LogExpression<Scalar>>(x);
1468template <
typename Scalar>
1481 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1483 std::string_view
name()
const override {
return "log10"; }
1500template <
typename Scalar>
1502 using enum ExpressionType;
1506 if (x->is_constant(Scalar(0))) {
1512 if (x->type() == CONSTANT) {
1513 return constant_ptr(log10(x->val));
1516 return make_expression_ptr<Log10Expression<Scalar>>(x);
1524template <
typename Scalar>
1538 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1540 std::string_view
name()
const override {
return "max"; }
1574template <
typename Scalar>
1577 using enum ExpressionType;
1581 if (
a->type() == CONSTANT &&
b->type() == CONSTANT) {
1582 return constant_ptr(max(
a->val,
b->val));
1585 return make_expression_ptr<MaxExpression<Scalar>>(a, b);
1593template <
typename Scalar>
1607 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1609 std::string_view
name()
const override {
return "min"; }
1644template <
typename Scalar>
1647 using enum ExpressionType;
1651 if (
a->type() == CONSTANT &&
b->type() == CONSTANT) {
1652 return constant_ptr(min(
a->val,
b->val));
1655 return make_expression_ptr<MinExpression<Scalar>>(a, b);
1658template <
typename Scalar>
1659ExpressionPtr<Scalar> pow(
const ExpressionPtr<Scalar>& base,
1660 const ExpressionPtr<Scalar>& power);
1666template <
typename Scalar, ExpressionType T>
1680 ExpressionType
type()
const override {
return T; }
1682 std::string_view
name()
const override {
return "pow"; }
1715template <
typename Scalar>
1718 using enum ExpressionType;
1722 if (
base->is_constant(Scalar(0))) {
1725 }
else if (base->is_constant(Scalar(1))) {
1729 if (power->is_constant(Scalar(0))) {
1730 return constant_ptr(Scalar(1));
1731 }
else if (power->is_constant(Scalar(1))) {
1737 if (base->type() == CONSTANT && power->type() == CONSTANT) {
1738 return constant_ptr(pow(base->val, power->val));
1741 if (power->is_constant(Scalar(2))) {
1742 if (base->type() == LINEAR) {
1743 return make_expression_ptr<MultExpression<Scalar, QUADRATIC>>(base, base);
1745 return make_expression_ptr<MultExpression<Scalar, NONLINEAR>>(base, base);
1749 return make_expression_ptr<PowExpression<Scalar, NONLINEAR>>(base, power);
1755template <
typename Scalar>
1766 }
else if (x ==
Scalar(0)) {
1773 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1775 std::string_view
name()
const override {
return "sign"; }
1782template <
typename Scalar>
1784 using enum ExpressionType;
1787 if (x->type() == CONSTANT) {
1788 if (x->val < Scalar(0)) {
1789 return constant_ptr(Scalar(-1));
1790 }
else if (x->val == Scalar(0)) {
1794 return constant_ptr(Scalar(1));
1798 return make_expression_ptr<SignExpression<Scalar>>(x);
1804template <
typename Scalar>
1817 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1819 std::string_view
name()
const override {
return "sin"; }
1823 return this->
adjoint * cos(x);
1837template <
typename Scalar>
1839 using enum ExpressionType;
1843 if (x->is_constant(Scalar(0))) {
1849 if (x->type() == CONSTANT) {
1850 return constant_ptr(sin(x->val));
1853 return make_expression_ptr<SinExpression<Scalar>>(x);
1859template <
typename Scalar>
1872 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1874 std::string_view
name()
const override {
return "sinh"; }
1878 return this->
adjoint * cosh(x);
1892template <
typename Scalar>
1894 using enum ExpressionType;
1898 if (x->is_constant(Scalar(0))) {
1904 if (x->type() == CONSTANT) {
1905 return constant_ptr(sinh(x->val));
1908 return make_expression_ptr<SinhExpression<Scalar>>(x);
1914template <
typename Scalar>
1927 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1929 std::string_view
name()
const override {
return "sqrt"; }
1947template <
typename Scalar>
1949 using enum ExpressionType;
1953 if (x->type() == CONSTANT) {
1954 if (x->val == Scalar(0)) {
1957 }
else if (x->val == Scalar(1)) {
1960 return constant_ptr(sqrt(x->val));
1964 return make_expression_ptr<SqrtExpression<Scalar>>(x);
1970template <
typename Scalar>
1983 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
1985 std::string_view
name()
const override {
return "tan"; }
2006template <
typename Scalar>
2008 using enum ExpressionType;
2012 if (x->is_constant(Scalar(0))) {
2018 if (x->type() == CONSTANT) {
2019 return constant_ptr(tan(x->val));
2022 return make_expression_ptr<TanExpression<Scalar>>(x);
2028template <
typename Scalar>
2041 ExpressionType
type()
const override {
return ExpressionType::NONLINEAR; }
2043 std::string_view
name()
const override {
return "tanh"; }
2064template <
typename Scalar>
2066 using enum ExpressionType;
2070 if (x->is_constant(Scalar(0))) {
2076 if (x->type() == CONSTANT) {
2077 return constant_ptr(tanh(x->val));
2080 return make_expression_ptr<TanhExpression<Scalar>>(x);
Definition intrusive_shared_ptr.hpp:27
Definition expression.hpp:773
std::string_view name() const override
Definition expression.hpp:787
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:799
constexpr AbsExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:777
ExpressionType type() const override
Definition expression.hpp:785
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:789
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:780
Definition expression.hpp:833
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:840
constexpr AcosExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:837
ExpressionType type() const override
Definition expression.hpp:845
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:849
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:854
std::string_view name() const override
Definition expression.hpp:847
Definition expression.hpp:887
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:894
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:903
std::string_view name() const override
Definition expression.hpp:901
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:908
ExpressionType type() const override
Definition expression.hpp:899
constexpr AsinExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:891
Definition expression.hpp:996
std::string_view name() const override
Definition expression.hpp:1012
Scalar value(Scalar y, Scalar x) const override
Definition expression.hpp:1005
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &y, const ExpressionPtr< Scalar > &x) const override
Definition expression.hpp:1022
constexpr Atan2Expression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1001
Scalar grad_l(Scalar y, Scalar x) const override
Definition expression.hpp:1014
Scalar grad_r(Scalar y, Scalar x) const override
Definition expression.hpp:1018
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &y, const ExpressionPtr< Scalar > &x) const override
Definition expression.hpp:1028
ExpressionType type() const override
Definition expression.hpp:1010
Definition expression.hpp:942
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:962
std::string_view name() const override
Definition expression.hpp:956
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:958
constexpr AtanExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:946
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:949
ExpressionType type() const override
Definition expression.hpp:954
Definition expression.hpp:444
constexpr BinaryMinusExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:449
std::string_view name() const override
Definition expression.hpp:457
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:469
Scalar grad_r(Scalar, Scalar) const override
Definition expression.hpp:461
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:463
Scalar value(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:453
ExpressionType type() const override
Definition expression.hpp:455
Scalar grad_l(Scalar, Scalar) const override
Definition expression.hpp:459
Definition expression.hpp:481
Scalar value(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:490
Scalar grad_r(Scalar, Scalar) const override
Definition expression.hpp:498
Scalar grad_l(Scalar, Scalar) const override
Definition expression.hpp:496
constexpr BinaryPlusExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:486
ExpressionType type() const override
Definition expression.hpp:492
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:506
std::string_view name() const override
Definition expression.hpp:494
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:500
Definition expression.hpp:517
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:540
std::string_view name() const override
Definition expression.hpp:531
constexpr CbrtExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:521
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:533
ExpressionType type() const override
Definition expression.hpp:529
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:524
Definition expression.hpp:576
ExpressionType type() const override
Definition expression.hpp:585
Scalar value(Scalar, Scalar) const override
Definition expression.hpp:583
std::string_view name() const override
Definition expression.hpp:587
constexpr ConstantExpression(Scalar value)
Definition expression.hpp:580
Definition expression.hpp:1058
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1065
ExpressionType type() const override
Definition expression.hpp:1070
constexpr CosExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1062
std::string_view name() const override
Definition expression.hpp:1072
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1074
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1079
Definition expression.hpp:1112
ExpressionType type() const override
Definition expression.hpp:1124
std::string_view name() const override
Definition expression.hpp:1126
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1119
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1133
constexpr CoshExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1116
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1128
Definition expression.hpp:594
constexpr DecisionVariableExpression()=default
Constructs a decision variable expression with a value of zero.
std::string_view name() const override
Definition expression.hpp:608
Scalar value(Scalar, Scalar) const override
Definition expression.hpp:604
constexpr DecisionVariableExpression(Scalar value)
Definition expression.hpp:601
ExpressionType type() const override
Definition expression.hpp:606
Definition expression.hpp:616
constexpr DivExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:621
ExpressionType type() const override
Definition expression.hpp:626
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:644
std::string_view name() const override
Definition expression.hpp:628
Scalar value(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:624
Scalar grad_l(Scalar, Scalar rhs) const override
Definition expression.hpp:630
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:638
Scalar grad_r(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:634
Definition expression.hpp:1166
std::string_view name() const override
Definition expression.hpp:1180
constexpr ErfExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1170
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1182
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1187
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1173
ExpressionType type() const override
Definition expression.hpp:1178
Definition expression.hpp:1222
std::string_view name() const override
Definition expression.hpp:1236
ExpressionType type() const override
Definition expression.hpp:1234
constexpr ExpExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1226
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1238
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1243
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1229
Definition expression.hpp:89
Scalar val
The value of the expression node.
Definition expression.hpp:94
friend ExpressionPtr< Scalar > operator-(const ExpressionPtr< Scalar > &lhs)
Definition expression.hpp:327
virtual Scalar grad_r(Scalar lhs, Scalar rhs) const
Definition expression.hpp:393
std::array< ExpressionPtr< Scalar >, 2 > args
Expression arguments.
Definition expression.hpp:104
uint32_t ref_count
Reference count for intrusive shared pointer.
Definition expression.hpp:118
constexpr Expression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:131
Scalar adjoint
The adjoint of the expression node, used during autodiff.
Definition expression.hpp:97
friend ExpressionPtr< Scalar > operator+(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs)
Definition expression.hpp:243
int32_t scratch
Definition expression.hpp:115
virtual Scalar grad_l(Scalar lhs, Scalar rhs) const
Definition expression.hpp:383
constexpr bool is_constant(Scalar constant) const
Definition expression.hpp:147
constexpr Expression()=default
Constructs a constant expression with a value of zero.
constexpr Expression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:138
virtual ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const
Definition expression.hpp:403
ExpressionPtr< Scalar > adjoint_expr
Definition expression.hpp:101
Scalar_ Scalar
Scalar type alias.
Definition expression.hpp:91
virtual ExpressionType type() const =0
friend ExpressionPtr< Scalar > operator/(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs)
Definition expression.hpp:207
friend ExpressionPtr< Scalar > operator-(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs)
Definition expression.hpp:288
friend ExpressionPtr< Scalar > operator*(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs)
Definition expression.hpp:155
virtual Scalar value(Scalar lhs, Scalar rhs) const =0
virtual std::string_view name() const =0
virtual ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const
Definition expression.hpp:414
friend ExpressionPtr< Scalar > operator+=(ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs)
Definition expression.hpp:279
constexpr Expression(Scalar value)
Definition expression.hpp:126
friend ExpressionPtr< Scalar > operator+(const ExpressionPtr< Scalar > &lhs)
Definition expression.hpp:353
Definition expression.hpp:1280
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &y) const override
Definition expression.hpp:1308
constexpr HypotExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1285
ExpressionType type() const override
Definition expression.hpp:1294
std::string_view name() const override
Definition expression.hpp:1296
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &y) const override
Definition expression.hpp:1314
Scalar value(Scalar x, Scalar y) const override
Definition expression.hpp:1289
Scalar grad_r(Scalar x, Scalar y) const override
Definition expression.hpp:1303
Scalar grad_l(Scalar x, Scalar y) const override
Definition expression.hpp:1298
Definition expression.hpp:1351
ExpressionType type() const override
Definition expression.hpp:1362
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1358
std::string_view name() const override
Definition expression.hpp:1364
constexpr IsNonnegativeExpression(ExpressionPtr< Scalar > x)
Definition expression.hpp:1355
Definition expression.hpp:1384
std::string_view name() const override
Definition expression.hpp:1397
ExpressionType type() const override
Definition expression.hpp:1395
constexpr IsPositiveExpression(ExpressionPtr< Scalar > x)
Definition expression.hpp:1388
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1391
Definition expression.hpp:1469
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1476
std::string_view name() const override
Definition expression.hpp:1483
constexpr Log10Expression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1473
ExpressionType type() const override
Definition expression.hpp:1481
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1489
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1485
Definition expression.hpp:1417
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1435
ExpressionType type() const override
Definition expression.hpp:1429
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1433
std::string_view name() const override
Definition expression.hpp:1431
constexpr LogExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1421
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1424
Definition expression.hpp:1525
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1550
Scalar value(Scalar a, Scalar b) const override
Definition expression.hpp:1533
Scalar grad_r(Scalar a, Scalar b) const override
Definition expression.hpp:1546
constexpr MaxExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1530
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1558
Scalar grad_l(Scalar a, Scalar b) const override
Definition expression.hpp:1542
ExpressionType type() const override
Definition expression.hpp:1538
std::string_view name() const override
Definition expression.hpp:1540
Definition expression.hpp:1594
std::string_view name() const override
Definition expression.hpp:1609
Scalar value(Scalar a, Scalar b) const override
Definition expression.hpp:1602
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1620
ExpressionType type() const override
Definition expression.hpp:1607
constexpr MinExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1599
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1629
Scalar grad_l(Scalar a, Scalar b) const override
Definition expression.hpp:1611
Scalar grad_r(Scalar a, Scalar b) const override
Definition expression.hpp:1615
Definition expression.hpp:656
Scalar grad_l(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:670
Scalar value(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:664
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:678
constexpr MultExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:661
ExpressionType type() const override
Definition expression.hpp:666
std::string_view name() const override
Definition expression.hpp:668
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:684
Scalar grad_r(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:674
Definition expression.hpp:1667
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &base, const ExpressionPtr< Scalar > &power) const override
Definition expression.hpp:1696
Scalar value(Scalar base, Scalar power) const override
Definition expression.hpp:1675
ExpressionType type() const override
Definition expression.hpp:1680
Scalar grad_l(Scalar base, Scalar power) const override
Definition expression.hpp:1684
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &base, const ExpressionPtr< Scalar > &power) const override
Definition expression.hpp:1703
Scalar grad_r(Scalar base, Scalar power) const override
Definition expression.hpp:1689
std::string_view name() const override
Definition expression.hpp:1682
constexpr PowExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1672
Definition expression.hpp:1756
constexpr SignExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1760
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1763
ExpressionType type() const override
Definition expression.hpp:1773
std::string_view name() const override
Definition expression.hpp:1775
Definition expression.hpp:1805
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1812
std::string_view name() const override
Definition expression.hpp:1819
ExpressionType type() const override
Definition expression.hpp:1817
constexpr SinExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1809
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1826
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1821
Definition expression.hpp:1860
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1876
std::string_view name() const override
Definition expression.hpp:1874
ExpressionType type() const override
Definition expression.hpp:1872
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1881
constexpr SinhExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1864
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1867
Definition expression.hpp:1915
ExpressionType type() const override
Definition expression.hpp:1927
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1931
std::string_view name() const override
Definition expression.hpp:1929
constexpr SqrtExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1919
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1922
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1936
Definition expression.hpp:1971
std::string_view name() const override
Definition expression.hpp:1985
constexpr TanExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1975
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1987
ExpressionType type() const override
Definition expression.hpp:1983
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1994
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1978
Definition expression.hpp:2029
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:2045
ExpressionType type() const override
Definition expression.hpp:2041
std::string_view name() const override
Definition expression.hpp:2043
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:2036
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:2052
constexpr TanhExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:2033
Definition expression.hpp:696
Scalar grad_l(Scalar, Scalar) const override
Definition expression.hpp:709
constexpr UnaryMinusExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:700
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:711
ExpressionType type() const override
Definition expression.hpp:705
Scalar value(Scalar lhs, Scalar) const override
Definition expression.hpp:703
std::string_view name() const override
Definition expression.hpp:707