Run model inference and get outputs.
107 {
108 assert(input_names.size() == input_values.size());
109 assert(batch_size > 0);
110
111
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
141
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
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
158 FloatArrays outputs;
159 for (auto& output_tensor : output_tensors) {
160 assert(output_tensor.IsTensor());
161
162
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}