From ce4ac6d3280f61b8c30801d044f32d423f69559c Mon Sep 17 00:00:00 2001 From: ytaous <4484531+ytaous@users.noreply.github.com> Date: Thu, 2 Jun 2022 23:23:51 -0700 Subject: [PATCH] Optimizer - add missing supported version for BiasSoftmaxFusion (#11616) * add missing version * opset check * fix format * reject fusion if type not allowed * per comments * trigger new build Co-authored-by: Ethan Tao --- .../core/optimizer/bias_softmax_fusion.cc | 21 ++++++-- .../test/optimizer/graph_transform_test.cc | 31 +++++++---- .../fusion/bias_softmax_fusion_bfloat16.onnx | Bin 0 -> 227 bytes ...softmax_fusion_simple_no_axis_opset13.onnx | Bin 0 -> 214 bytes .../transform/fusion/bias_softmax_gen.py | 51 +++++++++++++++++- 5 files changed, 88 insertions(+), 15 deletions(-) mode change 100644 => 100755 onnxruntime/core/optimizer/bias_softmax_fusion.cc mode change 100644 => 100755 onnxruntime/test/optimizer/graph_transform_test.cc create mode 100644 onnxruntime/test/testdata/transform/fusion/bias_softmax_fusion_bfloat16.onnx create mode 100644 onnxruntime/test/testdata/transform/fusion/bias_softmax_fusion_simple_no_axis_opset13.onnx mode change 100644 => 100755 onnxruntime/test/testdata/transform/fusion/bias_softmax_gen.py diff --git a/onnxruntime/core/optimizer/bias_softmax_fusion.cc b/onnxruntime/core/optimizer/bias_softmax_fusion.cc old mode 100644 new mode 100755 index f44c2169fd..0ca2e8b13f --- a/onnxruntime/core/optimizer/bias_softmax_fusion.cc +++ b/onnxruntime/core/optimizer/bias_softmax_fusion.cc @@ -59,9 +59,23 @@ bool TryBiasSoftmaxSubgraphMatch(Graph& graph, Node& start, Node*& add, Node*& s return false; } + // BiasSoftmax supports only float/float16/double - see ./onnxruntime/core/graph/contrib_ops/contrib_defs.cc + auto type_allowed = [](NodeArg* input) { + auto data_type = input->TypeAsProto()->tensor_type().elem_type(); + if (data_type != ONNX_NAMESPACE::TensorProto_DataType_DOUBLE && + data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 && + data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + return false; + } + return true; + }; + if (!type_allowed(input1) || !type_allowed(input2)) { + return false; + } + // check add is only consumed by softmax with matching exec provider Node& softmax_node = *graph.GetNode(add_node.OutputNodesBegin()->Index()); - if (!graph_utils::IsSupportedOptypeVersionAndDomain(softmax_node, "Softmax", {1, 11}) || + if (!graph_utils::IsSupportedOptypeVersionAndDomain(softmax_node, "Softmax", {1, 11, 13}) || softmax_node.GetExecutionProviderType() != add_node.GetExecutionProviderType()) { return false; } @@ -107,11 +121,12 @@ bool TrySelectInputAndBiasWithAlignment( // confirm all dimensions starting from softmax axis match for input and mask bool singlebatch_shape_matches = true; - int axis = 1; + // default axis = -1 if opset >= 13 + int axis = graph_utils::MatchesOpSinceVersion(softmax_node, {1, 11}) ? 1 : -1; auto& softmax_attr = softmax_node.GetAttributes(); if (softmax_attr.find("axis") != softmax_attr.end()) { auto& axis_attr = softmax_attr.at("axis"); - axis = utils::HasInt(axis_attr) ? (int)axis_attr.i() : 1; + axis = utils::HasInt(axis_attr) ? (int)axis_attr.i() : axis; } int N1 = input1->Shape()->dim_size(); diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc old mode 100644 new mode 100755 index aa85e9029f..f24f317b18 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -3514,12 +3514,9 @@ struct BiasSoftmaxFusionTester { } } - void TestFusionOccurs(int expected_broadcast_axis) { + void TestFusionOccurs(int expected_broadcast_axis, int expected_softmax_axis) { ASSERT_STATUS_OK(model_load_); - int expected_softmax_axis = 1; - GetAxis("Softmax", "axis", &expected_softmax_axis); - ASSERT_STATUS_OK(graph_transformation_mgr_.ApplyTransformers(p_model_->MainGraph(), TransformerLevel::Level2, *logger_)); std::map op_to_count = CountOpsInGraph(p_model_->MainGraph()); @@ -3556,25 +3553,37 @@ TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_GpuOnly) { TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_Simple_Rocm) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_simple.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get(), kRocmExecutionProvider); - tester.TestFusionOccurs(1); + tester.TestFusionOccurs(1, 1); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_Simple_Cuda) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_simple.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(1); + tester.TestFusionOccurs(1, 1); +} + +TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_Simple_Opset13_DefaultAxis) { + auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_simple_no_axis_opset13.onnx"; + BiasSoftmaxFusionTester tester(model_uri, logger_.get()); + tester.TestFusionOccurs(1, 1); +} + +TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_BFloat16_Input) { + auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_bfloat16.onnx"; + BiasSoftmaxFusionTester tester(model_uri, logger_.get()); + tester.TestNoFusionOccurs(); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_MiddleOnes) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_middleones.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(3); + tester.TestFusionOccurs(3, 6); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_ReversedInputs) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_middleones_reversed.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(3); + tester.TestFusionOccurs(3, 6); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_BadAxis) { @@ -3586,19 +3595,19 @@ TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_BadAxis) { TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_AllLeadingOnes) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_allleadingones.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(0); + tester.TestFusionOccurs(0, 6); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_SomeLeadingOnes) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_someleadingones.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(0); + tester.TestFusionOccurs(0, 6); } TEST_F(GraphTransformationTests, BiasSoftmaxFusionTest_NoLeadingOnes) { auto model_uri = MODEL_FOLDER "fusion/bias_softmax_fusion_noleadingones.onnx"; BiasSoftmaxFusionTester tester(model_uri, logger_.get()); - tester.TestFusionOccurs(0); + tester.TestFusionOccurs(0, 6); } static void TestBiasDropoutFusion(const PathString& file_path, const logging::Logger& logger, const int add_count = 0) { diff --git a/onnxruntime/test/testdata/transform/fusion/bias_softmax_fusion_bfloat16.onnx b/onnxruntime/test/testdata/transform/fusion/bias_softmax_fusion_bfloat16.onnx new file mode 100644 index 0000000000000000000000000000000000000000..04b53417ecf11f38f1e3dd23714ce32b7d79b168 GIT binary patch literal 227 zcmd|7+MkXO4ptg9ZOuSoZab|vAlq}R9ArUSi4gn!P zE>Q