26 #ifndef CASADI_SX_FUNCTION_HPP
27 #define CASADI_SX_FUNCTION_HPP
29 #include "x_function.hpp"
42 struct {
int i1, i2; };
53 class CASADI_EXPORT SXFunction :
54 public XFunction<SXFunction, Matrix<SXElem>, SXNode>{
59 SXFunction(
const std::string& name,
60 const std::vector<Matrix<SXElem> >& inputv,
61 const std::vector<Matrix<SXElem> >& outputv,
62 const std::vector<std::string>& name_in,
63 const std::vector<std::string>& name_out);
68 ~SXFunction()
override;
73 int eval(
const double** arg,
double** res, casadi_int* iw,
double* w,
void* mem)
const override;
75 void trace_instruction(std::ostream& trace, casadi_int k,
const double* w,
81 int eval_sx(
const SXElem** arg, SXElem** res,
82 casadi_int* iw, SXElem* w,
void* mem,
83 bool always_inline,
bool never_inline)
const override;
89 bool always_inline,
bool never_inline)
const override;
92 bool should_inline(
bool with_sx,
bool always_inline,
bool never_inline)
const override;
97 void ad_forward(
const std::vector<std::vector<SX> >& fseed,
98 std::vector<std::vector<SX> >& fsens)
const;
103 void ad_reverse(
const std::vector<std::vector<SX> >& aseed,
104 std::vector<std::vector<SX> >& asens)
const;
109 bool is_smooth()
const;
112 std::string print(
const ScalarAtomic& a)
const;
115 void print_arg(std::ostream &stream, casadi_int k,
const ScalarAtomic& el,
116 const double* w)
const;
119 void print_arg(CodeGenerator& g, casadi_int k,
const ScalarAtomic& el)
const;
122 void print_res(std::ostream &stream, casadi_int k,
const ScalarAtomic& el,
123 const double* w)
const;
126 void print_res(CodeGenerator& g, casadi_int k,
const ScalarAtomic& el)
const;
131 void disp_more(std::ostream& stream)
const override;
136 std::string class_name()
const override {
return "SXFunction";}
141 bool is_a(
const std::string& type,
bool recursive)
const override;
147 const SX sx_in(casadi_int ind)
const override;
148 const std::vector<SX> sx_in()
const override;
152 std::vector<SX> free_sx()
const override {
153 std::vector<SX> ret(free_vars_.size());
154 std::copy(free_vars_.begin(), free_vars_.end(), ret.begin());
161 bool has_free()
const override {
return !free_vars_.empty();}
166 std::vector<std::string> get_free()
const override {
167 std::vector<std::string> ret;
168 for (
auto&& e : free_vars_) ret.push_back(e.name());
175 std::vector<std::string> get_function()
const override;
180 const Function& get_function(
const std::string &name)
const override;
185 SX hess(casadi_int iind=0, casadi_int oind=0);
190 casadi_int n_instructions()
const override {
return algorithm_.size();}
195 casadi_int instruction_id(casadi_int k)
const override {
return algorithm_.at(k).op;}
200 std::vector<casadi_int> instruction_input(casadi_int k)
const override {
201 auto e = algorithm_.at(k);
203 const ExtendedAlgEl& m = call_.el[e.i1];
204 return vector_static_cast<casadi_int>(m.dep);
205 }
else if (casadi_math<double>::ndeps(e.op)==2 || e.op==OP_INPUT) {
207 }
else if (casadi_math<double>::ndeps(e.op)==1) {
217 double instruction_constant(casadi_int k)
const override {
218 return algorithm_.at(k).d;
224 std::vector<casadi_int> instruction_output(casadi_int k)
const override {
225 auto e = algorithm_.at(k);
227 const ExtendedAlgEl& m = call_.el[e.i1];
228 return vector_static_cast<casadi_int>(m.res);
229 }
else if (e.op==OP_OUTPUT) {
239 casadi_int n_nodes()
const override {
return algorithm_.size() - nnz_out();}
248 typedef ScalarAtomic AlgEl;
261 std::vector<AlgEl> algorithm_;
267 std::vector<SXElem> free_vars_;
270 std::vector<SXElem> operations_;
273 std::vector<SXElem> constants_;
276 std::vector<double> default_in_;
279 std::vector<bool> copy_elision_;
282 bool print_instructions_;
283 bool dump_trace_ =
false;
288 void serialize_body(SerializingStream &s)
const override;
291 struct ExtendedAlgEl {
292 ExtendedAlgEl(
const Function& fun);
295 std::vector<int> dep;
297 std::vector<int> res;
299 std::vector<int> copy_elision_arg;
300 std::vector<int> copy_elision_offset;
307 std::vector<int> f_nnz_in;
308 std::vector<int> f_nnz_out;
314 size_t sz_arg = 0, sz_res = 0, sz_iw = 0, sz_w = 0;
315 size_t sz_w_arg = 0, sz_w_res = 0;
316 std::vector<ExtendedAlgEl> el;
322 static ProtoFunction* deserialize(DeserializingStream& s);
324 static std::vector<SX> order(
const std::vector<SX>& expr);
330 static const Options options_;
331 const Options& get_options()
const override {
return options_;}
335 Dict generate_options(
const std::string& target=
"clone")
const override;
340 void init(
const Dict& opts)
override;
345 void init_copy_elision();
350 size_t codegen_sz_w(
const CodeGenerator& g)
const override;
355 void codegen_declarations(CodeGenerator& g)
const override;
360 void codegen_body(CodeGenerator& g)
const override;
365 int sp_forward(
const bvec_t** arg, bvec_t** res,
366 casadi_int* iw, bvec_t* w,
void* mem)
const override;
371 int eval_activity(
const bvec_t** arg, bvec_t** res,
372 casadi_int* iw, bvec_t* w,
void* mem)
const override;
377 int sp_reverse(bvec_t** arg, bvec_t** res, casadi_int* iw, bvec_t* w,
void* mem)
const override;
382 SX instructions_sx()
const override;
385 void find(std::map<FunctionInternal*, std::pair<Function, size_t> >& all_fun,
386 casadi_int max_depth)
const override;
391 void change_option(
const std::string& option_name,
const GenericType& option_value)
override;
396 double get_default_in(casadi_int ind)
const override {
return default_in_.at(ind);}
401 void export_code_body(
const std::string& lang,
402 std::ostream &stream,
const Dict& options)
const override;
405 bool just_in_time_opencl_;
408 bool just_in_time_sparsity_;
411 bool live_variables_;
415 void call_fwd(
const AlgEl& e,
const T** arg, T** res, casadi_int* iw, T* w)
const;
418 void call_activity(
const AlgEl& e,
const bvec_t** arg, bvec_t** res,
419 casadi_int* iw, bvec_t* w)
const;
422 void call_rev(
const AlgEl& e, T** arg, T** res, casadi_int* iw, T* w)
const;
424 template<
typename T,
typename CT>
425 void call_setup(
const ExtendedAlgEl& m,
426 CT*** call_arg, T*** call_res, casadi_int** call_iw, T** call_w, T** nz_in, T** nz_out)
const;
431 explicit SXFunction(DeserializingStream& s);
std::vector< MX > MXVector
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.