26 #ifndef CASADI_KRON_HPP
27 #define CASADI_KRON_HPP
29 #include "mx_node.hpp"
53 static MX create(
const MX& a,
const MX& b);
62 virtual void eval_kernel(
const double** arg,
double** res)
const;
63 virtual void eval_kernel(
const SXElem** arg,
SXElem** res)
const;
67 int eval_gen(
const T** arg, T** res, casadi_int* iw, T* w)
const {
68 eval_kernel(arg, res);
75 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
76 return eval_gen<double>(arg, res, iw, w);
83 return eval_gen<SXElem>(arg, res, iw, w);
89 void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
90 const std::vector<bool>& unique={})
const override;
95 int sp_forward(
const bvec_t** arg,
bvec_t** res, casadi_int* iw,
bvec_t* w)
const override;
105 void ad_forward(
const std::vector<std::vector<MX> >& fseed,
106 std::vector<std::vector<MX> >& fsens)
const override;
111 void ad_reverse(
const std::vector<std::vector<MX> >& aseed,
112 std::vector<std::vector<MX> >& asens)
const override;
117 void generate(CodeGenerator& g,
118 const std::vector<casadi_int>& arg,
119 const std::vector<casadi_int>& res,
120 const std::vector<bool>& arg_is_ref,
121 std::vector<bool>& res_is_ref)
const override;
126 std::string disp(
const std::vector<std::string>& arg)
const override;
157 static MXNode* try_create(
const MX& a,
const MX& b);
162 void eval_kernel(
const double** arg,
double** res)
const override;
163 void eval_kernel(
const SXElem** arg,
SXElem** res)
const override;
166 const std::vector<casadi_int>& arg,
167 const std::vector<casadi_int>& res,
168 const std::vector<bool>& arg_is_ref,
169 std::vector<bool>& res_is_ref)
const override;
181 static MXNode* try_create(
const MX& a,
const MX& b);
186 void eval_kernel(
const double** arg,
double** res)
const override;
187 void eval_kernel(
const SXElem** arg,
SXElem** res)
const override;
190 const std::vector<casadi_int>& arg,
191 const std::vector<casadi_int>& res,
192 const std::vector<bool>& arg_is_ref,
193 std::vector<bool>& res_is_ref)
const override;
205 static MXNode* try_create(
const MX& a,
const MX& b);
210 void eval_kernel(
const double** arg,
double** res)
const override;
211 void eval_kernel(
const SXElem** arg,
SXElem** res)
const override;
214 const std::vector<casadi_int>& arg,
215 const std::vector<casadi_int>& res,
216 const std::vector<bool>& arg_is_ref,
217 std::vector<bool>& res_is_ref)
const override;
245 static MX create(
const MX& m,
const MX& x,
bool inner);
255 virtual void eval_kernel(
const double** arg,
double** res,
double* w)
const;
260 int eval_gen(
const T** arg, T** res, casadi_int* iw, T* w)
const {
261 eval_kernel(arg, res, w);
268 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w)
const override {
269 return eval_gen<double>(arg, res, iw, w);
276 return eval_gen<SXElem>(arg, res, iw, w);
282 void eval_mx(
const std::vector<MX>& arg, std::vector<MX>& res,
283 const std::vector<bool>& unique={})
const override;
288 int sp_forward(
const bvec_t** arg,
bvec_t** res, casadi_int* iw,
bvec_t* w)
const override;
298 void ad_forward(
const std::vector<std::vector<MX> >& fseed,
299 std::vector<std::vector<MX> >& fsens)
const override;
304 void ad_reverse(
const std::vector<std::vector<MX> >& aseed,
305 std::vector<std::vector<MX> >& asens)
const override;
310 size_t sz_w()
const override;
315 void generate(CodeGenerator& g,
316 const std::vector<casadi_int>& arg,
317 const std::vector<casadi_int>& res,
318 const std::vector<bool>& arg_is_ref,
319 std::vector<bool>& res_is_ref)
const override;
324 std::string disp(
const std::vector<std::string>& arg)
const override;
363 static MXNode* try_create(
const MX& m,
const MX& x,
bool inner);
368 void eval_kernel(
const double** arg,
double** res,
double* w)
const override;
375 const std::vector<casadi_int>& arg,
376 const std::vector<casadi_int>& res,
377 const std::vector<bool>& arg_is_ref,
378 std::vector<bool>& res_is_ref)
const override;
390 static MXNode* try_create(
const MX& m,
const MX& x,
bool inner);
395 void eval_kernel(
const double** arg,
double** res,
double* w)
const override;
401 const std::vector<casadi_int>& arg,
402 const std::vector<casadi_int>& res,
403 const std::vector<bool>& arg_is_ref,
404 std::vector<bool>& res_is_ref)
const override;
416 static MXNode* try_create(
const MX& m,
const MX& x,
bool inner);
421 void eval_kernel(
const double** arg,
double** res,
double* w)
const override;
427 const std::vector<casadi_int>& arg,
428 const std::vector<casadi_int>& res,
429 const std::vector<bool>& arg_is_ref,
430 std::vector<bool>& res_is_ref)
const override;
Helper class for C code generation.
KronContract specialization: M dense, X dense (=> Y dense)
~DenseKronContract() override
DenseKronContract(const MX &m, const MX &x, bool inner)
DenseKronContract(DeserializingStream &s)
Kron specialization: both operands dense.
DenseKron(DeserializingStream &s)
DenseKron(const MX &a, const MX &b)
KronContract specialization: M dense, X sparse (=> Y dense)
DenseSparseKronContract(DeserializingStream &s)
DenseSparseKronContract(const MX &m, const MX &x, bool inner)
~DenseSparseKronContract() override
Kron specialization: dense a + sparse b.
~DenseSparseKron() override
DenseSparseKron(const MX &a, const MX &b)
DenseSparseKron(DeserializingStream &s)
Helper class for Serialization.
int eval_gen(const T **arg, T **res, casadi_int *iw, T *w) const
Evaluate the function (template) — dispatches to eval_kernel.
int eval(const double **arg, double **res, casadi_int *iw, double *w) const override
Evaluate numerically.
casadi_int op() const override
Get the operation.
bool inner_
Which axes are contracted (true = inner (mB, nB); false = outer (mA, nA))
int eval_sx(const SXElem **arg, SXElem **res, casadi_int *iw, SXElem *w) const override
Evaluate symbolically (SX)
~KronContract() override
Destructor.
casadi_int op() const override
Get the operation.
Kron(DeserializingStream &s)
Deserializing constructor.
~Kron() override
Destructor.
int eval_sx(const SXElem **arg, SXElem **res, casadi_int *iw, SXElem *w) const override
Evaluate symbolically (SX)
int eval_gen(const T **arg, T **res, casadi_int *iw, T *w) const
Evaluate the function (template) — dispatches to eval_kernel.
int eval(const double **arg, double **res, casadi_int *iw, double *w) const override
Evaluate numerically.
Node class for MX objects.
The basic scalar symbolic class of CasADi.
Helper class for Serialization.
KronContract specialization: M sparse, X dense.
SparseDenseKronContract(const MX &m, const MX &x, bool inner)
~SparseDenseKronContract() override
SparseDenseKronContract(DeserializingStream &s)
Kron specialization: sparse a + dense b.
~SparseDenseKron() override
SparseDenseKron(const MX &a, const MX &b)
SparseDenseKron(DeserializingStream &s)
unsigned long long bvec_t