Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
backend_rep.py77 linesDownload Raw Back to backend
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 
codekingpro/portable-devtools · Team Ai