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>
46 class GraphBuilderInternal;
49 using AddNodeFn = std::function<onnx::NodeProto*()>;
62 class Onnx :
public GraphModelInternal {
64 explicit Onnx(
const std::vector<uint8_t>& model_data);
68 static GraphModelInternal* creator(
const std::vector<uint8_t>& model_data) {
69 return new Onnx(model_data);
72 const char* plugin_name()
const override {
return "onnx"; }
73 std::string class_name()
const override {
return "Onnx"; }
77 static const Options options_;
78 const Options& get_options()
const override {
return options_; }
82 void init(
const Dict& opts)
override;
85 void fill_metadata(GraphBuilderInternal& gb)
const override;
86 Function import_symbolic(
const GraphBuilderInternal& gb,
const std::string& name)
override;
87 std::vector<uint8_t> export_symbolic(
const Function& f,
const Dict& opts)
override;
90 static const std::string meta_doc;
95 void load(
const Function& f);
98 void set_dimension(
const std::string& name, casadi_int dim);
101 Function create(
const std::string& name);
104 void load_bytes(
const std::vector<uint8_t>& data);
107 std::vector<uint8_t> save_bytes()
const;
111 onnx::ModelProto model_;
114 std::map<std::string, casadi_int> dimension_overrides_;
117 bool has_model_ =
false;
120 std::set<std::string> exported_functions_;
123 std::string casadi_real_ =
"double";
126 using IntegerConstants = std::map<std::string, std::vector<int64_t>>;
129 void process_graph_initializers(
130 const onnx::GraphProto& graph,
131 std::map<std::string, MX>& value_map,
134 void process_graph_inputs(
135 const onnx::GraphProto& graph,
136 std::map<std::string, MX>& value_map,
137 std::vector<MX>& func_inputs,
138 std::vector<std::string>& input_names,
141 void process_graph_nodes(
142 const onnx::GraphProto& graph,
143 std::map<std::string, MX>& value_map,
146 void collect_graph_outputs(
147 const onnx::GraphProto& graph,
148 const std::map<std::string, MX>& value_map,
149 std::vector<MX>& func_outputs,
150 std::vector<std::string>& output_names,
154 const onnx::FunctionProto* find_function(
const std::string& name,
155 const std::string& domain)
const;
161 casadi_int get_dimension(
const onnx::TensorShapeProto& shape,
int idx)
const;
165 DM tensor_to_dm(
const onnx::TensorProto& tensor)
const;
171 DM sparse_tensor_to_dm(
const onnx::SparseTensorProto& st)
const;
179 MX process_node_operation(
180 const std::string& op_type,
181 const onnx::NodeProto& node,
182 const std::vector<MX>& node_inputs,
183 const IntegerConstants& integer_constants);
194 onnx::FunctionProto* function_to_function_proto(
196 const std::string& domain);
202 bool is_if_else_function(
const Function& f)
const;
208 bool is_mapaccum_function(
const Function& f)
const;
216 bool is_map_function(
const Function& f)
const;
222 bool is_reduce_map_function(
const Function& f)
const;
225 void assert_not_control_flow(
const Function& called_func)
const;
231 template<
typename Container>
232 void export_call(Container* container,
const Function& called_func,
233 const std::vector<casadi_int>& i_vec,
234 const std::vector<casadi_int>& o,
235 std::map<casadi_int, std::string>& work_to_onnx,
236 const std::string& out_prefix);
239 template<
typename Container>
240 void export_map(Container* container,
const Function& map_fn,
241 const std::vector<casadi_int>& i_vec,
242 const std::vector<casadi_int>& o,
243 std::map<casadi_int, std::string>& work_to_onnx,
244 const std::string& out_prefix);
250 onnx::GraphProto build_scan_body(
const Function& base);
256 template<
typename Container>
257 void export_reduce_map(Container* container,
const Function& wrapper,
258 const std::vector<casadi_int>& i_vec,
259 const std::vector<casadi_int>& o,
260 std::map<casadi_int, std::string>& work_to_onnx,
261 const std::string& out_prefix);
264 onnx::GraphProto build_reduce_scan_body(
const Function& base,
265 const std::vector<bool>& reduce_in,
266 const std::vector<bool>& reduce_out,
267 const std::vector<std::string>& capture_names);
270 Function function_from_function_proto(
271 const onnx::FunctionProto& fp,
272 const std::vector<std::pair<casadi_int, casadi_int>>& in_shapes,
273 const std::string& name);
276 template<
typename Container>
277 void emit_reshape(Container* container,
const std::string& data,
278 const std::vector<casadi_int>& shape,
279 const std::string& output,
const std::string& shape_name);
282 template<
typename Container>
283 void colmajor_reshape_into(Container* container,
const std::string& data,
284 const std::vector<casadi_int>& dims,
285 const std::string& output,
const std::string& uniq);
289 template<
typename Container>
290 void emit_output_node(Container* container,
const std::string& data,
291 casadi_int src_rows, casadi_int src_cols,
const Sparsity& out_sp,
292 const std::string& output,
const std::string& uniq);
295 Function function_from_graph(
const onnx::GraphProto& graph,
const std::string& name);
298 template<
typename Container>
299 void export_if(Container* container,
const Function& switch_fn,
300 const std::vector<casadi_int>& i_vec,
301 const std::vector<casadi_int>& o,
302 std::map<casadi_int, std::string>& work_to_onnx,
303 const std::string& out_prefix);
306 onnx::GraphProto build_if_branch(
const Function& f,
307 const std::vector<std::string>& arg_names,
308 const std::string& prefix);
311 std::vector<MX> eval_captured_subgraph(
const onnx::GraphProto& graph,
312 std::map<std::string, MX> scope);
317 onnx::TensorProto::DataType real_type()
const {
318 return casadi_real_ ==
"float" ? onnx::TensorProto::FLOAT : onnx::TensorProto::DOUBLE;
322 void set_casadi_real(
const std::string& v) {
323 casadi_assert(v ==
"double" || v ==
"float",
324 "casadi_real must be \"double\" or \"float\", got \"" + v +
"\".");
329 void set_real_tensor_type(onnx::ValueInfoProto* value,
const Sparsity& sp);
332 void add_graph_inputs(onnx::GraphProto* graph,
const Function& f,
333 const std::string& name_prefix =
"");
334 void add_graph_outputs(onnx::GraphProto* graph,
const Function& f,
335 const std::string& name_prefix =
"");
338 void add_real_constant(AddNodeFn add_node,
const std::string& name,
339 const std::vector<double>& data,
340 const std::vector<casadi_int>& dims = {});
347 void emit_blockdiag(AddNodeFn add_node,
const std::vector<std::string>& names,
348 const std::vector<casadi_int>& row_off,
349 const std::vector<casadi_int>& col_off,
350 const std::vector<casadi_int>& br,
const std::vector<casadi_int>& bc,
351 casadi_int R, casadi_int C,
const std::string& output,
352 const std::string& uniq);
361 void emit_nonzero_remap(AddNodeFn add_node,
const std::string& data,
362 const Sparsity& sp_in,
const Sparsity& out_sp,
363 std::vector<casadi_int> idx,
const std::string& uniq,
364 const std::string& node_output);
379 std::string emit_sparsity_restore(AddNodeFn add_node,
const std::string& value,
380 const Sparsity& value_sp,
const Sparsity& target_sp,
381 const std::string& uniq,
const std::string& final_output);
389 void add_sparse_constant(AddNodeFn add_node,
const std::string& name,
const DM& dm);
393 void fill_sparse_tensor(onnx::SparseTensorProto* st,
const std::string& name,
399 Sparsity input_pattern(
const onnx::GraphProto& graph,
const std::string& name)
const;
402 bool process_operation(AddNodeFn add_node,
const Function& f, casadi_int op, casadi_int k,
403 const std::vector<casadi_int>& i_vec,
404 const std::vector<casadi_int>& o_vec,
405 std::map<casadi_int, std::string>& work_to_onnx,
406 const std::string& node_output);
407 bool process_operation(onnx::GraphProto* graph,
const Function& f, casadi_int op, casadi_int k,
408 const std::vector<casadi_int>& i_vec,
409 const std::vector<casadi_int>& o_vec,
410 std::map<casadi_int, std::string>& work_to_onnx,
411 const std::string& node_output);
417 std::string onnx_input_name(
const Function& f, casadi_int i);
418 std::string onnx_output_name(
const Function& f, casadi_int i);
425 casadi_int casadi_op;
426 const char* onnx_name;
433 const OpMapping* get_op_mapping(casadi_int op);
438 const OpMapping* get_op_mapping_by_name(
const std::string& onnx_name);
441 casadi_int get_int_attribute(
const onnx::NodeProto& node,
const std::string& name,
442 casadi_int default_value);
445 double get_float_attribute(
const onnx::NodeProto& node,
const std::string& name,
446 double default_value);
449 std::string get_string_attribute(
const onnx::NodeProto& node,
const std::string& name);
452 const onnx::GraphProto* get_graph_attribute(
const onnx::NodeProto& node,
453 const std::string& name);
456 void add_int_attribute(onnx::NodeProto* node,
const std::string& name, casadi_int value);
459 void add_ints_attribute(onnx::NodeProto* node,
const std::string& name,
460 const std::vector<casadi_int>& values);
463 void add_int_constant(onnx::GraphProto* graph,
const std::string& name,
464 const std::vector<casadi_int>& data, std::vector<casadi_int> dims = {});
467 onnx::NodeProto* create_binary_node(
469 const std::string& op_type,
470 const std::string& input1,
471 const std::string& input2,
472 const std::string& output);
475 onnx::NodeProto* create_unary_node(
477 const std::string& op_type,
478 const std::string& input,
479 const std::string& output);
482 onnx::NodeProto* create_binary_node(
483 onnx::GraphProto* graph,
484 const std::string& op_type,
485 const std::string& input1,
486 const std::string& input2,
487 const std::string& output);
490 onnx::NodeProto* create_unary_node(
491 onnx::GraphProto* graph,
492 const std::string& op_type,
493 const std::string& input,
494 const std::string& output);
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.