47 #ifndef CASADI_ORT_RUNTIME_HPP
48 #define CASADI_ORT_RUNTIME_HPP
50 #include <onnxruntime_c_api.h>
98 static const OrtApi* casadi_onnxruntime_api(
void) {
99 return OrtGetApiBase()->GetApi(ORT_API_VERSION);
103 static int casadi_onnxruntime_is_float(
long long et) {
104 return et == ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT
105 || et == ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE
106 || et == ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16
107 || et == ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16;
111 static double casadi_onnxruntime_half2d(uint16_t h) {
112 uint32_t sign = (
static_cast<uint32_t
>(h & 0x8000u)) << 16;
113 uint32_t e =
static_cast<uint32_t
>((h >> 10) & 0x1f), m =
static_cast<uint32_t
>(h & 0x3ff), bits;
120 while ((m & 0x400u) == 0) { m <<= 1; --e; }
122 bits = sign | (e << 23) | (m << 13);
124 }
else if (e == 31) {
125 bits = sign | 0x7f800000u | (m << 13);
127 bits = sign | ((e + 127 - 15) << 23) | (m << 13);
129 memcpy(&f, &bits, 4);
130 return static_cast<double>(f);
132 static uint16_t casadi_onnxruntime_d2half(
double d) {
133 float f =
static_cast<float>(d);
138 s =
static_cast<uint16_t
>((b >> 16) & 0x8000u);
139 ee =
static_cast<int32_t
>((b >> 23) & 0xff) - 127 + 15;
140 if (ee <= 0)
return s;
141 if (ee >= 31)
return static_cast<uint16_t
>(s | 0x7c00u);
142 e =
static_cast<uint16_t
>(ee);
143 return static_cast<uint16_t
>(s | (e << 10) | static_cast<uint16_t>((b >> 13) & 0x3ffu));
145 static double casadi_onnxruntime_bf162d(uint16_t v) {
146 uint32_t u = (
static_cast<uint32_t
>(v)) << 16;
149 return static_cast<double>(f);
151 static uint16_t casadi_onnxruntime_d2bf16(
double d) {
152 float f =
static_cast<float>(d);
155 return static_cast<uint16_t
>(u >> 16);
159 static void casadi_onnxruntime_to_row(
long long ndim,
const long long* dims,
160 const double* col,
double* row,
long long n) {
162 long long r, c, d0 = dims[0], d1 = dims[1];
163 for (r = 0; r < d0; ++r)
for (c = 0; c < d1; ++c) row[r * d1 + c] = col[r + c * d0];
165 long long j;
for (j = 0; j < n; ++j) row[j] = col[j];
169 static void casadi_onnxruntime_from_row(
long long ndim,
const long long* dims,
170 const double* row,
double* col,
long long n) {
172 long long r, c, d0 = dims[0], d1 = dims[1];
173 for (r = 0; r < d0; ++r)
for (c = 0; c < d1; ++c) col[r + c * d0] = row[r * d1 + c];
175 long long j;
for (j = 0; j < n; ++j) col[j] = row[j];
180 static size_t casadi_onnxruntime_elem_size(
long long et) {
182 case ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE:
183 case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64:
184 case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64:
return 8;
185 case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT:
186 case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32:
187 case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32:
return 4;
188 case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16:
189 case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16:
190 case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16:
191 case ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16:
return 2;
192 case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8:
193 case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8:
194 case ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL:
return 1;
201 #define CASADI_ORT_PACK(ETYPE, CTYPE) \
202 case ETYPE: { CTYPE* b = static_cast<CTYPE*>(dst); \
203 for (k = 0; k < n; ++k) b[k] = static_cast<CTYPE>(row[k]); \
205 static int casadi_onnxruntime_pack_into(
long long et,
const double* row,
long long n,
void* dst) {
208 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
float)
209 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE,
double)
210 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8, int8_t)
211 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8, uint8_t)
212 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16, int16_t)
213 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16, uint16_t)
214 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, int32_t)
215 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32, uint32_t)
216 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64, int64_t)
217 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64, uint64_t)
218 CASADI_ORT_PACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL, uint8_t)
219 case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16: {
220 uint16_t* b =
static_cast<uint16_t*
>(dst);
221 for (k = 0; k < n; ++k) b[k] = casadi_onnxruntime_d2half(row[k]);
223 case ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16: {
224 uint16_t* b =
static_cast<uint16_t*
>(dst);
225 for (k = 0; k < n; ++k) b[k] = casadi_onnxruntime_d2bf16(row[k]);
230 #undef CASADI_ORT_PACK
233 #define CASADI_ORT_UNPACK(ETYPE, CTYPE) \
234 case ETYPE: { const CTYPE* b = static_cast<const CTYPE*>(td); \
235 for (k = 0; k < n; ++k) row[k] = static_cast<double>(b[k]); \
237 static int casadi_onnxruntime_unpack(
long long et,
const void* td,
long long n,
double* row) {
240 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
float)
241 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE,
double)
242 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8, int8_t)
243 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8, uint8_t)
244 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16, int16_t)
245 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16, uint16_t)
246 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, int32_t)
247 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32, uint32_t)
248 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64, int64_t)
249 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64, uint64_t)
250 CASADI_ORT_UNPACK(ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL, uint8_t)
251 case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16:
252 {
const uint16_t* b =
static_cast<const uint16_t*
>(td);
253 for (k = 0; k < n; ++k) row[k] = casadi_onnxruntime_half2d(b[k]);
255 case ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16:
256 {
const uint16_t* b =
static_cast<const uint16_t*
>(td);
257 for (k = 0; k < n; ++k) row[k] = casadi_onnxruntime_bf162d(b[k]);
262 #undef CASADI_ORT_UNPACK
267 const OrtApi* api = casadi_onnxruntime_api();
268 OrtSessionOptions* so = 0;
271 if (api->CreateEnv(ORT_LOGGING_LEVEL_WARNING,
"casadi", &d->
env))
return 1;
272 if (api->CreateSessionOptions(&so))
return 1;
275 api->ReleaseSessionOptions(so);
278 api->ReleaseSessionOptions(so);
279 if (api->CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &d->
mem))
return 1;
290 const OrtApi* api = casadi_onnxruntime_api();
291 long long i, off, voff, k, maxnel = 1;
297 d->
row =
static_cast<double*
>(malloc(
sizeof(
double) *
static_cast<size_t>(maxnel)));
299 static_cast<OrtValue**
>(calloc(
static_cast<size_t>(p->
n_in),
sizeof(OrtValue*)));
301 static_cast<OrtValue**
>(calloc(
static_cast<size_t>(p->
n_out),
sizeof(OrtValue*)));
302 d->
buf =
static_cast<void**
>(calloc(
static_cast<size_t>(p->
n_in),
sizeof(
void*)));
304 for (i = 0, off = 0, voff = 0; i < p->
n_in;
307 size_t esz = casadi_onnxruntime_elem_size(p->
in_elem_type[i]);
309 if (esz == 0)
return 1;
310 d->
buf[i] = malloc(esz *
static_cast<size_t>(nel));
311 if (!d->
buf[i])
return 1;
314 casadi_onnxruntime_to_row(nd, p->
in_dims + off, p->
in_val + voff, d->
row, nel);
316 }
else if (src < 0) {
318 double fill = casadi_onnxruntime_is_float(p->
in_elem_type[i]) ? NAN : 0.0;
319 for (k = 0; k < nel; ++k) d->
row[k] = fill;
322 shp =
static_cast<int64_t*
>(malloc(
323 sizeof(int64_t) *
static_cast<size_t>(nd > 0 ? nd : 1)));
324 for (k = 0; k < nd; ++k) shp[k] = static_cast<int64_t>(p->
in_dims[off + k]);
325 if (api->CreateTensorWithDataAsOrtValue(d->
mem, d->
buf[i], esz *
static_cast<size_t>(nel),
326 shp,
static_cast<size_t>(nd),
327 static_cast<ONNXTensorElementDataType
>(p->
in_elem_type[i]), &d->
inv[i])) {
342 const double** arg,
double** res) {
343 const OrtApi* api = casadi_onnxruntime_api();
344 long long i, off, k, ret = 1;
345 if (!api || casadi_onnxruntime_prepare(d, p))
return 1;
348 for (i = 0, off = 0; i < p->
n_in; off += p->
in_ndim[i], ++i) {
350 if (src < 0)
continue;
352 casadi_onnxruntime_to_row(nd, p->
in_dims + off, arg[src], d->
row, nel);
355 double fill = casadi_onnxruntime_is_float(p->
in_elem_type[i]) ? NAN : 0.0;
356 for (k = 0; k < nel; ++k) d->
row[k] = fill;
362 const_cast<const OrtValue* const*
>(d->
inv),
static_cast<size_t>(p->
n_in),
365 for (i = 0, off = 0; i < p->
n_out; off += p->
out_ndim[i], ++i) {
368 if (api->GetTensorMutableData(d->
outv[i], &td))
goto cleanup;
373 api->ReleaseValue(d->
outv[i]); d->
outv[i] = 0;
378 for (i = 0; i < p->
n_out; ++i)
if (d->
outv[i]) { api->ReleaseValue(d->
outv[i]); d->
outv[i] = 0; }
379 return static_cast<int>(ret);
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