From 2fc3984e70c83ad8c326afb6f75d5ed4dcd843e6 Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Sat, 2 May 2020 07:36:21 +1000 Subject: [PATCH] Add test that C is unidirectionally broadcast-able before fusing the MatMul with Add. (#3780) Addresses #3764 --- .../core/framework/tensorprotoutils.cc | 16 +++++++-- onnxruntime/core/framework/tensorprotoutils.h | 4 +++ .../core/optimizer/matmul_add_fusion.cc | 32 ++++++++++++------ .../test/optimizer/graph_transform_test.cc | 25 +++++++++++--- .../matmul_add_not_broadcastable.onnx | Bin 0 -> 304 bytes 5 files changed, 60 insertions(+), 17 deletions(-) create mode 100644 onnxruntime/test/testdata/transform/matmul_add_fusion/matmul_add_not_broadcastable.onnx diff --git a/onnxruntime/core/framework/tensorprotoutils.cc b/onnxruntime/core/framework/tensorprotoutils.cc index 8ec6fb02ab..4021be3c2d 100644 --- a/onnxruntime/core/framework/tensorprotoutils.cc +++ b/onnxruntime/core/framework/tensorprotoutils.cc @@ -44,6 +44,19 @@ TensorProto ToTensor(const std::vector using namespace ONNX_NAMESPACE; @@ -29,6 +30,10 @@ Status MatMulAddFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, continue; } + if (!graph.GetNodeOutputsInGraphOutputs(node).empty()) { + continue; + } + auto next_node_itr = node.OutputNodesBegin(); if (next_node_itr == node.OutputNodesEnd()) { continue; @@ -69,25 +74,32 @@ Status MatMulAddFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, continue; } - auto matmul_output_name = matmul_node.OutputDefs()[0]->Name(); + const auto& matmul_output = *matmul_node.OutputDefs()[0]; + + auto matmul_output_name = matmul_output.Name(); auto gemm_input_defs = matmul_input_defs; if (matmul_output_name == add_input_defs[0]->Name()) { // matmul output as Add_A, should use Add_B as input C for gemm - // Gemm only support unidirectional broadcast on C - if (add_input_defs[1]->Shape()->dim_size() > 2) { - continue; - } gemm_input_defs.push_back(add_input_defs[1]); } else { // matmul output as Add_B, should use Add_A as input C for gemm - // Gemm only support unidirectional broadcast on C - if (add_input_defs[0]->Shape()->dim_size() > 2) { - continue; - } gemm_input_defs.push_back(add_input_defs[0]); } - if (!graph.GetNodeOutputsInGraphOutputs(node).empty()) { + // valid bias_shapes are (N) or (1, N) or (M, 1) or (M, N) as + // GEMM only supports unidirectional broadcast on the bias input C + const auto& bias_shape = *gemm_input_defs.back()->Shape(); + const auto& M = matmul_output.Shape()->dim()[0]; + const auto& N = matmul_output.Shape()->dim()[1]; + auto dim_has_value_1 = [](const TensorShapeProto_Dimension& dim) { + return dim.has_dim_value() && dim.dim_value() == 1; + }; + + bool valid = ((bias_shape.dim_size() == 1 && bias_shape.dim()[0] == N) || + (bias_shape.dim_size() == 2 && dim_has_value_1(bias_shape.dim()[0]) && bias_shape.dim()[1] == N) || + (bias_shape.dim_size() == 2 && bias_shape.dim()[0] == M && + (dim_has_value_1(bias_shape.dim()[1]) || bias_shape.dim()[1] == N))); + if (!valid) { continue; } diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index 6aae2911c0..368947d6d6 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -601,6 +601,25 @@ TEST_F(GraphTransformationTests, MatMulAddFusion_negitive_case) { ASSERT_TRUE(op_to_count["Gemm"] == 0); } +// Matmul+Add with shape [M,k]*[k,N]+[1,4], won't do the fusion +// 1,4 is not uni-directionally broadcast +TEST_F(GraphTransformationTests, MatMulAddFusion_NotBroadcastable) { + auto model_uri = MODEL_FOLDER "matmul_add_fusion/matmul_add_not_broadcastable.onnx"; + + std::shared_ptr p_model; + ASSERT_STATUS_OK(Model::Load(model_uri, p_model, nullptr, *logger_)); + Graph& graph = p_model->MainGraph(); + + onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; + graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level1); + ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); + + std::map op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["MatMul"] == 1); + ASSERT_TRUE(op_to_count["Add"] == 1); + ASSERT_TRUE(op_to_count["Gemm"] == 0); +} + #ifndef DISABLE_CONTRIB_OPS TEST_F(GraphTransformationTests, Gemm_Relu_three_input) { auto model_uri = MODEL_FOLDER "matmul_add_fusion/3Input/gemm_relu.onnx"; @@ -1116,7 +1135,6 @@ TEST_F(GraphTransformationTests, ReshapeFusionInternalReuseTest) { } } - TEST_F(GraphTransformationTests, ReshapeFusionGraphInputsTest) { auto model_uri = MODEL_FOLDER "fusion/reshape_fusion_with_graph_inputs.onnx"; std::shared_ptr p_model; @@ -1136,7 +1154,6 @@ TEST_F(GraphTransformationTests, ReshapeFusionGraphInputsTest) { ASSERT_EQ(op_to_count["Reshape"], 1); } - TEST_F(GraphTransformationTests, ExpandElimination) { auto model_uri = MODEL_FOLDER "expand_elimination.onnx"; std::shared_ptr model; @@ -1771,9 +1788,9 @@ TEST_F(GraphTransformationTests, SkipLayerNormFusion_Input_Output_Check) { std::vector& output_defs = node.MutableOutputDefs(); #ifdef ENABLE_TRAINING EXPECT_EQ(node.OutputDefs().size(), 3u) << "SkipLayerNormalization number of outputs does not equal to 3. Got:" << node.OutputDefs().size(); -#else +#else EXPECT_EQ(node.OutputDefs().size(), 1u) << "SkipLayerNormalization number of outputs does not equal to 1. Got:" << node.OutputDefs().size(); -#endif +#endif EXPECT_EQ(output_defs[0]->Name(), "19"); } else { EXPECT_EQ(node.OpType(), "MatMul") << "Unexpected node: " << node.OpType() << "," << node.Name(); diff --git a/onnxruntime/test/testdata/transform/matmul_add_fusion/matmul_add_not_broadcastable.onnx b/onnxruntime/test/testdata/transform/matmul_add_fusion/matmul_add_not_broadcastable.onnx new file mode 100644 index 0000000000000000000000000000000000000000..bd4d38c9e162f390682a0fc73715812de6271236 GIT binary patch literal 304 zcmd;J6Jjq(Gs@4)tB_(f)HBmFu$s=qrOU;hnO9I+Vr9U^0cIFl83=LsCYJb?=2#g> zu|Zf$P?}4Ni_pA6;Lg$#7GQKjnD52Hz`)=T#8Q%4ToNUPWE&qB4+onH c$T5TDqa@+}Lz7}o0vZW(44NE^6O#Zp0NSHL)&Kwi literal 0 HcmV?d00001