26 #ifndef CASADI_ONNX_FUNCTION_IMPL_HPP
27 #define CASADI_ONNX_FUNCTION_IMPL_HPP
29 #include "onnx_function.hpp"
30 #include "function_internal.hpp"
31 #include "plugin_interface.hpp"
38 class GraphBuilderInternal;
67 const std::vector<std::string>& inputs,
68 const std::vector<std::string>& outputs);
74 const std::vector<std::string>& inputs,
75 const std::vector<std::string>& outputs,
90 size_t get_n_in()
override {
return in_.size(); }
92 std::string
get_name_in(casadi_int i)
override {
return in_.at(i).name; }
93 std::string
get_name_out(casadi_int i)
override {
return out_.at(i).name; }
100 void init(
const Dict& opts)
override;
105 int eval(
const double** arg,
double** res,
106 casadi_int* iw,
double* w,
void* mem)
const override = 0;
113 bool has_forward(casadi_int nfwd)
const override;
114 Function get_forward(casadi_int nfwd,
const std::string& name,
115 const std::vector<std::string>& inames,
116 const std::vector<std::string>& onames,
117 const Dict& opts)
const override;
118 bool has_reverse(casadi_int nadj)
const override;
119 Function get_reverse(casadi_int nadj,
const std::string& name,
120 const std::vector<std::string>& inames,
121 const std::vector<std::string>& onames,
122 const Dict& opts)
const override;
123 bool has_jacobian()
const override;
124 Function get_jacobian(
const std::string& name,
125 const std::vector<std::string>& inames,
126 const std::vector<std::string>& onames,
127 const Dict& opts)
const override;
150 static Function create(
const std::string& solver,
151 const std::string& name,
153 const std::vector<std::string>& inputs,
154 const std::vector<std::string>& outputs,
164 #ifdef CASADI_WITH_THREADSAFE_SYMBOLICS
165 static std::mutex mutex_solvers_;
173 static Sparsity tensor_sparsity(
const std::vector<casadi_int>& shape);
179 bool diff_in(casadi_int i)
const {
return is_diff_in_.empty() || is_diff_in_.at(i); }
180 bool diff_out(casadi_int i)
const {
return is_diff_out_.empty() || is_diff_out_.at(i); }
183 static Function from_model_data(
const std::string& solver,
const std::string& name,
184 const std::vector<uint8_t>& model_data,
185 const std::vector<std::string>& inputs,
186 const std::vector<std::string>& outputs,
191 Function wrap_derivative(
const std::string& name,
192 const std::vector<std::string>& inames,
193 const std::vector<std::string>& onames,
194 const std::vector<Sparsity>& in_sp,
195 const std::vector<Sparsity>& out_sp,
196 const Dict& dim_bind,
const Dict& opts)
const;
202 std::vector<OnnxTensorInfo>
in_, out_;
217 std::string fwd_dim_ =
"nfwd";
218 std::string adj_dim_ =
"nadj";
Helper class for Serialization.
Internal class for Function.
Internal class for GraphBuilder.
Black-box ONNX function base; backends (e.g. onnxruntime) derive from this.
std::string serialize_base_function() const override
String used to identify the immediate FunctionInternal subclass.
static const std::string infix_
static std::map< std::string, Plugin > solvers_
Plugin registry.
void free_mem(void *mem) const override
Free memory block.
const Options & get_options() const override
Options.
std::string get_name_out(casadi_int i) override
Names of function input and outputs.
int eval(const double **arg, double **res, casadi_int *iw, double *w, void *mem) const override=0
Evaluate numerically.
Sparsity get_sparsity_in(casadi_int i) override
Get sparsity of a given input.
std::vector< casadi_int > in_src_
Per all_in_ entry: exposed-arg index (>=0), -2 baked value, or -1 unwired (default)
std::vector< OnnxTensorInfo > in_
Metadata for the exposed inputs/outputs (the selection)
Sparsity get_sparsity_out(casadi_int i) override
Get sparsity of a given output.
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
size_t get_n_out() override
Are all inputs and outputs scalar.
size_t get_n_in() override
Number of function inputs and outputs.
std::string get_name_in(casadi_int i) override
Names of function input and outputs.
static const Options options_
Options.
static std::string meta_doc
Documentation string.
void * alloc_mem() const override
Create memory block.
bool uses_output() const override
Do the derivative functions need nondifferentiated outputs?
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::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)
std::vector< double > in_val_
Baked input values, flat over all_in_ (numel each; placeholder block when not baked)
Interface for accessing input and output data structures.
Base class for FunctionInternal and LinsolInternal.
Helper class for Serialization.
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)
GenericType::Dict Dict
C++ equivalent of Python's dict or MATLAB's struct.
Function memory with temporary work vectors.
No statically-exposed plugin functions.
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.
Options metadata for a class.