25 #include "ort_interface.hpp"
26 #include <ort_runtime_str.h>
31 int CASADI_ONNX_ORT_EXPORT
36 plugin->version = CASADI_VERSION;
39 #ifdef ONNXRUNTIME_ADAPTOR
41 int ret = onnxruntime_adaptor_load(buffer,
sizeof(buffer));
43 casadi_warning(
"Failed to load ONNX Runtime adaptor: " + std::string(buffer) +
".");
56 "Black-box ONNX model evaluation through Microsoft's ONNX Runtime.\n"
57 #ifdef ONNXRUNTIME_ADAPTOR
59 "Needs the environmental variable CASADI_ONNXRUNTIME_LIB, holding the full path\n"
60 "of an ONNX Runtime shared library -- or, with no path separator, the module name\n"
61 "of one this process has already loaded.\n"
68 {
OT_STRING,
"Execution provider ('CPU', 'CUDA') [CPU]"}}
74 const std::vector<std::string>& inputs,
75 const std::vector<std::string>& outputs)
77 ort_api_(OrtGetApiBase()->GetApi(ORT_API_VERSION)), provider_(
"CPU") {
78 casadi_assert(ort_api_ !=
nullptr,
"Failed to obtain the ONNX Runtime API");
81 void OnnxRuntimeInterface::ort_check(OrtStatus* status,
const std::string& what)
const {
83 std::string msg = ort_api_->GetErrorMessage(status);
84 ort_api_->ReleaseStatus(status);
85 casadi_error(
"ONNX Runtime error (" + what +
"): " + msg);
88 void OnnxRuntimeInterface::build_prob() {
90 in_names_c_.clear(); out_names_c_.clear();
91 in_elem_.clear(); out_elem_.clear(); in_ndim_.clear(); out_ndim_.clear();
92 in_dims_.clear(); out_dims_.clear(); in_numel_.clear(); out_numel_.clear();
94 for (
const OnnxTensorInfo& t :
all_in_) {
95 in_names_c_.push_back(t.name.c_str());
96 in_elem_.push_back(t.elem_type);
97 in_ndim_.push_back(
static_cast<casadi_int
>(t.shape.size()));
98 for (casadi_int d : t.shape) in_dims_.push_back(d);
99 in_numel_.push_back(t.numel);
102 for (
const OnnxTensorInfo& t :
out_) {
103 out_names_c_.push_back(t.name.c_str());
104 out_elem_.push_back(t.elem_type);
105 out_ndim_.push_back(
static_cast<casadi_int
>(t.shape.size()));
106 for (casadi_int d : t.shape) out_dims_.push_back(d);
107 out_numel_.push_back(t.numel);
117 prob_.
in_ndim = in_ndim_.data();
119 prob_.
in_dims = in_dims_.data();
130 for (
auto&& op : opts) {
131 if (op.first ==
"provider") provider_ = op.second.to_string();
140 casadi_assert(casadi_onnxruntime_init(&m->d, &prob_) == 0,
141 "Failed to create ONNX Runtime session for '" +
name_ +
"'");
150 for (casadi_int i = 0; i < prob_.
n_in; ++i)
if (d.
inv[i]) ort_api_->ReleaseValue(d.
inv[i]);
154 for (casadi_int i = 0; i < prob_.
n_in; ++i)
if (d.
buf[i]) free(d.
buf[i]);
160 if (d.
mem) ort_api_->ReleaseMemoryInfo(d.
mem);
162 if (d.
env) ort_api_->ReleaseEnv(d.
env);
174 s.
version(
"OnnxRuntimeInterface", 1);
175 s.
pack(
"OnnxRuntimeInterface::provider", provider_);
180 ort_api_(OrtGetApiBase()->GetApi(ORT_API_VERSION)), provider_(
"CPU") {
181 casadi_assert(ort_api_ !=
nullptr,
"Failed to obtain the ONNX Runtime API");
182 s.
version(
"OnnxRuntimeInterface", 1);
183 s.
unpack(
"OnnxRuntimeInterface::provider", provider_);
188 casadi_int* iw,
double* w,
void* mem)
const {
190 return casadi_onnxruntime_solve(&m->d, &prob_, arg, res);
200 std::vector<std::string> in_nm, out_nm;
207 std::string in_elem = g.
constant(in_elem_), out_elem = g.
constant(out_elem_);
208 std::string in_ndim = g.
constant(in_ndim_), out_ndim = g.
constant(out_ndim_);
209 std::string in_dims = g.
constant(in_dims_.empty() ? std::vector<casadi_int> {0} : in_dims_);
210 std::string out_dims = g.
constant(out_dims_.empty() ? std::vector<casadi_int> {0} : out_dims_);
211 std::string in_numel = g.
constant(in_numel_), out_numel = g.
constant(out_numel_);
213 g <<
"static struct casadi_onnxruntime_prob prob = {"
215 << in_names <<
", " << out_names <<
", "
216 << in_src <<
", " << in_val <<
", "
217 << in_elem <<
", " << out_elem <<
", "
218 << in_ndim <<
", " << out_ndim <<
", "
219 << in_dims <<
", " << out_dims <<
", "
220 << in_numel <<
", " << out_numel <<
", "
221 <<
"(const unsigned char*)" << model <<
", " <<
model_data_.size() <<
"};\n";
222 g <<
"static struct casadi_onnxruntime_data data = {0};\n";
223 g <<
"if (casadi_onnxruntime_init(&data, &prob)) return 1;\n";
224 g <<
"return casadi_onnxruntime_solve(&data, &prob, arg, res);\n";
Helper class for C code generation.
std::string constant(const std::vector< casadi_int > &v)
Represent an array constant; adding it when new.
std::string sanitize_source(const std::string &src, const std::vector< std::string > &inst, bool add_shorthand=true)
Sanitize source files for codegen.
void add_include(const std::string &new_include, bool relative_path=false, const std::string &use_ifdef=std::string())
Add an include file optionally using a relative path "..." instead of an absolute path <....
std::stringstream auxiliaries
Helper class for Serialization.
void unpack(Sparsity &e)
Reconstruct an object from the input stream.
void version(const std::string &name, int v)
Internal class for GraphBuilder.
Black-box ONNX function base; backends (e.g. onnxruntime) derive from this.
std::vector< casadi_int > in_src_
Per all_in_ entry: exposed-arg index (>=0), -2 baked value, or -1 unwired (default)
void serialize_body(SerializingStream &s) const override
Serialize an object without type information.
static const Options options_
Options.
std::vector< OnnxTensorInfo > out_
std::vector< OnnxTensorInfo > all_in_
Metadata for every model input (a runtime backend must feed all of them)
std::vector< uint8_t > model_data_
Serialized ONNX model.
std::vector< double > in_val_
Baked input values, flat over all_in_ (numel each; placeholder block when not baked)
void init(const Dict &opts) override
Initialize.
static const std::string meta_doc
Documentation.
int init_mem(void *mem) const override
Initialize memory block: create the ONNX Runtime session.
~OnnxRuntimeInterface() override
static OnnxFunction * creator(const std::string &name, const GraphBuilderInternal *gb, const std::vector< std::string > &inputs, const std::vector< std::string > &outputs, const Dict &opts)
Plugin factory.
void free_mem(void *mem) const override
Free memory block: tear down the session and reusable scaffolding.
int eval(const double **arg, double **res, casadi_int *iw, double *w, void *mem) const override
Evaluate numerically.
void init(const Dict &opts) override
Initialize.
void codegen_declarations(CodeGenerator &g) const override
Generate code for the declarations.
static const Options options_
Options.
void codegen_body(CodeGenerator &g) const override
Generate code for the body.
void serialize_body(SerializingStream &s) const override
Serialize an object without type information.
static ProtoFunction * deserialize(DeserializingStream &s)
Deserialize into an OnnxRuntimeInterface.
OnnxRuntimeInterface(const std::string &name, const GraphBuilderInternal *gb, const std::vector< std::string > &inputs, const std::vector< std::string > &outputs)
static void registerPlugin(const Plugin &plugin, bool needs_lock=true)
Register an integrator in the factory.
virtual int init_mem(void *mem) const
Initalize memory block.
void clear_mem()
Clear all memory (called from destructor)
Helper class for Serialization.
void version(const std::string &name, int v)
void pack(const Sparsity &e)
Serializes an object to the output stream.
int CASADI_ONNX_ORT_EXPORT casadi_register_onnx_ort(OnnxFunction::Plugin *plugin)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
void CASADI_ONNX_ORT_EXPORT casadi_load_onnx_ort()
Per-checkout ONNX Runtime state (session + reusable eval scaffolding)
Metadata for a single ONNX tensor (input or output); shape is resolved by the builder.
const long long * in_ndim
const long long * out_ndim
const long long * out_elem_type
const char ** output_names
const unsigned char * model_data
const long long * in_numel
const long long * in_dims
const long long * in_elem_type
const long long * out_dims
const long long * out_numel
const char ** input_names