2626#include < onnxruntime_c_api.h>
2727#include < onnxruntime_cxx_api.h>
2828
29- #include < algorithm>
3029#include < cassert>
3130#include < cstddef>
3231#include < cstdint>
33- #include < iterator>
3432#include < memory>
3533#include < sstream>
3634#include < string>
@@ -97,6 +95,8 @@ void OnnxModel::initModel(const std::string& localPath, const bool enableOptimiz
9795 mSession = std::make_shared<Ort::Session>(*mEnv , modelPath.c_str (), sessionOptions);
9896 mMemInfo = Ort::MemoryInfo::CreateCpu (OrtAllocatorType::OrtArenaAllocator, OrtMemType::OrtMemTypeDefault);
9997
98+ mInputNamesChar .clear ();
99+ mOutputNamesChar .clear ();
100100 mInputNames .clear ();
101101 mInputShapes .clear ();
102102 mOutputNames .clear ();
@@ -114,6 +114,14 @@ void OnnxModel::initModel(const std::string& localPath, const bool enableOptimiz
114114 for (std::size_t i = 0 ; i < mSession ->GetOutputCount (); ++i) {
115115 mOutputShapes .emplace_back (mSession ->GetOutputTypeInfo (i).GetTensorTypeAndShapeInfo ().GetShape ());
116116 }
117+ mInputNamesChar .reserve (mInputNames .size ());
118+ for (const auto & name : mInputNames ) {
119+ mInputNamesChar .push_back (name.c_str ());
120+ }
121+ mOutputNamesChar .reserve (mOutputNames .size ());
122+ for (const auto & name : mOutputNames ) {
123+ mOutputNamesChar .push_back (name.c_str ());
124+ }
117125 LOG (info) << " Input Nodes:" ;
118126 for (std::size_t i = 0 ; i < mInputNames .size (); i++) {
119127 LOG (info) << " \t " << mInputNames [i] << " : " << printShape (mInputShapes [i]);
@@ -173,7 +181,7 @@ std::vector<int64_t> OnnxModel::inferInputShape(const std::size_t iinput, const
173181 return inputShape;
174182}
175183
176- std::vector<Ort::Value> OnnxModel::evalModelRaw ( std::vector<Ort::Value>& input)
184+ void OnnxModel::checkInput ( const std::vector<Ort::Value>& input) const
177185{
178186 if (!mSession ) {
179187 LOG (fatal) << " OnnxModel::evalModel called before initModel()" ;
@@ -184,19 +192,15 @@ std::vector<Ort::Value> OnnxModel::evalModelRaw(std::vector<Ort::Value>& input)
184192 for (std::size_t i = 0 ; i < input.size (); i++) {
185193 LOG (debug) << " Input tensor " << i << " shape: " << printShape (input[i].GetTensorTypeAndShapeInfo ().GetShape ());
186194 }
195+ }
187196
188- std::vector<const char *> inputNamesChar (mInputNames .size (), nullptr );
189- std::transform (std::begin (mInputNames ), std::end (mInputNames ), std::begin (inputNamesChar),
190- [](const std::string& str) { return str.c_str (); });
191-
192- std::vector<const char *> outputNamesChar (mOutputNames .size (), nullptr );
193- std::transform (std::begin (mOutputNames ), std::end (mOutputNames ), std::begin (outputNamesChar),
194- [](const std::string& str) { return str.c_str (); });
195-
197+ std::vector<Ort::Value> OnnxModel::evalModelRaw (std::vector<Ort::Value>& input)
198+ {
199+ checkInput (input);
196200 std::vector<Ort::Value> outputTensors;
197201 try {
198- const Ort::RunOptions runOptions;
199- outputTensors = mSession ->Run (runOptions, inputNamesChar .data (), input.data (), input.size (), outputNamesChar .data (), outputNamesChar .size ());
202+ const Ort::RunOptions runOptions{ nullptr } ;
203+ outputTensors = mSession ->Run (runOptions, mInputNamesChar .data (), input.data (), input.size (), mOutputNamesChar .data (), mOutputNamesChar .size ());
200204 } catch (const Ort::Exception& exception) {
201205 LOG (fatal) << " Error running model inference: " << exception.what ();
202206 }
@@ -206,21 +210,44 @@ std::vector<Ort::Value> OnnxModel::evalModelRaw(std::vector<Ort::Value>& input)
206210 LOG (fatal) << " Number of output tensors: " << outputTensors.size () << " does not agree with the model specified size: " << mOutputNames .size ();
207211 }
208212 for (std::size_t i = 0 ; i < outputTensors.size (); i++) {
209- const std::vector<int64_t > shape = outputTensors[i].GetTensorTypeAndShapeInfo ().GetShape ();
210- LOG (debug) << " Output tensor " << i << " shape: " << printShape (shape);
211- bool shapeOk = (shape.size () == mOutputShapes [i].size ());
212- for (std::size_t idim = 0 ; shapeOk && idim < shape.size (); idim++) {
213- // dynamic dimensions of the model (< 0) can take any value
214- shapeOk = (mOutputShapes [i][idim] < 0 ) || (shape[idim] == mOutputShapes [i][idim]);
215- }
216- if (!shapeOk) {
217- LOG (fatal) << " Shape of output tensor " << i << " does not agree with model specification! Output: " << printShape (shape) << " model: " << printShape (mOutputShapes [i]);
218- }
213+ checkOutput (outputTensors[i], i);
219214 }
220215
221216 return outputTensors;
222217}
223218
219+ Ort::Value OnnxModel::evalModelLast (std::vector<Ort::Value>& input)
220+ {
221+ checkInput (input);
222+ if (mOutputNamesChar .empty ()) {
223+ LOG (fatal) << " Model has no outputs" ;
224+ }
225+ Ort::Value output{nullptr };
226+ try {
227+ // A null RunOptions uses the runtime defaults without allocating options per call.
228+ const Ort::RunOptions runOptions{nullptr };
229+ mSession ->Run (runOptions, mInputNamesChar .data (), input.data (), input.size (), &mOutputNamesChar .back (), &output, 1 );
230+ } catch (const Ort::Exception& exception) {
231+ LOG (fatal) << " Error running model inference: " << exception.what ();
232+ }
233+ checkOutput (output, mOutputShapes .size () - 1 );
234+ return output;
235+ }
236+
237+ void OnnxModel::checkOutput (const Ort::Value& tensor, const std::size_t index) const
238+ {
239+ const std::vector<int64_t > shape = tensor.GetTensorTypeAndShapeInfo ().GetShape ();
240+ LOG (debug) << " Output tensor " << index << " shape: " << printShape (shape);
241+ bool shapeOk = (shape.size () == mOutputShapes [index].size ());
242+ for (std::size_t idim = 0 ; shapeOk && idim < shape.size (); idim++) {
243+ // Dynamic dimensions of the model (< 0) can take any value.
244+ shapeOk = (mOutputShapes [index][idim] < 0 ) || (shape[idim] == mOutputShapes [index][idim]);
245+ }
246+ if (!shapeOk) {
247+ LOG (fatal) << " Shape of output tensor " << index << " does not agree with model specification! Output: " << printShape (shape) << " model: " << printShape (mOutputShapes [index]);
248+ }
249+ }
250+
224251void OnnxModel::setActiveThreads (const int threads)
225252{
226253 activeThreads = threads;
0 commit comments