codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5"""
6Implements ONNX's backend API.
7"""
8
9from onnx.backend.base import BackendRep
10
11from onnxruntime import RunOptions
12
13# Allowlist of RunOptions attributes that are safe to set via the backend API.
14# 'terminate' excluded: setting it True would deny the current inference call.
15# 'training_mode' excluded: silently switches inference behavior in training builds.
16_ALLOWED_RUN_OPTIONS = frozenset(
17 {
18 "log_severity_level",
19 "log_verbosity_level",
20 "logid",
21 "only_execute_path_to_fetches",
22 }
23)
24
25
26class OnnxRuntimeBackendRep(BackendRep):
27 """
28 Wraps an :class:`onnxruntime.InferenceSession` to implement ONNX's
29 :class:`onnx.backend.base.BackendRep` interface for running predictions.
30 """
31
32 def __init__(self, session):
33 """
34 :param session: :class:`onnxruntime.InferenceSession`
35 """
36 self._session = session
37
38 def run(self, inputs, **kwargs): # type: (Any, **Any) -> Tuple[Any, ...]
39 """
40 Computes the prediction.
41 See :meth:`onnxruntime.InferenceSession.run`.
42
43 :param inputs: a list of input arrays (one per model input) or a single
44 array when the model has exactly one input
45 :param kwargs: only a safe subset of :class:`onnxruntime.RunOptions` attributes are
46 accepted; see ``_ALLOWED_RUN_OPTIONS`` for the list
47 :return: list of output arrays
48 """
49
50 options = RunOptions()
51 for k, v in kwargs.items():
52 if k in _ALLOWED_RUN_OPTIONS:
53 setattr(options, k, v)
54 elif hasattr(options, k):
55 raise RuntimeError(
56 f"RunOptions attribute '{k}' is not permitted via the backend API. "
57 f"Allowed attributes: {', '.join(sorted(_ALLOWED_RUN_OPTIONS))}"
58 )
59 # else: silently ignore unknown keys
60
61 if isinstance(inputs, list):
62 inps = {}
63 for i, inp in enumerate(self._session.get_inputs()):
64 inps[inp.name] = inputs[i]
65 outs = self._session.run(None, inps, options)
66 if isinstance(outs, list):
67 return outs
68 else:
69 output_names = [o.name for o in self._session.get_outputs()]
70 return [outs[name] for name in output_names]
71 else:
72 inp = self._session.get_inputs()
73 if len(inp) != 1:
74 raise RuntimeError(f"Model expect {len(inp)} inputs")
75 inps = {inp[0].name: inputs}
76 return self._session.run(None, inps, options)
77 