26 #ifndef CASADI_MX_NODE_HPP
27 #define CASADI_MX_NODE_HPP
30 #include "shared_object.hpp"
31 #include "sx_elem.hpp"
32 #include "calculus.hpp"
33 #include "code_generator.hpp"
41 class SerializingStream;
42 class DeserializingStream;
66 virtual bool __nonzero__()
const;
71 virtual bool is_zero()
const {
return false;}
76 virtual bool is_one()
const {
return false;}
86 virtual bool is_half()
const {
return false;}
91 virtual bool is_inf()
const {
return false;}
111 virtual bool is_value(
double val)
const {
return false;}
116 virtual bool is_eye()
const {
return false;}
131 void can_inline(std::map<const MXNode*, casadi_int>& nodeind)
const;
136 std::string print_compact(std::map<const MXNode*, casadi_int>& nodeind,
137 std::vector<std::string>& intermed)
const;
142 virtual std::string
disp(
const std::vector<std::string>& arg)
const = 0;
168 const std::vector<casadi_int>& arg,
169 const std::vector<casadi_int>& res,
170 const std::vector<bool>& arg_is_ref,
171 std::vector<bool>& res_is_ref)
const;
174 const std::vector<casadi_int>& arg,
175 const std::vector<casadi_int>& res,
176 const std::vector<bool>& arg_is_ref,
177 std::vector<bool>& res_is_ref,
183 virtual int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const;
188 virtual int eval_sx(
const SXElem** arg,
SXElem** res, casadi_int* iw,
SXElem* w)
const;
193 virtual void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
194 const std::vector<bool>& unique={})
const;
199 virtual void eval_linear(
const std::vector<std::array<MX, 3> >& arg,
200 std::vector<std::array<MX, 3> >& res)
const;
206 std::vector<std::array<MX, 3> >& res)
const;
214 void eval_linear_rearrange(
const std::vector<std::array<MX, 3> >& arg,
215 std::vector<std::array<MX, 3> >& res)
const;
220 virtual void ad_forward(
const std::vector<std::vector<MX> >& fseed,
221 std::vector<std::vector<MX> >& fsens)
const;
226 virtual void ad_reverse(
const std::vector<std::vector<MX> >& aseed,
227 std::vector<std::vector<MX> >& asens)
const;
232 virtual int sp_forward(
const bvec_t** arg,
bvec_t** res, casadi_int* iw,
bvec_t* w)
const;
238 for (casadi_int k=0; k<nout(); ++k) {
241 for (casadi_int i=0; i<sparsity(k).nnz(); ++i) v[i] = ~
static_cast<bvec_t>(0);
254 virtual const std::string& name()
const;
259 std::string class_name()
const override;
264 void disp(std::ostream& stream,
bool more)
const override;
274 virtual casadi_int n_primitives()
const;
279 virtual void primitives(std::vector<MX>::iterator& it)
const;
285 virtual void split_primitives(
const MX& x, std::vector<MX>::iterator& it)
const;
286 virtual void split_primitives(
const SX& x, std::vector<SX>::iterator& it)
const;
287 virtual void split_primitives(
const DM& x, std::vector<DM>::iterator& it)
const;
292 T join_primitives_gen(
typename std::vector<T>::const_iterator& it)
const;
298 virtual MX join_primitives(std::vector<MX>::const_iterator& it)
const;
299 virtual SX join_primitives(std::vector<SX>::const_iterator& it)
const;
300 virtual DM join_primitives(std::vector<DM>::const_iterator& it)
const;
308 virtual bool has_duplicates()
const;
315 virtual void reset_input()
const;
330 virtual casadi_int which_output()
const;
335 virtual const Function& which_function()
const;
340 virtual casadi_int
op()
const = 0;
343 virtual Dict info()
const;
375 virtual bool is_equal(
const MXNode* node, casadi_int depth)
const {
return false;}
387 bool sameOpAndDeps(
const MXNode* node, casadi_int depth)
const;
392 const MX&
dep(casadi_int ind=0)
const {
return dep_.at(ind);}
397 casadi_int n_dep()
const;
402 virtual casadi_int
nout()
const {
return 1;}
407 virtual MX get_output(casadi_int oind)
const;
413 virtual const Sparsity& sparsity(casadi_int oind)
const;
417 for (casadi_int i=0;i<dep_.size();++i) {
418 if (dep_[i].sparsity()!=arg[i].sparsity()) {
426 casadi_int
numel()
const {
return sparsity().numel(); }
427 casadi_int
nnz(casadi_int i=0)
const {
return sparsity(i).nnz(); }
428 casadi_int
size1()
const {
return sparsity().size1(); }
429 casadi_int
size2()
const {
return sparsity().size2(); }
430 std::pair<casadi_int, casadi_int>
size()
const {
return sparsity().size();}
433 virtual casadi_int ind()
const;
436 virtual casadi_int segment()
const;
439 virtual casadi_int offset()
const;
442 void set_sparsity(
const Sparsity& sparsity);
447 virtual size_t sz_arg()
const {
return n_dep();}
452 virtual size_t sz_res()
const {
return nout();}
457 virtual size_t sz_iw()
const {
return 0;}
462 virtual size_t sz_w()
const {
return 0;}
474 void set_dep(
const MX& dep);
477 void set_dep(
const MX& dep1,
const MX& dep2);
480 void set_dep(
const MX& dep1,
const MX& dep2,
const MX& dep3);
483 void set_dep(
const std::vector<MX>& dep);
486 void check_dep()
const;
498 virtual double to_double()
const;
501 virtual casadi_int
to_int()
const;
504 virtual DM get_DM()
const;
513 virtual MX get_horzcat(
const std::vector<MX>& x)
const;
516 virtual std::vector<MX> get_horzsplit(
const std::vector<casadi_int>& output_offset)
const;
519 virtual MX get_repmat(casadi_int m, casadi_int n)
const;
522 virtual MX get_repsum(casadi_int m, casadi_int n)
const;
525 virtual MX get_kron(
const MX& b)
const;
528 virtual MX get_kron_contract(
const MX& x,
bool inner)
const;
531 virtual MX get_vertcat(
const std::vector<MX>& x)
const;
534 virtual std::vector<MX> get_vertsplit(
const std::vector<casadi_int>& output_offset)
const;
537 virtual MX get_diagcat(
const std::vector<MX>& x)
const;
540 virtual std::vector<MX> get_diagsplit(
const std::vector<casadi_int>& offset1,
541 const std::vector<casadi_int>& offset2)
const;
544 virtual MX get_transpose()
const;
547 virtual MX get_reshape(
const Sparsity& sp)
const;
550 virtual MX get_sparsity_cast(
const Sparsity& sp)
const;
555 virtual MX get_mac(
const MX& y,
const MX& z,
556 const std::string& blas =
"reference")
const;
561 virtual MX get_einstein(
const MX& A,
const MX& B,
562 const std::vector<casadi_int>& dim_c,
const std::vector<casadi_int>& dim_a,
563 const std::vector<casadi_int>& dim_b,
564 const std::vector<casadi_int>& c,
const std::vector<casadi_int>& a,
565 const std::vector<casadi_int>& b)
const;
570 virtual MX get_bilin(
const MX& x,
const MX& y)
const;
575 virtual MX get_rank1(
const MX& alpha,
const MX& x,
const MX& y)
const;
580 virtual MX get_logsumexp()
const;
589 virtual MX get_solve(
const MX& r,
bool tr,
const Linsol& linear_solver)
const;
598 virtual MX get_solve_triu(
const MX& r,
bool tr)
const;
607 virtual MX get_solve_tril(
const MX& r,
bool tr)
const;
616 virtual MX get_solve_triu_unity(
const MX& r,
bool tr)
const;
625 virtual MX get_solve_tril_unity(
const MX& r,
bool tr)
const;
634 virtual MX get_nzref(
const Sparsity& sp,
const std::vector<casadi_int>& nz,
635 bool unique=
false)
const;
640 virtual MX get_nz_ref(
const MX& nz)
const;
645 virtual MX get_nz_ref(
const MX& inner,
const Slice& outer)
const;
650 virtual MX get_nz_ref(
const Slice& inner,
const MX& outer)
const;
655 virtual MX get_nz_ref(
const MX& inner,
const MX& outer)
const;
663 virtual MX get_nzassign(
const MX& y,
const std::vector<casadi_int>& nz)
const;
671 virtual MX get_nzadd(
const MX& y,
const std::vector<casadi_int>& nz)
const;
679 virtual MX get_nzassign(
const MX& y,
const MX& nz)
const;
687 virtual MX get_nzassign(
const MX& y,
const MX& inner,
const Slice& outer)
const;
695 virtual MX get_nzassign(
const MX& y,
const Slice& inner,
const MX& outer)
const;
703 virtual MX get_nzassign(
const MX& y,
const MX& inner,
const MX& outer)
const;
711 virtual MX get_nzadd(
const MX& y,
const MX& nz)
const;
719 virtual MX get_nzadd(
const MX& y,
const MX& inner,
const Slice& outer)
const;
727 virtual MX get_nzadd(
const MX& y,
const Slice& inner,
const MX& outer)
const;
735 virtual MX get_nzadd(
const MX& y,
const MX& inner,
const MX& outer)
const;
738 virtual MX get_subref(
const Slice& i,
const Slice& j)
const;
741 virtual MX get_subassign(
const MX& y,
const Slice& i,
const Slice& j)
const;
744 virtual MX get_project(
const Sparsity& sp,
bool unique=
false)
const;
747 virtual MX get_unary(casadi_int op,
bool unique=
false)
const;
750 MX get_binary(casadi_int op,
const MX& y,
bool unique_x=
false,
bool unique_y=
false)
const;
753 virtual MX _get_binary(casadi_int op,
const MX& y,
bool scX,
bool scY,
754 bool unique_x=
false,
bool unique_y=
false)
const;
757 virtual MX get_det(
const Linsol& linear_solver)
const;
760 virtual MX get_inv()
const;
763 virtual MX get_dot(
const MX& y)
const;
766 virtual MX get_norm_fro()
const;
769 virtual MX get_norm_2()
const;
772 virtual MX get_norm_inf()
const;
775 virtual MX get_norm_1()
const;
778 virtual MX get_mmin()
const;
781 virtual MX get_mmax()
const;
784 MX get_assert(
const MX& y,
const std::string& fail_message)
const;
787 MX get_monitor(
const std::string& comment)
const;
790 MX get_dump(
const std::string& base_filename,
const Dict& opts)
const;
796 MX get_low(
const MX& v,
const Dict& options)
const;
799 MX get_bspline(
const std::vector<double>& knots,
800 const std::vector<casadi_int>& offset,
801 const std::vector<double>& coeffs,
802 const std::vector<casadi_int>& degree,
804 const std::vector<casadi_int>& lookup_mode)
const;
806 MX get_bspline(
const MX& coeffs,
const std::vector<double>& knots,
807 const std::vector<casadi_int>& offset,
808 const std::vector<casadi_int>& degree,
810 const std::vector<casadi_int>& lookup_mode)
const;
813 MX get_convexify(
const Dict& opts)
const;
834 static void copy_fwd(
const bvec_t* arg,
bvec_t* res, casadi_int len);
839 static void copy_rev(
bvec_t* arg,
bvec_t* res, casadi_int len);
Helper class for C code generation.
Helper class for Serialization.
std::pair< casadi_int, casadi_int > size() const
Get the shape.
Node class for MX objects.
virtual bool has_output() const
Check if a multiple output node.
void eval_linear_unary(const std::vector< std::array< MX, 3 > > &arg, std::vector< std::array< MX, 3 > > &res) const
Evaluate the MX node on a const/linear/nonlinear partition.
virtual bool is_zero() const
Check if identically zero.
virtual size_t sz_arg() const
Get required length of arg field.
virtual int eval_activity(const bvec_t **arg, bvec_t **res, casadi_int *iw, bvec_t *w) const
Propagate signal activity forward (bit set = active)
virtual size_t codegen_sz_w() const
Length of w the node's GENERATED code needs (may exceed sz_w)
virtual bool is_valid_input() const
Check if valid function input.
static bool maxDepth()
Get equality checking depth.
virtual bool is_minus_inf() const
Check if identically -inf.
virtual size_t sz_w() const
Get required length of w field.
virtual bool is_one() const
Check if identically one.
virtual bool is_binary() const
Check if binary operation.
virtual void add_dependency(CodeGenerator &g) const
Add a dependent function.
virtual casadi_int n_inplace() const
Can the operation be performed inplace (i.e. overwrite the result)
std::pair< casadi_int, casadi_int > size() const
Sparsity sparsity_
The sparsity pattern.
casadi_int numel() const
Get shape.
static std::map< casadi_int, MXNode *(*)(DeserializingStream &)> deserialize_map
virtual bool is_nonnegative() const
Check if not negative.
const Sparsity & sparsity() const
Get the sparsity.
virtual size_t sz_res() const
Get required length of res field.
casadi_int nnz(casadi_int i=0) const
bool matches_sparsity(const std::vector< T > &arg) const
virtual void codegen_incref(CodeGenerator &g, std::set< void * > &added) const
Codegen incref.
virtual casadi_int nout() const
Number of outputs.
virtual bool is_value(double val) const
Check if a certain value.
const MX & dep(casadi_int ind=0) const
dependencies - functions that have to be evaluated before this one
std::vector< MX > dep_
dependencies - functions that have to be evaluated before this one
virtual bool is_unary() const
Check if unary operation.
virtual void codegen_decref(CodeGenerator &g, std::set< void * > &added) const
Codegen decref.
virtual casadi_int op() const =0
Get the operation.
virtual bool is_inf() const
Check if identically inf.
virtual bool has_refcount() const
Is reference counting needed in codegen?
virtual bool is_minus_one() const
Check if identically minus one.
virtual bool is_integer() const
Check if integer.
virtual std::string disp(const std::vector< std::string > &arg) const =0
Print expression.
virtual bool is_output() const
Check if evaluation output.
virtual bool is_equal(const MXNode *node, casadi_int depth) const
virtual size_t sz_iw() const
Get required length of iw field.
virtual bool is_eye() const
Check if identity matrix.
static MX to_matrix(const MX &x, const Sparsity &sp)
Convert scalar to matrix.
virtual bool is_half() const
Check if identically 0.5.
static casadi_int get_max_depth()
Get the depth to which equalities are being checked for simplifications.
Sparse matrix class. SX and DM are specializations.
The basic scalar symbolic class of CasADi.
Helper class for Serialization.
Class representing a Slice.
std::pair< casadi_int, casadi_int > size() const
Get the shape.
bool is_equal(double x, double y, casadi_int depth=0)
unsigned long long bvec_t
int to_int(casadi_int rhs)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.