ort_runtime.hpp
1 //
2 // MIT No Attribution
3 //
4 // Copyright (C) 2010-2023 Joel Andersson, Joris Gillis, Moritz Diehl, KU Leuven.
5 //
6 // Permission is hereby granted, free of charge, to any person obtaining a copy of this
7 // software and associated documentation files (the "Software"), to deal in the Software
8 // without restriction, including without limitation the rights to use, copy, modify,
9 // merge, publish, distribute, sublicense, and/or sell copies of the Software, and to
10 // permit persons to whom the Software is furnished to do so.
11 //
12 // THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED,
13 // INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A
14 // PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT
15 // HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION
16 // OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE
17 // SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
18 //
19 
20 
21 /* Runtime support for the black-box ONNX Runtime interface. Shared by the plugin's eval()
22  * and by generated C code after applying C-REPLACE directives. Both need the ONNX Runtime C API
23  * and link libonnxruntime. Inputs/outputs are double (casadi_real), indices long long.
24  * Every numeric ONNX tensor type is supported by converting to/from double. */
25 
26 // FILTER-MACROS OFF
27 // C-REPLACE "const_cast<const OrtValue* const*>" "(const OrtValue* const*) "
28 // C-REPLACE "static_cast<CTYPE>" "(CTYPE) "
29 // C-REPLACE "static_cast<CTYPE*>" "(CTYPE*) "
30 // C-REPLACE "static_cast<ONNXTensorElementDataType>" "(ONNXTensorElementDataType) "
31 // C-REPLACE "static_cast<OrtValue**>" "(OrtValue**) "
32 // C-REPLACE "static_cast<const CTYPE*>" "(const CTYPE*) "
33 // C-REPLACE "static_cast<const uint16_t*>" "(const uint16_t*) "
34 // C-REPLACE "static_cast<double>" "(double) "
35 // C-REPLACE "static_cast<double*>" "(double*) "
36 // C-REPLACE "static_cast<float>" "(float) "
37 // C-REPLACE "static_cast<int>" "(int) "
38 // C-REPLACE "static_cast<int32_t>" "(int32_t) "
39 // C-REPLACE "static_cast<int64_t>" "(int64_t) "
40 // C-REPLACE "static_cast<int64_t*>" "(int64_t*) "
41 // C-REPLACE "static_cast<size_t>" "(size_t) "
42 // C-REPLACE "static_cast<uint16_t>" "(uint16_t) "
43 // C-REPLACE "static_cast<uint16_t*>" "(uint16_t*) "
44 // C-REPLACE "static_cast<uint32_t>" "(uint32_t) "
45 // C-REPLACE "static_cast<void**>" "(void**) "
46 
47 #ifndef CASADI_ORT_RUNTIME_HPP
48 #define CASADI_ORT_RUNTIME_HPP
49 
50 #include <onnxruntime_c_api.h>
51 #include <math.h>
52 #include <stdint.h>
53 #include <stdlib.h>
54 #include <string.h>
55 
56 /* Model metadata; dims are flattened, with in_ndim/out_ndim giving the per-tensor counts.
57  * ORT requires every model input to be fed, so n_in/inputs cover ALL model inputs; in_src[i]
58  * selects what feeds input i: >=0 a caller arg index, -1 the default (NaN/0), -2 a baked
59  * constant taken from in_val (flat over all inputs, numel each). outputs cover only the
60  * selected (exposed) outputs. */
62  long long n_in, n_out;
63  const char** input_names;
64  const char** output_names;
65  const long long* in_src;
66  const double* in_val;
67  const long long* in_elem_type;
68  const long long* out_elem_type;
69  const long long* in_ndim;
70  const long long* out_ndim;
71  const long long* in_dims;
72  const long long* out_dims;
73  const long long* in_numel;
74  const long long* out_numel;
75  const unsigned char* model_data;
76  long long model_size;
77 };
78 
79 /* Per-instance ONNX Runtime state. The session/env/mem handles are created once by
80  * casadi_onnxruntime_init (the "prepare" step of the ONNX backend contract). The remaining
81  * fields hold the value-independent per-eval scaffolding built once by casadi_onnxruntime_prepare
82  * and reused across every casadi_onnxruntime_solve, so repeated evaluation does no allocation:
83  * - buf[i] : persistent typed input buffer (the OrtValue inv[i] is a non-owning view of it)
84  * - inv[i] : input tensor, created once over buf[i]; solve() only overwrites buf[i] contents
85  * - outv : scratch handles owned by Run() and released each solve()
86  * - row : conversion scratch sized to the largest input/output tensor */
88  OrtSession* session;
89  OrtEnv* env;
90  OrtMemoryInfo* mem;
91  OrtValue** inv;
92  OrtValue** outv;
93  void** buf;
94  double* row;
95  int prepared;
96 };
97 
98 static const OrtApi* casadi_onnxruntime_api(void) {
99  return OrtGetApiBase()->GetApi(ORT_API_VERSION);
100 }
101 
102 /* Floating-point ONNX types can hold a NaN sentinel; integer types cannot */
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;
108 }
109 
110 /* IEEE half / bfloat16 <-> double (stored as uint16_t) */
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;
114  float f;
115  if (e == 0) {
116  if (m == 0) {
117  bits = sign; // signed zero
118  } else { // subnormal -> normalize
119  e = 127 - 15 + 1;
120  while ((m & 0x400u) == 0) { m <<= 1; --e; }
121  m &= 0x3ffu;
122  bits = sign | (e << 23) | (m << 13);
123  }
124  } else if (e == 31) {
125  bits = sign | 0x7f800000u | (m << 13); /* inf / nan */
126  } else {
127  bits = sign | ((e + 127 - 15) << 23) | (m << 13); /* normal */
128  }
129  memcpy(&f, &bits, 4);
130  return static_cast<double>(f);
131 }
132 static uint16_t casadi_onnxruntime_d2half(double d) {
133  float f = static_cast<float>(d);
134  uint32_t b;
135  uint16_t s, e;
136  int32_t ee;
137  memcpy(&b, &f, 4);
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; /* underflow -> signed zero */
141  if (ee >= 31) return static_cast<uint16_t>(s | 0x7c00u);/* overflow -> inf */
142  e = static_cast<uint16_t>(ee);
143  return static_cast<uint16_t>(s | (e << 10) | static_cast<uint16_t>((b >> 13) & 0x3ffu));
144 }
145 static double casadi_onnxruntime_bf162d(uint16_t v) {
146  uint32_t u = (static_cast<uint32_t>(v)) << 16;
147  float f;
148  memcpy(&f, &u, 4);
149  return static_cast<double>(f);
150 }
151 static uint16_t casadi_onnxruntime_d2bf16(double d) {
152  float f = static_cast<float>(d);
153  uint32_t u;
154  memcpy(&u, &f, 4);
155  return static_cast<uint16_t>(u >> 16);
156 }
157 
158 /* CasADi matrices are column-major, ONNX tensors row-major: rank-2 differs by a transpose */
159 static void casadi_onnxruntime_to_row(long long ndim, const long long* dims,
160  const double* col, double* row, long long n) {
161  if (ndim == 2) {
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];
164  } else {
165  long long j; for (j = 0; j < n; ++j) row[j] = col[j];
166  }
167 }
168 
169 static void casadi_onnxruntime_from_row(long long ndim, const long long* dims,
170  const double* row, double* col, long long n) {
171  if (ndim == 2) {
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];
174  } else {
175  long long j; for (j = 0; j < n; ++j) col[j] = row[j];
176  }
177 }
178 
179 /* Byte size of one element of ONNX type et; 0 if et is non-numeric (unsupported). */
180 static size_t casadi_onnxruntime_elem_size(long long et) {
181  switch (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;
195  default: return 0;
196  }
197 }
198 
199 /* Pack a row-major double buffer into the caller-owned tensor buffer dst of ONNX type et.
200  * dst must hold casadi_onnxruntime_elem_size(et)*n bytes. Returns 0 on success, 1 if non-numeric. */
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]); \
204  return 0; }
205 static int casadi_onnxruntime_pack_into(long long et, const double* row, long long n, void* dst) {
206  long long k;
207  switch (et) {
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]);
222  return 0; }
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]);
226  return 0; }
227  default: return 1;
228  }
229 }
230 #undef CASADI_ORT_PACK
231 
232 /* Unpack an ONNX tensor of type et into a row-major double buffer. Returns 0 on success. */
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]); \
236  return 0; }
237 static int casadi_onnxruntime_unpack(long long et, const void* td, long long n, double* row) {
238  long long k;
239  switch (et) {
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]);
254  return 0; }
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]);
258  return 0; }
259  default: return 1;
260  }
261 }
262 #undef CASADI_ORT_UNPACK
263 
264 /* Create the session (idempotent). Returns 0 on success. */
265 static int casadi_onnxruntime_init(struct casadi_onnxruntime_data* d,
266  const struct casadi_onnxruntime_prob* p) {
267  const OrtApi* api = casadi_onnxruntime_api();
268  OrtSessionOptions* so = 0;
269  if (d->session) return 0;
270  if (!api) return 1;
271  if (api->CreateEnv(ORT_LOGGING_LEVEL_WARNING, "casadi", &d->env)) return 1;
272  if (api->CreateSessionOptions(&so)) return 1;
273  if (api->CreateSessionFromArray(d->env, p->model_data, static_cast<size_t>(p->model_size),
274  so, &d->session)) {
275  api->ReleaseSessionOptions(so);
276  return 1;
277  }
278  api->ReleaseSessionOptions(so);
279  if (api->CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &d->mem)) return 1;
280  return 0;
281 }
282 
283 /* Build the value-independent per-eval scaffolding once (idempotent): the typed input buffers,
284  * a reusable input tensor (OrtValue) viewing each buffer, the conversion scratch, and the
285  * constant (baked/unwired) input values. Requires the session/mem from casadi_onnxruntime_init.
286  * This is the per-instance half of the ONNX "prepare" step; solve() then does no allocation.
287  * Returns 0 on success. */
288 static int casadi_onnxruntime_prepare(struct casadi_onnxruntime_data* d,
289  const struct casadi_onnxruntime_prob* p) {
290  const OrtApi* api = casadi_onnxruntime_api();
291  long long i, off, voff, k, maxnel = 1;
292  if (d->prepared) return 0;
293  if (!api) return 1;
294  /* one scratch row, sized to the largest input or output tensor */
295  for (i = 0; i < p->n_in; ++i) if (p->in_numel[i] > maxnel) maxnel = p->in_numel[i];
296  for (i = 0; i < p->n_out; ++i) if (p->out_numel[i] > maxnel) maxnel = p->out_numel[i];
297  d->row = static_cast<double*>(malloc(sizeof(double) * static_cast<size_t>(maxnel)));
298  d->inv =
299  static_cast<OrtValue**>(calloc(static_cast<size_t>(p->n_in), sizeof(OrtValue*)));
300  d->outv =
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*)));
303  if (!d->row || !d->inv || !d->outv || !d->buf) return 1;
304  for (i = 0, off = 0, voff = 0; i < p->n_in;
305  off += p->in_ndim[i], voff += p->in_numel[i], ++i) {
306  long long nel = p->in_numel[i], nd = p->in_ndim[i], src = p->in_src[i];
307  size_t esz = casadi_onnxruntime_elem_size(p->in_elem_type[i]);
308  int64_t* shp;
309  if (esz == 0) return 1;
310  d->buf[i] = malloc(esz * static_cast<size_t>(nel));
311  if (!d->buf[i]) return 1;
312  // Pack constant inputs once; fill arg-fed inputs per solve().
313  if (src == -2) { /* baked value */
314  casadi_onnxruntime_to_row(nd, p->in_dims + off, p->in_val + voff, d->row, nel);
315  casadi_onnxruntime_pack_into(p->in_elem_type[i], d->row, nel, d->buf[i]);
316  } else if (src < 0) {
317  // Unwired inputs use NaN where the type allows, else 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;
320  casadi_onnxruntime_pack_into(p->in_elem_type[i], d->row, nel, d->buf[i]);
321  }
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])) {
328  free(shp);
329  return 1;
330  }
331  free(shp);
332  }
333  d->prepared = 1;
334  return 0;
335 }
336 
337 /* Evaluate the model. arg holds the exposed inputs; res the exposed outputs. Returns 0 on success.
338  * Allocation and tensor creation happen once in casadi_onnxruntime_prepare; per call this only
339  * refills the arg-fed input buffers, runs the pre-created input tensors, and unpacks the outputs. */
340 static int casadi_onnxruntime_solve(struct casadi_onnxruntime_data* d,
341  const struct casadi_onnxruntime_prob* p,
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;
346 
347  /* Refill only the inputs fed by the caller; baked/unwired buffers keep their prepared values */
348  for (i = 0, off = 0; i < p->n_in; off += p->in_ndim[i], ++i) {
349  long long nel = p->in_numel[i], nd = p->in_ndim[i], src = p->in_src[i];
350  if (src < 0) continue;
351  if (arg[src]) {
352  casadi_onnxruntime_to_row(nd, p->in_dims + off, arg[src], d->row, nel);
353  } else {
354  /* declared input not supplied -> NaN/0 sentinel, matching the unwired convention */
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;
357  }
358  casadi_onnxruntime_pack_into(p->in_elem_type[i], d->row, nel, d->buf[i]);
359  }
360 
361  if (api->Run(d->session, 0, p->input_names,
362  const_cast<const OrtValue* const*>(d->inv), static_cast<size_t>(p->n_in),
363  p->output_names, static_cast<size_t>(p->n_out), d->outv)) goto cleanup;
364 
365  for (i = 0, off = 0; i < p->n_out; off += p->out_ndim[i], ++i) {
366  void* td;
367  if (res[i]) {
368  if (api->GetTensorMutableData(d->outv[i], &td)) goto cleanup;
369  if (casadi_onnxruntime_unpack(p->out_elem_type[i], td, p->out_numel[i], d->row)) goto cleanup;
370  casadi_onnxruntime_from_row(p->out_ndim[i], p->out_dims + off,
371  d->row, res[i], p->out_numel[i]);
372  }
373  api->ReleaseValue(d->outv[i]); d->outv[i] = 0; /* Run owns these; release after reading */
374  }
375  ret = 0;
376 
377 cleanup:
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);
380 }
381 
382 #endif // CASADI_ORT_RUNTIME_HPP
383 
384 // FILTER-MACROS ON
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