25 #include "graph_builder_internal.hpp"
26 #include "onnx_function_impl.hpp"
27 #include "filesystem_impl.hpp"
36 stream <<
io <<
" " <<
name <<
": " <<
dtype <<
"[";
37 for (
size_t k = 0; k <
dimension.size(); ++k) {
44 if (
baked) stream <<
" (baked)";
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 +
"'");
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);
83 const std::string& format,
const Dict& opts) {
98 static std::string
null =
"null";
99 return is_null() ? null : (*this)->name_;
107 return (*this)->node(
name).dimension;
110 return (*this)->node(
name).dtype;
113 return (*this)->node(
name).dim_params;
116 return (*this)->dynamic_params();
120 (*this)->bind_dim(param, value);
123 const std::vector<casadi_int>& shape) {
124 (*this)->bind_shape(input_name, shape);
127 (*this)->set_value(input_name, value);
130 (*this)->set_value(input_name, std::vector<double>(1, value));
134 const std::vector<std::string>& name_in,
135 const std::vector<std::string>& name_out,
136 const Dict& opts)
const {
140 return (*this)->create_function(
name, {}, {}, opts);
143 (*this)->export_onnx(
filename, opts);
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) {
158 : opts_(opts), name_(name), format_(
"onnx"), fun_(f) {
159 populate_from_function();
165 void GraphBuilderInternal::populate_from_function() {
167 for (casadi_int i = 0; i <
fun_.
n_in(); ++i) {
175 for (casadi_int i = 0; i <
fun_.
n_out(); ++i) {
180 n.dim_params = {
"",
""};
186 if (n.
io ==
"input") {
190 std::vector<casadi_int> shape;
191 for (
size_t k = 0; k < n.
dimension.size(); ++k) {
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_ +
"'");
209 for (
const Node& n :
nodes_)
if (n.io ==
"input") ++c;
214 for (
const Node& n :
nodes_)
if (n.io ==
"output") ++c;
218 std::vector<std::string> r;
219 for (
const Node& n :
nodes_)
if (n.io ==
"input") r.push_back(n.name);
223 std::vector<std::string> r;
224 for (
const Node& n :
nodes_)
if (n.io ==
"output") r.push_back(n.name);
235 std::vector<std::string> r;
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]);
248 for (
const Node& n :
nodes_)
if (n.name == name)
return n;
249 casadi_error(
"Graph tensor '" + name +
"' not found in model '" +
name_ +
"'");
253 const std::vector<double>& value) {
254 find(input_name,
"input");
256 for (
Node& n :
nodes_)
if (n.io ==
"input" && n.name == 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 +
"'");
269 for (
size_t k = 0; k < t.
dimension.size(); ++k) {
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";
281 for (
auto&& op : opts) {
282 if (op.first ==
"symbolic") symbolic = op.second;
283 else if (op.first ==
"backend") backend = op.second.to_string();
285 o[op.first] = op.second;
290 "GraphBuilder: symbolic create requires a parsed model (build from a file)");
292 if (o.empty())
return f;
293 std::vector<MX> args = f.
mx_in(), res;
294 f.
call(args, res,
true);
299 casadi_assert(!
model_data_.empty(),
"GraphBuilder: numeric create requires model bytes");
304 std::vector<uint8_t> bytes;
309 casadi_assert(!
model_data_.empty(),
"GraphBuilder: nothing to export");
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());
318 stream <<
"GraphBuilder '" <<
name_ <<
"': " <<
n_in() <<
" input(s), "
319 <<
n_out() <<
" output(s)";
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();
327 stream <<
"\nDynamic dimensions:";
328 for (
const std::string& p : dp) stream <<
" " << p;
static std::string absolute(const std::string &path)
casadi_int size2_out(casadi_int ind) const
Get output dimension.
const MX mx_in(casadi_int ind) const
Get symbolic primitives equivalent to the input expressions.
casadi_int size1_in(casadi_int ind) const
Get input dimension.
const std::vector< std::string > & name_in() const
Get input scheme.
const std::string & name() const
Name of the function.
casadi_int n_out() const
Get the number of function outputs.
casadi_int n_in() const
Get the number of function inputs.
casadi_int size1_out(casadi_int ind) const
Get output dimension.
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.
casadi_int size2_in(casadi_int ind) const
Get input dimension.
const std::vector< std::string > & name_out() const
Get output scheme.
SharedObjectInternal * get() const
Get a const pointer to the node.
bool is_null() const
Is a null pointer?
void own(SharedObjectInternal *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)
~GraphBuilderInternal() override
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.
std::vector< uint8_t > model_data_
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 ¶m, 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.
void fill_metadata(GraphBuilderInternal &gb) const
Populate a builder's Node metadata from the parsed model.
Function import_symbolic(const GraphBuilderInternal &gb, const std::string &name) const
Symbolic import: rebuild the graph as a CasADi Function.
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.
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)
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.