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 <functional>
38 #include <map>
39 #include <set>
40 #include <string>
41 
43 namespace casadi {
44 
45  class GraphBuilderInternal;
46 
48  using AddNodeFn = std::function<onnx::NodeProto*()>;
49 
61  class Onnx : public GraphModelInternal {
62  public:
63  explicit Onnx(const std::vector<uint8_t>& model_data);
64  ~Onnx() override;
65 
67  static GraphModelInternal* creator(const std::vector<uint8_t>& model_data) {
68  return new Onnx(model_data);
69  }
70 
71  const char* plugin_name() const override { return "onnx"; }
72  std::string class_name() const override { return "Onnx"; }
73 
75 
76  static const Options options_;
77  const Options& get_options() const override { return options_; }
79 
81  void init(const Dict& opts) override;
82 
83  // ---- GraphModel interface ----
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;
87 
89  static const std::string meta_doc;
90 
91  // ---- symbolic graph engine ----
92 
94  void load(const Function& f);
95 
97  void set_dimension(const std::string& name, casadi_int dim);
98 
100  Function create(const std::string& name);
101 
103  void load_bytes(const std::vector<uint8_t>& data);
104 
106  std::vector<uint8_t> save_bytes() const;
107 
108  protected:
110  onnx::ModelProto model_;
111 
113  std::map<std::string, casadi_int> dimension_overrides_;
114 
116  bool has_model_ = false;
117 
119  std::set<std::string> exported_functions_;
120 
122  std::string casadi_real_ = "double";
123 
124  private:
125  // Import helpers (graph -> MX), operating on this translator
126  void process_graph_initializers(
127  const onnx::GraphProto& graph,
128  std::map<std::string, MX>& value_map,
129  bool verbose) const;
130 
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,
136  bool verbose) const;
137 
138  void process_graph_nodes(
139  const onnx::GraphProto& graph,
140  std::map<std::string, MX>& value_map,
141  bool verbose);
142 
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,
148  bool verbose) const;
149 
150  // Look up a model-level FunctionProto by name+domain, or nullptr if absent
151  const onnx::FunctionProto* find_function(const std::string& name,
152  const std::string& domain) const;
153 
158  casadi_int get_dimension(const onnx::TensorShapeProto& shape, int idx) const;
159 
162  DM tensor_to_dm(const onnx::TensorProto& tensor) const;
163 
168  DM sparse_tensor_to_dm(const onnx::SparseTensorProto& st) const;
169 
176  MX process_node_operation(
177  const std::string& op_type,
178  const onnx::NodeProto& node,
179  const std::vector<MX>& node_inputs);
180 
190  onnx::FunctionProto* function_to_function_proto(
191  const Function& f,
192  const std::string& domain);
193 
198  bool is_if_else_function(const Function& f) const;
199 
204  bool is_mapaccum_function(const Function& f) const;
205 
212  bool is_map_function(const Function& f) const;
213 
218  bool is_reduce_map_function(const Function& f) const;
219 
221  void assert_not_control_flow(const Function& called_func) const;
222 
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);
233 
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);
241 
246  onnx::GraphProto build_scan_body(const Function& base);
247 
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);
258 
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);
264 
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);
270 
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);
276 
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);
282 
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);
289 
291  Function function_from_graph(const onnx::GraphProto& graph, const std::string& name);
292 
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);
300 
302  onnx::GraphProto build_if_branch(const Function& f,
303  const std::vector<std::string>& arg_names,
304  const std::string& prefix);
305 
307  std::vector<MX> eval_captured_subgraph(const onnx::GraphProto& graph,
308  std::map<std::string, MX> scope);
309 
310  // --- Export helpers that depend on configuration (the real type) ---
311 
313  onnx::TensorProto::DataType real_type() const {
314  return casadi_real_ == "float" ? onnx::TensorProto::FLOAT : onnx::TensorProto::DOUBLE;
315  }
316 
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 + "\".");
321  casadi_real_ = v;
322  }
323 
325  void set_real_tensor_type(onnx::ValueInfoProto* value, const Sparsity& sp);
326 
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 = "");
332 
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 = {});
337 
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);
349 
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);
361 
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);
378 
385  void add_sparse_constant(AddNodeFn add_node, const std::string& name, const DM& dm);
386 
389  void fill_sparse_tensor(onnx::SparseTensorProto* st, const std::string& name,
390  const DM& dm) const;
391 
395  Sparsity input_pattern(const onnx::GraphProto& graph, const std::string& name) const;
396 
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);
408  };
409 
410  // ========== Export Helper Functions ==========
411 
413  std::string onnx_input_name(const Function& f, casadi_int i);
414  std::string onnx_output_name(const Function& f, casadi_int i);
415 
420  struct OpMapping {
421  casadi_int casadi_op;
422  const char* onnx_name;
423  int arity;
424  };
425 
429  const OpMapping* get_op_mapping(casadi_int op);
430 
434  const OpMapping* get_op_mapping_by_name(const std::string& onnx_name);
435 
437  casadi_int get_int_attribute(const onnx::NodeProto& node, const std::string& name,
438  casadi_int default_value);
439 
441  double get_float_attribute(const onnx::NodeProto& node, const std::string& name,
442  double default_value);
443 
445  std::string get_string_attribute(const onnx::NodeProto& node, const std::string& name);
446 
448  const onnx::GraphProto* get_graph_attribute(const onnx::NodeProto& node,
449  const std::string& name);
450 
452  void add_int_attribute(onnx::NodeProto* node, const std::string& name, casadi_int value);
453 
455  void add_ints_attribute(onnx::NodeProto* node, const std::string& name,
456  const std::vector<casadi_int>& values);
457 
459  void add_int_constant(onnx::GraphProto* graph, const std::string& name,
460  const std::vector<casadi_int>& data, std::vector<casadi_int> dims = {});
461 
463  onnx::NodeProto* create_binary_node(
464  AddNodeFn add_node,
465  const std::string& op_type,
466  const std::string& input1,
467  const std::string& input2,
468  const std::string& output);
469 
471  onnx::NodeProto* create_unary_node(
472  AddNodeFn add_node,
473  const std::string& op_type,
474  const std::string& input,
475  const std::string& output);
476 
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);
484 
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);
491 
492 } // namespace casadi
494 
495 #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