From 9a65a12633604eefc4bc3146dc4614b86fa7eac4 Mon Sep 17 00:00:00 2001 From: Tianlei WU Date: Thu, 4 Jun 2026 14:23:05 -0700 Subject: [PATCH] update qdq test to avoid asan oom --- .../qdq_transformer_fastmath_test.cc | 57 +++++++------------ .../test/optimizer/qdq_transformer_test.cc | 53 +++++++---------- 2 files changed, 43 insertions(+), 67 deletions(-) diff --git a/onnxruntime/test/optimizer/qdq_transformer_fastmath_test.cc b/onnxruntime/test/optimizer/qdq_transformer_fastmath_test.cc index 6b431c10f978a..7f09cea87cb3a 100644 --- a/onnxruntime/test/optimizer/qdq_transformer_fastmath_test.cc +++ b/onnxruntime/test/optimizer/qdq_transformer_fastmath_test.cc @@ -324,7 +324,7 @@ TEST(QDQTransformerTests, MatMul_S8S8U8_DisableFastMath) { template void QDQTransformerGemmTests(bool has_output_q, bool has_bias, bool beta_not_one = false, - bool disable_fastmath = false, bool alpha_not_one = false) { + bool disable_fastmath = false, bool alpha_not_one = false, int opset_version = 0) { auto test_case = [&](const std::vector& input1_shape, const std::vector& input2_shape, bool use_contrib_qdq = false) { auto build_test_case = [&](ModelTestBuilder& builder) { @@ -435,33 +435,19 @@ void QDQTransformerGemmTests(bool has_output_q, bool has_bias, bool beta_not_one kOrtSessionOptionsMlasGemmFastMathArm64Bfloat16, "1")); }; - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 12 /*opset_version*/, - NAN /*per_sample_tolerance*/, - NAN /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed()), - add_session_options); - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 18 /*opset_version*/, - NAN /*per_sample_tolerance*/, - NAN /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed()), - add_session_options); - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 19 /*opset_version*/, - NAN /*per_sample_tolerance*/, - NAN /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed()), - add_session_options); + const auto opset_versions = opset_version == 0 ? std::vector{12, 18, 19} + : std::vector{opset_version}; + for (int current_opset_version : opset_versions) { + TransformerTester(build_test_case, + check_binary_op_graph, + TransformerLevel::Level1, + TransformerLevel::Level2, + current_opset_version, + NAN /*per_sample_tolerance*/, + NAN /*relative_per_sample_tolerance*/, + std::make_unique(QDQIsInt8Allowed()), + add_session_options); + } if (disable_fastmath) { auto add_session_options = [&](SessionOptions& so) { @@ -498,17 +484,18 @@ void QDQTransformerGemmTests() { QDQTransformerGemmTests(false, true, true); QDQTransformerGemmTests(true, false, true); QDQTransformerGemmTests(true, true, true); - if constexpr (std::is_same_v && std::is_same_v && - std::is_same_v && std::is_same_v) { - QDQTransformerGemmTests(false, false, false, false, true); - QDQTransformerGemmTests(false, true, false, false, true); - QDQTransformerGemmTests(true, false, false, false, true); - QDQTransformerGemmTests(true, true, false, false, true); - } // dummy test to disable the fastmath session QDQTransformerGemmTests(true, true, true, true); } +TEST(QDQTransformerTests, Gemm_AlphaNotOne_U8U8U8_FastMath) { + constexpr int opset_version = 19; + QDQTransformerGemmTests(false, false, false, false, true, opset_version); + QDQTransformerGemmTests(false, true, false, false, true, opset_version); + QDQTransformerGemmTests(true, false, false, false, true, opset_version); + QDQTransformerGemmTests(true, true, false, false, true, opset_version); +} + TEST(QDQTransformerTests, Gemm_U8U8U8_FastMath) { QDQTransformerGemmTests(); QDQTransformerGemmTests(); diff --git a/onnxruntime/test/optimizer/qdq_transformer_test.cc b/onnxruntime/test/optimizer/qdq_transformer_test.cc index 6a545bc7a720a..f60a85ff45efb 100644 --- a/onnxruntime/test/optimizer/qdq_transformer_test.cc +++ b/onnxruntime/test/optimizer/qdq_transformer_test.cc @@ -719,7 +719,7 @@ TEST(QDQTransformerTests, MatMul_S8S8U8) { template void QDQTransformerGemmTests(bool has_output_q, bool has_bias, bool beta_not_one = false, - bool alpha_not_one = false) { + bool alpha_not_one = false, int opset_version = 0) { auto test_case = [&](const std::vector& input1_shape, const std::vector& input2_shape, bool use_contrib_qdq = false) { auto build_test_case = [&](ModelTestBuilder& builder) { @@ -825,30 +825,18 @@ void QDQTransformerGemmTests(bool has_output_q, bool has_bias, bool beta_not_one } }; - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 12 /*opset_version*/, - 0.01 /*per_sample_tolerance*/, - 0.01 /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed())); - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 18 /*opset_version*/, - 0.01 /*per_sample_tolerance*/, - 0.01 /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed())); - TransformerTester(build_test_case, - check_binary_op_graph, - TransformerLevel::Level1, - TransformerLevel::Level2, - 19 /*opset_version*/, - 0.01 /*per_sample_tolerance*/, - 0.01 /*relative_per_sample_tolerance*/, - std::make_unique(QDQIsInt8Allowed())); + const auto opset_versions = opset_version == 0 ? std::vector{12, 18, 19} + : std::vector{opset_version}; + for (int current_opset_version : opset_versions) { + TransformerTester(build_test_case, + check_binary_op_graph, + TransformerLevel::Level1, + TransformerLevel::Level2, + current_opset_version, + 0.01 /*per_sample_tolerance*/, + 0.01 /*relative_per_sample_tolerance*/, + std::make_unique(QDQIsInt8Allowed())); + } }; test_case({2, 2}, {2, 4}); @@ -868,13 +856,14 @@ void QDQTransformerGemmTests() { QDQTransformerGemmTests(false, true, true); QDQTransformerGemmTests(true, false, true); QDQTransformerGemmTests(true, true, true); - if constexpr (std::is_same_v && std::is_same_v && - std::is_same_v && std::is_same_v) { - QDQTransformerGemmTests(false, false, false, true); - QDQTransformerGemmTests(false, true, false, true); - QDQTransformerGemmTests(true, false, false, true); - QDQTransformerGemmTests(true, true, false, true); - } +} + +TEST(QDQTransformerTests, Gemm_AlphaNotOne_U8U8U8) { + constexpr int opset_version = 19; + QDQTransformerGemmTests(false, false, false, true, opset_version); + QDQTransformerGemmTests(false, true, false, true, opset_version); + QDQTransformerGemmTests(true, false, false, true, opset_version); + QDQTransformerGemmTests(true, true, false, true, opset_version); } TEST(QDQTransformerTests, Gemm_U8U8U8) {