18 :
m_env (
std::make_unique<
Ort::Env>(ORT_LOGGING_LEVEL_FATAL,
""))
21 Ort::SessionOptions session_options;
22 session_options.SetIntraOpNumThreads(1);
26 session_options.SetLogSeverityLevel(4);
27 session_options.SetGraphOptimizationLevel(
28 GraphOptimizationLevel::ORT_ENABLE_EXTENDED);
36 session_options.DisableCpuMemArena();
40 if (opts.execution_provider ==
"CUDA") {
41 Ort::CUDAProviderOptions cuda_options;
43 {
"device_id", std::to_string(opts.device_id)},
46 {
"use_tf32", opts.use_tf32 ?
"1" :
"0"},
48 session_options.AppendExecutionProvider_CUDA_V2(*cuda_options);
49 }
else if (opts.execution_provider !=
"CPU") {
50 throw std::runtime_error(
51 "unknown execution provider '" + opts.execution_provider +
"'");
55 Ort::AllocatorWithDefaultOptions allocator;
58 m_session = std::make_unique<Ort::Session>(
59 *m_env, path_to_onnx.c_str(), session_options);
62 m_metadata = loadMetadata(
"gnn_config");
63 m_num_inputs = m_session->GetInputCount();
64 m_num_outputs = m_session->GetOutputCount();
67 if (m_metadata.contains(
"onnx_model_version")) {
68 m_onnx_model_version = m_metadata[
"onnx_model_version"].get<
SaltModelVersion>();
70 throw std::runtime_error(
"Unknown Onnx model version!");
76 throw std::runtime_error(
"Onnx model version not found in metadata");
81 m_model_name = determineModelName();
84 for (
size_t i = 0;
i < m_num_inputs;
i++) {
85 m_input_node_names.push_back(m_session->GetInputNameAllocated(i, allocator).get());
89 for (
size_t i = 0;
i < m_num_outputs;
i++) {
90 const auto name = std::string(m_session->GetOutputNameAllocated(i, allocator).get());
91 const auto type = m_session->GetOutputTypeInfo(i).GetTensorTypeAndShapeInfo().GetElementType();
92 const int rank = m_session->GetOutputTypeInfo(i).GetTensorTypeAndShapeInfo().GetShape().size();
94 m_output_nodes.emplace_back(name, type, m_model_name);
96 m_output_nodes.emplace_back(name, type, rank);
109 Ort::AllocatorWithDefaultOptions allocator;
112 return std::string(
m_metadata[
"outputs"].begin().key());
116 std::set<std::string> model_names;
118 const auto name = std::string(
m_session->GetOutputNameAllocated(i, allocator).get());
119 size_t underscore_pos = name.find(
'_');
120 if (underscore_pos != std::string::npos) {
121 model_names.insert(name.substr(0, underscore_pos));
123 return std::string(
"");
126 if (model_names.size() != 1) {
127 throw std::runtime_error(
"SaltModel: model names are not consistent between outputs");
129 return *model_names.begin();
153 std::vector<float> input_tensor_values;
156 auto memory_info = Ort::MemoryInfo::CreateCpu(
157 OrtArenaAllocator, OrtMemTypeDefault
159 std::vector<Ort::Value> input_tensors;
161 input_tensors.push_back(Ort::Value::CreateTensor<float>(
162 memory_info, gnn_inputs.at(node_name).first.data(), gnn_inputs.at(node_name).first.size(),
163 gnn_inputs.at(node_name).second.data(), gnn_inputs.at(node_name).second.size())
168 std::vector<const char*> input_node_names;
171 input_node_names.push_back(name.c_str());
173 std::vector<const char*> output_node_names;
176 output_node_names.push_back(
node.name_in_model.c_str());
184 auto output_tensors = session.Run(Ort::RunOptions{
nullptr},
185 input_node_names.data(), input_tensors.data(), input_node_names.size(),
186 output_node_names.data(), output_node_names.size()
191 for (
size_t node_idx = 0; node_idx <
m_output_nodes.size(); ++node_idx) {
193 const auto& tensor = output_tensors[node_idx];
194 auto tensor_type = tensor.GetTypeInfo().GetTensorTypeAndShapeInfo().GetElementType();
195 auto tensor_shape = tensor.GetTypeInfo().GetTensorTypeAndShapeInfo().GetShape();
196 int length = tensor.GetTensorTypeAndShapeInfo().GetElementCount();
197 if (tensor_type == ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT) {
198 if (tensor_shape.size() == 0) {
199 output.singleFloat[output_node.name] = *tensor.GetTensorData<
float>();
200 }
else if (tensor_shape.size() == 1) {
201 const float* data = tensor.GetTensorData<
float>();
202 output.vecFloat[output_node.name] = std::vector<float>(data, data +
length);
204 throw std::runtime_error(
"Unsupported tensor shape for FLOAT type");
206 }
else if (tensor_type == ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8) {
207 if (tensor_shape.size() == 1) {
208 const char* data = tensor.GetTensorData<
char>();
209 output.vecChar[output_node.name] = std::vector<char>(data, data +
length);
211 throw std::runtime_error(
"Unsupported tensor shape for INT8 type");
214 throw std::runtime_error(
"Unsupported tensor type");