From 0e48187b4e8b7e7f629e725283751e06c11103d3 Mon Sep 17 00:00:00 2001 From: Yufeng Li Date: Mon, 17 May 2021 10:48:20 -0700 Subject: [PATCH] Add type checks for QDQ transformer (#7715) --- .../qdq_transformer/qdq_average_pool.cc | 13 + .../qdq_transformer/qdq_binary_op.cc | 13 + .../optimizer/qdq_transformer/qdq_concat.cc | 21 +- .../optimizer/qdq_transformer/qdq_conv.cc | 17 +- .../optimizer/qdq_transformer/qdq_matmul.cc | 13 +- .../optimizer/qdq_transformer/qdq_util.cc | 3 +- .../optimizer/graph_transform_test_builder.cc | 16 +- .../optimizer/graph_transform_test_builder.h | 3 +- .../test/optimizer/qdq_transformer_test.cc | 459 +++++++++++++++--- 9 files changed, 471 insertions(+), 87 deletions(-) diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_average_pool.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_average_pool.cc index edaf7900f3..04dca7d96e 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_average_pool.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_average_pool.cc @@ -12,6 +12,7 @@ class QDQAveragePoolTransformer : public QDQOperatorTransformer { public: QDQAveragePoolTransformer(Node& node, Graph& graph) : QDQOperatorTransformer(node, graph) {} + protected: bool TransformImpl(const std::vector& dq_nodes, const std::vector& q_nodes) override { std::vector input_defs(graph_.GetNode(dq_nodes[0]->Index())->MutableInputDefs()); @@ -29,6 +30,18 @@ class QDQAveragePoolTransformer : public QDQOperatorTransformer { .SetExecutionProviderType(kCpuExecutionProvider); return true; } + + bool Check(const std::vector& dq_nodes, const std::vector& q_nodes) const override { + if (!QDQOperatorTransformer::Check(dq_nodes, q_nodes)) { + return false; + } + + // Currently QLinearAveragePool only support activation type uint8_t + int32_t dt_input = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + return dt_input == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8 && + dt_output == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8; + } }; DEFINE_QDQ_CREATOR(AveragePool, QDQAveragePoolTransformer) diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_binary_op.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_binary_op.cc index 1a3c8183a6..d547d5ca4e 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_binary_op.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_binary_op.cc @@ -32,6 +32,19 @@ class QDQBinaryOpTransformer : public QDQOperatorTransformer { .SetExecutionProviderType(kCpuExecutionProvider); return true; } + + bool Check(const std::vector& dq_nodes, const std::vector& q_nodes) const override { + if (!QDQOperatorTransformer::Check(dq_nodes, q_nodes)) { + return false; + } + + // Currently QLinearConv only support activation type uint8_t + int32_t dt_input_1 = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + int32_t dt_input_2 = dq_nodes[1]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + return dt_input_1 == dt_input_2 && + dt_input_1 == dt_output; + } }; DEFINE_QDQ_CREATOR(Add, QDQBinaryOpTransformer) diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_concat.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_concat.cc index 64c209ee60..29cf236295 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_concat.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_concat.cc @@ -24,8 +24,8 @@ class QDQConcatTransformer : public QDQOperatorTransformer { input_defs.push_back(output_qnode->MutableInputDefs()[2]); for (size_t input_index = 0; input_index < input_count; ++input_index) { - auto qinput_defs = graph_.GetNode(dq_nodes[input_index]->Index())->MutableInputDefs(); - input_defs.insert(input_defs.end(), qinput_defs.begin(), qinput_defs.end()); + auto qinput_defs = graph_.GetNode(dq_nodes[input_index]->Index())->MutableInputDefs(); + input_defs.insert(input_defs.end(), qinput_defs.begin(), qinput_defs.end()); } graph_.AddNode(node_.Name(), @@ -39,6 +39,23 @@ class QDQConcatTransformer : public QDQOperatorTransformer { return true; } + + bool Check(const std::vector& dq_nodes, const std::vector& q_nodes) const override { + if (!QDQOperatorTransformer::Check(dq_nodes, q_nodes)) { + return false; + } + + // All DQs' inputs and Q's output should have same data type + int32_t dt_input = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + for (size_t dq_idx = 1; dq_idx < dq_nodes.size(); dq_idx++) { + if (dt_input != dq_nodes[dq_idx]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type()) { + return false; + } + } + + int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + return dt_input == dt_output; + } }; DEFINE_QDQ_CREATOR(Concat, QDQConcatTransformer) diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_conv.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_conv.cc index 45942796c0..d8feac5d10 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_conv.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_conv.cc @@ -42,9 +42,20 @@ class QDQConvTransformer : public QDQOperatorTransformer { return false; } - // Currently QLinearConv only support activation type uint8_t - int32_t dt = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); - return dt == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8; + // Currently QLinearConv only support activation type uint8_t and output type uint8_t + int32_t dt_input = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + if (dt_input != ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8 || + dt_output != ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8) { + return false; + } + + if (dq_nodes.size() < 3) { // no bias + return true; + } + + int32_t dt_bias = dq_nodes[2]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + return dt_bias == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_INT32; } }; diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_matmul.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_matmul.cc index 206900b64c..3e0b47c026 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_matmul.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_matmul.cc @@ -36,8 +36,17 @@ class QDQMatMulTransformer : public QDQOperatorTransformer { } // Currently Quant MatMul only support activation type uint8_t - int32_t dt = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); - return dt == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8; + int32_t dt_input1 = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + if (dt_input1 != ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8) { + return false; + } + + if (q_nodes.size() == 0) { + return true; + } + + int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type(); + return dt_output == ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_UINT8; } private: diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc index 45709a35b5..e0cc37d378 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc @@ -50,7 +50,8 @@ bool IsQDQPairSupported(const Graph& graph, const Node& q_node, const Node& dq_n Initializer dq_zp(*dq_zp_tensor_proto, graph.ModelPath()); Initializer dq_scale(*dq_scale_tensor_proto, graph.ModelPath()); - return *q_zp.data() == *dq_zp.data() && + return q_zp.data_type() == dq_zp.data_type() && + *q_zp.data() == *dq_zp.data() && *q_scale.data() == *dq_scale.data(); } diff --git a/onnxruntime/test/optimizer/graph_transform_test_builder.cc b/onnxruntime/test/optimizer/graph_transform_test_builder.cc index 3f4f11f4de..5731668957 100644 --- a/onnxruntime/test/optimizer/graph_transform_test_builder.cc +++ b/onnxruntime/test/optimizer/graph_transform_test_builder.cc @@ -22,7 +22,8 @@ void TransformerTester(const std::function& buil TransformerLevel target_level, int opset_version, double per_sample_tolerance, - double relative_per_sample_tolerance) { + double relative_per_sample_tolerance, + std::unique_ptr transformer) { // Build the model for this test. std::unordered_map domain_to_version; domain_to_version[kOnnxDomain] = opset_version; @@ -39,11 +40,18 @@ void TransformerTester(const std::function& buil std::string model_data; model.ToProto().SerializeToString(&model_data); - auto run_model = [&](TransformerLevel level, std::vector& fetches) { + auto run_model = [&](TransformerLevel level, std::vector& fetches, std::unique_ptr transformer = nullptr) { SessionOptions session_options; - session_options.graph_optimization_level = level; + session_options.graph_optimization_level = transformer ? baseline_level : level; +#if 0 // enable to dump model for debugging + session_options.optimized_model_filepath = L"model" + std::to_wstring(static_cast(level)) + L".onnx"; +#endif InferenceSessionWrapper session{session_options, GetEnvironment()}; ASSERT_TRUE(session.Load(model_data.data(), static_cast(model_data.size())).IsOK()); + if (transformer) { + ASSERT_TRUE(session.RegisterGraphTransformer(std::move(transformer), level).IsOK()); + } + auto status = session.Initialize(); if (!status.IsOK()) { std::cout << "Model initialized failed with status message: " << status.ErrorMessage() << std::endl; @@ -69,7 +77,7 @@ void TransformerTester(const std::function& buil run_model(baseline_level, baseline_fetches); std::vector target_fetches; - run_model(target_level, target_fetches); + run_model(target_level, target_fetches, std::move(transformer)); size_t num_outputs = baseline_fetches.size(); ASSERT_TRUE(num_outputs == target_fetches.size()); diff --git a/onnxruntime/test/optimizer/graph_transform_test_builder.h b/onnxruntime/test/optimizer/graph_transform_test_builder.h index c5f403161a..6410f8f6cd 100644 --- a/onnxruntime/test/optimizer/graph_transform_test_builder.h +++ b/onnxruntime/test/optimizer/graph_transform_test_builder.h @@ -266,7 +266,8 @@ void TransformerTester(const std::function& buil TransformerLevel target_level, int opset_version = 12, double per_sample_tolerance = 0.0, - double relative_per_sample_tolerance = 0.0); + double relative_per_sample_tolerance = 0.0, + std::unique_ptr transformer = nullptr); } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/optimizer/qdq_transformer_test.cc b/onnxruntime/test/optimizer/qdq_transformer_test.cc index c35b31b043..d046bec794 100644 --- a/onnxruntime/test/optimizer/qdq_transformer_test.cc +++ b/onnxruntime/test/optimizer/qdq_transformer_test.cc @@ -4,6 +4,7 @@ #include "core/graph/model.h" #include "core/graph/onnx_protobuf.h" #include "core/mlas/inc/mlas.h" +#include "core/optimizer/qdq_transformer/qdq_transformer.h" #include "core/session/environment.h" #include "core/session/inference_session.h" #include "test/compare_ortvalue.h" @@ -39,45 +40,107 @@ AddQDQNodePair(ModelTestBuilder& builder, NodeArg* q_input, float scale) { #ifndef DISABLE_CONTRIB_OPS -TEST(QDQTransformerTests, Conv) { - // TODO: enable fully use_default_zp tests after fixing inference bug in QuantizeLinear - auto test_case = [&](const std::vector& input_shape, const std::vector& weights_shape, bool use_default_zp) { +template +void QDQTransformerConvTests() { + auto test_case = [&](const std::vector& input_shape, const std::vector& weights_shape) { auto build_test_case = [&](ModelTestBuilder& builder) { auto* input_arg = builder.MakeInput(input_shape, -1.f, 1.f); auto* output_arg = builder.MakeOutput(); - auto* conv_output = builder.MakeIntermediate(); - auto* weight = builder.MakeInitializer(weights_shape, 0, 255); + typedef std::numeric_limits InputLimits; + typedef std::numeric_limits WeightLimits; + typedef std::numeric_limits OutputLimits; - auto* dq_w_output = builder.MakeIntermediate(); - auto* dq_output = AddQDQNodePair(builder, input_arg, .004f, 129); + InputType input_min_value = InputLimits::min(); + InputType input_max_value = InputLimits::max(); - if (use_default_zp) { - builder.AddDequantizeLinearNode(weight, .003f, dq_w_output); - } else { - builder.AddDequantizeLinearNode(weight, .003f, 118, dq_w_output); + WeightType weight_min_value = WeightLimits::min(); + WeightType weight_max_value = WeightLimits::max(); + if (std::is_same::value) { + weight_min_value /= 2; + weight_max_value /= 2; } - builder.AddConvNode(dq_output, dq_w_output, conv_output); - builder.AddQuantizeLinearNode(conv_output, .0039f, 135, output_arg); + auto* dq_w_output = builder.MakeIntermediate(); + auto* weight = builder.MakeInitializer(weights_shape, weight_min_value, weight_max_value); + builder.AddDequantizeLinearNode(weight, .03f, + (weight_min_value + weight_max_value) / 2 + 1, + dq_w_output); + + auto* dq_bias_output = builder.MakeIntermediate(); + auto* bias = builder.MakeInitializer({weights_shape[0]}, static_cast(0), static_cast(127)); + builder.AddDequantizeLinearNode(bias, .0012f, + 0, + dq_bias_output); + + auto* conv_output = builder.MakeIntermediate(); + auto* dq_output = AddQDQNodePair(builder, input_arg, .04f, + (input_min_value + input_max_value) / 2 + 1); + builder.AddNode("Conv", {dq_output, dq_w_output, dq_bias_output}, {conv_output}); + + auto* q_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(conv_output, .039f, + (OutputLimits::min() + OutputLimits::max()) / 2 + 1, + q_output); + + builder.AddDequantizeLinearNode(q_output, .039f, + (OutputLimits::min() + OutputLimits::max()) / 2 + 1, + output_arg); }; auto check_conv_graph = [&](InferenceSessionWrapper& session) { auto op_to_count = CountOpsInGraph(session.GetGraph()); - EXPECT_EQ(op_to_count["QLinearConv"], 1); - EXPECT_EQ(op_to_count["QuantizeLinear"], 1); - EXPECT_EQ(op_to_count["DequantizeLinear"], 0); + if (std::is_same::value && + std::is_same::value && + std::is_same::value) { + EXPECT_EQ(op_to_count["QLinearConv"], 1); + EXPECT_EQ(op_to_count["QuantizeLinear"], 1); + EXPECT_EQ(op_to_count["DequantizeLinear"], 1); + } else { + EXPECT_EQ(op_to_count["Conv"], 1); + EXPECT_EQ(op_to_count["QLinearConv"], 0); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 4); + } }; - TransformerTester(build_test_case, check_conv_graph, TransformerLevel::Level1, TransformerLevel::Level2); + TransformerTester(build_test_case, + check_conv_graph, + TransformerLevel::Level1, + TransformerLevel::Level2, + 12 /*opset_version*/, + 0.01 /*per_sample_tolerance*/, + 0.01 /*relative_per_sample_tolerance*/, + std::make_unique()); }; - test_case({1, 12, 37}, {32, 12, 5}, true); - test_case({1, 12, 37}, {32, 12, 5}, false); - test_case({1, 23, 13, 13}, {30, 23, 3, 3}, true); - test_case({1, 23, 13, 13}, {30, 23, 3, 3}, false); - test_case({1, 22, 11, 13, 15}, {30, 22, 5, 3, 3}, true); - test_case({1, 22, 11, 13, 15}, {30, 22, 5, 3, 3}, false); + test_case({1, 12, 37}, {32, 12, 5}); + test_case({1, 12, 37}, {32, 12, 5}); + test_case({1, 23, 13, 13}, {30, 23, 3, 3}); + test_case({1, 23, 13, 13}, {30, 23, 3, 3}); + test_case({1, 22, 11, 13, 15}, {30, 22, 5, 3, 3}); + test_case({1, 22, 11, 13, 15}, {30, 22, 5, 3, 3}); +} + +TEST(QDQTransformerTests, Conv) { + QDQTransformerConvTests(); + QDQTransformerConvTests(); + + // bias not int32_t + QDQTransformerConvTests(); + QDQTransformerConvTests(); + + // output not uint8_t + QDQTransformerConvTests(); + QDQTransformerConvTests(); + + // input not uint8_t + QDQTransformerConvTests(); + QDQTransformerConvTests(); + + // input not uint8_t and output not uint8_t + QDQTransformerConvTests(); + QDQTransformerConvTests(); } TEST(QDQTransformerTests, ConvMaxPoolReshape_UInt8) { @@ -187,31 +250,57 @@ TEST(QDQTransformerTests, ConvMaxPoolReshape_Int8) { test_case({1, 22, 11, 13, 15}, {30, 22, 5, 3, 3}); } -TEST(QDQTransformerTests, Add) { +template +void QDQTransformerAveragePoolTests() { auto test_case = [&](const std::vector& input_shape) { auto build_test_case = [&](ModelTestBuilder& builder) { - auto* input1_arg = builder.MakeInput(input_shape, -1.f, 1.f); - auto* input2_arg = builder.MakeInput(input_shape, -1.f, 1.f); + auto* input_arg = builder.MakeInput(input_shape, -1.f, 1.f); auto* output_arg = builder.MakeOutput(); + // add QDQ + AveragePool + auto* dq_output = AddQDQNodePair(builder, input_arg, .0035f, 7); + auto* averagepool_output = builder.MakeIntermediate(); + Node& pool_node = builder.AddNode("AveragePool", {dq_output}, {averagepool_output}); + std::vector pads((input_shape.size() - 2) * 2, 1); + pool_node.AddAttribute("pads", pads); + std::vector kernel_shape(input_shape.size() - 2, 3); + pool_node.AddAttribute("kernel_shape", kernel_shape); - // add QDQ + Add - auto* add_output = builder.MakeIntermediate(); - auto* dq_add_output1 = AddQDQNodePair(builder, input1_arg, .004f, 129); - auto* dq_add_output2 = AddQDQNodePair(builder, input2_arg, .004f, 129); - builder.AddNode("Add", {dq_add_output1, dq_add_output2}, {add_output}); - - // add Q - builder.AddQuantizeLinearNode(add_output, .0039f, 135, output_arg); + // add QDQ output + auto* q_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(averagepool_output, + .0038f, + std::numeric_limits::max() / 2, + q_output); + builder.AddDequantizeLinearNode(q_output, + .0039f, + std::numeric_limits::max() / 2, + output_arg); }; - auto check_add_graph = [&](InferenceSessionWrapper& session) { + auto check_binary_op_graph = [&](InferenceSessionWrapper& session) { auto op_to_count = CountOpsInGraph(session.GetGraph()); - EXPECT_EQ(op_to_count["com.microsoft.QLinearAdd"], 1); - EXPECT_EQ(op_to_count["QuantizeLinear"], 2); - EXPECT_EQ(op_to_count["DequantizeLinear"], 0); + if (std::is_same::value && + std::is_same::value) { + EXPECT_EQ(op_to_count["com.microsoft.QLinearAveragePool"], 1); + EXPECT_EQ(op_to_count["AveragePool"], 0); + EXPECT_EQ(op_to_count["QuantizeLinear"], 1); + EXPECT_EQ(op_to_count["DequantizeLinear"], 1); + } else { + EXPECT_EQ(op_to_count["com.microsoft.QLinearAveragePool"], 0); + EXPECT_EQ(op_to_count["AveragePool"], 1); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 2); + } }; - TransformerTester(build_test_case, check_add_graph, TransformerLevel::Level1, TransformerLevel::Level2); + 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()); }; test_case({1, 12, 37}); @@ -219,31 +308,86 @@ TEST(QDQTransformerTests, Add) { test_case({1, 22, 11, 13, 15}); } -TEST(QDQTransformerTests, Mul) { +TEST(QDQTransformerTests, AveragePool) { + QDQTransformerAveragePoolTests(); + QDQTransformerAveragePoolTests(); + QDQTransformerAveragePoolTests(); + QDQTransformerAveragePoolTests(); +} + +template +void QDQTransformerBinaryOpTests(const std::string& op_type, bool does_input1_support_int8 = true) { auto test_case = [&](const std::vector& input_shape) { auto build_test_case = [&](ModelTestBuilder& builder) { auto* input1_arg = builder.MakeInput(input_shape, -1.f, 1.f); auto* input2_arg = builder.MakeInput(input_shape, -1.f, 1.f); auto* output_arg = builder.MakeOutput(); - // add QDQ + Mul - auto* mul_output = builder.MakeIntermediate(); - auto* dq_mul_output1 = AddQDQNodePair(builder, input1_arg, .004f, 129); - auto* dq_mul_output2 = AddQDQNodePair(builder, input2_arg, .004f, 129); - builder.AddNode("Mul", {dq_mul_output1, dq_mul_output2}, {mul_output}); + // add QDQ 1 + auto* q1_output = builder.MakeIntermediate(); + auto* dq1_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(input1_arg, + .004f, + std::numeric_limits::max() / 2, + q1_output); + builder.AddDequantizeLinearNode(q1_output, + .0039f, + std::numeric_limits::max() / 2, + dq1_output); - // add Q - builder.AddQuantizeLinearNode(mul_output, .0039f, 135, output_arg); + // add QDQ 2 + auto* q2_output = builder.MakeIntermediate(); + auto* dq2_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(input2_arg, + .004f, + std::numeric_limits::max() / 2, + q2_output); + builder.AddDequantizeLinearNode(q2_output, + .0039f, + std::numeric_limits::max() / 2, + dq2_output); + + // add binary operator + auto* binary_op_output = builder.MakeIntermediate(); + builder.AddNode(op_type, {dq1_output, dq2_output}, {binary_op_output}); + + // add QDQ output + auto* q3_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(binary_op_output, + .0038f, + std::numeric_limits::max() / 2, + q3_output); + builder.AddDequantizeLinearNode(q3_output, + .0039f, + std::numeric_limits::max() / 2, + output_arg); }; - auto check_mul_graph = [&](InferenceSessionWrapper& session) { + auto check_binary_op_graph = [&](InferenceSessionWrapper& session) { auto op_to_count = CountOpsInGraph(session.GetGraph()); - EXPECT_EQ(op_to_count["com.microsoft.QLinearMul"], 1); - EXPECT_EQ(op_to_count["QuantizeLinear"], 2); - EXPECT_EQ(op_to_count["DequantizeLinear"], 0); + if (std::is_same::value && + std::is_same::value && + (does_input1_support_int8 || std::is_same::value)) { + EXPECT_EQ(op_to_count["com.microsoft.QLinear" + op_type], 1); + EXPECT_EQ(op_to_count[op_type], 0); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 1); + } else { + EXPECT_EQ(op_to_count["com.microsoft.QLinear" + op_type], 0); + EXPECT_EQ(op_to_count[op_type], 1); + EXPECT_EQ(op_to_count["QuantizeLinear"], 3); + EXPECT_EQ(op_to_count["DequantizeLinear"], 3); + } }; - TransformerTester(build_test_case, check_mul_graph, TransformerLevel::Level1, TransformerLevel::Level2); + 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()); }; test_case({1, 12, 37}); @@ -251,6 +395,158 @@ TEST(QDQTransformerTests, Mul) { test_case({1, 22, 11, 13, 15}); } +TEST(QDQTransformerTests, Add) { + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); +} + +TEST(QDQTransformerTests, Add_Have_Different_Types) { + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); + QDQTransformerBinaryOpTests("Add"); +} + +TEST(QDQTransformerTests, Mul) { + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); +} + +TEST(QDQTransformerTests, Mul_Have_Different_Types) { + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); + QDQTransformerBinaryOpTests("Mul"); +} + +template +void QDQTransformerMatMulTests(bool has_output_q) { + auto test_case = [&](const std::vector& input1_shape, const std::vector& input2_shape) { + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input1_arg = builder.MakeInput(input1_shape, -1.f, 1.f); + auto* input2_arg = builder.MakeInput(input2_shape, -1.f, 1.f); + auto* output_arg = builder.MakeOutput(); + + typedef std::numeric_limits Input1Limits; + typedef std::numeric_limits Input2Limits; + typedef std::numeric_limits OutputTypeLimits; + + // add QDQ 1 + auto* q1_output = builder.MakeIntermediate(); + auto* dq1_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(input1_arg, + .039f, + (Input1Limits::max() + Input1Limits::min()) / 2 + 1, + q1_output); + builder.AddDequantizeLinearNode(q1_output, + .039f, + (Input2Limits::max() + Input1Limits::min()) / 2 + 1, + dq1_output); + + // add QDQ 2 + auto* q2_output = builder.MakeIntermediate(); + auto* dq2_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(input2_arg, + .04f, + (Input2Limits::max() + Input2Limits::min()) / 2 + 1, + q2_output); + builder.AddDequantizeLinearNode(q2_output, + .04f, + (Input2Limits::max() + Input2Limits::min()) / 2 + 1, + dq2_output); + + if (has_output_q) { + // add binary operator + auto* matmul_op_output = builder.MakeIntermediate(); + builder.AddNode("MatMul", {dq1_output, dq2_output}, {matmul_op_output}); + + // add QDQ output + auto* q3_output = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(matmul_op_output, + .039f, + (OutputTypeLimits::max() + OutputTypeLimits::min()) / 2 + 1, + q3_output); + builder.AddDequantizeLinearNode(q3_output, + .039f, + (OutputTypeLimits::max() + OutputTypeLimits::min()) / 2 + 1, + output_arg); + } else { + builder.AddNode("MatMul", {dq1_output, dq2_output}, {output_arg}); + } + }; + + auto check_binary_op_graph = [&](InferenceSessionWrapper& session) { + auto op_to_count = CountOpsInGraph(session.GetGraph()); + if (has_output_q) { + if (std::is_same::value && + std::is_same::value) { + EXPECT_EQ(op_to_count["QLinearMatMul"], 1); + EXPECT_EQ(op_to_count["MatMul"], 0); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 1); + } else { + EXPECT_EQ(op_to_count["QLinearMatMul"], 0); + EXPECT_EQ(op_to_count["MatMul"], 1); + EXPECT_EQ(op_to_count["QuantizeLinear"], 3); + EXPECT_EQ(op_to_count["DequantizeLinear"], 3); + } + } else { + if (std::is_same::value) { + EXPECT_EQ(op_to_count["com.microsoft.MatMulIntegerToFloat"], 1); + EXPECT_EQ(op_to_count["MatMul"], 0); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 0); + } else { + EXPECT_EQ(op_to_count["com.microsoft.MatMulIntegerToFloat"], 0); + EXPECT_EQ(op_to_count["MatMul"], 1); + EXPECT_EQ(op_to_count["QuantizeLinear"], 2); + EXPECT_EQ(op_to_count["DequantizeLinear"], 2); + } + } + }; + + 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()); + }; + + test_case({1, 2, 2}, {1, 2, 4}); + test_case({1, 23, 13, 13}, {13, 13}); + test_case({1, 22, 11, 13, 15}, {1, 22, 11, 15, 15}); +} + +TEST(QDQTransformerTests, MatMul) { + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(true); +} + +TEST(QDQTransformerTests, MatMul_Have_Different_Types) { + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(false); + QDQTransformerMatMulTests(false); + + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(true); + QDQTransformerMatMulTests(true); +} + TEST(QDQTransformerTests, Gather) { auto test_case = [&](const std::vector& input1_shape, const std::vector& weights_shape) { auto build_test_case = [&](ModelTestBuilder& builder) { @@ -553,16 +849,16 @@ TEST(QDQTransformerTests, ConvAveragePoolReshape_UInt8) { builder.AddConvNode(dq_conv_output, dq_w_output, conv_output); // add QDQ + AveragePool - auto* dq_maxpool_output = AddQDQNodePair(builder, conv_output, .0035f, 135); - auto* maxpool_output = builder.MakeIntermediate(); - Node& pool_node = builder.AddNode("AveragePool", {dq_maxpool_output}, {maxpool_output}); + auto* dq_averagepool_output = AddQDQNodePair(builder, conv_output, .0035f, 135); + auto* averagepool_output = builder.MakeIntermediate(); + Node& pool_node = builder.AddNode("AveragePool", {dq_averagepool_output}, {averagepool_output}); std::vector pads((weights_shape.size() - 2) * 2, 1); pool_node.AddAttribute("pads", pads); std::vector kernel_shape(weights_shape.size() - 2, 3); pool_node.AddAttribute("kernel_shape", kernel_shape); // add QDQ + Reshape - auto* dq_reshape_output = AddQDQNodePair(builder, maxpool_output, .0035f, 135); + auto* dq_reshape_output = AddQDQNodePair(builder, averagepool_output, .0035f, 135); auto* reshape_shape = builder.Make1DInitializer({-1}); auto* reshape_output = builder.MakeIntermediate(); builder.AddNode("Reshape", {dq_reshape_output, reshape_shape}, {reshape_output}); @@ -611,16 +907,16 @@ TEST(QDQTransformerTests, ConvAveragePoolReshape_Int8) { builder.AddConvNode(dq_conv_output, dq_w_output, conv_output); // add QDQ + AveragePool - auto* dq_maxpool_output = AddQDQNodePair(builder, conv_output, .0035f, 7); - auto* maxpool_output = builder.MakeIntermediate(); - Node& pool_node = builder.AddNode("AveragePool", {dq_maxpool_output}, {maxpool_output}); + auto* dq_averagepool_output = AddQDQNodePair(builder, conv_output, .0035f, 7); + auto* averagepool_output = builder.MakeIntermediate(); + Node& pool_node = builder.AddNode("AveragePool", {dq_averagepool_output}, {averagepool_output}); std::vector pads((weights_shape.size() - 2) * 2, 1); pool_node.AddAttribute("pads", pads); std::vector kernel_shape(weights_shape.size() - 2, 3); pool_node.AddAttribute("kernel_shape", kernel_shape); // add QDQ + Reshape - auto* dq_reshape_output = AddQDQNodePair(builder, maxpool_output, .0035f, 7); + auto* dq_reshape_output = AddQDQNodePair(builder, averagepool_output, .0035f, 7); auto* reshape_shape = builder.Make1DInitializer({-1}); auto* reshape_output = builder.MakeIntermediate(); builder.AddNode("Reshape", {dq_reshape_output, reshape_shape}, {reshape_output}); @@ -670,16 +966,16 @@ TEST(QDQTransformerTests, ConvAveragePoolReshape_Int8_Fail) { builder.AddConvNode(dq_output, dq_w_output, conv_output); // add QDQ + AveragePool - auto* dq_maxpool_output = AddQDQNodePair(builder, conv_output, .0035f, 7); - auto* maxpool_output = builder.MakeIntermediate(); - Node& pool_node = builder.AddNode("AveragePool", {dq_maxpool_output}, {maxpool_output}); + auto* dq_averagepool_output = AddQDQNodePair(builder, conv_output, .0035f, 7); + auto* averagepool_output = builder.MakeIntermediate(); + Node& pool_node = builder.AddNode("AveragePool", {dq_averagepool_output}, {averagepool_output}); std::vector pads((weights_shape.size() - 2) * 2, 1); pool_node.AddAttribute("pads", pads); std::vector kernel_shape(weights_shape.size() - 2, 3); pool_node.AddAttribute("kernel_shape", kernel_shape); // add QDQ + Reshape - auto* dq_reshape_output = AddQDQNodePair(builder, maxpool_output, .0035f, 7); + auto* dq_reshape_output = AddQDQNodePair(builder, averagepool_output, .0035f, 7); auto* reshape_shape = builder.Make1DInitializer({-1}); auto* reshape_output = builder.MakeIntermediate(); builder.AddNode("Reshape", {dq_reshape_output, reshape_shape}, {reshape_output}); @@ -1143,15 +1439,21 @@ TEST(QDQTransformerTests, QDQPropagation_DQ_Q) { } TEST(QDQTransformerTests, Concat_UInt8) { - auto test_case = [&](const std::vector>& input_shapes, int64_t axis, bool can_trans = true) { + auto test_case = [&](const std::vector>& input_shapes, + int64_t axis, + bool has_input_float = false, + bool has_input_int8 = false, + bool has_output_int8 = false) { auto build_test_case = [&](ModelTestBuilder& builder) { auto input_count = input_shapes.size(); std::vector input_args; std::vector q_input_args; for (size_t i = 0; i < input_count; i++) { input_args.push_back(builder.MakeInput(input_shapes[i], -1.f, 1.f)); - if (!can_trans && i == 0) { + if (i == 0 && has_input_float) { q_input_args.push_back(input_args.back()); + } else if (i == 0 && has_input_int8) { + q_input_args.push_back(AddQDQNodePair(builder, input_args.back(), 0.05f, 1)); } else { q_input_args.push_back(AddQDQNodePair(builder, input_args.back(), 0.05f, 128)); } @@ -1161,15 +1463,22 @@ TEST(QDQTransformerTests, Concat_UInt8) { concat_node.AddAttribute("axis", axis); auto* q_concat_output = builder.MakeIntermediate(); - builder.AddQuantizeLinearNode(concat_output, 0.05f, 128, q_concat_output); + if (has_output_int8) { + builder.AddQuantizeLinearNode(concat_output, 0.05f, 1, q_concat_output); - auto* output_arg = builder.MakeOutput(); - builder.AddDequantizeLinearNode(q_concat_output, 0.05f, 128, output_arg); + auto* output_arg = builder.MakeOutput(); + builder.AddDequantizeLinearNode(q_concat_output, 0.05f, 1, output_arg); + } else { + builder.AddQuantizeLinearNode(concat_output, 0.05f, 128, q_concat_output); + + auto* output_arg = builder.MakeOutput(); + builder.AddDequantizeLinearNode(q_concat_output, 0.05f, 128, output_arg); + } }; - auto check_mp_reshape_graph = [&, can_trans](InferenceSessionWrapper& session) { + auto check_mp_reshape_graph = [&input_shapes, &has_input_float, &has_input_int8, &has_output_int8](InferenceSessionWrapper& session) { auto op_to_count = CountOpsInGraph(session.GetGraph()); - if (!can_trans) { + if (has_input_float || has_input_int8 || has_output_int8) { EXPECT_EQ(op_to_count["com.microsoft.QLinearConcat"], 0); } else { EXPECT_EQ(op_to_count["QuantizeLinear"], static_cast(input_shapes.size())); @@ -1184,13 +1493,15 @@ TEST(QDQTransformerTests, Concat_UInt8) { TransformerLevel::Level2, 12 /*opset_version*/, 0.01f /*per_sample_tolerance*/, - 0.01f /*relative_per_sample_tolerance*/); + 0.01f /*relative_per_sample_tolerance*/, + std::make_unique()); }; test_case({{1, 6, 36}, {1, 3, 36}}, 1); - test_case({{1, 6, 36}, {1, 3, 36}}, 1, false); test_case({{1, 6, 36}, {1, 6, 8}, {1, 6, 2}}, 2); - test_case({{1, 6, 36}, {1, 6, 8}, {1, 6, 2}}, 2, false); + test_case({{1, 6, 36}, {1, 6, 8}, {1, 6, 2}}, 2, true); + test_case({{1, 6, 36}, {1, 6, 8}, {1, 6, 2}}, 2, false, true); + test_case({{1, 6, 36}, {1, 6, 8}, {1, 6, 2}}, 2, false, false, true); } #endif // DISABLE_CONTRIB_OPS