26 #ifndef CASADI_CONSTANT_MX_HPP
27 #define CASADI_CONSTANT_MX_HPP
29 #include "mx_node.hpp"
32 #include "serializing_stream.hpp"
48 class CASADI_EXPORT ConstantMX :
public MXNode {
51 explicit ConstantMX(
const Sparsity& sp);
54 ~ConstantMX()
override = 0;
57 static ConstantMX* create(
const Sparsity& sp, casadi_int val);
58 static ConstantMX* create(
const Sparsity& sp,
int val) {
59 return create(sp,
static_cast<casadi_int
>(val));
63 static ConstantMX* create(
const Sparsity& sp,
double val);
66 static ConstantMX* create(
const Matrix<double>& val);
69 static ConstantMX* create(
const Sparsity& sp,
const std::string& fname);
72 static ConstantMX* create(
const Matrix<double>& val,
const std::string& name);
75 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override = 0;
78 int eval_sx(
const SXElem** arg, SXElem** res,
79 casadi_int* iw, SXElem* w)
const override = 0;
84 void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
85 const std::vector<bool>& unique={})
const override;
90 void ad_forward(
const std::vector<std::vector<MX> >& fseed,
91 std::vector<std::vector<MX> >& fsens)
const override;
96 void ad_reverse(
const std::vector<std::vector<MX> >& aseed,
97 std::vector<std::vector<MX> >& asens)
const override;
102 int sp_forward(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override;
107 int eval_activity(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override;
112 void nonzeros_to_activity(
const double* v, bvec_t* res)
const;
117 int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override;
122 casadi_int op()
const override {
return OP_CONST;}
125 double to_double()
const override = 0;
128 casadi_int to_int()
const override = 0;
131 Matrix<double> get_DM()
const override = 0;
137 MX get_dot(
const MX& y)
const override;
140 bool __nonzero__()
const override;
145 bool is_valid_input()
const override;
150 casadi_int n_primitives()
const override;
155 void primitives(std::vector<MX>::iterator& it)
const override;
159 void split_primitives_gen(
const T& x,
typename std::vector<T>::iterator& it)
const;
165 void split_primitives(
const MX& x, std::vector<MX>::iterator& it)
const override;
166 void split_primitives(
const SX& x, std::vector<SX>::iterator& it)
const override;
167 void split_primitives(
const DM& x, std::vector<DM>::iterator& it)
const override;
172 T join_primitives_gen(
typename std::vector<T>::const_iterator& it)
const;
178 MX join_primitives(std::vector<MX>::const_iterator& it)
const override;
179 SX join_primitives(std::vector<SX>::const_iterator& it)
const override;
180 DM join_primitives(std::vector<DM>::const_iterator& it)
const override;
186 bool has_duplicates()
const override {
return false;}
191 void reset_input()
const override {}
196 static MXNode* deserialize(DeserializingStream& s);
201 explicit ConstantMX(DeserializingStream& s) : MXNode(s) {}
205 class CASADI_EXPORT ConstantDM :
public ConstantMX {
211 explicit ConstantDM(
const Matrix<double>& x) : ConstantMX(x.sparsity()), x_(x) {}
214 ~ConstantDM()
override {}
219 std::string disp(
const std::vector<std::string>& arg)
const override {
226 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
227 std::copy(x_->begin(), x_->end(), res[0]);
234 int eval_sx(
const SXElem** arg, SXElem** res,
235 casadi_int* iw, SXElem* w)
const override {
236 std::copy(x_->begin(), x_->end(), res[0]);
243 void generate(CodeGenerator& g,
244 const std::vector<casadi_int>& arg,
245 const std::vector<casadi_int>& res,
246 const std::vector<bool>& arg_is_ref,
247 std::vector<bool>& res_is_ref)
const override;
253 bool is_one()
const override;
254 bool is_minus_one()
const override;
255 bool is_inf()
const override;
256 bool is_minus_inf()
const override;
257 bool is_half()
const override;
258 bool is_value(
double val)
const override;
259 bool is_nonnegative()
const override;
260 bool is_integer()
const override;
261 bool is_eye()
const override;
264 double to_double()
const override {
return x_.scalar();}
267 casadi_int to_int()
const override {
return static_cast<casadi_int
>(x_.scalar());}
270 Matrix<double> get_DM()
const override {
return x_;}
275 int eval_activity(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override {
276 nonzeros_to_activity(x_->data(), res[0]);
283 bool is_equal(
const MXNode* node, casadi_int depth)
const override;
293 void serialize_body(SerializingStream& s)
const override;
297 void serialize_type(SerializingStream& s)
const override;
302 explicit ConstantDM(DeserializingStream& s);
306 class CASADI_EXPORT ConstantFile :
public ConstantMX {
312 explicit ConstantFile(
const Sparsity& x,
const std::string& fname);
315 ~ConstantFile()
override {}
320 bool has_refcount()
const override {
return true; }
325 void codegen_incref(CodeGenerator& g, std::set<void*>& added)
const override;
330 std::string disp(
const std::vector<std::string>& arg)
const override;
333 double to_double()
const override;
336 casadi_int to_int()
const override;
339 Matrix<double> get_DM()
const override;
344 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
345 std::copy(x_.begin(), x_.end(), res[0]);
352 int eval_sx(
const SXElem** arg, SXElem** res,
353 casadi_int* iw, SXElem* w)
const override {
354 std::copy(x_.begin(), x_.end(), res[0]);
361 int eval_activity(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override {
362 nonzeros_to_activity(x_.data(), res[0]);
369 void generate(CodeGenerator& g,
370 const std::vector<casadi_int>& arg,
371 const std::vector<casadi_int>& res,
372 const std::vector<bool>& arg_is_ref,
373 std::vector<bool>& res_is_ref)
const override;
378 void add_dependency(CodeGenerator& g)
const override;
388 std::vector<double> x_;
393 void serialize_body(SerializingStream& s)
const override;
397 void serialize_type(SerializingStream& s)
const override;
402 explicit ConstantFile(DeserializingStream& s);
406 class CASADI_EXPORT ConstantPool :
public ConstantMX {
412 explicit ConstantPool(
const DM& x,
const std::string& name);
415 ~ConstantPool()
override {}
420 std::string disp(
const std::vector<std::string>& arg)
const override;
423 double to_double()
const override;
426 casadi_int to_int()
const override;
429 Matrix<double> get_DM()
const override;
434 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
435 if (res[0]) std::copy(x_.begin(), x_.end(), res[0]);
442 int eval_activity(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override {
443 nonzeros_to_activity(x_.data(), res[0]);
450 int eval_sx(
const SXElem** arg, SXElem** res,
451 casadi_int* iw, SXElem* w)
const override {
452 casadi_error(
"eval_sx not supported");
459 void generate(CodeGenerator& g,
460 const std::vector<casadi_int>& arg,
461 const std::vector<casadi_int>& res,
462 const std::vector<bool>& arg_is_ref,
463 std::vector<bool>& res_is_ref)
const override;
468 void add_dependency(CodeGenerator& g)
const override;
478 std::vector<double> x_;
483 void serialize_body(SerializingStream& s)
const override;
488 void serialize_type(SerializingStream& s)
const override;
493 explicit ConstantPool(DeserializingStream& s);
497 class CASADI_EXPORT ZeroByZero :
public ConstantMX {
502 explicit ZeroByZero() : ConstantMX(Sparsity(0, 0)) {
510 static ZeroByZero* getInstance() {
511 static ZeroByZero instance;
516 ~ZeroByZero()
override {
523 std::string disp(
const std::vector<std::string>& arg)
const override;
529 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
534 int eval_sx(
const SXElem** arg, SXElem** res,
535 casadi_int* iw, SXElem* w)
const override {
542 void generate(CodeGenerator& g,
543 const std::vector<casadi_int>& arg,
544 const std::vector<casadi_int>& res,
545 const std::vector<bool>& arg_is_ref,
546 std::vector<bool>& res_is_ref)
const override {}
549 double to_double()
const override {
return 0;}
552 casadi_int to_int()
const override {
return 0;}
555 DM get_DM()
const override {
return DM(); }
558 MX get_project(
const Sparsity& sp,
bool unique=
false)
const override;
561 MX get_nzref(
const Sparsity& sp,
const std::vector<casadi_int>& nz,
562 bool unique=
false)
const override;
565 MX get_nzassign(
const MX& y,
const std::vector<casadi_int>& nz)
const override;
568 MX get_transpose()
const override;
571 MX get_unary(casadi_int op,
bool unique)
const override;
574 MX _get_binary(casadi_int op,
const MX& y,
bool ScX,
bool ScY,
575 bool unique_x=
false,
bool unique_y=
false)
const override;
578 MX get_reshape(
const Sparsity& sp)
const override;
583 bool is_valid_input()
const override {
return true;}
588 const std::string& name()
const override {
589 static std::string dummyname;
596 void serialize_type(SerializingStream& s)
const override;
600 void serialize_body(SerializingStream& s)
const override;
608 struct RuntimeConst {
611 RuntimeConst(T v) : value(v) {}
612 static char type_char();
613 void serialize_type(SerializingStream& s)
const {
614 s.pack(
"Constant::value", value);
616 static RuntimeConst deserialize(DeserializingStream& s) {
618 s.unpack(
"Constant::value", v);
619 return RuntimeConst(v);
624 inline char RuntimeConst<T>::type_char() {
return 'u'; }
627 inline char RuntimeConst<casadi_int>::type_char() {
return 'I'; }
630 inline char RuntimeConst<double>::type_char() {
return 'D'; }
633 struct CompiletimeConst {
634 static const int value = v;
635 static char type_char();
636 void serialize_type(SerializingStream& s)
const {}
637 static CompiletimeConst deserialize(DeserializingStream& s) {
638 return CompiletimeConst();
643 inline char CompiletimeConst<v>::type_char() {
return 'u'; }
646 inline char CompiletimeConst<0>::type_char() {
return '0'; }
648 inline char CompiletimeConst<(-1)>::type_char() {
return 'm'; }
650 inline char CompiletimeConst<1>::type_char() {
return '1'; }
653 template<
typename Value>
654 class CASADI_EXPORT Constant :
public ConstantMX {
660 explicit Constant(
const Sparsity& sp, Value v = Value()) : ConstantMX(sp), v_(v) {}
665 explicit Constant(DeserializingStream& s,
const Value& v);
668 ~Constant()
override {}
673 std::string disp(
const std::vector<std::string>& arg)
const override;
679 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override;
682 int eval_sx(
const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w)
const override;
687 int eval_activity(
const bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w)
const override {
688 std::fill_n(res[0], nnz(), v_.value!=0 ? ~
static_cast<bvec_t
>(0) : 0);
695 void generate(CodeGenerator& g,
696 const std::vector<casadi_int>& arg,
697 const std::vector<casadi_int>& res,
698 const std::vector<bool>& arg_is_ref,
699 std::vector<bool>& res_is_ref)
const override;
705 bool is_one()
const override;
706 bool is_minus_one()
const override;
707 bool is_half()
const override;
708 bool is_inf()
const override;
709 bool is_minus_inf()
const override;
710 bool is_nonnegative()
const override;
711 bool is_integer()
const override;
712 bool is_eye()
const override;
713 bool is_value(
double val)
const override;
716 double to_double()
const override {
717 return static_cast<double>(v_.value);
721 casadi_int to_int()
const override {
722 return static_cast<casadi_int
>(v_.value);
726 Matrix<double> get_DM()
const override {
727 return Matrix<double>(sparsity(), to_double(),
false);
731 MX get_project(
const Sparsity& sp,
bool unique=
false)
const override;
734 MX get_nzref(
const Sparsity& sp,
const std::vector<casadi_int>& nz,
735 bool unique=
false)
const override;
738 MX get_nzassign(
const MX& y,
const std::vector<casadi_int>& nz)
const override;
741 MX get_transpose()
const override;
744 MX get_unary(casadi_int op,
bool unique=
false)
const override;
747 MX _get_binary(casadi_int op,
const MX& y,
bool ScX,
bool ScY,
748 bool unique_x=
false,
bool unique_y=
false)
const override;
751 MX get_reshape(
const Sparsity& sp)
const override;
754 MX get_horzcat(
const std::vector<MX>& x)
const override;
757 MX get_vertcat(
const std::vector<MX>& x)
const override;
762 bool is_equal(
const MXNode* node, casadi_int depth)
const override;
767 void serialize_body(SerializingStream& s)
const override;
771 void serialize_type(SerializingStream& s)
const override;
776 template<
typename Value>
777 bool Constant<Value>::is_zero()
const {
781 template<
typename Value>
782 bool Constant<Value>::is_one()
const {
783 return sparsity().is_dense() && v_.value==1;
786 template<
typename Value>
787 bool Constant<Value>::is_minus_one()
const {
788 return sparsity().is_dense() && v_.value==-1;
791 template<
typename Value>
792 bool Constant<Value>::is_half()
const {
793 return sparsity().is_dense() && v_.value==0.5;
796 template<
typename Value>
797 bool Constant<Value>::is_inf()
const {
801 template<
typename Value>
802 bool Constant<Value>::is_minus_inf()
const {
806 template<
typename Value>
807 bool Constant<Value>::is_nonnegative()
const {
811 template<
typename Value>
812 bool Constant<Value>::is_integer()
const {
816 template<
typename Value>
817 bool Constant<Value>::is_eye()
const {
818 return v_.value==1 && sparsity().is_diag();
821 template<
typename Value>
822 bool Constant<Value>::is_value(
double val)
const {
824 return sparsity().is_dense() && v_.value==val;
827 template<
typename Value>
828 void Constant<Value>::serialize_type(SerializingStream& s)
const {
830 s.pack(
"ConstantMX::type", Value::type_char());
831 v_.serialize_type(s);
834 template<
typename Value>
835 void Constant<Value>::serialize_body(SerializingStream& s)
const {
839 template<
typename Value>
840 Constant<Value>::Constant(DeserializingStream& s,
const Value& v) : ConstantMX(s), v_(v) {
843 template<
typename Value>
844 MX Constant<Value>::get_horzcat(
const std::vector<MX>& x)
const {
847 if (!i->is_value(to_double())) {
849 return ConstantMX::get_horzcat(x);
854 std::vector<Sparsity> sp;
855 for (
auto&& i : x) sp.push_back(i.sparsity());
856 return MX(horzcat(sp), v_.value,
false);
859 template<
typename Value>
860 MX Constant<Value>::get_vertcat(
const std::vector<MX>& x)
const {
863 if (!i->is_value(to_double())) {
865 return ConstantMX::get_vertcat(x);
870 std::vector<Sparsity> sp;
871 for (
auto&& i : x) sp.push_back(i.sparsity());
872 return MX(vertcat(sp), v_.value,
false);
875 template<
typename Value>
876 MX Constant<Value>::get_reshape(
const Sparsity& sp)
const {
877 return MX::create(
new Constant<Value>(sp, v_));
880 template<
typename Value>
881 MX Constant<Value>::get_transpose()
const {
882 return MX::create(
new Constant<Value>(sparsity().
T(), v_));
885 template<
typename Value>
886 MX Constant<Value>::get_unary(casadi_int op,
bool unique)
const {
889 casadi_math<double>::fun(op, to_double(), 0.0, ret);
890 if (operation_checker<F0XChecker>(op) || sparsity().is_dense()) {
891 return MX(sparsity(), ret);
894 if (
is_zero() && operation_checker<F0XChecker>(op)) {
895 return MX(sparsity(), ret,
false);
897 return repmat(MX(ret), size1(), size2());
901 casadi_math<double>::fun(op, 0, 0.0, ret2);
902 return DM(sparsity(), ret,
false)
903 +
DM(sparsity().pattern_inverse(), ret2,
false);
907 template<
typename Value>
908 MX Constant<Value>::_get_binary(casadi_int op,
const MX& y,
bool ScX,
bool ScY,
909 bool unique_x,
bool unique_y)
const {
910 casadi_assert_dev(sparsity()==y.sparsity() || ScX || ScY);
912 if (ScX && !operation_checker<FX0Checker>(op)) {
914 casadi_math<double>::fun(op, nnz()> 0 ? to_double(): 0.0, 0, ret);
917 Sparsity f = Sparsity::dense(y.size1(), y.size2());
918 MX yy = project(y, f);
919 return MX(f, shared_from_this<MX>())->_get_binary(op, yy,
false,
false, unique_x, unique_y);
921 }
else if (ScY && !operation_checker<F0XChecker>(op)) {
923 if (y->op()==OP_CONST &&
dynamic_cast<const ConstantDM*
>(y.get())==
nullptr) {
925 casadi_math<double>::fun(op, 0, y.nnz()>0 ? y->to_double() : 0, ret);
929 Sparsity f = Sparsity::dense(size1(), size2());
930 MX xx = project(shared_from_this<MX>(), f);
931 return xx->_get_binary(op, MX(f, y),
false,
false, unique_x, unique_y);
937 if (v_.value==0)
return ScY && !y->is_zero() ? repmat(y, size1(), size2()) : y;
940 if (v_.value==0)
return ScY && !y->is_zero() ? repmat(-y, size1(), size2()) : -y;
943 if (v_.value==1)
return y;
944 if (v_.value==-1)
return -y;
945 if (v_.value==2)
return y->get_unary(OP_TWICE);
948 if (v_.value==1)
return y->get_unary(OP_INV);
949 if (v_.value==-1)
return -y->get_unary(OP_INV);
953 if (v_.value==1)
return MX::ones(y.sparsity());
954 if (v_.value==std::exp(1.0))
return y->get_unary(OP_EXP);
961 if (y->op()==OP_CONST &&
dynamic_cast<const ConstantDM*
>(y.get())==
nullptr) {
962 double y_value = y.nnz()>0 ? y->to_double() : 0;
964 casadi_math<double>::fun(op, nnz()> 0.0 ? to_double(): 0, y_value, ret);
966 return MX(y.sparsity(), ret,
false);
970 return MXNode::_get_binary(op, y, ScX, ScY, unique_x, unique_y);
973 template<
typename Value>
974 int Constant<Value>::eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const {
975 std::fill(res[0], res[0]+nnz(), to_double());
979 template<
typename Value>
980 int Constant<Value>::
981 eval_sx(
const SXElem** arg, SXElem** res, casadi_int* iw, SXElem* w)
const {
982 std::fill(res[0], res[0]+nnz(), SXElem(v_.value));
986 template<
typename Value>
987 void Constant<Value>::generate(CodeGenerator& g,
988 const std::vector<casadi_int>& arg,
989 const std::vector<casadi_int>& res,
990 const std::vector<bool>& arg_is_ref,
991 std::vector<bool>& res_is_ref)
const {
994 }
else if (nnz()==1) {
995 g << g.workel(res[0]) <<
" = " << g.constant(to_double()) <<
";\n";
997 if (to_double()==0) {
998 g << g.clear(g.work(res[0], nnz(),
false), nnz()) <<
'\n';
1000 g << g.fill(g.work(res[0], nnz(),
false), nnz(), g.constant(to_double())) <<
'\n';
1005 template<
typename Value>
1006 MX Constant<Value>::get_nzref(
const Sparsity& sp,
const std::vector<casadi_int>& nz,
1007 bool unique)
const {
1010 for (std::vector<casadi_int>::const_iterator k=nz.begin(); k!=nz.end(); ++k) {
1013 return MXNode::get_nzref(sp, nz);
1017 return MX::create(
new Constant<Value>(sp, v_));
1020 template<
typename Value>
1021 MX Constant<Value>::get_nzassign(
const MX& y,
const std::vector<casadi_int>& nz)
const {
1022 if (y.is_constant() && y->is_zero() && v_.value==0) {
1027 return MXNode::get_nzassign(y, nz);
1030 template<
typename Value>
1031 MX Constant<Value>::get_project(
const Sparsity& sp,
bool unique)
const {
1033 return MX::create(
new Constant<Value>(sp, v_));
1034 }
else if (sp.is_dense()) {
1035 return densify(get_DM());
1037 return MXNode::get_project(sp, unique);
1041 template<
typename Value>
1043 Constant<Value>::disp(
const std::vector<std::string>& arg)
const {
1044 std::stringstream ss;
1048 ss.setf(std::ios::scientific);
1050 ss.unsetf(std::ios::scientific);
1052 if (sparsity().is_scalar()) {
1054 if (sparsity().nnz()==0) {
1059 }
else if (sparsity().is_empty()) {
1061 sparsity().disp(ss);
1066 }
else if (v_.value==1) {
1068 }
else if (v_.value!=v_.value) {
1070 }
else if (v_.value==std::numeric_limits<double>::infinity()) {
1072 }
else if (v_.value==-std::numeric_limits<double>::infinity()) {
1075 ss <<
"all_" << v_.value <<
"(";
1079 sparsity().disp(ss);
1085 template<
typename Value>
1086 bool Constant<Value>::is_equal(
const MXNode* node, casadi_int depth)
const {
1087 return node->is_value(to_double()) && sparsity()==node->sparsity();
virtual void serialize_type(SerializingStream &s) const
Serialize type information.
virtual void serialize_body(SerializingStream &s) const
Serialize an object without type information.
static casadi_int get_precision()
Get the 'precision, width & scientific' used in printing and serializing to streams.
static casadi_int get_width()
static bool get_scientific()
static bool is_inf(const T &val)
static bool is_nonnegative(const T &val)
static bool is_minus_inf(const T &val)
static bool is_integer(const T &val)