25 #include "onnx_function_impl.hpp"
26 #include "graph_builder_internal.hpp"
27 #include "casadi_misc.hpp"
28 #include "filesystem_impl.hpp"
47 std::vector<std::string> ret;
58 #ifdef CASADI_WITH_THREADSAFE_SYMBOLICS
59 std::mutex OnnxFunction::mutex_solvers_;
69 {
OT_STRING,
"Execution provider for the ONNX runtime backend"}},
71 {
OT_DICT,
"Sizes for symbolic/dynamic tensor dimensions (name -> size)"}},
73 {
OT_DICT,
"Explicit shapes for inputs (name -> shape), overriding the model's"}},
75 {
OT_DICT,
"Baked-in input values (name -> value); these inputs are not exposed"}},
77 {
OT_STRING,
"Symbolic dimension naming the forward seed count [nfwd]"}},
79 {
OT_STRING,
"Symbolic dimension naming the adjoint seed count [nadj]"}}
85 case 1:
return "FLOAT";
case 2:
return "UINT8";
case 3:
return "INT8";
86 case 4:
return "UINT16";
case 5:
return "INT16";
case 6:
return "INT32";
87 case 7:
return "INT64";
case 8:
return "STRING";
case 9:
return "BOOL";
88 case 10:
return "FLOAT16";
case 11:
return "DOUBLE";
case 12:
return "UINT32";
89 case 13:
return "UINT64";
case 16:
return "BFLOAT16";
90 default:
return "TYPE" +
str(t);
95 static const std::map<std::string, casadi_int> m = {
96 {
"FLOAT", 1}, {
"UINT8", 2}, {
"INT8", 3}, {
"UINT16", 4}, {
"INT16", 5},
97 {
"INT32", 6}, {
"INT64", 7}, {
"STRING", 8}, {
"BOOL", 9}, {
"FLOAT16", 10},
98 {
"DOUBLE", 11}, {
"UINT32", 12}, {
"UINT64", 13}, {
"BFLOAT16", 16}};
99 auto it = m.find(name);
100 return it != m.end() ? it->second : 0;
105 const std::vector<std::string>& req) {
106 if (req.empty())
return all;
107 std::vector<OnnxTensorInfo> sel;
108 for (
const std::string& name : req) {
111 if (t.name == name) { sel.push_back(t); found =
true;
break; }
112 casadi_assert(found,
"ONNX tensor '" + name +
"' not found in model");
118 const std::vector<std::string>& inputs,
119 const std::vector<std::string>& outputs)
129 std::vector<OnnxTensorInfo> all_out;
137 if (n.io ==
"input") {
141 all_out.push_back(t);
155 const std::vector<OnnxTensorInfo>& v) {
156 std::vector<std::string> names;
157 std::vector<std::vector<casadi_int>> shapes;
158 std::vector<casadi_int> elem_types, numels;
160 names.push_back(t.name);
161 shapes.push_back(t.shape);
162 elem_types.push_back(t.elem_type);
163 numels.push_back(t.numel);
165 s.
pack(d +
"::names", names);
166 s.
pack(d +
"::shapes", shapes);
167 s.
pack(d +
"::elem_types", elem_types);
168 s.
pack(d +
"::numels", numels);
171 std::vector<OnnxTensorInfo>& v) {
172 std::vector<std::string> names;
173 std::vector<std::vector<casadi_int>> shapes;
174 std::vector<casadi_int> elem_types, numels;
175 s.
unpack(d +
"::names", names);
176 s.
unpack(d +
"::shapes", shapes);
177 s.
unpack(d +
"::elem_types", elem_types);
178 s.
unpack(d +
"::numels", numels);
180 for (
size_t k = 0; k < names.size(); ++k)
181 v.push_back(
OnnxTensorInfo{names[k], shapes[k], elem_types[k], numels[k]});
198 s.
pack(
"OnnxFunction::model_inputs",
200 s.
pack(
"OnnxFunction::model_outputs",
213 int version = s.
version(
"OnnxFunction", 1, 2);
215 s.
unpack(
"OnnxFunction::model_data", bytes);
222 std::vector<std::string> mi, mo;
223 s.
unpack(
"OnnxFunction::model_inputs", mi);
225 s.
unpack(
"OnnxFunction::model_outputs", mo);
244 for (
auto&& op : opts) {
245 if (op.first ==
"fwd_dim")
fwd_dim_ = op.second.to_string();
246 else if (op.first ==
"adj_dim")
adj_dim_ = op.second.to_string();
247 if (op.first ==
"provider" || op.first ==
"fwd_dim" || op.first ==
"adj_dim")
251 std::vector<OnnxTensorInfo> exposed;
255 Dict inferred_opts = opts;
256 if (!opts.count(
"is_diff_in")) {
257 std::vector<bool> inferred;
259 auto infer = [&](
const std::string& kind,
const std::set<std::string>& inputs,
260 const std::set<std::string>& outputs) ->
bool {
262 for (
const auto& y :
out_) {
263 signature |= kind ==
"fwd" ? outputs.count(
"fwd_" + y.name)
264 : inputs.count(
"adj_" + y.name);
267 std::vector<bool> mask;
268 for (
const auto& x :
in_) {
269 mask.push_back(kind ==
"fwd" ? inputs.count(
"fwd_" + x.name)
270 : outputs.count(
"adj_" + x.name));
272 casadi_assert(!found || mask == inferred,
273 "ONNX derivative signatures disagree on differentiable inputs; specify is_diff_in");
278 for (
const std::string kind : {
"fwd",
"adj"}) {
283 std::set<std::string> inputs, outputs;
285 (node.io ==
"input" ? inputs : outputs).insert(node.name);
287 infer(kind, inputs, outputs);
289 if (found) inferred_opts[
"is_diff_in"] = inferred;
300 for (casadi_int j = 0; j < static_cast<casadi_int>(
in_.size()); ++j)
301 if (
in_[j].name == t.name) { src = j;
break; }
304 casadi_assert(
static_cast<casadi_int
>(bv->second.size()) == t.numel,
305 "Baked value for '" + t.name +
"' has " +
str(bv->second.size())
306 +
" elements, expected " +
str(t.numel));
308 for (
double v : bv->second)
in_val_.push_back(v);
321 casadi_int numel = 1;
322 for (casadi_int d : shape) numel *= d;
327 const std::vector<std::string>& inames,
const std::vector<std::string>& onames,
328 const std::vector<Sparsity>& in_sp,
const std::vector<Sparsity>& out_sp,
329 const Dict& dim_bind,
const Dict& opts)
const {
336 for (
auto&& d : dim_bind) bi->
dim_bindings_[d.first] = d.second;
339 std::set<std::string> inputs, outputs;
341 (n.io ==
"input" ? inputs : outputs).insert(n.name);
344 "ONNX derivative '" + bi->
name_ +
"' has an incomplete " + kind +
" signature");
345 std::vector<std::string> cin, con;
346 for (
const std::string& nm : inames)
if (inputs.count(nm)) cin.push_back(nm);
347 for (
const std::string& nm : onames)
if (outputs.count(nm)) con.push_back(nm);
350 if (n.dimension.size() != 2)
continue;
351 const auto& names = n.io ==
"input" ? inames : onames;
352 const auto& sparsities = n.io ==
"input" ? in_sp : out_sp;
353 for (
size_t i = 0; i < names.size(); ++i) {
354 if (names[i] != n.name)
continue;
355 for (
size_t k = 0; k < 2; ++k) {
356 const std::string& param = n.dim_params[k];
357 if (n.dimension[k] < 0 && !param.empty() && !bi->
dim_bindings_.count(param))
358 bi->
dim_bindings_[param] = k == 0 ? sparsities[i].size1() : sparsities[i].size2();
363 std::vector<bool> din(inames.size(),
true), dout(onames.size(),
true);
364 auto di = opts.find(
"is_diff_in"), do_ = opts.find(
"is_diff_out");
365 if (di != opts.end()) {
366 din = di->second.to_bool_vector();
368 for (
size_t i = 0; i <
in_.size(); ++i) din[i] =
diff_in(i);
369 for (
size_t j = 0; j <
out_.size(); ++j) din[
in_.size() + j] =
diff_out(j);
371 if (do_ != opts.end()) {
372 dout = do_->second.to_bool_vector();
373 }
else if (kind ==
"jac") {
374 for (
size_t j = 0; j <
out_.size(); ++j)
375 for (
size_t i = 0; i <
in_.size(); ++i)
378 std::vector<bool> cdi, cdo;
379 for (
size_t i = 0; i < inames.size(); ++i)
if (inputs.count(inames[i])) cdi.push_back(din[i]);
380 for (
size_t i = 0; i < onames.size(); ++i)
if (outputs.count(onames[i])) cdo.push_back(dout[i]);
382 copts[
"is_diff_in"] = cdi;
383 copts[
"is_diff_out"] = cdo;
385 for (
size_t i = 0; i < inames.size(); ++i) {
386 if (!inputs.count(inames[i]))
continue;
388 "ONNX derivative input '" + inames[i] +
"' has an incompatible shape");
390 for (
size_t i = 0; i < onames.size(); ++i) {
391 if (!outputs.count(onames[i]))
continue;
393 "ONNX derivative output '" + onames[i] +
"' has an incompatible shape");
396 std::map<std::string, MX> m;
397 std::vector<MX> args(inames.size());
398 for (
size_t i = 0; i < inames.size(); ++i) {
400 Sparsity sp = inputs.count(inames[i]) ? in_sp[i] :
Sparsity(in_sp[i].size());
401 args[i] =
MX::sym(inames[i], sp);
402 m[inames[i]] = args[i];
405 for (
const std::string& nm : cin) gin.push_back(m[nm]);
406 std::vector<MX> gout = g(gin);
407 std::map<std::string, MX> mo;
408 for (
size_t k = 0; k < con.size(); ++k) mo[con[k]] = gout[k];
409 std::vector<MX> outs(onames.size());
410 for (
size_t j = 0; j < onames.size(); ++j) {
411 auto it = mo.find(onames[j]);
412 outs[j] = it != mo.end() ? it->second :
MX::zeros(out_sp[j]);
415 wopts[
"is_diff_in"] = din;
416 wopts[
"is_diff_out"] = dout;
417 auto it = opts.find(
"derivative_of");
418 if (it != opts.end()) wopts[
"derivative_of"] = it->second;
419 return Function(name, args, outs, inames, onames, wopts);
429 const std::set<std::string>& inputs,
const std::set<std::string>& outputs)
const {
431 bool any_in =
false, any_out =
false;
432 for (
size_t i = 0; i <
in_.size(); ++i) {
435 if (kind ==
"fwd" && !inputs.count(pref +
in_[i].name))
return false;
436 if (kind ==
"adj" && !outputs.count(pref +
in_[i].name))
return false;
438 for (
size_t j = 0; j <
out_.size(); ++j) {
439 if (
diff_out(j) && !outputs.count(
"jac_" +
out_[j].name +
"_" +
in_[i].name))
444 for (
size_t j = 0; j <
out_.size(); ++j) {
447 if (kind ==
"fwd" && !outputs.count(pref +
out_[j].name))
return false;
448 if (kind ==
"adj" && !inputs.count(pref +
out_[j].name))
return false;
450 return any_in && any_out;
460 const std::vector<std::string>& inames,
461 const std::vector<std::string>& onames,
462 const Dict& opts)
const {
463 std::vector<Sparsity> isp, osp;
477 Dict db; db[dim] = nfwd;
478 return wrap_derivative(
"fwd", name, inames, onames, isp, osp, db, opts);
488 const std::vector<std::string>& inames,
489 const std::vector<std::string>& onames,
490 const Dict& opts)
const {
491 std::vector<Sparsity> isp, osp;
505 Dict db; db[dim] = nadj;
506 return wrap_derivative(
"adj", name, inames, onames, isp, osp, db, opts);
516 const std::vector<std::string>& inames,
517 const std::vector<std::string>& onames,
518 const Dict& opts)
const {
519 std::vector<Sparsity> isp, osp;
534 const std::vector<std::string>& inputs,
535 const std::vector<std::string>& outputs,
541 const std::vector<uint8_t>& model_data,
542 const std::vector<std::string>& inputs,
543 const std::vector<std::string>& outputs,
549 for (
auto&& op : opts) {
550 if (op.first ==
"dim_bindings") {
551 for (
auto&& d :
static_cast<Dict>(op.second)) bi->
dim_bindings_[d.first] = d.second;
552 }
else if (op.first ==
"input_shapes") {
553 for (
auto&& d :
static_cast<Dict>(op.second))
555 }
else if (op.first ==
"input_values") {
556 for (
auto&& d :
static_cast<Dict>(op.second))
559 fopts[op.first] = op.second;
Helper class for Serialization.
void unpack(Sparsity &e)
Reconstruct an object from the input stream.
void version(const std::string &name, int v)
static std::string parent_path(const std::string &path)
static std::string filename(const std::string &path)
static std::string ensure_trailing_slash(const std::string &path)
static bool exists(const std::string &path)
Internal class for Function.
std::string diff_prefix(const std::string &prefix) const
Determine prefix for differentiated functions.
void init(const Dict &opts) override
Initialize.
void serialize_body(SerializingStream &s) const override
Serialize an object without type information.
bool has_derivative() const
Can derivatives be calculated in any way?
static const Options options_
Options.
void serialize_type(SerializingStream &s) const override
Serialize type information.
virtual Dict info() const
std::string signature(const std::string &fname) const
Code generate the function.
const Sparsity & sparsity_out(casadi_int ind) const
Get sparsity of a given output.
static Function create(FunctionInternal *node)
Create from node.
const Sparsity & sparsity_in(casadi_int ind) const
Get sparsity of a given input.
static MX sym(const std::string &name, casadi_int nrow=1, casadi_int ncol=1)
Create an nrow-by-ncol symbolic primitive.
static MX zeros(casadi_int nrow=1, casadi_int ncol=1)
Create a dense matrix or a matrix with specified sparsity with all entries zero.
Internal class for GraphBuilder.
std::map< std::string, std::vector< double > > input_values_
Dict opts_
Original constructor options, retained for lazy derivative model loading.
const std::vector< Node > & node_list() const
std::string model_path_
Absolute source filename, empty for in-memory models.
std::vector< casadi_int > resolved_shape(const Node &n) const
std::map< std::string, casadi_int > dim_bindings_
Pending configuration carried into create()
std::map< std::string, std::vector< casadi_int > > input_shapes_
std::vector< uint8_t > model_data_
A mutable, format-neutral interface to a computational-graph model.
GraphBuilderInternal * get() const
Function get_forward(casadi_int nfwd, const std::string &name, const std::vector< std::string > &inames, const std::vector< std::string > &onames, const Dict &opts) const override
Return function that calculates forward derivatives.
static Sparsity tensor_sparsity(const std::vector< casadi_int > &shape)
Map an ONNX N-D shape to a 2-D CasADi sparsity (rank<=2 direct, rank>2 flattened column)
std::string derivative_path(const std::string &kind) const
Construct the sibling filename for a derivative of this entry point.
std::string fwd_dim_
Symbolic dimensions naming the forward/adjoint seed counts (bound in get_forward/reverse)
static const std::string infix_
static std::map< std::string, Plugin > solvers_
Plugin registry.
Function get_jacobian(const std::string &name, const std::vector< std::string > &inames, const std::vector< std::string > &onames, const Dict &opts) const override
Return Jacobian of all input elements with respect to all output elements.
std::map< std::string, casadi_int > dim_bindings_
std::set< std::string > model_outputs_
Function get_reverse(casadi_int nadj, const std::string &name, const std::vector< std::string > &inames, const std::vector< std::string > &onames, const Dict &opts) const override
Return function that calculates adjoint derivatives.
std::string model_path_
Filename of this entry point, used for lazy, recursive sibling discovery.
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.
std::map< std::string, std::vector< casadi_int > > input_shapes_
std::vector< OnnxTensorInfo > in_
Metadata for the exposed inputs/outputs (the selection)
Function wrap_derivative(const std::string &kind, const std::string &name, const std::vector< std::string > &inames, const std::vector< std::string > &onames, const std::vector< Sparsity > &in_sp, const std::vector< Sparsity > &out_sp, const Dict &dim_bind, const Dict &opts) const
Wrap an embedded or sibling derivative with CasADi's full signature.
bool has_jacobian() const override
Return Jacobian of all input elements with respect to all output elements.
void serialize_type(SerializingStream &s) const override
Serialize type information.
std::map< std::string, std::vector< double > > input_values_
Baked input values: input name -> value; such inputs are not exposed as Function inputs.
bool diff_out(casadi_int i) const
static ProtoFunction * deserialize(DeserializingStream &s)
Deserialize into a plugin instance (dispatches on the plugin name)
Dict info() const override
void build_io_map()
Compute the per-model-input feed map: in_src_ (arg index / -2 baked / -1 default) + in_val_.
bool has_forward(casadi_int nfwd) const override
Return function that calculates forward derivatives.
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 const Options options_
Options.
static std::string meta_doc
Documentation string.
Dict builder_opts_
Original GraphBuilder constructor options for derivative model loading.
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.
static Function from_model_data(const std::string &solver, const std::string &name, const std::vector< uint8_t > &model_data, const std::vector< std::string > &inputs, const std::vector< std::string > &outputs, const Dict &opts)
Freeze a function from raw model bytes: stage a transient GraphBuilder, then create()
std::set< std::string > model_inputs_
Names of every input/output tensor in the model (for derivative detection)
bool diff_in(casadi_int i) const
True if input/output index is differentiable (is_diff_in/out, default true)
OnnxFunction(const std::string &name, const GraphBuilderInternal *gb, const std::vector< std::string > &inputs, const std::vector< std::string > &outputs)
Construct by freezing a snapshot of a builder's metadata + config (exposed selection)
std::vector< double > in_val_
Baked input values, flat over all_in_ (numel each; placeholder block when not baked)
bool has_reverse(casadi_int nadj) const override
Return function that calculates adjoint derivatives.
void init(const Dict &opts) override
Initialize.
static bool has_plugin(const std::string &pname, bool verbose=false)
Check if a plugin is available or can be loaded.
void serialize_type(SerializingStream &s) const
Serialize type information.
static Plugin & getPlugin(const std::string &pname)
Load and get the creator function.
static ProtoFunction * deserialize(DeserializingStream &s)
Deserialize with type disambiguation.
virtual const char * plugin_name() const=0
static Plugin load_plugin(const std::string &pname, bool register_plugin=true, bool needs_lock=true)
Load a plugin dynamically.
Base class for FunctionInternal and LinsolInternal.
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.
casadi_int size1() const
Get the number of rows.
static Sparsity dense(casadi_int nrow, casadi_int ncol=1)
Create a dense rectangular sparsity pattern *.
casadi_int size2() const
Get the number of columns.
std::pair< casadi_int, casadi_int > size() const
Get the shape.
std::string onnxbackend_doc(const std::string &solver)
Get documentation for an ONNX runtime backend.
void load_onnxbackend(const std::string &solver)
Load an ONNX runtime backend.
bool has_onnxbackend(const std::string &solver)
Check if a given ONNX runtime backend is available.
std::vector< std::string > onnxbackend_solvers()
List available ONNX runtime backends.
std::string onnx_dtype_name(casadi_int t)
Human-readable name of an ONNX element-type enum (1=FLOAT, 11=DOUBLE, 7=INT64, ......
casadi_int onnx_dtype_enum(const std::string &name)
ONNX element-type enum for a human-readable name (inverse of onnx_dtype_name; 0 if unknown)
static std::vector< OnnxTensorInfo > select_tensors(const std::vector< OnnxTensorInfo > &all, const std::vector< std::string > &req)
static void unpack_tensors(DeserializingStream &s, const std::string &d, std::vector< OnnxTensorInfo > &v)
std::string str(const T &v)
String representation, any type.
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
bool all(const std::vector< bool > &v)
Check if all arguments are true.
static void pack_tensors(SerializingStream &s, const std::string &d, const std::vector< OnnxTensorInfo > &v)
std::vector< casadi_int > path(const std::vector< casadi_int > &map, casadi_int i_start)
Metadata for one graph tensor (graph input or output)
Metadata for a single ONNX tensor (input or output); shape is resolved by the builder.
casadi_int elem_type
ONNX element type enum (1=float, 11=double, 7=int64)
std::vector< casadi_int > shape
Resolved shape (dynamic dims bound or set to 1)
std::string name
ONNX tensor name.
casadi_int numel
Number of elements in the resolved shape.