Sleipnir C++ API
Loading...
Searching...
No Matches
expression.hpp
1// Copyright (c) Sleipnir contributors
2
3#pragma once
4
5#include <stdint.h>
6
7#include <algorithm>
8#include <array>
9#include <cmath>
10#include <memory>
11#include <numbers>
12#include <string_view>
13#include <utility>
14
15#include <gch/small_vector.hpp>
16
17#include "sleipnir/autodiff/expression_type.hpp"
18#include "sleipnir/util/intrusive_shared_ptr.hpp"
19#include "sleipnir/util/pool.hpp"
20
21namespace slp::detail {
22
23// The global pool allocator uses a thread-local static pool resource, which
24// isn't guaranteed to be initialized properly across DLL boundaries on Windows
25#ifdef _WIN32
26inline constexpr bool USE_POOL_ALLOCATOR = false;
27#else
28inline constexpr bool USE_POOL_ALLOCATOR = true;
29#endif
30
31template <typename Scalar>
32struct Expression;
33
34template <typename Scalar>
35constexpr void inc_ref_count(Expression<Scalar>* expr);
36template <typename Scalar>
37constexpr void dec_ref_count(Expression<Scalar>* expr);
38
42template <typename Scalar>
43using ExpressionPtr = IntrusiveSharedPtr<Expression<Scalar>>;
44
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)...);
55 } else {
56 return make_intrusive_shared<T>(std::forward<Args>(args)...);
57 }
58}
59
60template <typename Scalar, ExpressionType T>
61struct BinaryMinusExpression;
62
63template <typename Scalar, ExpressionType T>
64struct BinaryPlusExpression;
65
66template <typename Scalar>
67struct ConstantExpression;
68
69template <typename Scalar, ExpressionType T>
70struct DivExpression;
71
72template <typename Scalar, ExpressionType T>
73struct MultExpression;
74
75template <typename Scalar, ExpressionType T>
76struct UnaryMinusExpression;
77
82template <typename Scalar>
83ExpressionPtr<Scalar> constant_ptr(Scalar value);
84
88template <typename Scalar_>
89struct Expression {
91 using Scalar = Scalar_;
92
95
98
102
104 std::array<ExpressionPtr<Scalar>, 2> args{nullptr, nullptr};
105
116
119
121 constexpr Expression() = default;
122
126 explicit constexpr Expression(Scalar value) : val{value} {}
127
132 : args{std::move(lhs), nullptr} {}
133
140
141 virtual ~Expression() = default;
142
147 constexpr bool is_constant(Scalar constant) const {
148 return type() == ExpressionType::CONSTANT && val == constant;
149 }
150
156 const ExpressionPtr<Scalar>& rhs) {
157 using enum ExpressionType;
158
159 // Prune expression
160 if (lhs->is_constant(Scalar(0))) {
161 // Return zero, which lhs currently is
162 return lhs;
163 } else if (rhs->is_constant(Scalar(0))) {
164 // Return zero, which rhs currently is
165 return rhs;
166 } else if (lhs->is_constant(Scalar(1))) {
167 // Return rhs unmodified
168 return rhs;
169 } else if (rhs->is_constant(Scalar(1))) {
170 // Return lhs unmodified
171 return lhs;
172 }
173
174 // Evaluate constant
175 if (lhs->type() == CONSTANT && rhs->type() == CONSTANT) {
176 return constant_ptr(lhs->val * rhs->val);
177 }
178
179 // Evaluate expression type
180 if (lhs->type() == CONSTANT) {
181 if (rhs->type() == LINEAR) {
183 } else if (rhs->type() == QUADRATIC) {
185 } else {
187 }
188 } else if (rhs->type() == CONSTANT) {
189 if (lhs->type() == LINEAR) {
191 } else if (lhs->type() == QUADRATIC) {
193 } else {
195 }
196 } else if (lhs->type() == LINEAR && rhs->type() == LINEAR) {
198 } else {
200 }
201 }
202
208 const ExpressionPtr<Scalar>& rhs) {
209 using enum ExpressionType;
210
211 // Prune expression
212 if (lhs->is_constant(Scalar(0))) {
213 // Return zero, which lhs currently is
214 return lhs;
215 } else if (rhs->is_constant(Scalar(1))) {
216 // Return lhs unmodified
217 return lhs;
218 }
219
220 // Evaluate constant
221 if (lhs->type() == CONSTANT && rhs->type() == CONSTANT) {
222 return constant_ptr(lhs->val / rhs->val);
223 }
224
225 // Evaluate expression type
226 if (rhs->type() == CONSTANT) {
227 if (lhs->type() == LINEAR) {
229 } else if (lhs->type() == QUADRATIC) {
231 } else {
233 }
234 } else {
236 }
237 }
238
244 const ExpressionPtr<Scalar>& rhs) {
245 using enum ExpressionType;
246
247 // Prune expression. We check for nullptr because operator+ is used in
248 // adjoint accumulation, and child nodes can be null.
249 if (lhs == nullptr || lhs->is_constant(Scalar(0))) {
250 // Return rhs unmodified
251 return rhs;
252 } else if (rhs == nullptr || rhs->is_constant(Scalar(0))) {
253 // Return lhs unmodified
254 return lhs;
255 }
256
257 // Evaluate constant
258 if (lhs->type() == CONSTANT && rhs->type() == CONSTANT) {
259 return constant_ptr(lhs->val + rhs->val);
260 }
261
262 auto type = std::max(lhs->type(), rhs->type());
263 if (type == LINEAR) {
265 rhs);
266 } else if (type == QUADRATIC) {
268 rhs);
269 } else {
271 rhs);
272 }
273 }
274
283
289 const ExpressionPtr<Scalar>& rhs) {
290 using enum ExpressionType;
291
292 // Prune expression
293 if (lhs->is_constant(Scalar(0))) {
294 if (rhs->is_constant(Scalar(0))) {
295 // Return zero, which rhs currently is
296 return rhs;
297 } else {
298 // Return rhs negated
299 return -rhs;
300 }
301 } else if (rhs->is_constant(Scalar(0))) {
302 // Return lhs unmodified
303 return lhs;
304 }
305
306 // Evaluate constant
307 if (lhs->type() == CONSTANT && rhs->type() == CONSTANT) {
308 return constant_ptr(lhs->val - rhs->val);
309 }
310
311 auto type = std::max(lhs->type(), rhs->type());
312 if (type == LINEAR) {
314 rhs);
315 } else if (type == QUADRATIC) {
317 rhs);
318 } else {
320 rhs);
321 }
322 }
323
328 using enum ExpressionType;
329
330 // Prune expression
331 if (lhs->is_constant(Scalar(0))) {
332 // Return zero, which lhs currently is
333 return lhs;
334 }
335
336 // Evaluate constant
337 if (lhs->type() == CONSTANT) {
338 return constant_ptr(-lhs->val);
339 }
340
341 if (lhs->type() == LINEAR) {
343 } else if (lhs->type() == QUADRATIC) {
345 } else {
347 }
348 }
349
354 return lhs;
355 }
356
365 [[maybe_unused]] Scalar rhs) const = 0;
366
371 virtual ExpressionType type() const = 0;
372
376 virtual std::string_view name() const = 0;
377
384 [[maybe_unused]] Scalar rhs) const {
385 return Scalar(0);
386 }
387
394 [[maybe_unused]] Scalar rhs) const {
395 return Scalar(0);
396 }
397
405 [[maybe_unused]] const ExpressionPtr<Scalar>& rhs) const {
406 return constant_ptr(Scalar(0));
407 }
408
416 [[maybe_unused]] const ExpressionPtr<Scalar>& rhs) const {
417 return constant_ptr(Scalar(0));
418 }
419};
420
421template <typename Scalar>
422ExpressionPtr<Scalar> constant_ptr(Scalar value) {
424}
425
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);
438
443template <typename Scalar, ExpressionType T>
452
453 Scalar value(Scalar lhs, Scalar rhs) const override { return lhs - rhs; }
454
455 ExpressionType type() const override { return T; }
456
457 std::string_view name() const override { return "binary minus"; }
458
459 Scalar grad_l(Scalar, Scalar) const override { return this->adj; }
460
461 Scalar grad_r(Scalar, Scalar) const override { return -this->adj; }
462
465 const ExpressionPtr<Scalar>&) const override {
466 return this->adj_expr;
467 }
468
471 const ExpressionPtr<Scalar>&) const override {
472 return -this->adj_expr;
473 }
474};
475
480template <typename Scalar, ExpressionType T>
489
490 Scalar value(Scalar lhs, Scalar rhs) const override { return lhs + rhs; }
491
492 ExpressionType type() const override { return T; }
493
494 std::string_view name() const override { return "binary plus"; }
495
496 Scalar grad_l(Scalar, Scalar) const override { return this->adj; }
497
498 Scalar grad_r(Scalar, Scalar) const override { return this->adj; }
499
502 const ExpressionPtr<Scalar>&) const override {
503 return this->adj_expr;
504 }
505
508 const ExpressionPtr<Scalar>&) const override {
509 return this->adj_expr;
510 }
511};
512
516template <typename Scalar>
522 : Expression<Scalar>{std::move(lhs)} {}
523
524 Scalar value(Scalar x, Scalar) const override {
525 using std::cbrt;
526 return cbrt(x);
527 }
528
529 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
530
531 std::string_view name() const override { return "cbrt"; }
532
533 Scalar grad_l(Scalar x, Scalar) const override {
534 using std::cbrt;
535
536 Scalar c = cbrt(x);
537 return this->adj / (Scalar(3) * c * c);
538 }
539
541 const ExpressionPtr<Scalar>& x,
542 const ExpressionPtr<Scalar>&) const override {
543 auto c = cbrt(x);
544 return this->adj_expr / (constant_ptr(Scalar(3)) * c * c);
545 }
546};
547
552template <typename Scalar>
554 using enum ExpressionType;
555 using std::cbrt;
556
557 // Evaluate constant
558 if (x->type() == CONSTANT) {
559 if (x->val == Scalar(0)) {
560 // Return zero
561 return x;
562 } else if (x->val == Scalar(-1) || x->val == Scalar(1)) {
563 return x;
564 } else {
565 return constant_ptr(cbrt(x->val));
566 }
567 }
568
569 return make_expression_ptr<CbrtExpression<Scalar>>(x);
570}
571
575template <typename Scalar>
580 explicit constexpr ConstantExpression(Scalar value)
581 : Expression<Scalar>{value} {}
582
583 Scalar value(Scalar, Scalar) const override { return this->val; }
584
585 ExpressionType type() const override { return ExpressionType::CONSTANT; }
586
587 std::string_view name() const override { return "constant"; }
588};
589
593template <typename Scalar>
596 constexpr DecisionVariableExpression() = default;
597
603
604 Scalar value(Scalar, Scalar) const override { return this->val; }
605
606 ExpressionType type() const override { return ExpressionType::LINEAR; }
607
608 std::string_view name() const override { return "decision variable"; }
609};
610
615template <typename Scalar, ExpressionType T>
623
624 Scalar value(Scalar lhs, Scalar rhs) const override { return lhs / rhs; }
625
626 ExpressionType type() const override { return T; }
627
628 std::string_view name() const override { return "division"; }
629
630 Scalar grad_l(Scalar, Scalar rhs) const override { return this->adj / rhs; };
631
632 Scalar grad_r(Scalar lhs, Scalar rhs) const override {
633 return this->adj * -lhs / (rhs * rhs);
634 }
635
638 const ExpressionPtr<Scalar>& rhs) const override {
639 return this->adj_expr / rhs;
640 }
641
644 const ExpressionPtr<Scalar>& rhs) const override {
645 return this->adj_expr * -lhs / (rhs * rhs);
646 }
647};
648
653template <typename Scalar, ExpressionType T>
661
662 Scalar value(Scalar lhs, Scalar rhs) const override { return lhs * rhs; }
663
664 ExpressionType type() const override { return T; }
665
666 std::string_view name() const override { return "multiplication"; }
667
669 return this->adj * rhs;
670 }
671
673 return this->adj * lhs;
674 }
675
678 const ExpressionPtr<Scalar>& rhs) const override {
679 return this->adj_expr * rhs;
680 }
681
684 [[maybe_unused]] const ExpressionPtr<Scalar>& rhs) const override {
685 return this->adj_expr * lhs;
686 }
687};
688
693template <typename Scalar, ExpressionType T>
699 : Expression<Scalar>{std::move(lhs)} {}
700
701 Scalar value(Scalar lhs, Scalar) const override { return -lhs; }
702
703 ExpressionType type() const override { return T; }
704
705 std::string_view name() const override { return "unary minus"; }
706
707 Scalar grad_l(Scalar, Scalar) const override { return -this->adj; }
708
711 const ExpressionPtr<Scalar>&) const override {
712 return -this->adj_expr;
713 }
714};
715
720template <typename Scalar>
721constexpr void inc_ref_count(Expression<Scalar>* expr) {
722 ++expr->ref_count;
723}
724
729template <typename Scalar>
730constexpr void dec_ref_count(Expression<Scalar>* expr) {
731 // If a deeply nested tree is being deallocated all at once, calling the
732 // Expression destructor when expr's refcount reaches zero can cause a stack
733 // overflow. Instead, we iterate over its children to decrement their
734 // refcounts and deallocate them.
735 gch::small_vector<Expression<Scalar>*> stack;
736 stack.emplace_back(expr);
737
738 while (!stack.empty()) {
739 auto elem = stack.back();
740 stack.pop_back();
741
742 // Decrement the current node's refcount. If it reaches zero, deallocate the
743 // node and enqueue its children so their refcounts are decremented too.
744 if (--elem->ref_count == 0) {
745 if (elem->adj_expr != nullptr) {
746 stack.emplace_back(elem->adj_expr.get());
747 }
748 for (auto& arg : elem->args) {
749 if (arg != nullptr) {
750 stack.emplace_back(arg.get());
751 }
752 }
753
754 // Not calling the destructor here is safe because it only decrements
755 // refcounts, which was already done above.
756 if constexpr (USE_POOL_ALLOCATOR) {
757 auto alloc = global_pool_allocator<Expression<Scalar>>();
758 std::allocator_traits<decltype(alloc)>::deallocate(
759 alloc, elem, sizeof(Expression<Scalar>));
760 } else {
761 operator delete(elem);
762 }
763 }
764 }
765}
766
770template <typename Scalar>
776 : Expression<Scalar>{std::move(lhs)} {}
777
778 Scalar value(Scalar x, Scalar) const override {
779 using std::abs;
780 return abs(x);
781 }
782
783 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
784
785 std::string_view name() const override { return "abs"; }
786
787 Scalar grad_l(Scalar x, Scalar) const override {
788 if (x < Scalar(0)) {
789 return -this->adj;
790 } else if (x > Scalar(0)) {
791 return this->adj;
792 } else {
793 return Scalar(0);
794 }
795 }
796
798 const ExpressionPtr<Scalar>& x,
799 const ExpressionPtr<Scalar>&) const override {
800 return this->adj_expr * sign(x);
801 }
802};
803
808template <typename Scalar>
810 using enum ExpressionType;
811 using std::abs;
812
813 // Prune expression
814 if (x->is_constant(Scalar(0))) {
815 // Return zero, which x currently is
816 return x;
817 }
818
819 // Evaluate constant
820 if (x->type() == CONSTANT) {
821 return constant_ptr(abs(x->val));
822 }
823
824 return make_expression_ptr<AbsExpression<Scalar>>(x);
825}
826
830template <typename Scalar>
836 : Expression<Scalar>{std::move(lhs)} {}
837
838 Scalar value(Scalar x, Scalar) const override {
839 using std::acos;
840 return acos(x);
841 }
842
843 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
844
845 std::string_view name() const override { return "acos"; }
846
847 Scalar grad_l(Scalar x, Scalar) const override {
848 using std::sqrt;
849 return -this->adj / sqrt(Scalar(1) - x * x);
850 }
851
853 const ExpressionPtr<Scalar>& x,
854 const ExpressionPtr<Scalar>&) const override {
855 return -this->adj_expr / sqrt(constant_ptr(Scalar(1)) - x * x);
856 }
857};
858
863template <typename Scalar>
865 using enum ExpressionType;
866 using std::acos;
867
868 // Prune expression
869 if (x->is_constant(Scalar(0))) {
870 return constant_ptr(Scalar(std::numbers::pi) / Scalar(2));
871 }
872
873 // Evaluate constant
874 if (x->type() == CONSTANT) {
875 return constant_ptr(acos(x->val));
876 }
877
878 return make_expression_ptr<AcosExpression<Scalar>>(x);
879}
880
884template <typename Scalar>
890 : Expression<Scalar>{std::move(lhs)} {}
891
892 Scalar value(Scalar x, Scalar) const override {
893 using std::asin;
894 return asin(x);
895 }
896
897 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
898
899 std::string_view name() const override { return "asin"; }
900
901 Scalar grad_l(Scalar x, Scalar) const override {
902 using std::sqrt;
903 return this->adj / sqrt(Scalar(1) - x * x);
904 }
905
907 const ExpressionPtr<Scalar>& x,
908 const ExpressionPtr<Scalar>&) const override {
909 return this->adj_expr / sqrt(constant_ptr(Scalar(1)) - x * x);
910 }
911};
912
917template <typename Scalar>
919 using enum ExpressionType;
920 using std::asin;
921
922 // Prune expression
923 if (x->is_constant(Scalar(0))) {
924 // Return zero, which x currently is
925 return x;
926 }
927
928 // Evaluate constant
929 if (x->type() == CONSTANT) {
930 return constant_ptr(asin(x->val));
931 }
932
933 return make_expression_ptr<AsinExpression<Scalar>>(x);
934}
935
939template <typename Scalar>
945 : Expression<Scalar>{std::move(lhs)} {}
946
947 Scalar value(Scalar x, Scalar) const override {
948 using std::atan;
949 return atan(x);
950 }
951
952 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
953
954 std::string_view name() const override { return "atan"; }
955
956 Scalar grad_l(Scalar x, Scalar) const override {
957 return this->adj / (Scalar(1) + x * x);
958 }
959
961 const ExpressionPtr<Scalar>& x,
962 const ExpressionPtr<Scalar>&) const override {
963 return this->adj_expr / (constant_ptr(Scalar(1)) + x * x);
964 }
965};
966
971template <typename Scalar>
973 using enum ExpressionType;
974 using std::atan;
975
976 // Prune expression
977 if (x->is_constant(Scalar(0))) {
978 // Return zero, which x currently is
979 return x;
980 }
981
982 // Evaluate constant
983 if (x->type() == CONSTANT) {
984 return constant_ptr(atan(x->val));
985 }
986
987 return make_expression_ptr<AtanExpression<Scalar>>(x);
988}
989
993template <typename Scalar>
1001 : Expression<Scalar>{std::move(lhs), std::move(rhs)} {}
1002
1003 Scalar value(Scalar y, Scalar x) const override {
1004 using std::atan2;
1005 return atan2(y, x);
1006 }
1007
1008 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1009
1010 std::string_view name() const override { return "atan2"; }
1011
1012 Scalar grad_l(Scalar y, Scalar x) const override {
1013 return this->adj * x / (y * y + x * x);
1014 }
1015
1016 Scalar grad_r(Scalar y, Scalar x) const override {
1017 return this->adj * -y / (y * y + x * x);
1018 }
1019
1021 const ExpressionPtr<Scalar>& y,
1022 const ExpressionPtr<Scalar>& x) const override {
1023 return this->adj_expr * x / (y * y + x * x);
1024 }
1025
1027 const ExpressionPtr<Scalar>& y,
1028 const ExpressionPtr<Scalar>& x) const override {
1029 return this->adj_expr * -y / (y * y + x * x);
1030 }
1031};
1032
1038template <typename Scalar>
1040 const ExpressionPtr<Scalar>& x) {
1041 using enum ExpressionType;
1042 using std::atan2;
1043
1044 // Evaluate constant
1045 if (y->type() == CONSTANT && x->type() == CONSTANT) {
1046 return constant_ptr(atan2(y->val, x->val));
1047 }
1048
1049 return make_expression_ptr<Atan2Expression<Scalar>>(y, x);
1050}
1051
1055template <typename Scalar>
1061 : Expression<Scalar>{std::move(lhs)} {}
1062
1063 Scalar value(Scalar x, Scalar) const override {
1064 using std::cos;
1065 return cos(x);
1066 }
1067
1068 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1069
1070 std::string_view name() const override { return "cos"; }
1071
1072 Scalar grad_l(Scalar x, Scalar) const override {
1073 using std::sin;
1074 return this->adj * -sin(x);
1075 }
1076
1078 const ExpressionPtr<Scalar>& x,
1079 const ExpressionPtr<Scalar>&) const override {
1080 return this->adj_expr * -sin(x);
1081 }
1082};
1083
1088template <typename Scalar>
1090 using enum ExpressionType;
1091 using std::cos;
1092
1093 // Prune expression
1094 if (x->is_constant(Scalar(0))) {
1095 return constant_ptr(Scalar(1));
1096 }
1097
1098 // Evaluate constant
1099 if (x->type() == CONSTANT) {
1100 return constant_ptr(cos(x->val));
1101 }
1102
1103 return make_expression_ptr<CosExpression<Scalar>>(x);
1104}
1105
1109template <typename Scalar>
1115 : Expression<Scalar>{std::move(lhs)} {}
1116
1117 Scalar value(Scalar x, Scalar) const override {
1118 using std::cosh;
1119 return cosh(x);
1120 }
1121
1122 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1123
1124 std::string_view name() const override { return "cosh"; }
1125
1126 Scalar grad_l(Scalar x, Scalar) const override {
1127 using std::sinh;
1128 return this->adj * sinh(x);
1129 }
1130
1132 const ExpressionPtr<Scalar>& x,
1133 const ExpressionPtr<Scalar>&) const override {
1134 return this->adj_expr * sinh(x);
1135 }
1136};
1137
1142template <typename Scalar>
1144 using enum ExpressionType;
1145 using std::cosh;
1146
1147 // Prune expression
1148 if (x->is_constant(Scalar(0))) {
1149 return constant_ptr(Scalar(1));
1150 }
1151
1152 // Evaluate constant
1153 if (x->type() == CONSTANT) {
1154 return constant_ptr(cosh(x->val));
1155 }
1156
1157 return make_expression_ptr<CoshExpression<Scalar>>(x);
1158}
1159
1163template <typename Scalar>
1169 : Expression<Scalar>{std::move(lhs)} {}
1170
1171 Scalar value(Scalar x, Scalar) const override {
1172 using std::erf;
1173 return erf(x);
1174 }
1175
1176 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1177
1178 std::string_view name() const override { return "erf"; }
1179
1180 Scalar grad_l(Scalar x, Scalar) const override {
1181 using std::exp;
1182 return this->adj * Scalar(2.0 * std::numbers::inv_sqrtpi) * exp(-x * x);
1183 }
1184
1186 const ExpressionPtr<Scalar>& x,
1187 const ExpressionPtr<Scalar>&) const override {
1188 return this->adj_expr *
1189 constant_ptr(Scalar(2.0 * std::numbers::inv_sqrtpi)) * exp(-x * x);
1190 }
1191};
1192
1197template <typename Scalar>
1199 using enum ExpressionType;
1200 using std::erf;
1201
1202 // Prune expression
1203 if (x->is_constant(Scalar(0))) {
1204 // Return zero, which x currently is
1205 return x;
1206 }
1207
1208 // Evaluate constant
1209 if (x->type() == CONSTANT) {
1210 return constant_ptr(erf(x->val));
1211 }
1212
1213 return make_expression_ptr<ErfExpression<Scalar>>(x);
1214}
1215
1219template <typename Scalar>
1225 : Expression<Scalar>{std::move(lhs)} {}
1226
1227 Scalar value(Scalar x, Scalar) const override {
1228 using std::exp;
1229 return exp(x);
1230 }
1231
1232 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1233
1234 std::string_view name() const override { return "exp"; }
1235
1236 Scalar grad_l(Scalar x, Scalar) const override {
1237 using std::exp;
1238 return this->adj * exp(x);
1239 }
1240
1242 const ExpressionPtr<Scalar>& x,
1243 const ExpressionPtr<Scalar>&) const override {
1244 return this->adj_expr * exp(x);
1245 }
1246};
1247
1252template <typename Scalar>
1254 using enum ExpressionType;
1255 using std::exp;
1256
1257 // Prune expression
1258 if (x->is_constant(Scalar(0))) {
1259 return constant_ptr(Scalar(1));
1260 }
1261
1262 // Evaluate constant
1263 if (x->type() == CONSTANT) {
1264 return constant_ptr(exp(x->val));
1265 }
1266
1267 return make_expression_ptr<ExpExpression<Scalar>>(x);
1268}
1269
1270template <typename Scalar>
1271ExpressionPtr<Scalar> hypot(const ExpressionPtr<Scalar>& x,
1272 const ExpressionPtr<Scalar>& y);
1273
1277template <typename Scalar>
1285 : Expression<Scalar>{std::move(lhs), std::move(rhs)} {}
1286
1287 Scalar value(Scalar x, Scalar y) const override {
1288 using std::hypot;
1289 return hypot(x, y);
1290 }
1291
1292 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1293
1294 std::string_view name() const override { return "hypot"; }
1295
1296 Scalar grad_l(Scalar x, Scalar y) const override {
1297 using std::hypot;
1298 return this->adj * x / hypot(x, y);
1299 }
1300
1301 Scalar grad_r(Scalar x, Scalar y) const override {
1302 using std::hypot;
1303 return this->adj * y / hypot(x, y);
1304 }
1305
1307 const ExpressionPtr<Scalar>& x,
1308 const ExpressionPtr<Scalar>& y) const override {
1309 return this->adj_expr * x / hypot(x, y);
1310 }
1311
1313 const ExpressionPtr<Scalar>& x,
1314 const ExpressionPtr<Scalar>& y) const override {
1315 return this->adj_expr * y / hypot(x, y);
1316 }
1317};
1318
1324template <typename Scalar>
1326 const ExpressionPtr<Scalar>& y) {
1327 using enum ExpressionType;
1328 using std::hypot;
1329
1330 // Prune expression
1331 if (x->is_constant(Scalar(0))) {
1332 return abs(y);
1333 } else if (y->is_constant(Scalar(0))) {
1334 return abs(x);
1335 }
1336
1337 // Evaluate constant
1338 if (x->type() == CONSTANT && y->type() == CONSTANT) {
1339 return constant_ptr(hypot(x->val, y->val));
1340 }
1341
1342 return make_expression_ptr<HypotExpression<Scalar>>(x, y);
1343}
1344
1348template <typename Scalar>
1354 : Expression<Scalar>{std::move(x)} {}
1355
1356 Scalar value(Scalar x, Scalar) const override {
1357 return x >= Scalar(0) ? Scalar(1) : Scalar(0);
1358 }
1359
1360 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1361
1362 std::string_view name() const override { return "is nonnegative"; }
1363};
1364
1369template <typename Scalar>
1370ExpressionPtr<Scalar> is_nonnegative(const ExpressionPtr<Scalar>& x) {
1371 if (x->type() == ExpressionType::CONSTANT) {
1372 return constant_ptr(x->val >= Scalar(0) ? Scalar(1) : Scalar(0));
1373 }
1374
1375 return make_expression_ptr<IsNonnegativeExpression<Scalar>>(x);
1376}
1377
1381template <typename Scalar>
1387 : Expression<Scalar>{std::move(x)} {}
1388
1389 Scalar value(Scalar x, Scalar) const override {
1390 return x > Scalar(0) ? Scalar(1) : Scalar(0);
1391 }
1392
1393 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1394
1395 std::string_view name() const override { return "is positive"; }
1396};
1397
1402template <typename Scalar>
1403ExpressionPtr<Scalar> is_positive(const ExpressionPtr<Scalar>& x) {
1404 if (x->type() == ExpressionType::CONSTANT) {
1405 return constant_ptr(x->val > Scalar(0) ? Scalar(1) : Scalar(0));
1406 }
1407
1408 return make_expression_ptr<IsPositiveExpression<Scalar>>(x);
1409}
1410
1414template <typename Scalar>
1420 : Expression<Scalar>{std::move(lhs)} {}
1421
1422 Scalar value(Scalar x, Scalar) const override {
1423 using std::log;
1424 return log(x);
1425 }
1426
1427 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1428
1429 std::string_view name() const override { return "log"; }
1430
1431 Scalar grad_l(Scalar x, Scalar) const override { return this->adj / x; }
1432
1434 const ExpressionPtr<Scalar>& x,
1435 const ExpressionPtr<Scalar>&) const override {
1436 return this->adj_expr / x;
1437 }
1438};
1439
1444template <typename Scalar>
1446 using enum ExpressionType;
1447 using std::log;
1448
1449 // Prune expression
1450 if (x->is_constant(Scalar(0))) {
1451 // Return zero, which x currently is
1452 return x;
1453 }
1454
1455 // Evaluate constant
1456 if (x->type() == CONSTANT) {
1457 return constant_ptr(log(x->val));
1458 }
1459
1460 return make_expression_ptr<LogExpression<Scalar>>(x);
1461}
1462
1466template <typename Scalar>
1472 : Expression<Scalar>{std::move(lhs)} {}
1473
1474 Scalar value(Scalar x, Scalar) const override {
1475 using std::log10;
1476 return log10(x);
1477 }
1478
1479 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1480
1481 std::string_view name() const override { return "log10"; }
1482
1483 Scalar grad_l(Scalar x, Scalar) const override {
1484 return this->adj / (Scalar(std::numbers::ln10) * x);
1485 }
1486
1488 const ExpressionPtr<Scalar>& x,
1489 const ExpressionPtr<Scalar>&) const override {
1490 return this->adj_expr / (constant_ptr(Scalar(std::numbers::ln10)) * x);
1491 }
1492};
1493
1498template <typename Scalar>
1500 using enum ExpressionType;
1501 using std::log10;
1502
1503 // Prune expression
1504 if (x->is_constant(Scalar(0))) {
1505 // Return zero, which x currently is
1506 return x;
1507 }
1508
1509 // Evaluate constant
1510 if (x->type() == CONSTANT) {
1511 return constant_ptr(log10(x->val));
1512 }
1513
1514 return make_expression_ptr<Log10Expression<Scalar>>(x);
1515}
1516
1522template <typename Scalar>
1530
1531 Scalar value(Scalar a, Scalar b) const override {
1532 using std::max;
1533 return max(a, b);
1534 }
1535
1536 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1537
1538 std::string_view name() const override { return "max"; }
1539
1540 Scalar grad_l(Scalar a, Scalar b) const override {
1541 return a >= b ? this->adj : Scalar(0);
1542 }
1543
1544 Scalar grad_r(Scalar a, Scalar b) const override {
1545 return a >= b ? Scalar(0) : this->adj;
1546 }
1547
1549 const ExpressionPtr<Scalar>& a,
1550 const ExpressionPtr<Scalar>& b) const override {
1551 // adjoint * (a >= b)
1552 // adjoint * (a - b >= 0)
1553 return this->adj_expr * is_nonnegative(a - b);
1554 }
1555
1557 const ExpressionPtr<Scalar>& a,
1558 const ExpressionPtr<Scalar>& b) const override {
1559 // adjoint * !(a >= b)
1560 // adjoint * (a < b)
1561 // adjoint * (b > a)
1562 // adjoint * (b - a > 0)
1563 return this->adj_expr * is_positive(b - a);
1564 }
1565};
1566
1572template <typename Scalar>
1574 const ExpressionPtr<Scalar>& b) {
1575 using enum ExpressionType;
1576 using std::max;
1577
1578 // Evaluate constant
1579 if (a->type() == CONSTANT && b->type() == CONSTANT) {
1580 return constant_ptr(max(a->val, b->val));
1581 }
1582
1583 return make_expression_ptr<MaxExpression<Scalar>>(a, b);
1584}
1585
1591template <typename Scalar>
1599
1600 Scalar value(Scalar a, Scalar b) const override {
1601 using std::min;
1602 return min(a, b);
1603 }
1604
1605 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1606
1607 std::string_view name() const override { return "min"; }
1608
1609 Scalar grad_l(Scalar a, Scalar b) const override {
1610 return a <= b ? this->adj : Scalar(0);
1611 }
1612
1614 [[maybe_unused]] Scalar b) const override {
1615 return a <= b ? Scalar(0) : this->adj;
1616 }
1617
1619 const ExpressionPtr<Scalar>& a,
1620 const ExpressionPtr<Scalar>& b) const override {
1621 // adjoint * (a <= b)
1622 // adjoint * (b >= a)
1623 // adjoint * (b - a >= 0)
1624 return this->adj_expr * is_nonnegative(b - a);
1625 }
1626
1628 const ExpressionPtr<Scalar>& a,
1629 const ExpressionPtr<Scalar>& b) const override {
1630 // adjoint * !(a <= b)
1631 // adjoint * (a > b)
1632 // adjoint * (a - b > 0)
1633 return this->adj_expr * is_positive(a - b);
1634 }
1635};
1636
1642template <typename Scalar>
1644 const ExpressionPtr<Scalar>& b) {
1645 using enum ExpressionType;
1646 using std::min;
1647
1648 // Evaluate constant
1649 if (a->type() == CONSTANT && b->type() == CONSTANT) {
1650 return constant_ptr(min(a->val, b->val));
1651 }
1652
1653 return make_expression_ptr<MinExpression<Scalar>>(a, b);
1654}
1655
1656template <typename Scalar>
1657ExpressionPtr<Scalar> pow(const ExpressionPtr<Scalar>& base,
1658 const ExpressionPtr<Scalar>& power);
1659
1664template <typename Scalar, ExpressionType T>
1672
1673 Scalar value(Scalar base, Scalar power) const override {
1674 using std::pow;
1675 return pow(base, power);
1676 }
1677
1678 ExpressionType type() const override { return T; }
1679
1680 std::string_view name() const override { return "pow"; }
1681
1683 using std::pow;
1684 return this->adj * pow(base, power - Scalar(1)) * power;
1685 }
1686
1688 using std::log;
1689 using std::pow;
1690
1691 return this->adj * pow(base, power) * log(base);
1692 }
1693
1696 const ExpressionPtr<Scalar>& power) const override {
1697 return this->adj_expr * pow(base, power - constant_ptr(Scalar(1))) * power;
1698 }
1699
1702 const ExpressionPtr<Scalar>& power) const override {
1703 return this->adj_expr * pow(base, power) * log(base);
1704 }
1705};
1706
1712template <typename Scalar>
1715 using enum ExpressionType;
1716 using std::pow;
1717
1718 // Prune expression
1719 if (base->is_constant(Scalar(0))) {
1720 // Return zero, which base currently is
1721 return base;
1722 } else if (base->is_constant(Scalar(1))) {
1723 // Return one, which base currently is
1724 return base;
1725 }
1726 if (power->is_constant(Scalar(0))) {
1727 return constant_ptr(Scalar(1));
1728 } else if (power->is_constant(Scalar(1))) {
1729 // Return base unmodified
1730 return base;
1731 }
1732
1733 // Evaluate constant
1734 if (base->type() == CONSTANT && power->type() == CONSTANT) {
1735 return constant_ptr(pow(base->val, power->val));
1736 }
1737
1738 if (power->is_constant(Scalar(2))) {
1739 if (base->type() == LINEAR) {
1740 return make_expression_ptr<MultExpression<Scalar, QUADRATIC>>(base, base);
1741 } else {
1742 return make_expression_ptr<MultExpression<Scalar, NONLINEAR>>(base, base);
1743 }
1744 }
1745
1746 return make_expression_ptr<PowExpression<Scalar, NONLINEAR>>(base, power);
1747}
1748
1752template <typename Scalar>
1758 : Expression<Scalar>{std::move(lhs)} {}
1759
1760 Scalar value(Scalar x, Scalar) const override {
1761 if (x < Scalar(0)) {
1762 return Scalar(-1);
1763 } else if (x == Scalar(0)) {
1764 return Scalar(0);
1765 } else {
1766 return Scalar(1);
1767 }
1768 }
1769
1770 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1771
1772 std::string_view name() const override { return "sign"; }
1773};
1774
1779template <typename Scalar>
1781 using enum ExpressionType;
1782
1783 // Evaluate constant
1784 if (x->type() == CONSTANT) {
1785 if (x->val < Scalar(0)) {
1786 return constant_ptr(Scalar(-1));
1787 } else if (x->val == Scalar(0)) {
1788 // Return zero
1789 return x;
1790 } else {
1791 return constant_ptr(Scalar(1));
1792 }
1793 }
1794
1795 return make_expression_ptr<SignExpression<Scalar>>(x);
1796}
1797
1801template <typename Scalar>
1807 : Expression<Scalar>{std::move(lhs)} {}
1808
1809 Scalar value(Scalar x, Scalar) const override {
1810 using std::sin;
1811 return sin(x);
1812 }
1813
1814 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1815
1816 std::string_view name() const override { return "sin"; }
1817
1818 Scalar grad_l(Scalar x, Scalar) const override {
1819 using std::cos;
1820 return this->adj * cos(x);
1821 }
1822
1824 const ExpressionPtr<Scalar>& x,
1825 const ExpressionPtr<Scalar>&) const override {
1826 return this->adj_expr * cos(x);
1827 }
1828};
1829
1834template <typename Scalar>
1836 using enum ExpressionType;
1837 using std::sin;
1838
1839 // Prune expression
1840 if (x->is_constant(Scalar(0))) {
1841 // Return zero, which x currently is
1842 return x;
1843 }
1844
1845 // Evaluate constant
1846 if (x->type() == CONSTANT) {
1847 return constant_ptr(sin(x->val));
1848 }
1849
1850 return make_expression_ptr<SinExpression<Scalar>>(x);
1851}
1852
1856template <typename Scalar>
1862 : Expression<Scalar>{std::move(lhs)} {}
1863
1864 Scalar value(Scalar x, Scalar) const override {
1865 using std::sinh;
1866 return sinh(x);
1867 }
1868
1869 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1870
1871 std::string_view name() const override { return "sinh"; }
1872
1873 Scalar grad_l(Scalar x, Scalar) const override {
1874 using std::cosh;
1875 return this->adj * cosh(x);
1876 }
1877
1879 const ExpressionPtr<Scalar>& x,
1880 const ExpressionPtr<Scalar>&) const override {
1881 return this->adj_expr * cosh(x);
1882 }
1883};
1884
1889template <typename Scalar>
1891 using enum ExpressionType;
1892 using std::sinh;
1893
1894 // Prune expression
1895 if (x->is_constant(Scalar(0))) {
1896 // Return zero, which x currently is
1897 return x;
1898 }
1899
1900 // Evaluate constant
1901 if (x->type() == CONSTANT) {
1902 return constant_ptr(sinh(x->val));
1903 }
1904
1905 return make_expression_ptr<SinhExpression<Scalar>>(x);
1906}
1907
1911template <typename Scalar>
1917 : Expression<Scalar>{std::move(lhs)} {}
1918
1919 Scalar value(Scalar x, Scalar) const override {
1920 using std::sqrt;
1921 return sqrt(x);
1922 }
1923
1924 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1925
1926 std::string_view name() const override { return "sqrt"; }
1927
1928 Scalar grad_l(Scalar x, Scalar) const override {
1929 using std::sqrt;
1930 return this->adj / (Scalar(2) * sqrt(x));
1931 }
1932
1934 const ExpressionPtr<Scalar>& x,
1935 const ExpressionPtr<Scalar>&) const override {
1936 return this->adj_expr / (constant_ptr(Scalar(2)) * sqrt(x));
1937 }
1938};
1939
1944template <typename Scalar>
1946 using enum ExpressionType;
1947 using std::sqrt;
1948
1949 // Evaluate constant
1950 if (x->type() == CONSTANT) {
1951 if (x->val == Scalar(0)) {
1952 // Return zero
1953 return x;
1954 } else if (x->val == Scalar(1)) {
1955 return x;
1956 } else {
1957 return constant_ptr(sqrt(x->val));
1958 }
1959 }
1960
1961 return make_expression_ptr<SqrtExpression<Scalar>>(x);
1962}
1963
1967template <typename Scalar>
1973 : Expression<Scalar>{std::move(lhs)} {}
1974
1975 Scalar value(Scalar x, Scalar) const override {
1976 using std::tan;
1977 return tan(x);
1978 }
1979
1980 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
1981
1982 std::string_view name() const override { return "tan"; }
1983
1984 Scalar grad_l(Scalar x, Scalar) const override {
1985 using std::cos;
1986
1987 auto c = cos(x);
1988 return this->adj / (c * c);
1989 }
1990
1992 const ExpressionPtr<Scalar>& x,
1993 const ExpressionPtr<Scalar>&) const override {
1994 auto c = cos(x);
1995 return this->adj_expr / (c * c);
1996 }
1997};
1998
2003template <typename Scalar>
2005 using enum ExpressionType;
2006 using std::tan;
2007
2008 // Prune expression
2009 if (x->is_constant(Scalar(0))) {
2010 // Return zero, which x currently is
2011 return x;
2012 }
2013
2014 // Evaluate constant
2015 if (x->type() == CONSTANT) {
2016 return constant_ptr(tan(x->val));
2017 }
2018
2019 return make_expression_ptr<TanExpression<Scalar>>(x);
2020}
2021
2025template <typename Scalar>
2031 : Expression<Scalar>{std::move(lhs)} {}
2032
2033 Scalar value(Scalar x, Scalar) const override {
2034 using std::tanh;
2035 return tanh(x);
2036 }
2037
2038 ExpressionType type() const override { return ExpressionType::NONLINEAR; }
2039
2040 std::string_view name() const override { return "tanh"; }
2041
2042 Scalar grad_l(Scalar x, Scalar) const override {
2043 using std::cosh;
2044
2045 auto c = cosh(x);
2046 return this->adj / (c * c);
2047 }
2048
2050 const ExpressionPtr<Scalar>& x,
2051 const ExpressionPtr<Scalar>&) const override {
2052 auto c = cosh(x);
2053 return this->adj_expr / (c * c);
2054 }
2055};
2056
2061template <typename Scalar>
2063 using enum ExpressionType;
2064 using std::tanh;
2065
2066 // Prune expression
2067 if (x->is_constant(Scalar(0))) {
2068 // Return zero, which x currently is
2069 return x;
2070 }
2071
2072 // Evaluate constant
2073 if (x->type() == CONSTANT) {
2074 return constant_ptr(tanh(x->val));
2075 }
2076
2077 return make_expression_ptr<TanhExpression<Scalar>>(x);
2078}
2079
2080} // namespace slp::detail
Definition intrusive_shared_ptr.hpp:27
Definition expression.hpp:771
std::string_view name() const override
Definition expression.hpp:785
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:797
constexpr AbsExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:775
ExpressionType type() const override
Definition expression.hpp:783
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:787
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:778
Definition expression.hpp:831
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:838
constexpr AcosExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:835
ExpressionType type() const override
Definition expression.hpp:843
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:847
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:852
std::string_view name() const override
Definition expression.hpp:845
Definition expression.hpp:885
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:892
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:901
std::string_view name() const override
Definition expression.hpp:899
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:906
ExpressionType type() const override
Definition expression.hpp:897
constexpr AsinExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:889
Definition expression.hpp:994
std::string_view name() const override
Definition expression.hpp:1010
Scalar value(Scalar y, Scalar x) const override
Definition expression.hpp:1003
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &y, const ExpressionPtr< Scalar > &x) const override
Definition expression.hpp:1020
constexpr Atan2Expression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:999
Scalar grad_l(Scalar y, Scalar x) const override
Definition expression.hpp:1012
Scalar grad_r(Scalar y, Scalar x) const override
Definition expression.hpp:1016
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &y, const ExpressionPtr< Scalar > &x) const override
Definition expression.hpp:1026
ExpressionType type() const override
Definition expression.hpp:1008
Definition expression.hpp:940
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:960
std::string_view name() const override
Definition expression.hpp:954
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:956
constexpr AtanExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:944
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:947
ExpressionType type() const override
Definition expression.hpp:952
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:1056
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1063
ExpressionType type() const override
Definition expression.hpp:1068
constexpr CosExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1060
std::string_view name() const override
Definition expression.hpp:1070
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1072
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1077
Definition expression.hpp:1110
ExpressionType type() const override
Definition expression.hpp:1122
std::string_view name() const override
Definition expression.hpp:1124
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1117
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1131
constexpr CoshExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1114
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1126
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:642
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:636
Scalar grad_r(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:632
Definition expression.hpp:1164
std::string_view name() const override
Definition expression.hpp:1178
constexpr ErfExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1168
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1180
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1185
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1171
ExpressionType type() const override
Definition expression.hpp:1176
Definition expression.hpp:1220
std::string_view name() const override
Definition expression.hpp:1234
ExpressionType type() const override
Definition expression.hpp:1232
constexpr ExpExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1224
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1236
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1241
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1227
Definition expression.hpp:89
Scalar val
The value of the expression node.
Definition expression.hpp:94
ExpressionPtr< Scalar > adj_expr
Definition expression.hpp:101
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
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
Scalar adj
The adjoint of the expression node, used during autodiff.
Definition expression.hpp:97
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
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:1278
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &y) const override
Definition expression.hpp:1306
constexpr HypotExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1283
ExpressionType type() const override
Definition expression.hpp:1292
std::string_view name() const override
Definition expression.hpp:1294
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &y) const override
Definition expression.hpp:1312
Scalar value(Scalar x, Scalar y) const override
Definition expression.hpp:1287
Scalar grad_r(Scalar x, Scalar y) const override
Definition expression.hpp:1301
Scalar grad_l(Scalar x, Scalar y) const override
Definition expression.hpp:1296
Definition expression.hpp:1349
ExpressionType type() const override
Definition expression.hpp:1360
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1356
std::string_view name() const override
Definition expression.hpp:1362
constexpr IsNonnegativeExpression(ExpressionPtr< Scalar > x)
Definition expression.hpp:1353
Definition expression.hpp:1382
std::string_view name() const override
Definition expression.hpp:1395
ExpressionType type() const override
Definition expression.hpp:1393
constexpr IsPositiveExpression(ExpressionPtr< Scalar > x)
Definition expression.hpp:1386
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1389
Definition expression.hpp:1467
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1474
std::string_view name() const override
Definition expression.hpp:1481
constexpr Log10Expression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1471
ExpressionType type() const override
Definition expression.hpp:1479
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1487
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1483
Definition expression.hpp:1415
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1433
ExpressionType type() const override
Definition expression.hpp:1427
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1431
std::string_view name() const override
Definition expression.hpp:1429
constexpr LogExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1419
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1422
Definition expression.hpp:1523
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1548
Scalar value(Scalar a, Scalar b) const override
Definition expression.hpp:1531
Scalar grad_r(Scalar a, Scalar b) const override
Definition expression.hpp:1544
constexpr MaxExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1528
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1556
Scalar grad_l(Scalar a, Scalar b) const override
Definition expression.hpp:1540
ExpressionType type() const override
Definition expression.hpp:1536
std::string_view name() const override
Definition expression.hpp:1538
Definition expression.hpp:1592
std::string_view name() const override
Definition expression.hpp:1607
Scalar value(Scalar a, Scalar b) const override
Definition expression.hpp:1600
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1618
ExpressionType type() const override
Definition expression.hpp:1605
constexpr MinExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1597
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &a, const ExpressionPtr< Scalar > &b) const override
Definition expression.hpp:1627
Scalar grad_l(Scalar a, Scalar b) const override
Definition expression.hpp:1609
Scalar grad_r(Scalar a, Scalar b) const override
Definition expression.hpp:1613
Definition expression.hpp:654
Scalar grad_l(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:668
Scalar value(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:662
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:676
constexpr MultExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:659
ExpressionType type() const override
Definition expression.hpp:664
std::string_view name() const override
Definition expression.hpp:666
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &lhs, const ExpressionPtr< Scalar > &rhs) const override
Definition expression.hpp:682
Scalar grad_r(Scalar lhs, Scalar rhs) const override
Definition expression.hpp:672
Definition expression.hpp:1665
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &base, const ExpressionPtr< Scalar > &power) const override
Definition expression.hpp:1694
Scalar value(Scalar base, Scalar power) const override
Definition expression.hpp:1673
ExpressionType type() const override
Definition expression.hpp:1678
Scalar grad_l(Scalar base, Scalar power) const override
Definition expression.hpp:1682
ExpressionPtr< Scalar > grad_expr_r(const ExpressionPtr< Scalar > &base, const ExpressionPtr< Scalar > &power) const override
Definition expression.hpp:1700
Scalar grad_r(Scalar base, Scalar power) const override
Definition expression.hpp:1687
std::string_view name() const override
Definition expression.hpp:1680
constexpr PowExpression(ExpressionPtr< Scalar > lhs, ExpressionPtr< Scalar > rhs)
Definition expression.hpp:1670
Definition expression.hpp:1753
constexpr SignExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1757
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1760
ExpressionType type() const override
Definition expression.hpp:1770
std::string_view name() const override
Definition expression.hpp:1772
Definition expression.hpp:1802
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1809
std::string_view name() const override
Definition expression.hpp:1816
ExpressionType type() const override
Definition expression.hpp:1814
constexpr SinExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1806
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1823
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1818
Definition expression.hpp:1857
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1873
std::string_view name() const override
Definition expression.hpp:1871
ExpressionType type() const override
Definition expression.hpp:1869
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1878
constexpr SinhExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1861
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1864
Definition expression.hpp:1912
ExpressionType type() const override
Definition expression.hpp:1924
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1928
std::string_view name() const override
Definition expression.hpp:1926
constexpr SqrtExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1916
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1919
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1933
Definition expression.hpp:1968
std::string_view name() const override
Definition expression.hpp:1982
constexpr TanExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:1972
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:1984
ExpressionType type() const override
Definition expression.hpp:1980
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:1991
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:1975
Definition expression.hpp:2026
Scalar grad_l(Scalar x, Scalar) const override
Definition expression.hpp:2042
ExpressionType type() const override
Definition expression.hpp:2038
std::string_view name() const override
Definition expression.hpp:2040
Scalar value(Scalar x, Scalar) const override
Definition expression.hpp:2033
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &x, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:2049
constexpr TanhExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:2030
Definition expression.hpp:694
Scalar grad_l(Scalar, Scalar) const override
Definition expression.hpp:707
constexpr UnaryMinusExpression(ExpressionPtr< Scalar > lhs)
Definition expression.hpp:698
ExpressionPtr< Scalar > grad_expr_l(const ExpressionPtr< Scalar > &, const ExpressionPtr< Scalar > &) const override
Definition expression.hpp:709
ExpressionType type() const override
Definition expression.hpp:703
Scalar value(Scalar lhs, Scalar) const override
Definition expression.hpp:701
std::string_view name() const override
Definition expression.hpp:705