LDMX Software
ONNXRuntime.cxx
1
2#include "Tools/ONNXRuntime.h"
3
4#include <algorithm>
5#include <functional>
6#include <numeric>
7
8namespace ldmx {
9namespace ort {
10using namespace ::Ort;
11#if ORT_API_VERSION == 2
12// version used when first integrated onnx into ldmx-sw
13// and version downloaded by cmake infrastructure
14// only support x86_64 architectures
15std::string get_input_name(std::unique_ptr<Session>& s, size_t i,
16 AllocatorWithDefaultOptions a) {
17 return s->GetInputName(i, a);
18}
19std::string get_output_name(std::unique_ptr<Session>& s, size_t i,
20 AllocatorWithDefaultOptions a) {
21 return s->GetOutputName(i, a);
22}
23#else
24// latest version with prebuilds for both x86_64 and arm64
25// architectures but contains a slight API change
26std::string getInputName(std::unique_ptr<Session>& s, size_t i,
27 AllocatorWithDefaultOptions a) {
28 return s->GetInputNameAllocated(i, a).get();
29}
30std::string getOutputName(std::unique_ptr<Session>& s, size_t i,
31 AllocatorWithDefaultOptions a) {
32 return s->GetOutputNameAllocated(i, a).get();
33}
34#if ORT_API_VERSION != 15
35#pragma warning( \
36 "Untested ONNX version, not certain of API, assuming API version 15.")
37#endif
38#endif
39
40Env ONNXRuntime::env(ORT_LOGGING_LEVEL_WARNING, "");
41
42ONNXRuntime::ONNXRuntime(const std::string& model_path,
43 const SessionOptions* session_options) {
44 // create session
45 if (session_options) {
46 session_.reset(new Session(env, model_path.c_str(), *session_options));
47 } else {
48 SessionOptions sess_opts;
49 sess_opts.SetIntraOpNumThreads(1);
50 session_.reset(new Session(env, model_path.c_str(), sess_opts));
51 }
52 AllocatorWithDefaultOptions allocator;
53
54 // get input names and shapes
55 size_t num_input_nodes = session_->GetInputCount();
56 input_node_strings_.resize(num_input_nodes);
57 input_node_names_.resize(num_input_nodes);
58 input_node_dims_.clear();
59
60 for (size_t i = 0; i < num_input_nodes; i++) {
61 // get input node names
62 std::string input_name(getInputName(session_, i, allocator));
63 input_node_strings_[i] = input_name;
64 input_node_names_[i] = input_node_strings_[i].c_str();
65
66 // get input shapes
67 auto type_info = session_->GetInputTypeInfo(i);
68 auto tensor_info = type_info.GetTensorTypeAndShapeInfo();
69 size_t num_dims = tensor_info.GetDimensionsCount();
70 input_node_dims_[input_name].resize(num_dims);
71 const auto input_shape = tensor_info.GetShape();
72 std::copy(input_shape.begin(), input_shape.end(),
73 input_node_dims_[input_name].begin());
74
75 // set the batch size to 1 by default
76 input_node_dims_[input_name].at(0) = 1;
77 }
78
79 size_t num_output_nodes = session_->GetOutputCount();
80 output_node_strings_.resize(num_output_nodes);
81 output_node_names_.resize(num_output_nodes);
82 output_node_dims_.clear();
83
84 for (size_t i = 0; i < num_output_nodes; i++) {
85 // get output node names
86 std::string output_name(getOutputName(session_, i, allocator));
87 output_node_strings_[i] = output_name;
88 output_node_names_[i] = output_node_strings_[i].c_str();
89
90 // get output node types
91 auto type_info = session_->GetOutputTypeInfo(i);
92 auto tensor_info = type_info.GetTensorTypeAndShapeInfo();
93 size_t num_dims = tensor_info.GetDimensionsCount();
94 output_node_dims_[output_name].resize(num_dims);
95 const auto output_shape = tensor_info.GetShape();
96 std::copy(output_shape.begin(), output_shape.end(),
97 output_node_dims_[output_name].begin());
98
99 // the 0th dim depends on the batch size
100 output_node_dims_[output_name].at(0) = -1;
101 }
102}
103
104FloatArrays ONNXRuntime::run(const std::vector<std::string>& input_names,
105 FloatArrays& input_values,
106 const std::vector<std::string>& output_names,
107 int64_t batch_size) const {
108 assert(input_names.size() == input_values.size());
109 assert(batch_size > 0);
110
111 // create input tensor objects from data values
112 std::vector<Value> input_tensors;
113 auto memory_info =
114 MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault);
115 for (const auto& name : input_node_strings_) {
116 auto iter = std::find(input_names.begin(), input_names.end(), name);
117 if (iter == input_names.end()) {
118 throw std::runtime_error("Input '" + name + "' is not provided!");
119 }
120 auto value = input_values.begin() + (iter - input_names.begin());
121 auto input_dims = input_node_dims_.at(name);
122 if (input_dims.size() > 0) {
123 input_dims[0] = batch_size;
124 }
125 auto expected_len = std::accumulate(input_dims.begin(), input_dims.end(), 1,
126 std::multiplies<int64_t>());
127 if (expected_len != (int64_t)value->size()) {
128 throw std::runtime_error("Input array '" + name +
129 "' has a wrong size of " +
130 std::to_string(value->size()) + ", expected " +
131 std::to_string(expected_len));
132 }
133 auto input_tensor =
134 Value::CreateTensor<float>(memory_info, value->data(), value->size(),
135 input_dims.data(), input_dims.size());
136 assert(input_tensor.IsTensor());
137 input_tensors.emplace_back(std::move(input_tensor));
138 }
139
140 // set output node names; will get all outputs if `output_names` is not
141 // provided
142 std::vector<const char*> run_output_node_names;
143 if (output_names.empty()) {
144 run_output_node_names = output_node_names_;
145 } else {
146 for (const auto& name : output_names) {
147 run_output_node_names.push_back(name.c_str());
148 }
149 }
150
151 // run
152 auto output_tensors =
153 session_->Run(RunOptions{nullptr}, input_node_names_.data(),
154 input_tensors.data(), input_tensors.size(),
155 run_output_node_names.data(), run_output_node_names.size());
156
157 // convert output to floats
158 FloatArrays outputs;
159 for (auto& output_tensor : output_tensors) {
160 assert(output_tensor.IsTensor());
161
162 // get output shape
163 auto tensor_info = output_tensor.GetTensorTypeAndShapeInfo();
164 auto length = tensor_info.GetElementCount();
165
166 auto floatarr = output_tensor.GetTensorMutableData<float>();
167 outputs.emplace_back(floatarr, floatarr + length);
168 }
169 assert(outputs.size() == run_output_node_names.size());
170
171 return outputs;
172}
173
174const std::vector<std::string>& ONNXRuntime::getOutputNames() const {
175 if (session_) {
176 return output_node_strings_;
177 } else {
178 throw std::runtime_error("ONNXRuntime session is not initialized!");
179 }
180}
181
182const std::vector<int64_t>& ONNXRuntime::getOutputShape(
183 const std::string& output_name) const {
184 auto iter = output_node_dims_.find(output_name);
185 if (iter == output_node_dims_.end()) {
186 throw std::runtime_error("Output name '" + output_name + "' is invalid!");
187 } else {
188 return iter->second;
189 }
190}
191
192} // namespace ort
193} // namespace ldmx
const std::vector< int64_t > & getOutputShape(const std::string &output_name) const
Get the shape of a output node.
ONNXRuntime(const std::string &model_path, const ::Ort::SessionOptions *session_options=nullptr)
Class constructor.
FloatArrays run(const std::vector< std::string > &input_names, FloatArrays &input_values, const std::vector< std::string > &output_names={}, int64_t batch_size=1) const
Run model inference and get outputs.
const std::vector< std::string > & getOutputNames() const
Get the names of all the output nodes.