From cb46d79108c953d25db29432245fa74544d4b44c Mon Sep 17 00:00:00 2001 From: Chi Lo <54722500+chilo-ms@users.noreply.github.com> Date: Wed, 20 Apr 2022 10:14:28 -0700 Subject: [PATCH] Model tests refactor (#11194) * Update model test * update comment * create map to hold OnnxModelInfo so test doesn't need to reload the model again * revert the code and use GTEST_SKIP() to skip test * fix bug * revert LATEST_ONNX_OPSET_SUPPORTED_BY_TENSORRT --- onnxruntime/test/providers/cpu/model_tests.cc | 26 ++++++++++++------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/onnxruntime/test/providers/cpu/model_tests.cc b/onnxruntime/test/providers/cpu/model_tests.cc index 19145da702..f2b5fdc5b4 100644 --- a/onnxruntime/test/providers/cpu/model_tests.cc +++ b/onnxruntime/test/providers/cpu/model_tests.cc @@ -52,6 +52,11 @@ struct BrokenTest { #ifdef GTEST_ALLOW_UNINSTANTIATED_PARAMETERIZED_TEST GTEST_ALLOW_UNINSTANTIATED_PARAMETERIZED_TEST(ModelTest); #endif + +void SkipTest() { + GTEST_SKIP() << "Skipping single test"; +} + TEST_P(ModelTest, Run) { std::basic_string param = GetParam(); size_t pos = param.find(ORT_TSTR("_")); @@ -72,16 +77,19 @@ TEST_P(ModelTest, Run) { // them is enabled here to save CI build time. // Besides saving CI build time, TRT isn’t able to support full ONNX ops spec and therefore some testcases will fail. // That's one of reasons we skip those testcases and only test latest ONNX opsets. + SkipTest(); return; } if (model_info->GetONNXOpSetVersion() == 10 && provider_name == "dnnl") { // DNNL can run most of the model tests, but only part of // them is enabled here to save CI build time. + SkipTest(); return; } #ifndef ENABLE_TRAINING if (model_info->HasDomain(ONNX_NAMESPACE::AI_ONNX_TRAINING_DOMAIN) || model_info->HasDomain(ONNX_NAMESPACE::AI_ONNX_PREVIEW_TRAINING_DOMAIN)) { + SkipTest(); return; } #endif @@ -396,7 +404,6 @@ TEST_P(ModelTest, Run) { // TensorRT EP CI uses Nvidia Tesla M60 which doesn't support fp16. broken_tests_keyword_set.insert({"FLOAT16"}); - } if (provider_name == "dml") { @@ -547,14 +554,16 @@ TEST_P(ModelTest, Run) { if (iter != broken_tests.end() && (model_version == TestModelInfo::unknown_version || iter->broken_versions_.empty() || iter->broken_versions_.find(model_version) != iter->broken_versions_.end())) { + SkipTest(); return; } for (auto iter2 = broken_tests_keyword_set.begin(); iter2 != broken_tests_keyword_set.end(); ++iter2) { - std::string keyword = *iter2; - if (ToUTF8String(test_case_name).find(keyword) != std::string::npos) { - return; - } + std::string keyword = *iter2; + if (ToUTF8String(test_case_name).find(keyword) != std::string::npos) { + SkipTest(); + return; + } } } bool is_single_node = !model_info->GetNodeName().empty(); @@ -567,7 +576,6 @@ TEST_P(ModelTest, Run) { if (provider_name == "cpu" && is_single_node) use_single_thread.push_back(true); - std::unique_ptr l = CreateOnnxTestCase(ToUTF8String(test_case_name), std::move(model_info), per_sample_tolerance, relative_per_sample_tolerance); @@ -599,7 +607,7 @@ TEST_P(ModelTest, Run) { 1000, 1, 1 << 30, - 1, // enable fp16 + 1, // enable fp16 0, nullptr, 0, @@ -995,7 +1003,7 @@ TEST_P(ModelTest, Run) { return v; } -auto ExpandModelName = [](const ::testing::TestParamInfo& info) { +auto ExpandModelName = [](const ::testing::TestParamInfo& info) { // use info.param here to generate the test suffix std::basic_string name = info.param; @@ -1034,4 +1042,4 @@ auto ExpandModelName = [](const ::testing::TestParamInfo& INSTANTIATE_TEST_SUITE_P(ModelTests, ModelTest, testing::ValuesIn(GetParameterStrings()), ExpandModelName); } // namespace test -} // namespace onnxruntime +} // namespace onnxruntime \ No newline at end of file