Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
backend.py215 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
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 
codekingpro/portable-devtools · Team Ai