Project
Loading...
Searching...
No Matches
onnxruntime_ep_inference.cxx
Go to the documentation of this file.
1// Copyright 2019-2020 CERN and copyright holders of ALICE O2.
2// See https://alice-o2.web.cern.ch/copyright for details of the copyright holders.
3// All rights not expressly granted are reserved.
4//
5// This software is distributed under the terms of the GNU General Public
6// License v3 (GPL Version 3), copied verbatim in the file "COPYING".
7//
8// In applying this license CERN does not waive the privileges and immunities
9// granted to it by virtue of its status as an Intergovernmental Organization
10// or submit itself to any jurisdiction.
11
15
16#include <onnxruntime_cxx_api.h>
17
18#include <algorithm>
19#include <cctype>
20#include <cmath>
21#include <cstdlib>
22#include <iostream>
23#include <limits>
24#include <numeric>
25#include <sstream>
26#include <stdexcept>
27#include <string>
28#include <unordered_map>
29#include <vector>
30
31namespace
32{
33
34struct Arguments {
35 std::string modelPath;
36 std::string provider;
37 int deviceId = 0;
38 size_t expectedInputElements = 0;
39 size_t expectedOutputElements = 0;
40 bool requireProviderAssignment = true;
41};
42
43std::string toLower(std::string value)
44{
45 std::transform(value.begin(), value.end(), value.begin(), [](unsigned char c) {
46 return static_cast<char>(std::tolower(c));
47 });
48 return value;
49}
50
51bool hasProvider(const std::vector<std::string>& providers, const std::string& provider)
52{
53 return std::find(providers.begin(), providers.end(), provider) != providers.end();
54}
55
56std::string join(const std::vector<std::string>& values)
57{
58 std::ostringstream os;
59 for (size_t i = 0; i < values.size(); ++i) {
60 os << (i == 0 ? "" : ", ") << values[i];
61 }
62 return os.str();
63}
64
65void usage(const char* argv0)
66{
67 std::cerr << "usage: " << argv0
68 << " --model MODEL.onnx --provider cpu|migraphx|cuda|tensorrt "
69 "[--device-id N] [--expected-input-elements N] "
70 "[--expected-output-elements N] [--allow-cpu-fallback]\n";
71}
72
73Arguments parseArguments(int argc, char** argv)
74{
75 Arguments args;
76 for (int i = 1; i < argc; ++i) {
77 const std::string arg = argv[i];
78 auto needValue = [&](const char* name) -> std::string {
79 if (i + 1 >= argc) {
80 throw std::runtime_error(std::string("missing value for ") + name);
81 }
82 return argv[++i];
83 };
84
85 if (arg == "--model") {
86 args.modelPath = needValue("--model");
87 } else if (arg == "--provider") {
88 args.provider = toLower(needValue("--provider"));
89 } else if (arg == "--device-id") {
90 args.deviceId = std::stoi(needValue("--device-id"));
91 } else if (arg == "--expected-input-elements") {
92 args.expectedInputElements = std::stoull(needValue("--expected-input-elements"));
93 } else if (arg == "--expected-output-elements") {
94 args.expectedOutputElements = std::stoull(needValue("--expected-output-elements"));
95 } else if (arg == "--allow-cpu-fallback") {
96 args.requireProviderAssignment = false;
97 } else if (arg == "--help" || arg == "-h") {
98 usage(argv[0]);
99 std::exit(0);
100 } else {
101 throw std::runtime_error("unknown argument: " + arg);
102 }
103 }
104
105 if (args.modelPath.empty()) {
106 throw std::runtime_error("--model is required");
107 }
108 if (args.provider != "cpu" && args.provider != "migraphx" && args.provider != "cuda" && args.provider != "tensorrt") {
109 throw std::runtime_error("--provider must be one of: cpu, migraphx, cuda, tensorrt");
110 }
111 return args;
112}
113
114std::string ortProviderName(const std::string& provider)
115{
116 if (provider == "cpu") {
117 return "CPUExecutionProvider";
118 }
119 if (provider == "migraphx") {
120 return "MIGraphXExecutionProvider";
121 }
122 if (provider == "cuda") {
123 return "CUDAExecutionProvider";
124 }
125 if (provider == "tensorrt") {
126 return "TensorrtExecutionProvider";
127 }
128 throw std::runtime_error("unsupported provider: " + provider);
129}
130
131void appendProvider(Ort::SessionOptions& options, const Arguments& args)
132{
133 if (args.provider == "cpu") {
134 return;
135 }
136 if (args.provider == "cuda") {
137#ifdef ORT_CUDA_BUILD
138 OrtCUDAProviderOptions cudaOptions{};
139 cudaOptions.device_id = args.deviceId;
140 options.AppendExecutionProvider_CUDA(cudaOptions);
141 return;
142#else
143 throw std::runtime_error("CUDA execution provider support was not enabled at build time");
144#endif
145 }
146 if (args.provider == "migraphx") {
147#ifdef ORT_MIGRAPHX_BUILD
148 OrtMIGraphXProviderOptions migraphxOptions{};
149 migraphxOptions.device_id = args.deviceId;
150 migraphxOptions.migraphx_mem_limit = std::numeric_limits<size_t>::max();
151 options.AppendExecutionProvider_MIGraphX(migraphxOptions);
152 return;
153#else
154 throw std::runtime_error("MIGraphX execution provider support was not enabled at build time");
155#endif
156 }
157 if (args.provider == "tensorrt") {
158#ifdef ORT_TENSORRT_BUILD
159 Ort::TensorRTProviderOptions tensorrtOptions;
160 tensorrtOptions.Update({{"device_id", std::to_string(args.deviceId)}});
161 options.AppendExecutionProvider_TensorRT_V2(*tensorrtOptions);
162 return;
163#else
164 throw std::runtime_error("TensorRT execution provider support was not enabled at build time");
165#endif
166 }
167}
168
169std::vector<int64_t> concreteShape(std::vector<int64_t> shape)
170{
171 for (auto& dim : shape) {
172 if (dim <= 0) {
173 dim = 1;
174 }
175 }
176 return shape;
177}
178
179size_t elementCount(const std::vector<int64_t>& shape)
180{
181 if (shape.empty()) {
182 return 1;
183 }
184 return std::accumulate(shape.begin(), shape.end(), size_t{1}, [](size_t product, int64_t dim) {
185 if (dim <= 0) {
186 throw std::runtime_error("invalid concrete tensor dimension");
187 }
188 return product * static_cast<size_t>(dim);
189 });
190}
191
192std::string shapeString(const std::vector<int64_t>& shape)
193{
194 std::ostringstream os;
195 os << "[";
196 for (size_t i = 0; i < shape.size(); ++i) {
197 os << (i == 0 ? "" : ",") << shape[i];
198 }
199 os << "]";
200 return os.str();
201}
202
203bool assignedToProvider(const Ort::Session& session, const std::string& providerName, size_t& assignedNodes)
204{
205 assignedNodes = 0;
206 for (const auto& subgraph : session.GetEpGraphAssignmentInfo()) {
207 if (subgraph.GetEpName() == providerName) {
208 assignedNodes += subgraph.GetNodes().size();
209 }
210 }
211 return assignedNodes > 0;
212}
213
214} // namespace
215
216int main(int argc, char** argv)
217{
218 try {
219 const auto args = parseArguments(argc, argv);
220 const auto providerName = ortProviderName(args.provider);
221 const auto availableProviders = Ort::GetAvailableProviders();
222 if (!hasProvider(availableProviders, providerName)) {
223 throw std::runtime_error(providerName + " is not available in this ONNX Runtime build. Available providers: " + join(availableProviders));
224 }
225
226 Ort::Env env(ORT_LOGGING_LEVEL_WARNING, "onnxruntime-ep-inference");
227 Ort::SessionOptions options;
228 options.SetIntraOpNumThreads(1);
229 options.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL);
230 options.AddConfigEntry("session.record_ep_graph_assignment_info", "1");
231 appendProvider(options, args);
232
233 Ort::Session session(env, args.modelPath.c_str(), options);
234 size_t assignedNodes = 0;
235 if (args.provider != "cpu" && args.requireProviderAssignment && !assignedToProvider(session, providerName, assignedNodes)) {
236 throw std::runtime_error(providerName + " did not receive any graph nodes");
237 }
238
239 Ort::AllocatorWithDefaultOptions allocator;
240 Ort::MemoryInfo memoryInfo = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault);
241
242 std::vector<std::string> inputNames;
243 std::vector<const char*> inputNamePointers;
244 std::vector<std::vector<float>> inputBuffers;
245 std::vector<Ort::Value> inputValues;
246 size_t totalInputElements = 0;
247
248 const size_t inputCount = session.GetInputCount();
249 if (inputCount == 0) {
250 throw std::runtime_error("model has no inputs");
251 }
252 inputNames.reserve(inputCount);
253 inputNamePointers.reserve(inputCount);
254 inputBuffers.reserve(inputCount);
255 inputValues.reserve(inputCount);
256
257 for (size_t i = 0; i < inputCount; ++i) {
258 auto name = session.GetInputNameAllocated(i, allocator);
259 inputNames.emplace_back(name.get());
260 inputNamePointers.push_back(inputNames.back().c_str());
261
262 auto typeInfo = session.GetInputTypeInfo(i);
263 if (typeInfo.GetONNXType() != ONNX_TYPE_TENSOR) {
264 throw std::runtime_error("input " + inputNames.back() + " is not a tensor");
265 }
266 auto tensorInfo = typeInfo.GetTensorTypeAndShapeInfo();
267 if (tensorInfo.GetElementType() != ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT) {
268 throw std::runtime_error("input " + inputNames.back() + " is not a float tensor");
269 }
270
271 const auto shape = concreteShape(tensorInfo.GetShape());
272 const auto elements = elementCount(shape);
273 totalInputElements += elements;
274 inputBuffers.emplace_back(elements);
275 for (size_t j = 0; j < elements; ++j) {
276 inputBuffers.back()[j] = static_cast<float>((static_cast<int>((i + j) % 23) - 11) * 0.03125f);
277 }
278 inputValues.emplace_back(Ort::Value::CreateTensor<float>(
279 memoryInfo, inputBuffers.back().data(), elements, shape.data(), shape.size()));
280 std::cout << "input[" << i << "] " << inputNames.back() << " shape=" << shapeString(shape)
281 << " elements=" << elements << "\n";
282 }
283
284 if (args.expectedInputElements != 0 && totalInputElements != args.expectedInputElements) {
285 throw std::runtime_error("model input element count is " + std::to_string(totalInputElements) +
286 ", expected " + std::to_string(args.expectedInputElements));
287 }
288
289 std::vector<std::string> outputNames;
290 std::vector<const char*> outputNamePointers;
291 const size_t outputCount = session.GetOutputCount();
292 if (outputCount == 0) {
293 throw std::runtime_error("model has no outputs");
294 }
295 outputNames.reserve(outputCount);
296 outputNamePointers.reserve(outputCount);
297 for (size_t i = 0; i < outputCount; ++i) {
298 auto name = session.GetOutputNameAllocated(i, allocator);
299 outputNames.emplace_back(name.get());
300 outputNamePointers.push_back(outputNames.back().c_str());
301 }
302
303 auto outputs = session.Run(Ort::RunOptions{nullptr},
304 inputNamePointers.data(),
305 inputValues.data(),
306 inputValues.size(),
307 outputNamePointers.data(),
308 outputNamePointers.size());
309
310 if (outputs.size() != outputCount) {
311 throw std::runtime_error("ONNX Runtime returned an unexpected number of outputs");
312 }
313
314 size_t totalOutputElements = 0;
315 for (size_t i = 0; i < outputs.size(); ++i) {
316 if (!outputs[i].IsTensor()) {
317 throw std::runtime_error("output " + outputNames[i] + " is not a tensor");
318 }
319 auto tensorInfo = outputs[i].GetTensorTypeAndShapeInfo();
320 if (tensorInfo.GetElementType() != ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT) {
321 throw std::runtime_error("output " + outputNames[i] + " is not a float tensor");
322 }
323 const auto shape = tensorInfo.GetShape();
324 const auto elements = tensorInfo.GetElementCount();
325 totalOutputElements += elements;
326 const float* data = outputs[i].GetTensorData<float>();
327 for (size_t j = 0; j < elements; ++j) {
328 if (!std::isfinite(data[j])) {
329 throw std::runtime_error("output " + outputNames[i] + " contains a non-finite value");
330 }
331 }
332 std::cout << "output[" << i << "] " << outputNames[i] << " shape=" << shapeString(shape)
333 << " elements=" << elements << "\n";
334 }
335
336 if (args.expectedOutputElements != 0 && totalOutputElements != args.expectedOutputElements) {
337 throw std::runtime_error("model output element count is " + std::to_string(totalOutputElements) +
338 ", expected " + std::to_string(args.expectedOutputElements));
339 }
340
341 std::cout << "provider=" << providerName << " assigned_nodes=" << assignedNodes
342 << " total_inputs=" << totalInputElements
343 << " total_outputs=" << totalOutputElements << "\n";
344 return 0;
345 } catch (const Ort::Exception& ex) {
346 std::cerr << "ONNX Runtime error: " << ex.what() << "\n";
347 } catch (const std::exception& ex) {
348 std::cerr << "error: " << ex.what() << "\n";
349 }
350
351 usage(argv[0]);
352 return 1;
353}
int32_t i
uint32_t j
Definition RawData.h:0
uint32_t c
Definition RawData.h:2
GLuint const GLchar * name
Definition glcorearb.h:781
GLsizei const GLfloat * value
Definition glcorearb.h:819
GLenum GLsizei GLsizei GLint * values
Definition glcorearb.h:1576
GLboolean * data
Definition glcorearb.h:298
GLsizeiptr const void GLenum usage
Definition glcorearb.h:659
std::string toLower(std::string const &s)
constexpr auto join(Ts const &... t)
Definition ASoA.h:3515
std::string to_string(gsl::span< T, Size > span)
Definition common.h:52
#define main