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));
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);
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");
239 Ort::AllocatorWithDefaultOptions allocator;
240 Ort::MemoryInfo memoryInfo = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault);
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;
248 const size_t inputCount = session.GetInputCount();
249 if (inputCount == 0) {
250 throw std::runtime_error(
"model has no inputs");
252 inputNames.reserve(inputCount);
253 inputNamePointers.reserve(inputCount);
254 inputBuffers.reserve(inputCount);
255 inputValues.reserve(inputCount);
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());
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");
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");
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);
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";
284 if (args.expectedInputElements != 0 && totalInputElements != args.expectedInputElements) {
285 throw std::runtime_error(
"model input element count is " +
std::to_string(totalInputElements) +
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");
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());
303 auto outputs = session.Run(Ort::RunOptions{
nullptr},
304 inputNamePointers.data(),
307 outputNamePointers.data(),
308 outputNamePointers.size());
310 if (outputs.size() != outputCount) {
311 throw std::runtime_error(
"ONNX Runtime returned an unexpected number of outputs");
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");
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");
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");
332 std::cout <<
"output[" <<
i <<
"] " << outputNames[
i] <<
" shape=" << shapeString(shape)
333 <<
" elements=" << elements <<
"\n";
336 if (args.expectedOutputElements != 0 && totalOutputElements != args.expectedOutputElements) {
337 throw std::runtime_error(
"model output element count is " +
std::to_string(totalOutputElements) +
341 std::cout <<
"provider=" << providerName <<
" assigned_nodes=" << assignedNodes
342 <<
" total_inputs=" << totalInputElements
343 <<
" total_outputs=" << totalOutputElements <<
"\n";
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";