mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
103 lines
4.2 KiB
C++
103 lines
4.2 KiB
C++
// Copyright (c) Microsoft Corporation. All rights reserved.
|
|
// Licensed under the MIT License.
|
|
|
|
#include "core/session/inference_session.h"
|
|
#include "test/providers/provider_test_utils.h"
|
|
#include "test/framework/test_utils.h"
|
|
#include "gtest/gtest.h"
|
|
#include "core/providers/tensorrt/tensorrt_execution_provider.h"
|
|
|
|
using namespace std;
|
|
using namespace ONNX_NAMESPACE;
|
|
using namespace ::onnxruntime::logging;
|
|
|
|
namespace onnxruntime {
|
|
|
|
namespace test {
|
|
void VerifyOutputs(const std::vector<MLValue>& fetches,
|
|
const std::vector<int64_t>& expected_dims,
|
|
const std::vector<float>& expected_values) {
|
|
ASSERT_EQ(1, fetches.size());
|
|
auto& rtensor = fetches.front().Get<Tensor>();
|
|
TensorShape expected_shape(expected_dims);
|
|
ASSERT_EQ(expected_shape, rtensor.Shape());
|
|
const std::vector<float> found(rtensor.template Data<float>(), rtensor.template Data<float>() + expected_values.size());
|
|
ASSERT_EQ(expected_values, found);
|
|
}
|
|
|
|
TEST(TensorrtExecutionProviderTest, FunctionTest) {
|
|
onnxruntime::Model model("graph_1");
|
|
auto& graph = model.MainGraph();
|
|
std::vector<onnxruntime::NodeArg*> inputs;
|
|
std::vector<onnxruntime::NodeArg*> outputs;
|
|
|
|
// FLOAT tensor.
|
|
ONNX_NAMESPACE::TypeProto float_tensor;
|
|
float_tensor.mutable_tensor_type()->set_elem_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT);
|
|
float_tensor.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(1);
|
|
float_tensor.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(3);
|
|
float_tensor.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(2);
|
|
|
|
auto& input_arg_1 = graph.GetOrCreateNodeArg("X", &float_tensor);
|
|
auto& input_arg_2 = graph.GetOrCreateNodeArg("Y", &float_tensor);
|
|
inputs.push_back(&input_arg_1);
|
|
inputs.push_back(&input_arg_2);
|
|
auto& output_arg = graph.GetOrCreateNodeArg("node_1_out_1", &float_tensor);
|
|
outputs.push_back(&output_arg);
|
|
graph.AddNode("node_1", "Add", "node 1.", inputs, outputs);
|
|
|
|
auto& input_arg_3 = graph.GetOrCreateNodeArg("Z", &float_tensor);
|
|
inputs.clear();
|
|
inputs.push_back(&output_arg);
|
|
inputs.push_back(&input_arg_3);
|
|
auto& output_arg_2 = graph.GetOrCreateNodeArg("M", &float_tensor);
|
|
outputs.clear();
|
|
outputs.push_back(&output_arg_2);
|
|
graph.AddNode("node_2", "Add", "node 2.", inputs, outputs);
|
|
|
|
auto status = graph.Resolve();
|
|
ASSERT_TRUE(status.IsOK());
|
|
std::string model_file_name = "trt_execution_provider_test_graph.onnx";
|
|
status = onnxruntime::Model::Save(model, model_file_name);
|
|
|
|
std::vector<int64_t> dims_mul_x = {1, 3, 2};
|
|
std::vector<float> values_mul_x = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f};
|
|
MLValue ml_value_x;
|
|
CreateMLValue<float>(TestTensorrtExecutionProvider()->GetAllocator(0, OrtMemTypeCPU), dims_mul_x, values_mul_x, &ml_value_x);
|
|
MLValue ml_value_y;
|
|
CreateMLValue<float>(TestTensorrtExecutionProvider()->GetAllocator(0, OrtMemTypeCPU), dims_mul_x, values_mul_x, &ml_value_y);
|
|
MLValue ml_value_z;
|
|
CreateMLValue<float>(TestTensorrtExecutionProvider()->GetAllocator(0, OrtMemTypeCPU), dims_mul_x, values_mul_x, &ml_value_z);
|
|
NameMLValMap feeds;
|
|
feeds.insert(std::make_pair("X", ml_value_x));
|
|
feeds.insert(std::make_pair("Y", ml_value_y));
|
|
feeds.insert(std::make_pair("Z", ml_value_z));
|
|
|
|
// prepare outputs
|
|
std::vector<std::string> output_names;
|
|
output_names.push_back("M");
|
|
std::vector<MLValue> fetches;
|
|
|
|
// prepare expected inputs and outputs
|
|
std::vector<int64_t> expected_dims_mul_m = {1, 3, 2};
|
|
std::vector<float> expected_values_mul_m = {3.0f, 6.0f, 9.0f, 12.0f, 15.0f, 18.0f};
|
|
|
|
SessionOptions so;
|
|
so.session_logid = "TensorrtExecutionProviderTest.FunctionTest";
|
|
RunOptions run_options;
|
|
run_options.run_tag = so.session_logid;
|
|
|
|
InferenceSession session_object{so};
|
|
session_object.RegisterExecutionProvider(std::make_unique<::onnxruntime::TensorrtExecutionProvider>());
|
|
status = session_object.Load(model_file_name);
|
|
ASSERT_TRUE(status.IsOK());
|
|
status = session_object.Initialize();
|
|
ASSERT_TRUE(status.IsOK());
|
|
|
|
// Now run
|
|
status = session_object.Run(run_options, feeds, output_names, &fetches);
|
|
ASSERT_TRUE(status.IsOK());
|
|
VerifyOutputs(fetches, expected_dims_mul_m, expected_values_mul_m);
|
|
}
|
|
} // namespace test
|
|
} // namespace onnxruntime
|