26 #ifndef CASADI_ONNX_FUNCTION_IMPL_HPP
27 #define CASADI_ONNX_FUNCTION_IMPL_HPP
29 #include "onnx_function.hpp"
30 #include "function_internal.hpp"
31 #include "plugin_interface.hpp"
38 class GraphBuilderInternal;
41 struct OnnxTensorInfo {
43 std::vector<casadi_int> shape;
48 struct CASADI_EXPORT OnnxMemory :
public FunctionMemory {
52 CASADI_EXPORT std::string onnx_dtype_name(casadi_int elem_type);
55 CASADI_EXPORT casadi_int onnx_dtype_enum(
const std::string& name);
62 class CASADI_EXPORT OnnxFunction :
public FunctionInternal,
public PluginInterface<OnnxFunction> {
65 OnnxFunction(
const std::string& name,
66 const GraphBuilderInternal* gb,
67 const std::vector<std::string>& inputs,
68 const std::vector<std::string>& outputs);
69 ~OnnxFunction()
override;
72 typedef OnnxFunction* (*Creator)(
const std::string& name,
73 const GraphBuilderInternal* gb,
74 const std::vector<std::string>& inputs,
75 const std::vector<std::string>& outputs,
85 static const Options options_;
86 const Options& get_options()
const override {
return options_;}
90 size_t get_n_in()
override {
return in_.size(); }
91 size_t get_n_out()
override {
return out_.size(); }
92 std::string get_name_in(casadi_int i)
override {
return in_.at(i).name; }
93 std::string get_name_out(casadi_int i)
override {
return out_.at(i).name; }
94 Sparsity get_sparsity_in(casadi_int i)
override {
return tensor_sparsity(in_.at(i).shape); }
95 Sparsity get_sparsity_out(casadi_int i)
override {
return tensor_sparsity(out_.at(i).shape); }
100 void init(
const Dict& opts)
override;
102 void* alloc_mem()
const override {
return new OnnxMemory(); }
103 void free_mem(
void* mem)
const override {
delete static_cast<OnnxMemory*
>(mem); }
105 int eval(
const double** arg,
double** res,
106 casadi_int* iw,
double* w,
void* mem)
const override = 0;
108 bool uses_output()
const override {
return false; }
113 bool has_forward(casadi_int nfwd)
const override;
114 Function get_forward(casadi_int nfwd,
const std::string& name,
115 const std::vector<std::string>& inames,
116 const std::vector<std::string>& onames,
117 const Dict& opts)
const override;
118 bool has_reverse(casadi_int nadj)
const override;
119 Function get_reverse(casadi_int nadj,
const std::string& name,
120 const std::vector<std::string>& inames,
121 const std::vector<std::string>& onames,
122 const Dict& opts)
const override;
123 bool has_jacobian()
const override;
124 Function get_jacobian(
const std::string& name,
125 const std::vector<std::string>& inames,
126 const std::vector<std::string>& onames,
127 const Dict& opts)
const override;
132 void serialize_body(SerializingStream &s)
const override;
136 void serialize_type(SerializingStream &s)
const override;
140 std::string serialize_base_function()
const override {
return "Onnx"; }
144 static ProtoFunction* deserialize(DeserializingStream& s);
147 static std::string meta_doc;
150 static Function create(
const std::string& solver,
151 const std::string& name,
152 const GraphBuilderInternal* gb,
153 const std::vector<std::string>& inputs,
154 const std::vector<std::string>& outputs,
158 static const std::string infix_;
161 static std::map<std::string, Plugin> solvers_;
164 #ifdef CASADI_WITH_THREADSAFE_SYMBOLICS
165 static std::mutex mutex_solvers_;
170 explicit OnnxFunction(DeserializingStream& s);
173 static Sparsity tensor_sparsity(
const std::vector<casadi_int>& shape);
179 bool diff_in(casadi_int i)
const {
return is_diff_in_.empty() || is_diff_in_.at(i); }
180 bool diff_out(casadi_int i)
const {
return is_diff_out_.empty() || is_diff_out_.at(i); }
183 static Function from_model_data(
const std::string& solver,
const std::string& name,
184 const std::vector<uint8_t>& model_data,
185 const std::vector<std::string>& inputs,
186 const std::vector<std::string>& outputs,
191 Function wrap_derivative(
const std::string& name,
192 const std::vector<std::string>& inames,
193 const std::vector<std::string>& onames,
194 const std::vector<Sparsity>& in_sp,
195 const std::vector<Sparsity>& out_sp,
196 const Dict& dim_bind,
const Dict& opts)
const;
199 std::vector<uint8_t> model_data_;
202 std::vector<OnnxTensorInfo> in_, out_;
205 std::vector<OnnxTensorInfo> all_in_;
208 std::vector<casadi_int> in_src_;
211 std::vector<double> in_val_;
214 std::set<std::string> model_inputs_, model_outputs_;
217 std::string fwd_dim_ =
"nfwd";
218 std::string adj_dim_ =
"nadj";
221 std::map<std::string, std::vector<double>> input_values_;
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.