Graph Framework
Loading...
Searching...
No Matches
math.hpp
Go to the documentation of this file.
1//------------------------------------------------------------------------------
4//------------------------------------------------------------------------------
5
6#ifndef math_h
7#define math_h
8
9#include <cmath>
10
11#include "node.hpp"
12
13namespace graph {
14//******************************************************************************
15// Sqrt node.
16//******************************************************************************
17//------------------------------------------------------------------------------
24//------------------------------------------------------------------------------
25 template<jit::float_scalar T, bool SAFE_MATH=false>
26 class sqrt_node final : public straight_node<T, SAFE_MATH> {
27 private:
28//------------------------------------------------------------------------------
33//------------------------------------------------------------------------------
34 static std::string to_string(leaf_node<T, SAFE_MATH> *a) {
35 return "sqrt" + jit::format_to_string(reinterpret_cast<size_t> (a));
36 }
37
38 public:
39//------------------------------------------------------------------------------
43//------------------------------------------------------------------------------
46
47//------------------------------------------------------------------------------
53//------------------------------------------------------------------------------
55 backend::buffer<T> result = this->arg->evaluate();
56 result.sqrt();
57 return result;
58 }
59
60//------------------------------------------------------------------------------
64//------------------------------------------------------------------------------
66 auto ac = constant_cast(this->arg);
67
68 if (ac.get()) {
69 if (ac->is(0) || ac->is(1)) {
70 return this->arg;
71 }
72 return constant<T, SAFE_MATH> (this->evaluate());
73 }
74
75 auto ap1 = piecewise_1D_cast(this->arg);
76 if (ap1.get()) {
77 return piecewise_1D(this->evaluate(),
78 ap1->get_arg(),
79 ap1->get_scale(),
80 ap1->get_offset());
81 }
82
83 auto ap2 = piecewise_2D_cast(this->arg);
84 if (ap2.get()) {
85 return piecewise_2D(this->evaluate(),
86 ap2->get_num_columns(),
87 ap2->get_left(), ap2->get_x_scale(), ap2->get_x_offset(),
88 ap2->get_right(), ap2->get_y_scale(), ap2->get_y_offset());
89 }
90
91// Handle cases like sqrt(c*x) where c is constant or cases like sqrt((x^a)*y).
92// Note that we need to disable this reduction C is a negative real.
93 auto am = multiply_cast(this->arg);
94 if (am.get()) {
95 if (pow_cast(am->get_left()).get() ||
96 am->get_left()->is_constant() ||
97 pow_cast(am->get_right()).get() ||
98 am->get_right()->is_constant()) {
99 if constexpr (jit::complex_scalar<T>) {
100 return sqrt(am->get_left()) *
101 sqrt(am->get_right());
102 } else {
103 if (am->get_left()->is_constant() &&
104 !am->get_left()->evaluate().is_negative()) {
105 return sqrt(am->get_left()) *
106 sqrt(am->get_right());
107 }
108 }
109 }
110 }
111
112 auto ad = divide_cast(this->arg);
113 if (ad.get()) {
114// sqrt((c1*x)/y) -> c2*sqrt(x/y)
115 auto alm = multiply_cast(ad->get_left());
116 if (alm.get() && alm->get_left()->is_constant()) {
117 return sqrt(alm->get_left()) *
118 sqrt(alm->get_right()/ad->get_right());
119 }
120
121// Handle cases like sqrt(x^a/b) and sqrt(a/x^b) or sqrt(c/b) and sqrt(a/c)
122// where c is a constant.
123 if (pow_cast(ad->get_left()).get() ||
124 ad->get_left()->is_constant() ||
125 pow_cast(ad->get_right()).get() ||
126 ad->get_right()->is_constant()) {
127 return sqrt(ad->get_left()) /
128 sqrt(ad->get_right());
129 }
130 }
131
132 return this->shared_from_this();
133 }
134
135//------------------------------------------------------------------------------
142//------------------------------------------------------------------------------
144 if (this->is_match(x)) {
145 return one<T, SAFE_MATH> ();
146 }
147
148 const size_t hash = reinterpret_cast<size_t> (x.get());
149 if (this->df_cache.find(hash) == this->df_cache.end()) {
150 this->df_cache[hash] = this->arg->df(x)
151 / (2.0*this->shared_from_this());
152 }
153 return this->df_cache[hash];
154 }
155
156//------------------------------------------------------------------------------
164//------------------------------------------------------------------------------
166 compile(std::ostringstream &stream,
167 jit::register_map &registers,
169 const jit::register_usage &usage) {
170 if (registers.find(this) == registers.end()) {
171 shared_leaf<T, SAFE_MATH> a = this->arg->compile(stream,
172 registers,
173 indices,
174 usage);
175
176 registers[this] = jit::to_string('r', this);
177 stream << " const ";
178 jit::add_type<T> (stream);
179 stream << " " << registers[this] << " = sqrt("
180 << registers[a.get()] << ")";
181 this->endline(stream, usage);
182 }
183
184 return this->shared_from_this();
185 }
186
187//------------------------------------------------------------------------------
192//------------------------------------------------------------------------------
194 if (this == x.get()) {
195 return true;
196 }
197
198 auto x_cast = sqrt_cast(x);
199 if (x_cast.get()) {
200 return this->arg->is_match(x_cast->get_arg());
201 }
202
203 return false;
204 }
205
206//------------------------------------------------------------------------------
208//------------------------------------------------------------------------------
209 virtual void to_latex() const {
210 std::cout << "\\sqrt{";
211 this->arg->to_latex();
212 std::cout << "}";
213 }
214
215//------------------------------------------------------------------------------
219//------------------------------------------------------------------------------
220 virtual bool is_power_like() const {
221 return true;
222 }
223
224//------------------------------------------------------------------------------
228//------------------------------------------------------------------------------
230 return this->arg;
231 }
232
233//------------------------------------------------------------------------------
237//------------------------------------------------------------------------------
239 return constant<T, SAFE_MATH> (static_cast<T> (0.5));
240 }
241
242//------------------------------------------------------------------------------
246//------------------------------------------------------------------------------
248 if (this->has_pseudo()) {
249 return sqrt(this->arg->remove_pseudo());
250 }
251 return this->shared_from_this();
252 }
253
254//------------------------------------------------------------------------------
260//------------------------------------------------------------------------------
261 virtual shared_leaf<T, SAFE_MATH> to_vizgraph(std::stringstream &stream,
262 jit::register_map &registers) {
263 if (registers.find(this) == registers.end()) {
264 const std::string name = jit::to_string('r', this);
265 registers[this] = name;
266 stream << " " << name
267 << " [label = \"sqrt\", shape = oval, style = filled, fillcolor = blue, fontcolor = white];" << std::endl;
268
269 auto a = this->arg->to_vizgraph(stream, registers);
270 stream << " " << name << " -- " << registers[a.get()] << ";" << std::endl;
271 }
272
273 return this->shared_from_this();
274 }
275 };
276
277//------------------------------------------------------------------------------
285//------------------------------------------------------------------------------
286 template<jit::float_scalar T, bool SAFE_MATH=false>
288 auto temp = std::make_shared<sqrt_node<T, SAFE_MATH>> (x)->reduce();
289// Test for hash collisions.
290 for (size_t i = temp->get_hash();
292 if (leaf_node<T, SAFE_MATH>::caches.nodes.find(i) ==
295 return temp;
296 } else if (temp->is_match(leaf_node<T, SAFE_MATH>::caches.nodes[i])) {
298 }
299 }
300#if defined(__clang__) || defined(__GNUC__)
302#else
303 assert(false && "Should never reach.");
304#endif
305 }
306
308 template<jit::float_scalar T, bool SAFE_MATH=false>
309 using shared_sqrt = std::shared_ptr<sqrt_node<T, SAFE_MATH>>;
310
311//------------------------------------------------------------------------------
319//------------------------------------------------------------------------------
320 template<jit::float_scalar T, bool SAFE_MATH=false>
322 return std::dynamic_pointer_cast<sqrt_node<T, SAFE_MATH>> (x);
323 }
324
325//******************************************************************************
326// Exp node.
327//******************************************************************************
328//------------------------------------------------------------------------------
335//------------------------------------------------------------------------------
336 template<jit::float_scalar T, bool SAFE_MATH=false>
337 class exp_node final : public straight_node<T, SAFE_MATH> {
338 private:
339//------------------------------------------------------------------------------
344//------------------------------------------------------------------------------
345 static std::string to_string(leaf_node<T, SAFE_MATH> *a) {
346 return "exp" + jit::format_to_string(reinterpret_cast<size_t> (a));
347 }
348
349 public:
350//------------------------------------------------------------------------------
354//------------------------------------------------------------------------------
357
358//------------------------------------------------------------------------------
364//------------------------------------------------------------------------------
366 backend::buffer<T> result = this->arg->evaluate();
367 result.exp();
368 return result;
369 }
370
371//------------------------------------------------------------------------------
375//------------------------------------------------------------------------------
377 if (constant_cast(this->arg).get()) {
378 return constant<T, SAFE_MATH> (this->evaluate());
379 }
380
381 auto ap1 = piecewise_1D_cast(this->arg);
382 if (ap1.get()) {
383 return piecewise_1D(this->evaluate(),
384 ap1->get_arg(),
385 ap1->get_scale(),
386 ap1->get_offset());
387 }
388
389 auto ap2 = piecewise_2D_cast(this->arg);
390 if (ap2.get()) {
391 return piecewise_2D(this->evaluate(),
392 ap2->get_num_columns(),
393 ap2->get_left(), ap2->get_x_scale(), ap2->get_x_offset(),
394 ap2->get_right(), ap2->get_y_scale(), ap2->get_y_offset());
395 }
396
397// Reduce exp(log(x)) -> x
398 auto a = log_cast(this->arg);
399 if (a.get()) {
400 return a->get_arg();
401 }
402
403 return this->shared_from_this();
404 }
405
406//------------------------------------------------------------------------------
413//------------------------------------------------------------------------------
415 if (this->is_match(x)) {
416 return one<T, SAFE_MATH> ();
417 }
418
419 const size_t hash = reinterpret_cast<size_t> (x.get());
420 if (this->df_cache.find(hash) == this->df_cache.end()) {
421 this->df_cache[hash] = this->shared_from_this()*this->arg->df(x);
422 }
423 return this->df_cache[hash];
424 }
425
426//------------------------------------------------------------------------------
434//------------------------------------------------------------------------------
436 compile(std::ostringstream &stream,
437 jit::register_map &registers,
439 const jit::register_usage &usage) {
440 if (registers.find(this) == registers.end()) {
441 shared_leaf<T, SAFE_MATH> a = this->arg->compile(stream,
442 registers,
443 indices,
444 usage);
445
446 registers[this] = jit::to_string('r', this);
447 stream << " const ";
448 jit::add_type<T> (stream);
449 stream << " " << registers[this] << " = ";
450 if constexpr (SAFE_MATH) {
451 if constexpr (jit::complex_scalar<T>) {
452 stream << "real(";
453 }
454 stream << registers[a.get()];
455 if constexpr (jit::complex_scalar<T>) {
456 stream << ")";
457 }
458 stream << " < 709.8 ? ";
459 }
460 stream << "exp(" << registers[a.get()] << ")";
461 if constexpr (SAFE_MATH) {
462 stream << " : ";
463 if constexpr (jit::complex_scalar<T>) {
464 jit::add_type<T> (stream);
465 stream << "(";
466 }
468 if constexpr (jit::complex_scalar<T>) {
469 stream << ")";
470 }
471 }
472 stream << "";
473 this->endline(stream, usage);
474 }
475
476 return this->shared_from_this();
477 }
478
479//------------------------------------------------------------------------------
484//------------------------------------------------------------------------------
486 if (this == x.get()) {
487 return true;
488 }
489
490 auto x_cast = exp_cast(x);
491 if (x_cast.get()) {
492 return this->arg->is_match(x_cast->get_arg());
493 }
494
495 return false;
496 }
497
498//------------------------------------------------------------------------------
500//------------------------------------------------------------------------------
501 virtual void to_latex() const {
502 std::cout << "e^{\\left(";
503 this->arg->to_latex();
504 std::cout << "\\right)}";
505 }
506
507//------------------------------------------------------------------------------
511//------------------------------------------------------------------------------
513 if (this->has_pseudo()) {
514 return exp(this->arg->remove_pseudo());
515 }
516 return this->shared_from_this();
517 }
518
519//------------------------------------------------------------------------------
525//------------------------------------------------------------------------------
526 virtual shared_leaf<T, SAFE_MATH> to_vizgraph(std::stringstream &stream,
527 jit::register_map &registers) {
528 if (registers.find(this) == registers.end()) {
529 const std::string name = jit::to_string('r', this);
530 registers[this] = name;
531 stream << " " << name
532 << " [label = \"exp\", shape = oval, style = filled, fillcolor = blue, fontcolor = white];" << std::endl;
533
534 auto a = this->arg->to_vizgraph(stream, registers);
535 stream << " " << name << " -- " << registers[a.get()] << ";" << std::endl;
536 }
537
538 return this->shared_from_this();
539 }
540 };
541
542//------------------------------------------------------------------------------
550//------------------------------------------------------------------------------
551 template<jit::float_scalar T, bool SAFE_MATH=false>
553 auto temp = std::make_shared<exp_node<T, SAFE_MATH>> (x)->reduce();
554// Test for hash collisions.
555 for (size_t i = temp->get_hash();
557 if (leaf_node<T, SAFE_MATH>::caches.nodes.find(i) ==
560 return temp;
561 } else if (temp->is_match(leaf_node<T, SAFE_MATH>::caches.nodes[i])) {
563 }
564 }
565#if defined(__clang__) || defined(__GNUC__)
567#else
568 assert(false && "Should never reach.");
569#endif
570 }
571
573 template<jit::float_scalar T, bool SAFE_MATH=false>
574 using shared_exp = std::shared_ptr<exp_node<T, SAFE_MATH>>;
575
576//------------------------------------------------------------------------------
584//------------------------------------------------------------------------------
585 template<jit::float_scalar T, bool SAFE_MATH=false>
587 return std::dynamic_pointer_cast<exp_node<T, SAFE_MATH>> (x);
588 }
589
590//******************************************************************************
591// Log node.
592//******************************************************************************
593//------------------------------------------------------------------------------
600//------------------------------------------------------------------------------
601 template<jit::float_scalar T, bool SAFE_MATH=false>
602 class log_node final : public straight_node<T, SAFE_MATH> {
603 private:
604//------------------------------------------------------------------------------
609//------------------------------------------------------------------------------
610 static std::string to_string(leaf_node<T, SAFE_MATH> *a) {
611 return "log" + jit::format_to_string(reinterpret_cast<size_t> (a));
612 }
613
614 public:
615//------------------------------------------------------------------------------
619//------------------------------------------------------------------------------
622
623//------------------------------------------------------------------------------
629//------------------------------------------------------------------------------
631 backend::buffer<T> result = this->arg->evaluate();
632 result.log();
633 return result;
634 }
635
636//------------------------------------------------------------------------------
640//------------------------------------------------------------------------------
642 if (constant_cast(this->arg).get()) {
643 return constant<T, SAFE_MATH> (this->evaluate());
644 }
645
646 auto ap1 = piecewise_1D_cast(this->arg);
647 if (ap1.get()) {
648 return piecewise_1D(this->evaluate(),
649 ap1->get_arg(),
650 ap1->get_scale(),
651 ap1->get_offset());
652 }
653
654 auto ap2 = piecewise_2D_cast(this->arg);
655 if (ap2.get()) {
656 return piecewise_2D(this->evaluate(),
657 ap2->get_num_columns(),
658 ap2->get_left(), ap2->get_x_scale(), ap2->get_x_offset(),
659 ap2->get_right(), ap2->get_y_scale(), ap2->get_y_offset());
660 }
661
662// Reduce log(exp(x)) -> x
663 auto a = exp_cast(this->arg);
664 if (a.get()) {
665 return a->get_arg();
666 }
667
668 return this->shared_from_this();
669 }
670
671//------------------------------------------------------------------------------
678//------------------------------------------------------------------------------
680 if (this->is_match(x)) {
681 return one<T, SAFE_MATH> ();
682 }
683
684 const size_t hash = reinterpret_cast<size_t> (x.get());
685 if (this->df_cache.find(hash) == this->df_cache.end()) {
686 this->df_cache[hash] = this->arg->df(x)/this->arg;
687 }
688 return this->df_cache[hash];
689 }
690
691//------------------------------------------------------------------------------
699//------------------------------------------------------------------------------
701 compile(std::ostringstream &stream,
702 jit::register_map &registers,
704 const jit::register_usage &usage) {
705 if (registers.find(this) == registers.end()) {
706 shared_leaf<T, SAFE_MATH> a = this->arg->compile(stream,
707 registers,
708 indices,
709 usage);
710
711 registers[this] = jit::to_string('r', this);
712 stream << " const ";
713 jit::add_type<T> (stream);
714 stream << " " << registers[this] << " = log("
715 << registers[a.get()] << ")";
716 this->endline(stream, usage);
717 }
718
719 return this->shared_from_this();
720 }
721
722//------------------------------------------------------------------------------
727//------------------------------------------------------------------------------
729 if (this == x.get()) {
730 return true;
731 }
732
733 auto x_cast = log_cast(x);
734 if (x_cast.get()) {
735 return this->arg->is_match(x_cast->get_arg());
736 }
737
738 return false;
739 }
740
741//------------------------------------------------------------------------------
743//------------------------------------------------------------------------------
744 virtual void to_latex() const {
745 std::cout << "\\ln{\\left(";
746 this->arg->to_latex();
747 std::cout << "\\right)}";
748 }
749
750//------------------------------------------------------------------------------
754//------------------------------------------------------------------------------
756 if (this->has_pseudo()) {
757 return log(this->arg->remove_pseudo());
758 }
759 return this->shared_from_this();
760 }
761
762//------------------------------------------------------------------------------
768//------------------------------------------------------------------------------
769 virtual shared_leaf<T, SAFE_MATH> to_vizgraph(std::stringstream &stream,
770 jit::register_map &registers) {
771 if (registers.find(this) == registers.end()) {
772 const std::string name = jit::to_string('r', this);
773 registers[this] = name;
774 stream << " " << name
775 << " [label = \"log\", shape = oval, style = filled, fillcolor = blue, fontcolor = white];" << std::endl;
776
777 auto a = this->arg->to_vizgraph(stream, registers);
778 stream << " " << name << " -- " << registers[a.get()] << ";" << std::endl;
779 }
780
781 return this->shared_from_this();
782 }
783 };
784
785//------------------------------------------------------------------------------
793//------------------------------------------------------------------------------
794 template<jit::float_scalar T, bool SAFE_MATH=false>
796 auto temp = std::make_shared<log_node<T, SAFE_MATH>> (x)->reduce();
797// Test for hash collisions.
798 for (size_t i = temp->get_hash(); i < std::numeric_limits<size_t>::max(); i++) {
799 if (leaf_node<T, SAFE_MATH>::caches.nodes.find(i) ==
802 return temp;
803 } else if (temp->is_match(leaf_node<T, SAFE_MATH>::caches.nodes[i])) {
805 }
806 }
807#if defined(__clang__) || defined(__GNUC__)
809#else
810 assert(false && "Should never reach.");
811#endif
812 }
813
815 template<jit::float_scalar T, bool SAFE_MATH=false>
816 using shared_log = std::shared_ptr<log_node<T, SAFE_MATH>>;
817
818//------------------------------------------------------------------------------
826//------------------------------------------------------------------------------
827 template<jit::float_scalar T, bool SAFE_MATH=false>
829 return std::dynamic_pointer_cast<log_node<T, SAFE_MATH>> (x);
830 }
831
832//******************************************************************************
833// Pow node.
834//******************************************************************************
835//------------------------------------------------------------------------------
842//------------------------------------------------------------------------------
843 template<jit::float_scalar T, bool SAFE_MATH=false>
844 class pow_node final : public branch_node<T, SAFE_MATH> {
845 private:
846//------------------------------------------------------------------------------
852//------------------------------------------------------------------------------
853 static std::string to_string(leaf_node<T, SAFE_MATH> *l,
855 return "pow" + jit::format_to_string(reinterpret_cast<size_t> (l))
856 + jit::format_to_string(reinterpret_cast<size_t> (r));
857 }
858
859 public:
860//------------------------------------------------------------------------------
865//------------------------------------------------------------------------------
870
871//------------------------------------------------------------------------------
877//------------------------------------------------------------------------------
879 backend::buffer<T> l_result = this->left->evaluate();
880 backend::buffer<T> r_result = this->right->evaluate();
881 return backend::pow(l_result, r_result);
882 }
883
884//------------------------------------------------------------------------------
888//------------------------------------------------------------------------------
890 auto lc = constant_cast(this->left);
891 auto rc = constant_cast(this->right);
892
893 if (rc.get() && rc->is(0)) {
894 return one<T, SAFE_MATH> ();
895 } else if (rc.get() && rc->is(1)) {
896 return this->left;
897 } else if (rc.get() && rc->is(0.5)) {
898 return sqrt(this->left);
899 } else if (rc.get() && rc->is(2)){
900 auto sq = sqrt_cast(this->left);
901 if (sq.get()) {
902 return sq->get_arg();
903 }
904 }
905
906 if (lc.get() && rc.get()) {
907 return constant<T, SAFE_MATH> (this->evaluate());
908 }
909
910 auto pl1 = piecewise_1D_cast(this->left);
911 auto pr1 = piecewise_1D_cast(this->right);
912 if (pl1.get() && (rc.get() || pl1->is_arg_match(this->right))) {
913 return piecewise_1D(this->evaluate(), pl1->get_arg(),
914 pl1->get_scale(), pl1->get_offset());
915 } else if (pr1.get() && (lc.get() || pr1->is_arg_match(this->left))) {
916 return piecewise_1D(this->evaluate(), pr1->get_arg(),
917 pr1->get_scale(), pr1->get_offset());
918 }
919
920 auto pl2 = piecewise_2D_cast(this->left);
921 auto pr2 = piecewise_2D_cast(this->right);
922 if (pl2.get() && (rc.get() || pl2->is_arg_match(this->right))) {
923 return piecewise_2D(this->evaluate(),
924 pl2->get_num_columns(),
925 pl2->get_left(), pl2->get_x_scale(), pl2->get_x_offset(),
926 pl2->get_right(), pl2->get_y_scale(), pl2->get_y_offset());
927 } else if (pr2.get() && (lc.get() || pr2->is_arg_match(this->left))) {
928 return piecewise_2D(this->evaluate(),
929 pr2->get_num_columns(),
930 pr2->get_left(), pr2->get_x_scale(), pr2->get_x_offset(),
931 pr2->get_right(), pr2->get_y_scale(), pr2->get_y_offset());
932 }
933
934// Combine 2D and 1D piecewise constants if a row or column matches.
935 if (pr2.get() && pr2->is_row_match(this->left)) {
936 backend::buffer<T> result = pl1->evaluate();
937 result.pow_row(pr2->evaluate());
938 return piecewise_2D(result,
939 pr2->get_num_columns(),
940 pr2->get_left(), pr2->get_x_scale(), pr2->get_x_offset(),
941 pr2->get_right(), pr2->get_y_scale(), pr2->get_y_offset());
942 } else if (pr2.get() && pr2->is_col_match(this->left)) {
943 backend::buffer<T> result = pl1->evaluate();
944 result.pow_col(pr2->evaluate());
945 return piecewise_2D(result,
946 pr2->get_num_columns(),
947 pr2->get_left(), pr2->get_x_scale(), pr2->get_x_offset(),
948 pr2->get_right(), pr2->get_y_scale(), pr2->get_y_offset());
949 } else if (pl2.get() && pl2->is_row_match(this->right)) {
950 backend::buffer<T> result = pl2->evaluate();
951 result.pow_row(pr1->evaluate());
952 return piecewise_2D(result,
953 pl2->get_num_columns(),
954 pl2->get_left(), pl2->get_x_scale(), pl2->get_x_offset(),
955 pl2->get_right(), pl2->get_y_scale(), pl2->get_y_offset());
956 } else if (pl2.get() && pl2->is_col_match(this->right)) {
957 backend::buffer<T> result = pl2->evaluate();
958 result.pow_col(pr1->evaluate());
959 return piecewise_2D(result,
960 pl2->get_num_columns(),
961 pl2->get_left(), pl2->get_x_scale(), pl2->get_x_offset(),
962 pl2->get_right(), pl2->get_y_scale(), pl2->get_y_offset());
963 }
964
965 auto lp = pow_cast(this->left);
966// Only run this reduction if the right is an integer constant value.
967 if (lp.get() && rc.get() && rc->is_integer()) {
968 return pow(lp->get_left(), lp->get_right()*this->right);
969 }
970
971// Handle cases where (c*x)^a, (x*c)^a, (a*sqrt(b))^c and (a*b^c)^2.
972// These reductions only make sense if the power is constant.
973 auto lm = multiply_cast(this->left);
974 if (lm.get() && rc.get()) {
975 if (lm->get_left()->is_constant() ||
976 lm->get_right()->is_constant() ||
977 sqrt_cast(lm->get_left()).get() ||
978 sqrt_cast(lm->get_right()).get() ||
979 pow_cast(lm->get_left()).get() ||
980 pow_cast(lm->get_right()).get()) {
981 return pow(lm->get_left(), this->right) *
982 pow(lm->get_right(), this->right);
983 }
984
985// ((Sqrt(a)*b)*c)^d -> a^(d/2)*(b*c)^d
986// ((b*Sqrt(a))*c)^d -> a^(d/2)*(b*c)^d
987 auto lmlm = multiply_cast(lm->get_left());
988 if (lmlm.get()) {
989 if (lmlm->get_left()->is_constant() ||
990 lmlm->get_right()->is_constant() ||
991 sqrt_cast(lmlm->get_left()).get() ||
992 sqrt_cast(lmlm->get_right()).get() ||
993 pow_cast(lmlm->get_left()).get() ||
994 pow_cast(lmlm->get_right()).get()) {
995 return pow(lmlm->get_left(), this->right) *
996 pow(lmlm->get_right(), this->right) *
997 pow(lm->get_right(), this->right);
998 }
999 }
1000 }
1001
1002// These reductions only make sense if the power is constant.
1003 auto ld = divide_cast(this->left);
1004 if (ld.get() && rc.get()) {
1005// For even exponents e.
1006// (-a/b)^e -> (a/b)^e
1007 auto ldlm = multiply_cast(ld->get_left());
1008 if (ldlm.get()) {
1009 if (rc.get() &&
1010 rc->evaluate().is_even()) {
1011 if (ldlm->get_left()->is_constant()) {
1012 return pow(ldlm->get_left(), this->right) *
1013 pow(ldlm->get_right()/ld->get_right(),
1014 this->right);
1015 }
1016 }
1017 if (ldlm->get_left()->is_constant() ||
1018 ldlm->get_right()->is_constant() ||
1019 sqrt_cast(ldlm->get_left()).get() ||
1020 sqrt_cast(ldlm->get_right()).get() ||
1021 pow_cast(ldlm->get_left()).get() ||
1022 pow_cast(ldlm->get_right()).get()) {
1023 return pow(ldlm->get_left(), this->right) *
1024 pow(ldlm->get_right(), this->right)/
1025 pow(ld->get_right(), this->right);
1026 }
1027
1028 auto ldlmlm = multiply_cast(ldlm->get_left());
1029 if (ldlmlm.get()) {
1030 if (ldlmlm->get_left()->is_constant() ||
1031 ldlmlm->get_right()->is_constant() ||
1032 sqrt_cast(ldlmlm->get_left()).get() ||
1033 sqrt_cast(ldlmlm->get_right()).get() ||
1034 pow_cast(ldlmlm->get_left()).get() ||
1035 pow_cast(ldlmlm->get_right()).get()) {
1036 return (pow(ldlmlm->get_left(), this->right) *
1037 pow(ldlmlm->get_right(), this->right) *
1038 pow(ldlm->get_right(), this->right)) /
1039 pow(ld->get_right(), this->right);
1040 }
1041 }
1042
1043 auto ldlmrm = multiply_cast(ldlm->get_right());
1044 if (ldlmrm.get()) {
1045 if (ldlmrm->get_left()->is_constant() ||
1046 ldlmrm->get_right()->is_constant() ||
1047 sqrt_cast(ldlmrm->get_left()).get() ||
1048 sqrt_cast(ldlmrm->get_right()).get() ||
1049 pow_cast(ldlmrm->get_left()).get() ||
1050 pow_cast(ldlmrm->get_right()).get()) {
1051 return (pow(ldlmrm->get_left(), this->right) *
1052 pow(ldlmrm->get_right(), this->right) *
1053 pow(ldlm->get_left(), this->right)) /
1054 pow(ld->get_right(), this->right);
1055 }
1056 }
1057 }
1058
1059// Handle cases where (c/x)^a, (x/c)^a, (a/sqrt(b))^c and (a/b^c)^2.
1060 if (ld->get_left()->is_constant() ||
1061 ld->get_right()->is_constant() ||
1062 sqrt_cast(ld->get_left()).get() ||
1063 sqrt_cast(ld->get_right()).get() ||
1064 pow_cast(ld->get_left()).get() ||
1065 pow_cast(ld->get_right()).get()) {
1066 return pow(ld->get_left(), this->right) /
1067 pow(ld->get_right(), this->right);
1068 }
1069
1070// Handle cases where (a/(b*sqrt(c))), (a/(sqrt(c)*b)), (a/(b*c^d)), (a/(c^d*b))
1071 auto ldrm = multiply_cast(ld->get_right());
1072 if (ldrm.get()) {
1073 if (ldrm->get_left()->is_constant() ||
1074 ldrm->get_right()->is_constant() ||
1075 sqrt_cast(ldrm->get_left()).get() ||
1076 sqrt_cast(ldrm->get_right()).get() ||
1077 pow_cast(ldrm->get_left()).get() ||
1078 pow_cast(ldrm->get_right()).get()) {
1079 return pow(ld->get_left(), this->right) /
1080 (pow(ldrm->get_left(), this->right) *
1081 pow(ldrm->get_right(), this->right));
1082 }
1083
1084 auto ldrmlm = multiply_cast(ldrm->get_left());
1085 if (ldrmlm.get()) {
1086 if (ldrmlm->get_left()->is_constant() ||
1087 ldrmlm->get_right()->is_constant() ||
1088 sqrt_cast(ldrmlm->get_left()).get() ||
1089 sqrt_cast(ldrmlm->get_right()).get() ||
1090 pow_cast(ldrmlm->get_left()).get() ||
1091 pow_cast(ldrmlm->get_right()).get()) {
1092 return pow(ld->get_left(), this->right) /
1093 (pow(ldrmlm->get_left(), this->right) *
1094 pow(ldrmlm->get_right(), this->right) *
1095 pow(ldrm->get_right(), this->right));
1096 }
1097 }
1098
1099 auto ldrmrm = multiply_cast(ldrm->get_right());
1100 if (ldrmrm.get()) {
1101 if (ldrmrm->get_left()->is_constant() ||
1102 ldrmrm->get_right()->is_constant() ||
1103 sqrt_cast(ldrmrm->get_left()).get() ||
1104 sqrt_cast(ldrmrm->get_right()).get() ||
1105 pow_cast(ldrmrm->get_left()).get() ||
1106 pow_cast(ldrmrm->get_right()).get()) {
1107 return pow(ld->get_left(), this->right) /
1108 (pow(ldrmrm->get_left(), this->right) *
1109 pow(ldrmrm->get_right(), this->right) *
1110 pow(ldrm->get_left(), this->right));
1111 }
1112 }
1113 }
1114
1115 if (is_variable_combinable(ld->get_left(),
1116 ld->get_right())) {
1117 return pow(ld->get_left()->get_power_base(),
1118 this->right*(ld->get_left()->get_power_exponent() -
1119 ld->get_right()->get_power_exponent()));
1120 }
1121
1122 if (ldrm.get()) {
1123 auto ldrmlm = multiply_cast(ldrm->get_left());
1124 if (ldrmlm.get()) {
1125 if (is_variable_combinable(ldrm->get_right(),
1126 ldrmlm->get_right()->get_power_base())) {
1127 return pow(ld->get_left()/ldrmlm->get_left(),
1128 this->right) /
1129 pow(ldrm->get_right()*ldrmlm->get_right(),
1130 this->right);
1131 } else if (is_variable_combinable(ldrm->get_right(),
1132 ldrmlm->get_left()->get_power_base())) {
1133 return pow(ld->get_left()/ldrmlm->get_right(),
1134 this->right) /
1135 pow(ldrm->get_right()*ldrmlm->get_left(),
1136 this->right);
1137 } else if (is_variable_combinable(ldrmlm->get_left(),
1138 ldrmlm->get_right()->get_power_base()) ||
1139 is_variable_combinable(ldrmlm->get_right(),
1140 ldrmlm->get_left()->get_power_base())) {
1141 return pow(ld->get_left()/ldrm->get_right(),
1142 this->right) /
1143 pow(ldrmlm->get_left()*ldrmlm->get_right(),
1144 this->right);
1145 }
1146 }
1147 }
1148 }
1149
1150// Reduce sqrt(a)^b
1151 auto lsq = sqrt_cast(this->left);
1152 if (lsq.get()) {
1153 return pow(lsq->get_arg(),
1154 this->right/2.0);
1155 }
1156
1157// Reduce exp(x)^n -> exp(n*x) when x is an integer.
1158 auto temp = exp_cast(this->left);
1159 if (temp.get() && rc.get() && rc->is_integer()) {
1160 return exp(this->right*temp->get_arg());
1161 }
1162
1163 return this->shared_from_this();
1164 }
1165
1166//------------------------------------------------------------------------------
1173//------------------------------------------------------------------------------
1176 if (this->is_match(x)) {
1177 return one<T, SAFE_MATH> ();
1178 }
1179
1180 const size_t hash = reinterpret_cast<size_t> (x.get());
1181 if (this->df_cache.find(hash) == this->df_cache.end()) {
1182 this->df_cache[hash] = pow(this->left, this->right - 1.0)
1183 * (this->right*this->left->df(x) +
1184 this->left*log(this->left)*this->right->df(x));
1185 }
1186 return this->df_cache[hash];
1187 }
1188
1189//------------------------------------------------------------------------------
1197//------------------------------------------------------------------------------
1199 compile(std::ostringstream &stream,
1200 jit::register_map &registers,
1202 const jit::register_usage &usage) {
1203 if (registers.find(this) == registers.end()) {
1204 shared_leaf<T, SAFE_MATH> l = this->left->compile(stream,
1205 registers,
1206 indices,
1207 usage);
1209 auto temp = constant_cast(this->right);
1210 if (!temp.get() || !temp->is_integer()) {
1211 r = this->right->compile(stream, registers, indices, usage);
1212 }
1213
1214 registers[this] = jit::to_string('r', this);
1215 stream << " const ";
1216 jit::add_type<T> (stream);
1217 stream << " " << registers[this] << " = ";
1218 if (temp.get() && temp->is_integer()) {
1219 stream << registers[l.get()];
1220 const size_t end = static_cast<size_t> (std::real(this->right->evaluate().at(0)));
1221 for (size_t i = 1; i < end; i++) {
1222 stream << "*" << registers[l.get()];
1223 }
1224 } else {
1225 stream << "pow("
1226 << registers[l.get()] << ", "
1227 << registers[r.get()] << ")";
1228 }
1229 this->endline(stream, usage);
1230 }
1231
1232 return this->shared_from_this();
1233 }
1234
1235//------------------------------------------------------------------------------
1240//------------------------------------------------------------------------------
1242 if (this == x.get()) {
1243 return true;
1244 }
1245
1246 auto x_cast = pow_cast(x);
1247 if (x_cast.get()) {
1248 return this->left->is_match(x_cast->get_left()) &&
1249 this->right->is_match(x_cast->get_right());
1250 }
1251
1252 return false;
1253 }
1254
1255//------------------------------------------------------------------------------
1257//------------------------------------------------------------------------------
1258 virtual void to_latex() const {
1259 auto use_brackets = !constant_cast(this->left).get() &&
1260 !variable_cast(this->left).get();
1261
1262 if (use_brackets) {
1263 std::cout << "\\left(";
1264 }
1265 this->left->to_latex();
1266 if (use_brackets) {
1267 std::cout << "\\right)";
1268 }
1269 std::cout << "^{";
1270 this->right->to_latex();
1271 std::cout << "}";
1272 }
1273
1274//------------------------------------------------------------------------------
1280//------------------------------------------------------------------------------
1281 virtual shared_leaf<T, SAFE_MATH> to_vizgraph(std::stringstream &stream,
1282 jit::register_map &registers) {
1283 if (registers.find(this) == registers.end()) {
1284 const std::string name = jit::to_string('r', this);
1285 registers[this] = name;
1286 stream << " " << name
1287 << " [label = \"pow\", shape = oval, style = filled, fillcolor = blue, fontcolor = white];" << std::endl;
1288
1289 auto l = this->left->to_vizgraph(stream, registers);
1290 stream << " " << name << " -- " << registers[l.get()] << ";" << std::endl;
1291 auto r = this->right->to_vizgraph(stream, registers);
1292 stream << " " << name << " -- " << registers[r.get()] << ";" << std::endl;
1293 }
1294
1295 return this->shared_from_this();
1296 }
1297
1298//------------------------------------------------------------------------------
1302//------------------------------------------------------------------------------
1303 virtual bool is_all_variables() const {
1304 return this->left->is_all_variables() &&
1305 (this->right->is_all_variables() ||
1306 constant_cast(this->right).get());
1307 }
1308
1309//------------------------------------------------------------------------------
1313//------------------------------------------------------------------------------
1314 virtual bool is_power_like() const {
1315 return true;
1316 }
1317
1318//------------------------------------------------------------------------------
1322//------------------------------------------------------------------------------
1324 return this->left;
1325 }
1326
1327//------------------------------------------------------------------------------
1331//------------------------------------------------------------------------------
1333 return this->right;
1334 }
1335
1336//------------------------------------------------------------------------------
1340//------------------------------------------------------------------------------
1342 if (this->has_pseudo()) {
1343 return pow(this->left->remove_pseudo(),
1344 this->right->remove_pseudo());
1345 }
1346 return this->shared_from_this();
1347 }
1348 };
1349
1350//------------------------------------------------------------------------------
1358//------------------------------------------------------------------------------
1359 template<jit::float_scalar T, bool SAFE_MATH=false>
1362 auto temp = std::make_shared<pow_node<T, SAFE_MATH>> (l, r)->reduce();
1363// Test for hash collisions.
1364 for (size_t i = temp->get_hash();
1366 if (leaf_node<T, SAFE_MATH>::caches.nodes.find(i) ==
1367 leaf_node<T, SAFE_MATH>::caches.nodes.end()) {
1369 return temp;
1370 } else if (temp->is_match(leaf_node<T, SAFE_MATH>::caches.nodes[i])) {
1371 return leaf_node<T, SAFE_MATH>::caches.nodes[i];
1372 }
1373 }
1374#if defined(__clang__) || defined(__GNUC__)
1376#else
1377 assert(false && "Should never reach.");
1378#endif
1379 }
1380
1381//------------------------------------------------------------------------------
1390//------------------------------------------------------------------------------
1391 template<jit::float_scalar T, jit::float_scalar L, bool SAFE_MATH=false>
1394 return pow(constant<T, SAFE_MATH> (static_cast<T> (l)), r);
1395 }
1396
1397//------------------------------------------------------------------------------
1406//------------------------------------------------------------------------------
1407 template<jit::float_scalar T, jit::float_scalar R, bool SAFE_MATH=false>
1409 const R r) {
1410 return pow(l, constant<T, SAFE_MATH> (static_cast<T> (r)));
1411 }
1412
1414 template<jit::float_scalar T, bool SAFE_MATH=false>
1415 using shared_pow = std::shared_ptr<pow_node<T, SAFE_MATH>>;
1416
1417//------------------------------------------------------------------------------
1422//------------------------------------------------------------------------------
1423 template<jit::float_scalar T, bool SAFE_MATH=false>
1425 return std::dynamic_pointer_cast<pow_node<T, SAFE_MATH>> (x);
1426 }
1427
1428//******************************************************************************
1429// Erfi node.
1430//******************************************************************************
1431//------------------------------------------------------------------------------
1438//------------------------------------------------------------------------------
1439 template<jit::complex_scalar T, bool SAFE_MATH=false>
1440 class erfi_node final : public straight_node<T, SAFE_MATH> {
1441 private:
1442//------------------------------------------------------------------------------
1447//------------------------------------------------------------------------------
1448 static std::string to_string(leaf_node<T, SAFE_MATH> *a) {
1449 return "erfi" + jit::format_to_string(reinterpret_cast<size_t> (a));
1450 }
1451
1452 public:
1453//------------------------------------------------------------------------------
1457//------------------------------------------------------------------------------
1460
1461//------------------------------------------------------------------------------
1467//------------------------------------------------------------------------------
1469 backend::buffer<T> result = this->arg->evaluate();
1470 result.erfi();
1471 return result;
1472 }
1473
1474//------------------------------------------------------------------------------
1478//------------------------------------------------------------------------------
1480 if (constant_cast(this->arg).get()) {
1481 return constant<T, SAFE_MATH> (this->evaluate());
1482 }
1483
1484 auto ap1 = piecewise_1D_cast(this->arg);
1485 if (ap1.get()) {
1486 return piecewise_1D(this->evaluate(),
1487 ap1->get_arg(),
1488 ap1->get_scale(),
1489 ap1->get_offset());
1490 }
1491
1492 auto ap2 = piecewise_2D_cast(this->arg);
1493 if (ap2.get()) {
1494 return piecewise_2D(this->evaluate(),
1495 ap2->get_num_columns(),
1496 ap2->get_left(), ap2->get_x_scale(), ap2->get_x_offset(),
1497 ap2->get_right(), ap2->get_y_scale(), ap2->get_y_offset());
1498 }
1499
1500 return this->shared_from_this();
1501 }
1502
1503//------------------------------------------------------------------------------
1510//------------------------------------------------------------------------------
1512 if (this->is_match(x)) {
1513 return one<T, SAFE_MATH> ();
1514 }
1515
1516 const size_t hash = reinterpret_cast<size_t> (x.get());
1517 if (this->df_cache.find(hash) == this->df_cache.end()) {
1518 this->df_cache[hash] = 2.0/std::sqrt(M_PI)
1519 * exp(this->arg*this->arg)*this->arg->df(x);
1520 }
1521 return this->df_cache[hash];
1522 }
1523
1524//------------------------------------------------------------------------------
1532//------------------------------------------------------------------------------
1534 compile(std::ostringstream &stream,
1535 jit::register_map &registers,
1537 const jit::register_usage &usage) {
1538 if (registers.find(this) == registers.end()) {
1539 shared_leaf<T, SAFE_MATH> a = this->arg->compile(stream,
1540 registers,
1541 indices,
1542 usage);
1543
1544 registers[this] = jit::to_string('r', this);
1545 stream << " const ";
1546 jit::add_type<T> (stream);
1547 stream << " " << registers[this] << " = special::erfi("
1548 << registers[a.get()] << ")";
1549 this->endline(stream, usage);
1550 }
1551
1552 return this->shared_from_this();
1553 }
1554
1555//------------------------------------------------------------------------------
1560//------------------------------------------------------------------------------
1562 if (this == x.get()) {
1563 return true;
1564 }
1565
1566 auto x_cast = erfi_cast(x);
1567 if (x_cast.get()) {
1568 return this->arg->is_match(x_cast->get_arg());
1569 }
1570
1571 return false;
1572 }
1573
1574//------------------------------------------------------------------------------
1576//------------------------------------------------------------------------------
1577 virtual void to_latex() const {
1578 std::cout << "erfi\\left(";
1579 this->arg->to_latex();
1580 std::cout << "\\right)";
1581 }
1582
1583//------------------------------------------------------------------------------
1587//------------------------------------------------------------------------------
1589 if (this->has_pseudo()) {
1590 return erfi(this->arg->remove_pseudo());
1591 }
1592 return this->shared_from_this();
1593 }
1594
1595//------------------------------------------------------------------------------
1601//------------------------------------------------------------------------------
1602 virtual shared_leaf<T, SAFE_MATH> to_vizgraph(std::stringstream &stream,
1603 jit::register_map &registers) {
1604 if (registers.find(this) == registers.end()) {
1605 const std::string name = jit::to_string('r', this);
1606 registers[this] = name;
1607 stream << " " << name
1608 << " [label = \"erfi\", shape = oval, style = filled, fillcolor = blue, fontcolor = white];" << std::endl;
1609
1610 auto a = this->arg->to_vizgraph(stream, registers);
1611 stream << " " << name << " -- " << registers[a.get()] << ";" << std::endl;
1612 }
1613
1614 return this->shared_from_this();
1615 }
1616 };
1617
1618//------------------------------------------------------------------------------
1626//------------------------------------------------------------------------------
1627 template<jit::complex_scalar T, bool SAFE_MATH=false>
1629 auto temp = std::make_shared<erfi_node<T, SAFE_MATH>> (x)->reduce();
1630// Test for hash collisions.
1631 for (size_t i = temp->get_hash();
1633 if (leaf_node<T, SAFE_MATH>::caches.nodes.find(i) ==
1634 leaf_node<T, SAFE_MATH>::caches.nodes.end()) {
1636 return temp;
1637 } else if (temp->is_match(leaf_node<T, SAFE_MATH>::caches.nodes[i])) {
1638 return leaf_node<T, SAFE_MATH>::caches.nodes[i];
1639 }
1640 }
1641#if defined(__clang__) || defined(__GNUC__)
1643#else
1644 assert(false && "Should never reach.");
1645#endif
1646 }
1647
1649 template<jit::complex_scalar T, bool SAFE_MATH=false>
1650 using shared_erfi = std::shared_ptr<erfi_node<T, SAFE_MATH>>;
1651
1652//------------------------------------------------------------------------------
1660//------------------------------------------------------------------------------
1661 template<jit::complex_scalar T, bool SAFE_MATH=false>
1663 return std::dynamic_pointer_cast<erfi_node<T, SAFE_MATH>> (x);
1664 }
1665}
1666
1667#endif /* math_h */
Class representing a generic buffer.
Definition backend.hpp:29
void erfi()
Take erfi.
Definition backend.hpp:259
void log()
Take log.
Definition backend.hpp:232
void sqrt()
Take sqrt.
Definition backend.hpp:214
void pow_col(const buffer< T > &x)
Pow col operation.
Definition backend.hpp:747
void pow_row(const buffer< T > &x)
Pow row operation.
Definition backend.hpp:711
void exp()
Take exp.
Definition backend.hpp:223
Class representing a branch node.
Definition node.hpp:1165
shared_leaf< T, SAFE_MATH > right
Right branch of the tree.
Definition node.hpp:1170
shared_leaf< T, SAFE_MATH > left
Left branch of the tree.
Definition node.hpp:1168
An imaginary error function node.
Definition math.hpp:1440
erfi_node(shared_leaf< T, SAFE_MATH > x)
Construct a exp node.
Definition math.hpp:1458
virtual shared_leaf< T, SAFE_MATH > reduce()
Reduce the erfi(x).
Definition math.hpp:1479
virtual shared_leaf< T, SAFE_MATH > compile(std::ostringstream &stream, jit::register_map &registers, jit::register_map &indices, const jit::register_usage &usage)
Compile the node.
Definition math.hpp:1534
virtual void to_latex() const
Convert the node to latex.
Definition math.hpp:1577
virtual shared_leaf< T, SAFE_MATH > remove_pseudo()
Remove pseudo variable nodes.
Definition math.hpp:1588
virtual shared_leaf< T, SAFE_MATH > df(shared_leaf< T, SAFE_MATH > x)
Transform node to derivative.
Definition math.hpp:1511
virtual bool is_match(shared_leaf< T, SAFE_MATH > x)
Query if the nodes match.
Definition math.hpp:1561
virtual backend::buffer< T > evaluate()
Evaluate the results of erfi.
Definition math.hpp:1468
virtual shared_leaf< T, SAFE_MATH > to_vizgraph(std::stringstream &stream, jit::register_map &registers)
Convert the node to vizgraph.
Definition math.hpp:1602
A exp node.
Definition math.hpp:337
virtual backend::buffer< T > evaluate()
Evaluate the results of exp.
Definition math.hpp:365
virtual shared_leaf< T, SAFE_MATH > df(shared_leaf< T, SAFE_MATH > x)
Transform node to derivative.
Definition math.hpp:414
virtual bool is_match(shared_leaf< T, SAFE_MATH > x)
Query if the nodes match.
Definition math.hpp:485
virtual shared_leaf< T, SAFE_MATH > remove_pseudo()
Remove pseudo variable nodes.
Definition math.hpp:512
exp_node(shared_leaf< T, SAFE_MATH > x)
Construct a exp node.
Definition math.hpp:355
virtual void to_latex() const
Convert the node to latex.
Definition math.hpp:501
virtual shared_leaf< T, SAFE_MATH > reduce()
Reduce the exp(x).
Definition math.hpp:376
virtual shared_leaf< T, SAFE_MATH > to_vizgraph(std::stringstream &stream, jit::register_map &registers)
Convert the node to vizgraph.
Definition math.hpp:526
virtual shared_leaf< T, SAFE_MATH > compile(std::ostringstream &stream, jit::register_map &registers, jit::register_map &indices, const jit::register_usage &usage)
Compile the node.
Definition math.hpp:436
Class representing a node leaf.
Definition node.hpp:364
virtual void endline(std::ostringstream &stream, const jit::register_usage &usage) const final
End a line in the kernel source.
Definition node.hpp:639
std::map< size_t, std::shared_ptr< leaf_node< T, SAFE_MATH > > > df_cache
Cache derivative terms.
Definition node.hpp:371
virtual bool has_pseudo() const
Query if the node contains pseudo variables.
Definition node.hpp:620
const size_t hash
Hash for node.
Definition node.hpp:367
A log node.
Definition math.hpp:602
virtual shared_leaf< T, SAFE_MATH > to_vizgraph(std::stringstream &stream, jit::register_map &registers)
Convert the node to vizgraph.
Definition math.hpp:769
virtual shared_leaf< T, SAFE_MATH > df(shared_leaf< T, SAFE_MATH > x)
Transform node to derivative.
Definition math.hpp:679
virtual backend::buffer< T > evaluate()
Evaluate the results of log.
Definition math.hpp:630
virtual bool is_match(shared_leaf< T, SAFE_MATH > x)
Query if the nodes match.
Definition math.hpp:728
virtual shared_leaf< T, SAFE_MATH > reduce()
Reduce the log(x).
Definition math.hpp:641
virtual void to_latex() const
Convert the node to latex.
Definition math.hpp:744
virtual shared_leaf< T, SAFE_MATH > compile(std::ostringstream &stream, jit::register_map &registers, jit::register_map &indices, const jit::register_usage &usage)
Compile the node.
Definition math.hpp:701
log_node(shared_leaf< T, SAFE_MATH > x)
Construct a log node.
Definition math.hpp:620
virtual shared_leaf< T, SAFE_MATH > remove_pseudo()
Remove pseudo variable nodes.
Definition math.hpp:755
An power node.
Definition math.hpp:844
virtual shared_leaf< T, SAFE_MATH > to_vizgraph(std::stringstream &stream, jit::register_map &registers)
Convert the node to vizgraph.
Definition math.hpp:1281
virtual bool is_all_variables() const
Test if node acts like a variable.
Definition math.hpp:1303
virtual shared_leaf< T, SAFE_MATH > get_power_base()
Get the base of a power.
Definition math.hpp:1323
virtual backend::buffer< T > evaluate()
Evaluate the results of addition.
Definition math.hpp:878
pow_node(shared_leaf< T, SAFE_MATH > l, shared_leaf< T, SAFE_MATH > r)
Construct an power node.
Definition math.hpp:866
virtual bool is_power_like() const
Test if the node acts like a power of variable.
Definition math.hpp:1314
virtual shared_leaf< T, SAFE_MATH > compile(std::ostringstream &stream, jit::register_map &registers, jit::register_map &indices, const jit::register_usage &usage)
Compile the node.
Definition math.hpp:1199
virtual shared_leaf< T, SAFE_MATH > remove_pseudo()
Remove pseudo variable nodes.
Definition math.hpp:1341
virtual bool is_match(shared_leaf< T, SAFE_MATH > x)
Query if the nodes match.
Definition math.hpp:1241
virtual shared_leaf< T, SAFE_MATH > get_power_exponent() const
Get the exponent of a power.
Definition math.hpp:1332
virtual void to_latex() const
Convert the node to latex.
Definition math.hpp:1258
virtual shared_leaf< T, SAFE_MATH > reduce()
Reduce a power node.
Definition math.hpp:889
virtual shared_leaf< T, SAFE_MATH > df(shared_leaf< T, SAFE_MATH > x)
Transform node to derivative.
Definition math.hpp:1175
A sqrt node.
Definition math.hpp:26
virtual shared_leaf< T, SAFE_MATH > get_power_exponent() const
Get the exponent of a power.
Definition math.hpp:238
virtual backend::buffer< T > evaluate()
Evaluate the results of sqrt.
Definition math.hpp:54
virtual shared_leaf< T, SAFE_MATH > reduce()
Reduce the sqrt(x).
Definition math.hpp:65
sqrt_node(shared_leaf< T, SAFE_MATH > x)
Construct a sqrt node.
Definition math.hpp:44
virtual bool is_power_like() const
Test if the node acts like a power of variable.
Definition math.hpp:220
virtual bool is_match(shared_leaf< T, SAFE_MATH > x)
Query if the nodes match.
Definition math.hpp:193
virtual shared_leaf< T, SAFE_MATH > compile(std::ostringstream &stream, jit::register_map &registers, jit::register_map &indices, const jit::register_usage &usage)
Compile the node.
Definition math.hpp:166
virtual shared_leaf< T, SAFE_MATH > remove_pseudo()
Remove pseudo variable nodes.
Definition math.hpp:247
virtual shared_leaf< T, SAFE_MATH > df(shared_leaf< T, SAFE_MATH > x)
Transform node to derivative.
Definition math.hpp:143
virtual shared_leaf< T, SAFE_MATH > to_vizgraph(std::stringstream &stream, jit::register_map &registers)
Convert the node to vizgraph.
Definition math.hpp:261
virtual void to_latex() const
Convert the node to latex.
Definition math.hpp:209
virtual shared_leaf< T, SAFE_MATH > get_power_base()
Get the base of a power.
Definition math.hpp:229
Class representing a straight node.
Definition node.hpp:1051
shared_leaf< T, SAFE_MATH > arg
Argument.
Definition node.hpp:1054
Complex scalar concept.
Definition register.hpp:24
subroutine assert(test, message)
Assert check.
Definition f_binding_test.f90:38
buffer< T > pow(buffer< T > &base, buffer< T > &exponent)
Take the power.
Definition backend.hpp:1057
Name space for graph nodes.
Definition arithmetic.hpp:13
shared_piecewise_2D< T, SAFE_MATH > piecewise_2D_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a piecewise 2D node.
Definition piecewise.hpp:1428
shared_leaf< T, SAFE_MATH > log(shared_leaf< T, SAFE_MATH > x)
Define log convenience function.
Definition math.hpp:795
shared_pow< T, SAFE_MATH > pow_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a power node.
Definition math.hpp:1424
constexpr shared_leaf< T, SAFE_MATH > zero()
Forward declare for zero.
Definition node.hpp:986
shared_sqrt< T, SAFE_MATH > sqrt_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a sqrt node.
Definition math.hpp:321
shared_leaf< T, SAFE_MATH > pow(shared_leaf< T, SAFE_MATH > l, shared_leaf< T, SAFE_MATH > r)
Build power node.
Definition math.hpp:1360
shared_divide< T, SAFE_MATH > divide_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a divide node.
Definition arithmetic.hpp:3720
std::shared_ptr< erfi_node< T, SAFE_MATH > > shared_erfi
Convenience type alias for shared exp nodes.
Definition math.hpp:1650
shared_piecewise_1D< T, SAFE_MATH > piecewise_1D_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a piecewise 1D node.
Definition piecewise.hpp:637
shared_multiply< T, SAFE_MATH > multiply_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a multiply node.
Definition arithmetic.hpp:2755
shared_leaf< T, SAFE_MATH > exp(shared_leaf< T, SAFE_MATH > x)
Define exp convenience function.
Definition math.hpp:552
shared_exp< T, SAFE_MATH > exp_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a exp node.
Definition math.hpp:586
shared_constant< T, SAFE_MATH > constant_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a constant node.
Definition node.hpp:1034
shared_leaf< T, SAFE_MATH > erfi(shared_leaf< T, SAFE_MATH > x)
Define erfi convenience function.
Definition math.hpp:1628
bool is_variable_combinable(shared_leaf< T, SAFE_MATH > a, shared_leaf< T, SAFE_MATH > b)
Check if the variable is combinable.
Definition arithmetic.hpp:75
shared_variable< T, SAFE_MATH > variable_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a variable node.
Definition node.hpp:1727
constexpr T i
Convenience type for imaginary constant.
Definition node.hpp:1018
shared_erfi< T, SAFE_MATH > erfi_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a exp node.
Definition math.hpp:1662
shared_leaf< T, SAFE_MATH > sqrt(shared_leaf< T, SAFE_MATH > x)
Define sqrt convenience function.
Definition math.hpp:287
std::shared_ptr< exp_node< T, SAFE_MATH > > shared_exp
Convenience type alias for shared exp nodes.
Definition math.hpp:574
shared_log< T, SAFE_MATH > log_cast(shared_leaf< T, SAFE_MATH > x)
Cast to a exp node.
Definition math.hpp:828
std::shared_ptr< leaf_node< T, SAFE_MATH > > shared_leaf
Convenience type alias for shared leaf nodes.
Definition node.hpp:676
std::shared_ptr< log_node< T, SAFE_MATH > > shared_log
Convenience type alias for shared log nodes.
Definition math.hpp:816
std::shared_ptr< pow_node< T, SAFE_MATH > > shared_pow
Convenience type alias for shared add nodes.
Definition math.hpp:1415
std::shared_ptr< sqrt_node< T, SAFE_MATH > > shared_sqrt
Convenience type alias for shared sqrt nodes.
Definition math.hpp:309
std::string format_to_string(const T value)
Convert a value to a string while avoiding locale.
Definition register.hpp:212
std::map< void *, size_t > register_usage
Type alias for counting register usage.
Definition register.hpp:259
std::map< void *, std::string > register_map
Type alias for mapping node pointers to register names.
Definition register.hpp:257
std::string to_string(const char prefix, const NODE *pointer)
Convert a graph::leaf_node pointer to a string.
Definition register.hpp:246
Base nodes of graph computation framework.
void piecewise_1D()
Tests for 1D piecewise nodes.
Definition piecewise_test.cpp:80
void piecewise_2D()
Tests for 2D piecewise nodes.
Definition piecewise_test.cpp:319