26 #ifndef CASADI_ONNX_MODEL_HPP
27 #define CASADI_ONNX_MODEL_HPP
29 #include <casadi/core/graph_model_internal.hpp>
30 #include <casadi/core/mx.hpp>
31 #include <casadi/interfaces/onnx/casadi_graphmodel_onnx_export.h>
34 #define ONNX_NAMESPACE onnx
35 #include <onnx/onnx_pb.h>
45 class GraphBuilderInternal;
48 using AddNodeFn = std::function<onnx::NodeProto*()>;
61 class Onnx :
public GraphModelInternal {
63 explicit Onnx(
const std::vector<uint8_t>& model_data);
67 static GraphModelInternal* creator(
const std::vector<uint8_t>& model_data) {
68 return new Onnx(model_data);
71 const char* plugin_name()
const override {
return "onnx"; }
72 std::string class_name()
const override {
return "Onnx"; }
76 static const Options options_;
77 const Options& get_options()
const override {
return options_; }
81 void init(
const Dict& opts)
override;
84 void fill_metadata(GraphBuilderInternal& gb)
const override;
85 Function import_symbolic(
const GraphBuilderInternal& gb,
const std::string& name)
override;
86 std::vector<uint8_t> export_symbolic(
const Function& f,
const Dict& opts)
override;
89 static const std::string meta_doc;
94 void load(
const Function& f);
97 void set_dimension(
const std::string& name, casadi_int dim);
100 Function create(
const std::string& name);
103 void load_bytes(
const std::vector<uint8_t>& data);
106 std::vector<uint8_t> save_bytes()
const;
110 onnx::ModelProto model_;
113 std::map<std::string, casadi_int> dimension_overrides_;
116 bool has_model_ =
false;
119 std::set<std::string> exported_functions_;
122 std::string casadi_real_ =
"double";
126 void process_graph_initializers(
127 const onnx::GraphProto& graph,
128 std::map<std::string, MX>& value_map,
131 void process_graph_inputs(
132 const onnx::GraphProto& graph,
133 std::map<std::string, MX>& value_map,
134 std::vector<MX>& func_inputs,
135 std::vector<std::string>& input_names,
138 void process_graph_nodes(
139 const onnx::GraphProto& graph,
140 std::map<std::string, MX>& value_map,
143 void collect_graph_outputs(
144 const onnx::GraphProto& graph,
145 const std::map<std::string, MX>& value_map,
146 std::vector<MX>& func_outputs,
147 std::vector<std::string>& output_names,
151 const onnx::FunctionProto* find_function(
const std::string& name,
152 const std::string& domain)
const;
158 casadi_int get_dimension(
const onnx::TensorShapeProto& shape,
int idx)
const;
162 DM tensor_to_dm(
const onnx::TensorProto& tensor)
const;
168 DM sparse_tensor_to_dm(
const onnx::SparseTensorProto& st)
const;
176 MX process_node_operation(
177 const std::string& op_type,
178 const onnx::NodeProto& node,
179 const std::vector<MX>& node_inputs);
190 onnx::FunctionProto* function_to_function_proto(
192 const std::string& domain);
198 bool is_if_else_function(
const Function& f)
const;
204 bool is_mapaccum_function(
const Function& f)
const;
212 bool is_map_function(
const Function& f)
const;
218 bool is_reduce_map_function(
const Function& f)
const;
221 void assert_not_control_flow(
const Function& called_func)
const;
227 template<
typename Container>
228 void export_call(Container* container,
const Function& called_func,
229 const std::vector<casadi_int>& i_vec,
230 const std::vector<casadi_int>& o,
231 std::map<casadi_int, std::string>& work_to_onnx,
232 const std::string& out_prefix);
235 template<
typename Container>
236 void export_map(Container* container,
const Function& map_fn,
237 const std::vector<casadi_int>& i_vec,
238 const std::vector<casadi_int>& o,
239 std::map<casadi_int, std::string>& work_to_onnx,
240 const std::string& out_prefix);
246 onnx::GraphProto build_scan_body(
const Function& base);
252 template<
typename Container>
253 void export_reduce_map(Container* container,
const Function& wrapper,
254 const std::vector<casadi_int>& i_vec,
255 const std::vector<casadi_int>& o,
256 std::map<casadi_int, std::string>& work_to_onnx,
257 const std::string& out_prefix);
260 onnx::GraphProto build_reduce_scan_body(
const Function& base,
261 const std::vector<bool>& reduce_in,
262 const std::vector<bool>& reduce_out,
263 const std::vector<std::string>& capture_names);
266 Function function_from_function_proto(
267 const onnx::FunctionProto& fp,
268 const std::vector<std::pair<casadi_int, casadi_int>>& in_shapes,
269 const std::string& name);
272 template<
typename Container>
273 void emit_reshape(Container* container,
const std::string& data,
274 const std::vector<casadi_int>& shape,
275 const std::string& output,
const std::string& shape_name);
278 template<
typename Container>
279 void colmajor_reshape_into(Container* container,
const std::string& data,
280 const std::vector<casadi_int>& dims,
281 const std::string& output,
const std::string& uniq);
285 template<
typename Container>
286 void emit_output_node(Container* container,
const std::string& data,
287 casadi_int src_rows, casadi_int src_cols,
const Sparsity& out_sp,
288 const std::string& output,
const std::string& uniq);
291 Function function_from_graph(
const onnx::GraphProto& graph,
const std::string& name);
294 template<
typename Container>
295 void export_if(Container* container,
const Function& switch_fn,
296 const std::vector<casadi_int>& i_vec,
297 const std::vector<casadi_int>& o,
298 std::map<casadi_int, std::string>& work_to_onnx,
299 const std::string& out_prefix);
302 onnx::GraphProto build_if_branch(
const Function& f,
303 const std::vector<std::string>& arg_names,
304 const std::string& prefix);
307 std::vector<MX> eval_captured_subgraph(
const onnx::GraphProto& graph,
308 std::map<std::string, MX> scope);
313 onnx::TensorProto::DataType real_type()
const {
314 return casadi_real_ ==
"float" ? onnx::TensorProto::FLOAT : onnx::TensorProto::DOUBLE;
318 void set_casadi_real(
const std::string& v) {
319 casadi_assert(v ==
"double" || v ==
"float",
320 "casadi_real must be \"double\" or \"float\", got \"" + v +
"\".");
325 void set_real_tensor_type(onnx::ValueInfoProto* value,
const Sparsity& sp);
328 void add_graph_inputs(onnx::GraphProto* graph,
const Function& f,
329 const std::string& name_prefix =
"");
330 void add_graph_outputs(onnx::GraphProto* graph,
const Function& f,
331 const std::string& name_prefix =
"");
334 void add_real_constant(AddNodeFn add_node,
const std::string& name,
335 const std::vector<double>& data,
336 const std::vector<casadi_int>& dims = {});
343 void emit_blockdiag(AddNodeFn add_node,
const std::vector<std::string>& names,
344 const std::vector<casadi_int>& row_off,
345 const std::vector<casadi_int>& col_off,
346 const std::vector<casadi_int>& br,
const std::vector<casadi_int>& bc,
347 casadi_int R, casadi_int C,
const std::string& output,
348 const std::string& uniq);
357 void emit_nonzero_remap(AddNodeFn add_node,
const std::string& data,
358 const Sparsity& sp_in,
const Sparsity& out_sp,
359 std::vector<casadi_int> idx,
const std::string& uniq,
360 const std::string& node_output);
375 std::string emit_sparsity_restore(AddNodeFn add_node,
const std::string& value,
376 const Sparsity& value_sp,
const Sparsity& target_sp,
377 const std::string& uniq,
const std::string& final_output);
385 void add_sparse_constant(AddNodeFn add_node,
const std::string& name,
const DM& dm);
389 void fill_sparse_tensor(onnx::SparseTensorProto* st,
const std::string& name,
395 Sparsity input_pattern(
const onnx::GraphProto& graph,
const std::string& name)
const;
398 bool process_operation(AddNodeFn add_node,
const Function& f, casadi_int op, casadi_int k,
399 const std::vector<casadi_int>& i_vec,
400 const std::vector<casadi_int>& o_vec,
401 std::map<casadi_int, std::string>& work_to_onnx,
402 const std::string& node_output);
403 bool process_operation(onnx::GraphProto* graph,
const Function& f, casadi_int op, casadi_int k,
404 const std::vector<casadi_int>& i_vec,
405 const std::vector<casadi_int>& o_vec,
406 std::map<casadi_int, std::string>& work_to_onnx,
407 const std::string& node_output);
413 std::string onnx_input_name(
const Function& f, casadi_int i);
414 std::string onnx_output_name(
const Function& f, casadi_int i);
421 casadi_int casadi_op;
422 const char* onnx_name;
429 const OpMapping* get_op_mapping(casadi_int op);
434 const OpMapping* get_op_mapping_by_name(
const std::string& onnx_name);
437 casadi_int get_int_attribute(
const onnx::NodeProto& node,
const std::string& name,
438 casadi_int default_value);
441 double get_float_attribute(
const onnx::NodeProto& node,
const std::string& name,
442 double default_value);
445 std::string get_string_attribute(
const onnx::NodeProto& node,
const std::string& name);
448 const onnx::GraphProto* get_graph_attribute(
const onnx::NodeProto& node,
449 const std::string& name);
452 void add_int_attribute(onnx::NodeProto* node,
const std::string& name, casadi_int value);
455 void add_ints_attribute(onnx::NodeProto* node,
const std::string& name,
456 const std::vector<casadi_int>& values);
459 void add_int_constant(onnx::GraphProto* graph,
const std::string& name,
460 const std::vector<casadi_int>& data, std::vector<casadi_int> dims = {});
463 onnx::NodeProto* create_binary_node(
465 const std::string& op_type,
466 const std::string& input1,
467 const std::string& input2,
468 const std::string& output);
471 onnx::NodeProto* create_unary_node(
473 const std::string& op_type,
474 const std::string& input,
475 const std::string& output);
478 onnx::NodeProto* create_binary_node(
479 onnx::GraphProto* graph,
480 const std::string& op_type,
481 const std::string& input1,
482 const std::string& input2,
483 const std::string& output);
486 onnx::NodeProto* create_unary_node(
487 onnx::GraphProto* graph,
488 const std::string& op_type,
489 const std::string& input,
490 const std::string& output);
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.