13#ifndef dealii_differentiation_sd_symengine_number_visitor_internal_h
14#define dealii_differentiation_sd_symengine_number_visitor_internal_h
18#ifdef DEAL_II_WITH_SYMENGINE
26# include <boost/serialization/split_member.hpp>
28# include <symengine/basic.h>
29# include <symengine/dict.h>
30# include <symengine/symengine_exception.h>
31# include <symengine/symengine_rcp.h>
32# include <symengine/visitor.h>
38#ifdef DEAL_II_WITH_SYMENGINE
52 template <
typename ReturnType,
typename ExpressionType>
56 std::vector<std::pair<SD::Expression, SD::Expression>>;
122 call(ReturnType *output_values,
124 const ReturnType *substitution_values);
131 template <
class Archive>
133 save(Archive &archive,
const unsigned int version)
const;
140 template <
class Archive>
142 load(Archive &archive,
const unsigned int version);
150 template <
class Archive>
152 serialize(Archive &archive,
const unsigned int version);
156 BOOST_SERIALIZATION_SPLIT_MEMBER()
164 template <
typename StreamType>
200 init(
const SymEngine::vec_basic &dependent_functions);
227 call(ReturnType *output_values,
228 const SymEngine::vec_basic &independent_symbols,
229 const ReturnType *substitution_values);
260 template <
typename ReturnType,
typename ExpressionType>
262 :
public SymEngine::BaseVisitor<
263 DictionarySubstitutionVisitor<ReturnType, ExpressionType>>
302 const bool use_cse =
false);
329 const SymEngine::Basic &dependent_function,
330 const bool use_cse =
false);
355 const bool use_cse =
false);
383 const bool use_cse =
false);
418 call(ReturnType *output_values,
const ReturnType *substitution_values);
438 call(
const std::vector<ReturnType> &substitution_values);
445 template <
class Archive>
447 save(Archive &archive,
const unsigned int version)
const;
454 template <
class Archive>
456 load(Archive &archive,
const unsigned int version);
464 template <
class Archive>
466 serialize(Archive &archive,
const unsigned int version);
470 BOOST_SERIALIZATION_SPLIT_MEMBER()
486 template <
typename StreamType>
489 const bool print_independent_symbols =
false,
490 const bool print_dependent_functions =
false,
491 const bool print_cse_reductions =
false)
const;
499# define IMPLEMENT_DSV_BVISIT(Argument) \
500 void bvisit(const Argument &) \
502 AssertThrow(false, ExcNotImplemented()); \
505 IMPLEMENT_DSV_BVISIT(SymEngine::Basic)
506 IMPLEMENT_DSV_BVISIT(SymEngine::Symbol)
507 IMPLEMENT_DSV_BVISIT(SymEngine::Constant)
508 IMPLEMENT_DSV_BVISIT(SymEngine::Integer)
509 IMPLEMENT_DSV_BVISIT(SymEngine::Rational)
510 IMPLEMENT_DSV_BVISIT(SymEngine::RealDouble)
511 IMPLEMENT_DSV_BVISIT(SymEngine::ComplexDouble)
512 IMPLEMENT_DSV_BVISIT(SymEngine::Add)
513 IMPLEMENT_DSV_BVISIT(SymEngine::Mul)
514 IMPLEMENT_DSV_BVISIT(SymEngine::Pow)
515 IMPLEMENT_DSV_BVISIT(SymEngine::Log)
516 IMPLEMENT_DSV_BVISIT(SymEngine::Sin)
517 IMPLEMENT_DSV_BVISIT(SymEngine::Cos)
518 IMPLEMENT_DSV_BVISIT(SymEngine::Tan)
519 IMPLEMENT_DSV_BVISIT(SymEngine::Csc)
520 IMPLEMENT_DSV_BVISIT(SymEngine::Sec)
521 IMPLEMENT_DSV_BVISIT(SymEngine::Cot)
522 IMPLEMENT_DSV_BVISIT(SymEngine::ASin)
523 IMPLEMENT_DSV_BVISIT(SymEngine::ACos)
524 IMPLEMENT_DSV_BVISIT(SymEngine::ATan)
525 IMPLEMENT_DSV_BVISIT(SymEngine::ATan2)
526 IMPLEMENT_DSV_BVISIT(SymEngine::ACsc)
527 IMPLEMENT_DSV_BVISIT(SymEngine::ASec)
528 IMPLEMENT_DSV_BVISIT(SymEngine::ACot)
529 IMPLEMENT_DSV_BVISIT(SymEngine::Sinh)
530 IMPLEMENT_DSV_BVISIT(SymEngine::Cosh)
531 IMPLEMENT_DSV_BVISIT(SymEngine::Tanh)
532 IMPLEMENT_DSV_BVISIT(SymEngine::Csch)
533 IMPLEMENT_DSV_BVISIT(SymEngine::Sech)
534 IMPLEMENT_DSV_BVISIT(SymEngine::Coth)
535 IMPLEMENT_DSV_BVISIT(SymEngine::ASinh)
536 IMPLEMENT_DSV_BVISIT(SymEngine::ACosh)
537 IMPLEMENT_DSV_BVISIT(SymEngine::ATanh)
538 IMPLEMENT_DSV_BVISIT(SymEngine::ACsch)
539 IMPLEMENT_DSV_BVISIT(SymEngine::ACoth)
540 IMPLEMENT_DSV_BVISIT(SymEngine::ASech)
541 IMPLEMENT_DSV_BVISIT(SymEngine::Abs)
542 IMPLEMENT_DSV_BVISIT(SymEngine::Gamma)
543 IMPLEMENT_DSV_BVISIT(SymEngine::LogGamma)
544 IMPLEMENT_DSV_BVISIT(SymEngine::Erf)
545 IMPLEMENT_DSV_BVISIT(SymEngine::Erfc)
546 IMPLEMENT_DSV_BVISIT(SymEngine::Max)
547 IMPLEMENT_DSV_BVISIT(SymEngine::Min)
549# undef IMPLEMENT_DSV_BVISIT
587 template <
typename ReturnType,
typename ExpressionType>
593 dependent_functions));
598 template <
typename ReturnType,
typename ExpressionType>
601 const SymEngine::vec_basic &dependent_functions)
619 SymEngine::vec_pair se_replacements;
620 SymEngine::vec_basic se_reduced_exprs;
621 SymEngine::cse(se_replacements, se_reduced_exprs, dependent_functions);
623 intermediate_symbols_exprs =
632 template <
typename ReturnType,
typename ExpressionType>
635 ReturnType *output_values,
637 const ReturnType *substitution_values)
641 independent_symbols),
642 substitution_values);
647 template <
typename ReturnType,
typename ExpressionType>
650 ReturnType *output_values,
651 const SymEngine::vec_basic &independent_symbols,
652 const ReturnType *substitution_values)
657 SymEngine::map_basic_basic substitution_value_map;
658 for (
unsigned i = 0; i < independent_symbols.size(); ++i)
659 substitution_value_map[independent_symbols[i]] =
660 static_cast<const SymEngine::RCP<const SymEngine::Basic> &
>(
661 ExpressionType(substitution_values[i]));
665 for (
const auto &expression : intermediate_symbols_exprs)
667 const SymEngine::RCP<const SymEngine::Basic> &cse_symbol =
669 const SymEngine::RCP<const SymEngine::Basic> &cse_expr =
671 Assert(substitution_value_map.find(cse_symbol) ==
672 substitution_value_map.end(),
674 "Reduced symbol already appears in substitution map. "
675 "Is there a clash between the reduced symbol name and "
676 "the symbol used for an independent variable?"));
677 substitution_value_map[cse_symbol] =
678 static_cast<const SymEngine::RCP<const SymEngine::Basic> &
>(
679 ExpressionType(ExpressionType(cse_expr)
680 .
template substitute_and_evaluate<ReturnType>(
681 substitution_value_map)));
685 for (
unsigned i = 0; i < reduced_exprs.size(); ++i)
686 output_values[i] = ExpressionType(reduced_exprs[i])
687 .template substitute_and_evaluate<ReturnType>(
688 substitution_value_map);
693 template <
typename ReturnType,
typename ExpressionType>
694 template <
class Archive>
698 const unsigned int )
const
702 ar &intermediate_symbols_exprs;
708 template <
typename ReturnType,
typename ExpressionType>
709 template <
class Archive>
720 ar &intermediate_symbols_exprs;
726 template <
typename ReturnType,
typename ExpressionType>
727 template <
typename StreamType>
730 StreamType &stream)
const
732 stream <<
"Common subexpression elimination: \n";
733 stream <<
" Intermediate reduced expressions: \n";
734 for (
unsigned i = 0; i < intermediate_symbols_exprs.size(); ++i)
736 const SymEngine::RCP<const SymEngine::Basic> &cse_symbol =
737 intermediate_symbols_exprs[i].first;
738 const SymEngine::RCP<const SymEngine::Basic> &cse_expr =
739 intermediate_symbols_exprs[i].second;
740 stream <<
" " << i <<
": " << cse_symbol <<
" = " << cse_expr
744 stream <<
" Final reduced expressions for dependent variables: \n";
745 for (
unsigned i = 0; i < reduced_exprs.size(); ++i)
746 stream <<
" " << i <<
": " << reduced_exprs[i] <<
'\n';
748 stream << std::flush;
753 template <
typename ReturnType,
typename ExpressionType>
761 return (n_reduced_expressions() > 0) ||
762 (n_intermediate_expressions() > 0);
767 template <
typename ReturnType,
typename ExpressionType>
769 CSEDictionaryVisitor<ReturnType,
770 ExpressionType>::n_intermediate_expressions()
const
772 return intermediate_symbols_exprs.size();
777 template <
typename ReturnType,
typename ExpressionType>
782 return reduced_exprs.size();
790 template <
typename ReturnType,
typename ExpressionType>
794 const SD::Expression &output,
802 template <
typename ReturnType,
typename ExpressionType>
805 const SymEngine::vec_basic &inputs,
806 const SymEngine::Basic &output,
810 SD::Expression(output.rcp_from_this()),
816 template <
typename ReturnType,
typename ExpressionType>
819 const SymEngine::vec_basic &inputs,
820 const SymEngine::vec_basic &outputs,
830 template <
typename ReturnType,
typename ExpressionType>
837 independent_symbols.clear();
838 dependent_functions.clear();
840 independent_symbols = inputs;
847 if (use_cse ==
false)
848 dependent_functions = outputs;
857 template <
typename ReturnType,
typename ExpressionType>
860 const std::vector<ReturnType> &substitution_values)
863 dependent_functions.size() == 1,
865 "Cannot use this call function when more than one symbolic expression is to be evaluated."));
867 substitution_values.size() == independent_symbols.size(),
869 "Input substitution vector does not match size of symbol vector."));
872 call(&out, substitution_values.data());
878 template <
typename ReturnType,
typename ExpressionType>
881 ReturnType *output_values,
882 const ReturnType *substitution_values)
887 cse.call(output_values, independent_symbols, substitution_values);
892 SymEngine::map_basic_basic substitution_value_map;
893 for (
unsigned i = 0; i < independent_symbols.size(); ++i)
894 substitution_value_map[independent_symbols[i]] =
895 static_cast<const SymEngine::RCP<const SymEngine::Basic> &
>(
896 ExpressionType(substitution_values[i]));
903 for (
unsigned i = 0; i < dependent_functions.size(); ++i)
905 ExpressionType(dependent_functions[i])
906 .template substitute_and_evaluate<ReturnType>(
907 substitution_value_map);
913 template <
typename ReturnType,
typename ExpressionType>
914 template <
class Archive>
918 const unsigned int version)
const
929 ar &independent_symbols;
930 cse.save(ar, version);
931 ar &dependent_functions;
936 template <
typename ReturnType,
typename ExpressionType>
937 template <
class Archive>
941 const unsigned int version)
951 ar &independent_symbols;
952 cse.load(ar, version);
953 ar &dependent_functions;
958 template <
typename ReturnType,
typename ExpressionType>
959 template <
typename StreamType>
963 const bool print_independent_symbols,
964 const bool print_dependent_functions,
965 const bool print_cse_reductions)
const
967 if (print_independent_symbols)
969 stream <<
"Independent variables: \n";
970 for (
unsigned i = 0; i < independent_symbols.size(); ++i)
971 stream <<
" " << i <<
": " << independent_symbols[i] <<
'\n';
973 stream << std::flush;
977 if (print_cse_reductions && cse.executed())
985 if (print_dependent_functions)
987 stream <<
"Dependent variables: \n";
988 for (
unsigned i = 0; i < dependent_functions.size(); ++i)
989 stream <<
" " << i << dependent_functions[i] <<
'\n';
991 stream << std::flush;
CSEDictionaryVisitor()=default
void serialize(Archive &archive, const unsigned int version)
unsigned int n_intermediate_expressions() const
void call(ReturnType *output_values, const SymEngine::vec_basic &independent_symbols, const ReturnType *substitution_values)
symbol_vector_pair intermediate_symbols_exprs
void call(ReturnType *output_values, const types::symbol_vector &independent_symbols, const ReturnType *substitution_values)
types::symbol_vector reduced_exprs
void save(Archive &archive, const unsigned int version) const
void print(StreamType &stream) const
void init(const SymEngine::vec_basic &dependent_functions)
virtual ~CSEDictionaryVisitor()=default
void init(const types::symbol_vector &dependent_functions)
void load(Archive &archive, const unsigned int version)
unsigned int n_reduced_expressions() const
std::vector< std::pair< SD::Expression, SD::Expression > > symbol_vector_pair
ReturnType call(const std::vector< ReturnType > &substitution_values)
void init(const types::symbol_vector &independent_symbols, const Expression &dependent_function, const bool use_cse=false)
void init(const SymEngine::vec_basic &independent_symbols, const SymEngine::Basic &dependent_function, const bool use_cse=false)
SD::types::symbol_vector dependent_functions
void init(const types::symbol_vector &independent_symbols, const types::symbol_vector &dependent_functions, const bool use_cse=false)
void call(ReturnType *output_values, const ReturnType *substitution_values)
DictionarySubstitutionVisitor()=default
void print(StreamType &stream, const bool print_independent_symbols=false, const bool print_dependent_functions=false, const bool print_cse_reductions=false) const
CSEDictionaryVisitor< ReturnType, ExpressionType > cse
virtual ~DictionarySubstitutionVisitor() override=default
void save(Archive &archive, const unsigned int version) const
void serialize(Archive &archive, const unsigned int version)
void init(const SymEngine::vec_basic &independent_symbols, const SymEngine::vec_basic &dependent_functions, const bool use_cse=false)
SD::types::symbol_vector independent_symbols
void load(Archive &archive, const unsigned int version)
#define DEAL_II_NAMESPACE_OPEN
#define DEAL_II_NAMESPACE_CLOSE
#define Assert(cond, exc)
static ::ExceptionBase & ExcInternalError()
static ::ExceptionBase & ExcMessage(std::string arg1)
std::vector< std::pair< Expression, Expression > > convert_basic_pair_vector_to_expression_pair_vector(const SymEngine::vec_pair &symbol_value_vector)
SD::types::symbol_vector convert_basic_vector_to_expression_vector(const SymEngine::vec_basic &symbol_vector)
SymEngine::vec_basic convert_expression_vector_to_basic_vector(const SD::types::symbol_vector &symbol_vector)
std::vector< SD::Expression > symbol_vector
static constexpr const T & value(const T &t)