ort_interface.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 "ort_interface.hpp"
26 #include <ort_runtime_str.h>
27 
28 namespace casadi {
29 
30  extern "C"
31  int CASADI_ONNX_ORT_EXPORT
32  casadi_register_onnx_ort(OnnxFunction::Plugin* plugin) {
33  plugin->creator = OnnxRuntimeInterface::creator;
34  plugin->name = "ort";
35  plugin->doc = OnnxRuntimeInterface::meta_doc.c_str();
36  plugin->version = CASADI_VERSION;
37  plugin->options = &OnnxRuntimeInterface::options_;
38  plugin->deserialize = &OnnxRuntimeInterface::deserialize;
39  #ifdef ONNXRUNTIME_ADAPTOR
40  char buffer[400];
41  int ret = onnxruntime_adaptor_load(buffer, sizeof(buffer));
42  if (ret!=0) {
43  casadi_warning("Failed to load ONNX Runtime adaptor: " + std::string(buffer) + ".");
44  return 1;
45  }
46  #endif
47  return 0;
48  }
49 
50  extern "C"
51  void CASADI_ONNX_ORT_EXPORT casadi_load_onnx_ort() {
53  }
54 
55  const std::string OnnxRuntimeInterface::meta_doc =
56  "Black-box ONNX model evaluation through Microsoft's ONNX Runtime.\n"
57  #ifdef ONNXRUNTIME_ADAPTOR
58  // No runtime ships with CasADi; the adaptor opens the one named here
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"
62  #endif
63  ;
64 
65  const Options OnnxRuntimeInterface::options_
67  {{"provider",
68  {OT_STRING, "Execution provider ('CPU', 'CUDA') [CPU]"}}
69  }
70  };
71 
73  const GraphBuilderInternal* gb,
74  const std::vector<std::string>& inputs,
75  const std::vector<std::string>& outputs)
76  : OnnxFunction(name, gb, inputs, outputs),
77  ort_api_(OrtGetApiBase()->GetApi(ORT_API_VERSION)), provider_("CPU") {
78  casadi_assert(ort_api_ != nullptr, "Failed to obtain the ONNX Runtime API");
79  }
80 
81  void OnnxRuntimeInterface::ort_check(OrtStatus* status, const std::string& what) const {
82  if (!status) return;
83  std::string msg = ort_api_->GetErrorMessage(status);
84  ort_api_->ReleaseStatus(status);
85  casadi_error("ONNX Runtime error (" + what + "): " + msg);
86  }
87 
88  void OnnxRuntimeInterface::build_prob() {
89  // Metadata (all_in_, in_, out_, in_src_, in_val_, resolved shapes) is frozen by the base
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();
93  // Inputs: ALL model inputs (the base in_src_/in_val_ feed map covers them)
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);
100  }
101  // Outputs: only the exposed selection (ORT runs just the needed subgraph)
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);
108  }
109  prob_.n_in = all_in_.size();
110  prob_.n_out = out_.size();
111  prob_.input_names = in_names_c_.data();
112  prob_.output_names = out_names_c_.data();
113  prob_.in_src = in_src_.data();
114  prob_.in_val = in_val_.data();
115  prob_.in_elem_type = in_elem_.data();
116  prob_.out_elem_type = out_elem_.data();
117  prob_.in_ndim = in_ndim_.data();
118  prob_.out_ndim = out_ndim_.data();
119  prob_.in_dims = in_dims_.data();
120  prob_.out_dims = out_dims_.data();
121  prob_.in_numel = in_numel_.data();
122  prob_.out_numel = out_numel_.data();
123  prob_.model_data = model_data_.data();
124  prob_.model_size = static_cast<casadi_int>(model_data_.size());
125  }
126 
127  void OnnxRuntimeInterface::init(const Dict& opts) {
128  OnnxFunction::init(opts);
129 
130  for (auto&& op : opts) {
131  if (op.first == "provider") provider_ = op.second.to_string();
132  }
133 
134  build_prob();
135  }
136 
137  int OnnxRuntimeInterface::init_mem(void* mem) const {
138  if (OnnxFunction::init_mem(mem)) return 1;
139  auto* m = static_cast<OnnxRuntimeMemory*>(mem);
140  casadi_assert(casadi_onnxruntime_init(&m->d, &prob_) == 0,
141  "Failed to create ONNX Runtime session for '" + name_ + "'");
142  return 0;
143  }
144 
145  void OnnxRuntimeInterface::free_mem(void* mem) const {
146  auto* m = static_cast<OnnxRuntimeMemory*>(mem);
147  casadi_onnxruntime_data& d = m->d;
148  // Release the reusable scaffolding (built lazily by the runtime's prepare step) ...
149  if (d.inv) {
150  for (casadi_int i = 0; i < prob_.n_in; ++i) if (d.inv[i]) ort_api_->ReleaseValue(d.inv[i]);
151  free(d.inv);
152  }
153  if (d.buf) {
154  for (casadi_int i = 0; i < prob_.n_in; ++i) if (d.buf[i]) free(d.buf[i]);
155  free(d.buf);
156  }
157  free(d.outv);
158  free(d.row);
159  // ... then the session handles created by init_mem
160  if (d.mem) ort_api_->ReleaseMemoryInfo(d.mem);
161  if (d.session) ort_api_->ReleaseSession(d.session);
162  if (d.env) ort_api_->ReleaseEnv(d.env);
163  delete m;
164  }
165 
167  // Free the memory pool while this is still the dynamic type, so the virtual
168  // free_mem() above runs (the base ~OnnxFunction::clear_mem() is too late).
169  clear_mem();
170  }
171 
174  s.version("OnnxRuntimeInterface", 1);
175  s.pack("OnnxRuntimeInterface::provider", provider_);
176  }
177 
179  : OnnxFunction(s),
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_);
184  build_prob();
185  }
186 
187  int OnnxRuntimeInterface::eval(const double** arg, double** res,
188  casadi_int* iw, double* w, void* mem) const {
189  auto* m = static_cast<OnnxRuntimeMemory*>(mem);
190  return casadi_onnxruntime_solve(&m->d, &prob_, arg, res);
191  }
192 
194  g.add_include("onnxruntime_c_api.h");
195  g.auxiliaries << g.sanitize_source(ort_runtime_str, {});
196  }
197 
199  // Embed the model and its metadata (all inputs + the in_src map), then call the runtime
200  std::vector<std::string> in_nm, out_nm;
201  for (const OnnxTensorInfo& t : all_in_) in_nm.push_back(t.name);
202  for (const OnnxTensorInfo& t : out_) out_nm.push_back(t.name);
203  std::string model = g.constant(std::vector<char>(model_data_.begin(), model_data_.end()));
204  std::string in_names = g.constant(in_nm), out_names = g.constant(out_nm);
205  std::string in_src = g.constant(in_src_);
206  std::string in_val = g.constant(in_val_.empty() ? std::vector<double> {0.0} : in_val_);
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_);
212 
213  g << "static struct casadi_onnxruntime_prob prob = {"
214  << all_in_.size() << ", " << out_.size() << ", "
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";
225  }
226 
227 } // namespace casadi
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.
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.
The casadi namespace.
Definition: archiver.cpp:28
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.
OrtMemoryInfo * mem
Definition: ort_runtime.hpp:90
const long long * in_ndim
Definition: ort_runtime.hpp:69
const long long * out_ndim
Definition: ort_runtime.hpp:70
const long long * out_elem_type
Definition: ort_runtime.hpp:68
const char ** output_names
Definition: ort_runtime.hpp:64
const unsigned char * model_data
Definition: ort_runtime.hpp:75
const long long * in_numel
Definition: ort_runtime.hpp:73
const long long * in_dims
Definition: ort_runtime.hpp:71
const long long * in_elem_type
Definition: ort_runtime.hpp:67
const long long * out_dims
Definition: ort_runtime.hpp:72
const double * in_val
Definition: ort_runtime.hpp:66
const long long * in_src
Definition: ort_runtime.hpp:65
const long long * out_numel
Definition: ort_runtime.hpp:74
const char ** input_names
Definition: ort_runtime.hpp:63