graph_builder.cpp
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 #include "graph_builder_internal.hpp"
26 #include "onnx_function_impl.hpp"
27 #include "filesystem_impl.hpp"
28 #include <fstream>
29 
30 namespace casadi {
31 
32  // ---------- Node ----------
33 
34  void Node::disp(std::ostream& stream, bool more) const {
35  (void)more;
36  stream << io << " " << name << ": " << dtype << "[";
37  for (size_t k = 0; k < dimension.size(); ++k) {
38  if (k) stream << "x";
39  if (dimension[k] >= 0) stream << dimension[k];
40  else
41  stream << (dim_params[k].empty() ? "?" : dim_params[k]);
42  }
43  stream << "]";
44  if (baked) stream << " (baked)";
45  }
46 
47  std::string Node::get_str(bool more) const {
48  std::stringstream ss;
49  disp(ss, more);
50  return ss.str();
51  }
52 
53  // Infer the model format from a file suffix
54  static std::string format_from_path(const std::string& path) {
55  auto dot = path.find_last_of('.');
56  std::string suffix = dot == std::string::npos ? "" : path.substr(dot + 1);
57  if (suffix == "onnx") return "onnx";
58  casadi_error("GraphBuilder: cannot infer format from '" + path + "'");
59  }
60 
61  // ---------- public GraphBuilder ----------
62 
64  }
65 
66  GraphBuilder::GraphBuilder(const std::string& model_path, const Dict& opts) {
67  std::ifstream file(model_path, std::ios::binary | std::ios::ate);
68  casadi_assert(file.is_open(), "Cannot open model file: " + model_path);
69  std::streamsize size = file.tellg();
70  file.seekg(0, std::ios::beg);
71  std::vector<uint8_t> data(static_cast<size_t>(size));
72  casadi_assert(file.read(reinterpret_cast<char*>(data.data()), size),
73  "Cannot read model file: " + model_path);
74  own(new GraphBuilderInternal(model_path, data, format_from_path(model_path), opts));
75  if (Filesystem::is_enabled()) (*this)->model_path_ = Filesystem::absolute(model_path);
76  }
77 
78  GraphBuilder::GraphBuilder(const Function& f, const Dict& opts) {
79  own(new GraphBuilderInternal(f.name(), f, opts));
80  }
81 
82  GraphBuilder::GraphBuilder(const std::string& name, const std::vector<uint8_t>& model_data,
83  const std::string& format, const Dict& opts) {
84  own(new GraphBuilderInternal(name, model_data, format, opts));
85  }
86 
88  return static_cast<GraphBuilderInternal*>(SharedObject::operator->());
89  }
91  return static_cast<const GraphBuilderInternal*>(SharedObject::operator->());
92  }
94  return static_cast<GraphBuilderInternal*>(SharedObject::get());
95  }
96 
97  const std::string& GraphBuilder::name() const {
98  static std::string null = "null";
99  return is_null() ? null : (*this)->name_;
100  }
101 
102  casadi_int GraphBuilder::n_in() const { return (*this)->n_in(); }
103  casadi_int GraphBuilder::n_out() const { return (*this)->n_out(); }
104  std::vector<std::string> GraphBuilder::name_in() const { return (*this)->name_in(); }
105  std::vector<std::string> GraphBuilder::name_out() const { return (*this)->name_out(); }
106  std::vector<casadi_int> GraphBuilder::dimension(const std::string& name) const {
107  return (*this)->node(name).dimension;
108  }
109  std::string GraphBuilder::dtype(const std::string& name) const {
110  return (*this)->node(name).dtype;
111  }
112  std::vector<std::string> GraphBuilder::dimension_param(const std::string& name) const {
113  return (*this)->node(name).dim_params;
114  }
115  std::vector<std::string> GraphBuilder::dynamic_params() const {
116  return (*this)->dynamic_params();
117  }
118 
119  void GraphBuilder::bind_dim(const std::string& param, casadi_int value) {
120  (*this)->bind_dim(param, value);
121  }
122  void GraphBuilder::bind_shape(const std::string& input_name,
123  const std::vector<casadi_int>& shape) {
124  (*this)->bind_shape(input_name, shape);
125  }
126  void GraphBuilder::set(const std::string& input_name, const std::vector<double>& value) {
127  (*this)->set_value(input_name, value);
128  }
129  void GraphBuilder::set(const std::string& input_name, double value) {
130  (*this)->set_value(input_name, std::vector<double>(1, value));
131  }
132 
133  Function GraphBuilder::create(const std::string& name,
134  const std::vector<std::string>& name_in,
135  const std::vector<std::string>& name_out,
136  const Dict& opts) const {
137  return (*this)->create_function(name, name_in, name_out, opts);
138  }
139  Function GraphBuilder::create(const std::string& name, const Dict& opts) const {
140  return (*this)->create_function(name, {}, {}, opts);
141  }
142  void GraphBuilder::export_onnx(const std::string& filename, const Dict& opts) {
143  (*this)->export_onnx(filename, opts);
144  }
145 
146  // ---------- GraphBuilderInternal ----------
147 
149  const std::vector<uint8_t>& model_data,
150  const std::string& format, const Dict& opts)
151  : opts_(opts), name_(name), format_(format), model_data_(model_data) {
152  model_ = GraphModel(format, model_data, opts);
153  model_.fill_metadata(*this);
154  }
155 
156  GraphBuilderInternal::GraphBuilderInternal(const std::string& name, const Function& f,
157  const Dict& opts)
158  : opts_(opts), name_(name), format_("onnx"), fun_(f) {
159  populate_from_function();
160  }
161 
163  }
164 
165  void GraphBuilderInternal::populate_from_function() {
166  nodes_.clear();
167  for (casadi_int i = 0; i < fun_.n_in(); ++i) {
168  Node n;
169  n.name = fun_.name_in(i);
170  n.io = "input";
171  n.dimension = {fun_.size1_in(i), fun_.size2_in(i)};
172  n.dim_params = {"", ""};
173  nodes_.push_back(n);
174  }
175  for (casadi_int i = 0; i < fun_.n_out(); ++i) {
176  Node n;
177  n.name = fun_.name_out(i);
178  n.io = "output";
179  n.dimension = {fun_.size1_out(i), fun_.size2_out(i)};
180  n.dim_params = {"", ""};
181  nodes_.push_back(n);
182  }
183  }
184 
185  std::vector<casadi_int> GraphBuilderInternal::resolved_shape(const Node& n) const {
186  if (n.io == "input") {
187  auto ov = input_shapes_.find(n.name);
188  if (ov != input_shapes_.end()) return ov->second; // explicit override
189  }
190  std::vector<casadi_int> shape;
191  for (size_t k = 0; k < n.dimension.size(); ++k) {
192  casadi_int d = n.dimension[k];
193  if (d < 0) { // dynamic: bound name else default 1
194  auto it = dim_bindings_.find(n.dim_params[k]);
195  d = (it != dim_bindings_.end()) ? it->second : 1;
196  }
197  shape.push_back(d);
198  }
199  return shape;
200  }
201 
202  const Node& GraphBuilderInternal::find(const std::string& name, const std::string& io) const {
203  for (const Node& n : nodes_) if (n.io == io && n.name == name) return n;
204  casadi_error("Graph tensor '" + name + "' (" + io + ") not found in model '" + name_ + "'");
205  }
206 
207  casadi_int GraphBuilderInternal::n_in() const {
208  casadi_int c = 0;
209  for (const Node& n : nodes_) if (n.io == "input") ++c;
210  return c;
211  }
212  casadi_int GraphBuilderInternal::n_out() const {
213  casadi_int c = 0;
214  for (const Node& n : nodes_) if (n.io == "output") ++c;
215  return c;
216  }
217  std::vector<std::string> GraphBuilderInternal::name_in() const {
218  std::vector<std::string> r;
219  for (const Node& n : nodes_) if (n.io == "input") r.push_back(n.name);
220  return r;
221  }
222  std::vector<std::string> GraphBuilderInternal::name_out() const {
223  std::vector<std::string> r;
224  for (const Node& n : nodes_) if (n.io == "output") r.push_back(n.name);
225  return r;
226  }
227  std::vector<casadi_int> GraphBuilderInternal::input_shape(const std::string& name) const {
228  return find(name, "input").dimension;
229  }
230  std::vector<casadi_int> GraphBuilderInternal::output_shape(const std::string& name) const {
231  return find(name, "output").dimension;
232  }
233 
234  std::vector<std::string> GraphBuilderInternal::dynamic_params() const {
235  std::vector<std::string> r;
236  for (const Node& n : nodes_) {
237  for (size_t k = 0; k < n.dimension.size(); ++k) {
238  if (n.dimension[k] < 0 && !n.dim_params[k].empty() &&
239  std::find(r.begin(), r.end(), n.dim_params[k]) == r.end()) {
240  r.push_back(n.dim_params[k]);
241  }
242  }
243  }
244  return r;
245  }
246 
247  Node GraphBuilderInternal::node(const std::string& name) const {
248  for (const Node& n : nodes_) if (n.name == name) return n;
249  casadi_error("Graph tensor '" + name + "' not found in model '" + name_ + "'");
250  }
251 
252  void GraphBuilderInternal::set_value(const std::string& input_name,
253  const std::vector<double>& value) {
254  find(input_name, "input"); // validate the name
255  input_values_[input_name] = value;
256  for (Node& n : nodes_) if (n.io == "input" && n.name == input_name) {
257  n.value = value;
258  n.baked = true;
259  }
260  }
261 
262  void GraphBuilderInternal::bind_shape(const std::string& input_name,
263  const std::vector<casadi_int>& shape) {
264  const Node& t = find(input_name, "input");
265  casadi_assert(shape.size() == t.dimension.size(),
266  "bind_shape: rank mismatch for '" + input_name + "'");
267  input_shapes_[input_name] = shape;
268  // Pinning a named dynamic axis also binds that dim everywhere it appears (e.g. outputs)
269  for (size_t k = 0; k < t.dimension.size(); ++k) {
270  if (t.dimension[k] < 0 && !t.dim_params[k].empty()) dim_bindings_[t.dim_params[k]] = shape[k];
271  }
272  }
273 
275  const std::vector<std::string>& inputs,
276  const std::vector<std::string>& outputs,
277  const Dict& opts) const {
278  bool symbolic = false;
279  std::string backend = "ort";
280  Dict o;
281  for (auto&& op : opts) {
282  if (op.first == "symbolic") symbolic = op.second;
283  else if (op.first == "backend") backend = op.second.to_string();
284  else
285  o[op.first] = op.second;
286  }
287 
288  if (symbolic) {
289  casadi_assert(!model_.is_null(),
290  "GraphBuilder: symbolic create requires a parsed model (build from a file)");
291  Function f = model_.import_symbolic(*this, name);
292  if (o.empty()) return f;
293  std::vector<MX> args = f.mx_in(), res;
294  f.call(args, res, true);
295  return Function(name, args, res, f.name_in(), f.name_out(), o);
296  }
297 
298  // Numeric path: OnnxFunction freezes a snapshot directly from this builder's config
299  casadi_assert(!model_data_.empty(), "GraphBuilder: numeric create requires model bytes");
300  return OnnxFunction::create(backend, name, this, inputs, outputs, o);
301  }
302 
303  void GraphBuilderInternal::export_onnx(const std::string& filename, const Dict& opts) {
304  std::vector<uint8_t> bytes;
305  if (!fun_.is_null()) {
306  GraphModel gm(format_);
307  bytes = gm.export_symbolic(fun_, opts);
308  } else {
309  casadi_assert(!model_data_.empty(), "GraphBuilder: nothing to export");
310  bytes = model_data_;
311  }
312  std::ofstream out(filename, std::ios::binary);
313  casadi_assert(out.good(), "Cannot open output file: " + filename);
314  out.write(reinterpret_cast<const char*>(bytes.data()), bytes.size());
315  }
316 
317  void GraphBuilderInternal::disp(std::ostream& stream, bool more) const {
318  stream << "GraphBuilder '" << name_ << "': " << n_in() << " input(s), "
319  << n_out() << " output(s)";
320  if (!more) return;
321  stream << "\nInputs:";
322  for (const Node& n : nodes_) if (n.io == "input") stream << "\n " << n.get_str();
323  stream << "\nOutputs:";
324  for (const Node& n : nodes_) if (n.io == "output") stream << "\n " << n.get_str();
325  std::vector<std::string> dp = dynamic_params();
326  if (!dp.empty()) {
327  stream << "\nDynamic dimensions:";
328  for (const std::string& p : dp) stream << " " << p;
329  }
330  }
331 
332 } // namespace casadi
static std::string absolute(const std::string &path)
Definition: filesystem.cpp:78
static bool is_enabled()
Definition: filesystem.cpp:83
Function object.
Definition: function.hpp:60
casadi_int size2_out(casadi_int ind) const
Get output dimension.
Definition: function.cpp:991
const MX mx_in(casadi_int ind) const
Get symbolic primitives equivalent to the input expressions.
Definition: function.cpp:1781
casadi_int size1_in(casadi_int ind) const
Get input dimension.
Definition: function.cpp:979
const std::vector< std::string > & name_in() const
Get input scheme.
Definition: function.cpp:1113
const std::string & name() const
Name of the function.
Definition: function.cpp:1504
casadi_int n_out() const
Get the number of function outputs.
Definition: function.cpp:975
casadi_int n_in() const
Get the number of function inputs.
Definition: function.cpp:971
casadi_int size1_out(casadi_int ind) const
Get output dimension.
Definition: function.cpp:987
void call(const std::vector< DM > &arg, std::vector< DM > &res, bool always_inline=false, bool never_inline=false) const
Evaluate the function symbolically or numerically.
Definition: function.cpp:509
casadi_int size2_in(casadi_int ind) const
Get input dimension.
Definition: function.cpp:983
const std::vector< std::string > & name_out() const
Get output scheme.
Definition: function.cpp:1117
SharedObjectInternal * get() const
Get a const pointer to the node.
SharedObjectInternal * operator->() const
Access a member function or object.
Internal class for GraphBuilder.
void export_onnx(const std::string &filename, const Dict &opts)
void bind_shape(const std::string &input_name, const std::vector< casadi_int > &shape)
std::map< std::string, std::vector< double > > input_values_
std::vector< Node > nodes_
Tensor metadata (inputs followed by outputs)
Function fun_
Source Function (export lifecycle); null when built from a model.
const Node & find(const std::string &name, const std::string &io) const
Locate a node by name in a given I/O role (throws if absent)
Node node(const std::string &name) const
std::vector< casadi_int > resolved_shape(const Node &n) const
std::vector< std::string > name_out() const
void set_value(const std::string &input_name, const std::vector< double > &value)
std::map< std::string, casadi_int > dim_bindings_
Pending configuration carried into create()
std::map< std::string, std::vector< casadi_int > > input_shapes_
GraphModel model_
Parsed model backend (import lifecycle); null when built from a Function.
GraphBuilderInternal(const std::string &name, const std::vector< uint8_t > &model_data, const std::string &format, const Dict &opts)
Construct from parsed model bytes of a given format.
std::vector< casadi_int > input_shape(const std::string &name) const
std::vector< casadi_int > output_shape(const std::string &name) const
std::vector< std::string > dynamic_params() const
std::vector< std::string > name_in() const
void disp(std::ostream &stream, bool more) const override
Print a description of the object.
Function create_function(const std::string &name, const std::vector< std::string > &name_in, const std::vector< std::string > &name_out, const Dict &opts) const
GraphBuilder()
Default constructor.
std::vector< std::string > dynamic_params() const
Names of the symbolic/dynamic dimensions in the model.
std::string dtype(const std::string &name) const
Element type name of a tensor (input or output) by name (FLOAT, INT64, ...)
GraphBuilderInternal * operator->()
void bind_dim(const std::string &param, casadi_int value)
Bind a symbolic/dynamic dimension to a concrete size.
casadi_int n_in() const
Declared shape of a tensor (input or output) by name; -1 for dynamic dimensions.
std::vector< std::string > name_in() const
Declared shape of a tensor (input or output) by name; -1 for dynamic dimensions.
Function create() const
Freeze into an evaluable Function, default naming.
std::vector< casadi_int > dimension(const std::string &name) const
Declared shape of a tensor (input or output) by name; -1 for dynamic dimensions.
void set(const std::string &input_name, const std::vector< double > &value)
Bake a fixed value into an input; it is fed at create() and not exposed as a Function input.
void bind_shape(const std::string &input_name, const std::vector< casadi_int > &shape)
Pin the full shape of an input.
GraphBuilderInternal * get() const
const std::string & name() const
Name of the model.
casadi_int n_out() const
Declared shape of a tensor (input or output) by name; -1 for dynamic dimensions.
void export_onnx(const std::string &filename, const Dict &opts=Dict())
Export to an ONNX model file.
std::vector< std::string > name_out() const
Declared shape of a tensor (input or output) by name; -1 for dynamic dimensions.
std::vector< std::string > dimension_param(const std::string &name) const
Per-axis symbolic dimension name of a tensor by name ("" for static axes)
Format-agnostic handle to a parsed computational-graph model.
std::vector< uint8_t > export_symbolic(const Function &f, const Dict &opts=Dict())
Symbolic export: serialize a CasADi Function to model bytes.
Definition: graph_model.cpp:62
void fill_metadata(GraphBuilderInternal &gb) const
Populate a builder's Node metadata from the parsed model.
Definition: graph_model.cpp:53
Function import_symbolic(const GraphBuilderInternal &gb, const std::string &name) const
Symbolic import: rebuild the graph as a CasADi Function.
Definition: graph_model.cpp:56
static Function create(const std::string &solver, const std::string &name, const GraphBuilderInternal *gb, const std::vector< std::string > &inputs, const std::vector< std::string > &outputs, const Dict &opts)
Plugin factory.
The casadi namespace.
Definition: archiver.cpp:28
static std::string format_from_path(const std::string &path)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
T dot(const std::vector< T > &a, const std::vector< T > &b)
std::vector< casadi_int > path(const std::vector< casadi_int > &map, casadi_int i_start)
std::string filename(const std::string &path)
Definition: ghc.cpp:55
Metadata for one graph tensor (graph input or output)
void disp(std::ostream &stream, bool more=false) const
std::string get_str(bool more=false) const
std::vector< casadi_int > dimension
Declared shape, -1 for dynamic dimensions.
std::string name
Tensor name.
std::string io
"input" or "output"
std::string dtype
Element type name (FLOAT, INT64, ...)
std::vector< std::string > dim_params
Symbolic name per axis ("" if static)
bool baked
True if a fixed value was set() for this input.