25 #include "onnx_model.hpp"
26 #include <casadi/core/graph_builder_internal.hpp>
34 int CASADI_GRAPHMODEL_ONNX_EXPORT
37 plugin->name =
"onnx";
39 plugin->version = CASADI_VERSION;
50 "ONNX backend for GraphModel: protobuf metadata, symbolic import and export.\n";
55 {
OT_STRING,
"Real type for exported tensors: 'double' (default) or 'float'"}}
68 for (
auto&& op : opts) {
69 if (op.first ==
"casadi_real") set_casadi_real(op.second.to_string());
74 casadi_assert(
model_.ParseFromArray(data.data(),
static_cast<int>(data.size())),
75 "Failed to parse ONNX model from memory");
80 casadi_assert(
has_model_,
"No ONNX model loaded.");
82 casadi_assert(
model_.SerializeToString(&s),
"Failed to serialize ONNX model");
83 return std::vector<uint8_t>(s.begin(), s.end());
87 static Node io_node(
const onnx::ValueInfoProto& vi,
const std::string& io) {
91 const onnx::TypeProto::Tensor& tt = vi.type().tensor_type();
93 const onnx::TensorShapeProto& sh = tt.shape();
94 for (
int k = 0; k < sh.dim_size(); ++k) {
95 const auto& d = sh.dim(k);
96 if (d.has_dim_value()) {
97 n.
dimension.push_back(
static_cast<casadi_int
>(d.dim_value()));
101 n.
dim_params.push_back(d.has_dim_param() ? d.dim_param() :
"");
109 const onnx::GraphProto& g =
model_.graph();
112 std::set<std::string> init_names;
113 for (
int i = 0; i < g.initializer_size(); ++i) init_names.insert(g.initializer(i).name());
115 for (
int i = 0; i < g.input_size(); ++i)
116 if (!init_names.count(g.input(i).name())) gb.
add_node(
io_node(g.input(i),
"input"));
117 for (
int i = 0; i < g.output_size(); ++i) gb.
add_node(
io_node(g.output(i),
"output"));
126 auto it = opts.find(
"casadi_real");
127 if (it != opts.end()) set_casadi_real(it->second.to_string());
Internal class for GraphBuilder.
void clear_nodes()
Drop all tensor descriptors (called by a backend before re-filling)
void add_node(const Node &n)
Append a tensor descriptor (called by a backend during fill_metadata)
const std::map< std::string, casadi_int > & dim_bindings() const
Base interface for format-specific graph-model backends.
static const Options options_
Options.
const std::vector< uint8_t > & model_data() const
Raw model bytes.
virtual void init(const Dict &opts)
Initialize.
std::vector< uint8_t > export_symbolic(const Function &f, const Dict &opts) override
Serialize a CasADi Function as model bytes (symbolic export; mutates the backend's engine)
static const std::string meta_doc
Documentation.
static const Options options_
Options.
Function create(const std::string &name)
Create a CasADi Function from the loaded ONNX graph.
void load_bytes(const std::vector< uint8_t > &data)
Load a graph from serialized ONNX bytes.
static GraphModelInternal * creator(const std::vector< uint8_t > &model_data)
Plugin factory.
std::vector< uint8_t > save_bytes() const
Serialize the loaded graph/model to ONNX bytes.
onnx::ModelProto model_
ONNX model protocol buffer.
Function import_symbolic(const GraphBuilderInternal &gb, const std::string &name) override
Rebuild the graph as a CasADi Function (symbolic import; mutates the backend's engine)
void init(const Dict &opts) override
Initialize.
bool has_model_
Whether a model has been loaded.
void load(const Function &f)
Load a CasADi Function and convert to the ONNX representation.
Onnx(const std::vector< uint8_t > &model_data)
void set_dimension(const std::string &name, casadi_int dim)
Set dimension for a symbolic variable.
void fill_metadata(GraphBuilderInternal &gb) const override
Populate a builder's Node metadata from the parsed model.
static void registerPlugin(const Plugin &plugin, bool needs_lock=true)
Register an integrator in the factory.
std::string onnx_dtype_name(casadi_int t)
Human-readable name of an ONNX element-type enum (1=FLOAT, 11=DOUBLE, 7=INT64, ......
static Node io_node(const onnx::ValueInfoProto &vi, const std::string &io)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
int CASADI_GRAPHMODEL_ONNX_EXPORT casadi_register_graphmodel_onnx(GraphModelInternal::Plugin *plugin)
void CASADI_GRAPHMODEL_ONNX_EXPORT casadi_load_graphmodel_onnx()
Metadata for one graph tensor (graph input or output)
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)