LDMX Software
ONNXRuntime.h
1
2#ifndef TOOLS_ONNXRUNTIME_H
3#define TOOLS_ONNXRUNTIME_H
4
5#include <cassert>
6#include <map>
7#include <memory>
8#include <string>
9#include <vector>
10
11#include "onnxruntime_cxx_api.h"
12
13namespace ldmx {
14namespace ort {
15
16typedef std::vector<std::vector<float>> FloatArrays;
17
23 public:
30 ONNXRuntime(const std::string& model_path,
31 const ::Ort::SessionOptions* session_options = nullptr);
32 ONNXRuntime(const ONNXRuntime&) = delete;
33 ONNXRuntime& operator=(const ONNXRuntime&) = delete;
34 ~ONNXRuntime() = default;
35
49 FloatArrays run(const std::vector<std::string>& input_names,
50 FloatArrays& input_values,
51 const std::vector<std::string>& output_names = {},
52 int64_t batch_size = 1) const;
53
58 const std::vector<std::string>& getOutputNames() const;
59
66 const std::vector<int64_t>& getOutputShape(
67 const std::string& output_name) const;
68
69 private:
70 static ::Ort::Env env;
71 std::unique_ptr<::Ort::Session> session_;
72
73 std::vector<std::string> input_node_strings_;
74 std::vector<const char*> input_node_names_;
75 std::map<std::string, std::vector<int64_t>> input_node_dims_;
76
77 std::vector<std::string> output_node_strings_;
78 std::vector<const char*> output_node_names_;
79 std::map<std::string, std::vector<int64_t>> output_node_dims_;
80};
81
82} // namespace ort
83} // namespace ldmx
84
85#endif /* TOOLS_ONNXRUNTIME_H_ */
A convenience wrapper of the ONNXRuntime C++ API.
Definition ONNXRuntime.h:22
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.