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
9import os
10import unittest
11
12import packaging.version
13from onnx import ModelProto, helper, version # noqa: F401
14from onnx.backend.base import Backend
15from onnx.checker import check_model
16
17from onnxruntime import InferenceSession, SessionOptions, get_available_providers, get_device
18from onnxruntime.backend.backend_rep import OnnxRuntimeBackendRep
19
20# Allowlist of SessionOptions attributes that are safe to set via the backend API.
21# Dangerous attributes intentionally excluded:
22# optimized_model_filepath — triggers Model::Save(), overwrites arbitrary files
23# profile_file_prefix — writes profiling JSON to arbitrary path
24# enable_profiling — causes uncontrolled file writes to cwd
25_ALLOWED_SESSION_OPTIONS = frozenset(
26 {
27 "enable_cpu_mem_arena",
28 "enable_mem_pattern",
29 "enable_mem_reuse",
30 "execution_mode",
31 "execution_order",
32 "graph_optimization_level",
33 "inter_op_num_threads",
34 "intra_op_num_threads",
35 "log_severity_level",
36 "log_verbosity_level",
37 "logid",
38 "use_deterministic_compute",
39 "use_per_session_threads",
40 }
41)
42
43
44class OnnxRuntimeBackend(Backend):
45 """
46 Implements
47 `ONNX's backend API <https://github.com/onnx/onnx/blob/main/docs/ImplementingAnOnnxBackend.md>`_
48 with *ONNX Runtime*.
49 The backend is mostly used when you need to switch between
50 multiple runtimes with the same API.
51 `Importing models from ONNX to Caffe2 <https://github.com/onnx/tutorials/blob/master/tutorials/OnnxCaffe2Import.ipynb>`_
52 shows how to use *caffe2* as a backend for a converted model.
53 Note: This is not the official Python API.
54 """
55
56 allowReleasedOpsetsOnly = bool(os.getenv("ALLOW_RELEASED_ONNX_OPSET_ONLY", "1") == "1") # noqa: N815
57
58 @classmethod
59 def is_compatible(cls, model, device=None, **kwargs):
60 """
61 Return whether the model is compatible with the backend.
62
63 :param model: unused
64 :param device: None to use the default device or a string (ex: `'CPU'`)
65 :return: boolean
66 """
67 if device is None:
68 device = get_device()
69 return cls.supports_device(device)
70
71 @classmethod
72 def is_opset_supported(cls, model):
73 """
74 Return whether the opset for the model is supported by the backend.
75 When By default only released onnx opsets are allowed by the backend
76 To test new opsets env variable ALLOW_RELEASED_ONNX_OPSET_ONLY should be set to 0
77
78 :param model: Model whose opsets needed to be verified.
79 :return: boolean and error message if opset is not supported.
80 """
81 if cls.allowReleasedOpsetsOnly:
82 for opset in model.opset_import:
83 domain = opset.domain if opset.domain else "ai.onnx"
84 try:
85 key = (domain, opset.version)
86 if key not in helper.OP_SET_ID_VERSION_MAP:
87 error_message = (
88 "Skipping this test as only released onnx opsets are supported."
89 "To run this test set env variable ALLOW_RELEASED_ONNX_OPSET_ONLY to 0."
90 f" Got Domain '{domain}' version '{opset.version}'."
91 )
92 return False, error_message
93 except AttributeError:
94 # for some CI pipelines accessing helper.OP_SET_ID_VERSION_MAP
95 # is generating attribute error. TODO investigate the pipelines to
96 # fix this error. Falling back to a simple version check when this error is encountered
97 if (domain == "ai.onnx" and opset.version > 12) or (domain == "ai.ommx.ml" and opset.version > 2):
98 error_message = (
99 "Skipping this test as only released onnx opsets are supported."
100 "To run this test set env variable ALLOW_RELEASED_ONNX_OPSET_ONLY to 0."
101 f" Got Domain '{domain}' version '{opset.version}'."
102 )
103 return False, error_message
104 return True, ""
105
106 @classmethod
107 def supports_device(cls, device):
108 """
109 Check whether the backend is compiled with particular device support.
110 In particular it's used in the testing suite.
111 """
112 if device == "CUDA":
113 device = "GPU"
114 return "-" + device in get_device() or device + "-" in get_device() or device == get_device()
115
116 @classmethod
117 def prepare(cls, model, device=None, **kwargs):
118 """
119 Load the model and creates an :class:`onnxruntime.backend.backend_rep.OnnxRuntimeBackendRep`
120 ready to be used as a backend.
121
122 :param model: the model to prepare — accepts a file path (str), serialized
123 model (bytes), :class:`onnx.ModelProto`, :class:`onnxruntime.InferenceSession`,
124 or :class:`onnxruntime.backend.backend_rep.OnnxRuntimeBackendRep` (returned as-is)
125 :param device: requested device for the computation,
126 None means the default one which depends on
127 the compilation settings
128 :param kwargs: only a safe subset of :class:`onnxruntime.SessionOptions` attributes are
129 accepted; see ``_ALLOWED_SESSION_OPTIONS`` for the list
130 :return: :class:`onnxruntime.backend.backend_rep.OnnxRuntimeBackendRep`
131 """
132 if isinstance(model, OnnxRuntimeBackendRep):
133 return model
134 elif isinstance(model, InferenceSession):
135 return OnnxRuntimeBackendRep(model)
136 elif isinstance(model, (str, bytes)):
137 options = SessionOptions()
138 for k, v in kwargs.items():
139 if k in _ALLOWED_SESSION_OPTIONS:
140 setattr(options, k, v)
141 elif hasattr(options, k):
142 raise RuntimeError(
143 f"SessionOptions attribute '{k}' is not permitted via the backend API. "
144 f"Allowed attributes: {', '.join(sorted(_ALLOWED_SESSION_OPTIONS))}"
145 )
146 # else: silently ignore unknown keys
147
148 excluded_providers = os.getenv("ORT_ONNX_BACKEND_EXCLUDE_PROVIDERS", default="").split(",")
149 providers = [x for x in get_available_providers() if (x not in excluded_providers)]
150
151 inf = InferenceSession(model, sess_options=options, providers=providers)
152 # backend API is primarily used for ONNX test/validation. As such, we should disable session.run() fallback
153 # which may hide test failures.
154 inf.disable_fallback()
155 if device is not None and not cls.supports_device(device):
156 raise RuntimeError(f"Incompatible device expected '{device}', got '{get_device()}'")
157 return cls.prepare(inf, device, **kwargs)
158 else:
159 # type: ModelProto
160 # check_model serializes the model anyways, so serialize the model once here
161 # and reuse it below in the cls.prepare call to avoid an additional serialization
162 # only works with onnx >= 1.10.0 hence the version check
163 onnx_version = packaging.version.parse(version.version) or packaging.version.Version("0")
164 onnx_supports_serialized_model_check = onnx_version.release >= (1, 10, 0)
165 bin_or_model = model.SerializeToString() if onnx_supports_serialized_model_check else model
166 check_model(bin_or_model)
167 opset_supported, error_message = cls.is_opset_supported(model)
168 if not opset_supported:
169 raise unittest.SkipTest(error_message)
170 # Now bin might be serialized, if it's not we need to serialize it otherwise we'll have
171 # an infinite recursive call
172 bin = bin_or_model
173 if not isinstance(bin, (str, bytes)):
174 bin = bin.SerializeToString()
175 return cls.prepare(bin, device, **kwargs)
176
177 @classmethod
178 def run_model(cls, model, inputs, device=None, **kwargs):
179 """
180 Compute the prediction.
181
182 :param model: the model to run — accepts a file path (str), serialized
183 model (bytes), :class:`onnx.ModelProto`, :class:`onnxruntime.InferenceSession`,
184 or :class:`onnxruntime.backend.backend_rep.OnnxRuntimeBackendRep`
185 :param inputs: inputs
186 :param device: requested device for the computation,
187 None means the default one which depends on
188 the compilation settings
189 :param kwargs: ``run_model()`` forwards kwargs to both ``prepare()`` and ``rep.run()``.
190 ``prepare()`` validates and applies ``_ALLOWED_SESSION_OPTIONS`` only when creating
191 a new session from a model path or bytes; if ``model`` is already an
192 ``InferenceSession`` or ``OnnxRuntimeBackendRep``, session-option kwargs are
193 silently ignored. ``rep.run()`` always validates against ``_ALLOWED_RUN_OPTIONS``
194 and raises ``RuntimeError`` for known-but-blocked run attributes.
195 Logging-related kwargs (``log_severity_level``, ``log_verbosity_level``, ``logid``)
196 appear in both allowlists.
197 :return: predictions
198 """
199 rep = cls.prepare(model, device, **kwargs)
200 return rep.run(inputs, **kwargs)
201
202 @classmethod
203 def run_node(cls, node, inputs, device=None, outputs_info=None, **kwargs):
204 """
205 This method is not implemented as it is much more efficient
206 to run a whole model than every node independently.
207 """
208 raise NotImplementedError("It is much more efficient to run a whole model than every node independently.")
209
210
211is_compatible = OnnxRuntimeBackend.is_compatible
212prepare = OnnxRuntimeBackend.prepare
213run = OnnxRuntimeBackend.run_model
214supports_device = OnnxRuntimeBackend.supports_device
215 