onnx_model.hpp
1 /*
2  * This file is part of CasADi.
3  *
4  * CasADi -- A symbolic framework for dynamic optimization.
5  * Copyright (C) 2010-2023 Joel Andersson, Joris Gillis, Moritz Diehl,
6  * KU Leuven. All rights reserved.
7  * Copyright (C) 2011-2014 Greg Horn
8  *
9  * CasADi is free software; you can redistribute it and/or
10  * modify it under the terms of the GNU Lesser General Public
11  * License as published by the Free Software Foundation; either
12  * version 3 of the License, or (at your option) any later version.
13  *
14  * CasADi is distributed in the hope that it will be useful,
15  * but WITHOUT ANY WARRANTY; without even the implied warranty of
16  * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
17  * Lesser General Public License for more details.
18  *
19  * You should have received a copy of the GNU Lesser General Public
20  * License along with CasADi; if not, write to the Free Software
21  * Foundation, Inc., 51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA
22  *
23  */
24 
25 
26 #ifndef CASADI_ONNX_MODEL_HPP
27 #define CASADI_ONNX_MODEL_HPP
28 
29 #include <casadi/core/graph_model_internal.hpp>
30 #include <casadi/core/mx.hpp>
31 #include <casadi/interfaces/onnx/casadi_graphmodel_onnx_export.h>
32 
33 #define ONNX_ML 1
34 #define ONNX_NAMESPACE onnx
35 #include <onnx/onnx_pb.h>
36 
37 #include <cstdint>
38 #include <functional>
39 #include <map>
40 #include <set>
41 #include <string>
42 
44 namespace casadi {
45 
46  class GraphBuilderInternal;
47 
49  using AddNodeFn = std::function<onnx::NodeProto*()>;
50 
62  class Onnx : public GraphModelInternal {
63  public:
64  explicit Onnx(const std::vector<uint8_t>& model_data);
65  ~Onnx() override;
66 
68  static GraphModelInternal* creator(const std::vector<uint8_t>& model_data) {
69  return new Onnx(model_data);
70  }
71 
72  const char* plugin_name() const override { return "onnx"; }
73  std::string class_name() const override { return "Onnx"; }
74 
76 
77  static const Options options_;
78  const Options& get_options() const override { return options_; }
80 
82  void init(const Dict& opts) override;
83 
84  // ---- GraphModel interface ----
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;
88 
90  static const std::string meta_doc;
91 
92  // ---- symbolic graph engine ----
93 
95  void load(const Function& f);
96 
98  void set_dimension(const std::string& name, casadi_int dim);
99 
101  Function create(const std::string& name);
102 
104  void load_bytes(const std::vector<uint8_t>& data);
105 
107  std::vector<uint8_t> save_bytes() const;
108 
109  protected:
111  onnx::ModelProto model_;
112 
114  std::map<std::string, casadi_int> dimension_overrides_;
115 
117  bool has_model_ = false;
118 
120  std::set<std::string> exported_functions_;
121 
123  std::string casadi_real_ = "double";
124 
125  private:
126  using IntegerConstants = std::map<std::string, std::vector<int64_t>>;
127 
128  // Import helpers (graph -> MX), operating on this translator
129  void process_graph_initializers(
130  const onnx::GraphProto& graph,
131  std::map<std::string, MX>& value_map,
132  bool verbose) const;
133 
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,
139  bool verbose) const;
140 
141  void process_graph_nodes(
142  const onnx::GraphProto& graph,
143  std::map<std::string, MX>& value_map,
144  bool verbose);
145 
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,
151  bool verbose) const;
152 
153  // Look up a model-level FunctionProto by name+domain, or nullptr if absent
154  const onnx::FunctionProto* find_function(const std::string& name,
155  const std::string& domain) const;
156 
161  casadi_int get_dimension(const onnx::TensorShapeProto& shape, int idx) const;
162 
165  DM tensor_to_dm(const onnx::TensorProto& tensor) const;
166 
171  DM sparse_tensor_to_dm(const onnx::SparseTensorProto& st) const;
172 
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);
184 
194  onnx::FunctionProto* function_to_function_proto(
195  const Function& f,
196  const std::string& domain);
197 
202  bool is_if_else_function(const Function& f) const;
203 
208  bool is_mapaccum_function(const Function& f) const;
209 
216  bool is_map_function(const Function& f) const;
217 
222  bool is_reduce_map_function(const Function& f) const;
223 
225  void assert_not_control_flow(const Function& called_func) const;
226 
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);
237 
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);
245 
250  onnx::GraphProto build_scan_body(const Function& base);
251 
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);
262 
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);
268 
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);
274 
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);
280 
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);
286 
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);
293 
295  Function function_from_graph(const onnx::GraphProto& graph, const std::string& name);
296 
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);
304 
306  onnx::GraphProto build_if_branch(const Function& f,
307  const std::vector<std::string>& arg_names,
308  const std::string& prefix);
309 
311  std::vector<MX> eval_captured_subgraph(const onnx::GraphProto& graph,
312  std::map<std::string, MX> scope);
313 
314  // --- Export helpers that depend on configuration (the real type) ---
315 
317  onnx::TensorProto::DataType real_type() const {
318  return casadi_real_ == "float" ? onnx::TensorProto::FLOAT : onnx::TensorProto::DOUBLE;
319  }
320 
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 + "\".");
325  casadi_real_ = v;
326  }
327 
329  void set_real_tensor_type(onnx::ValueInfoProto* value, const Sparsity& sp);
330 
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 = "");
336 
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 = {});
341 
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);
353 
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);
365 
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);
382 
389  void add_sparse_constant(AddNodeFn add_node, const std::string& name, const DM& dm);
390 
393  void fill_sparse_tensor(onnx::SparseTensorProto* st, const std::string& name,
394  const DM& dm) const;
395 
399  Sparsity input_pattern(const onnx::GraphProto& graph, const std::string& name) const;
400 
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);
412  };
413 
414  // ========== Export Helper Functions ==========
415 
417  std::string onnx_input_name(const Function& f, casadi_int i);
418  std::string onnx_output_name(const Function& f, casadi_int i);
419 
424  struct OpMapping {
425  casadi_int casadi_op;
426  const char* onnx_name;
427  int arity;
428  };
429 
433  const OpMapping* get_op_mapping(casadi_int op);
434 
438  const OpMapping* get_op_mapping_by_name(const std::string& onnx_name);
439 
441  casadi_int get_int_attribute(const onnx::NodeProto& node, const std::string& name,
442  casadi_int default_value);
443 
445  double get_float_attribute(const onnx::NodeProto& node, const std::string& name,
446  double default_value);
447 
449  std::string get_string_attribute(const onnx::NodeProto& node, const std::string& name);
450 
452  const onnx::GraphProto* get_graph_attribute(const onnx::NodeProto& node,
453  const std::string& name);
454 
456  void add_int_attribute(onnx::NodeProto* node, const std::string& name, casadi_int value);
457 
459  void add_ints_attribute(onnx::NodeProto* node, const std::string& name,
460  const std::vector<casadi_int>& values);
461 
463  void add_int_constant(onnx::GraphProto* graph, const std::string& name,
464  const std::vector<casadi_int>& data, std::vector<casadi_int> dims = {});
465 
467  onnx::NodeProto* create_binary_node(
468  AddNodeFn add_node,
469  const std::string& op_type,
470  const std::string& input1,
471  const std::string& input2,
472  const std::string& output);
473 
475  onnx::NodeProto* create_unary_node(
476  AddNodeFn add_node,
477  const std::string& op_type,
478  const std::string& input,
479  const std::string& output);
480 
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);
488 
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);
495 
496 } // namespace casadi
498 
499 #endif // CASADI_ONNX_MODEL_HPP
The casadi namespace.
Definition: archiver.hpp:32
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
Matrix< double > DM
Definition: dm_fwd.hpp:33