diff --git a/onnxruntime/core/optimizer/layer_norm_fusion.cc b/onnxruntime/core/optimizer/layer_norm_fusion.cc index 351c8552c0595..25fc7f3b39d31 100644 --- a/onnxruntime/core/optimizer/layer_norm_fusion.cc +++ b/onnxruntime/core/optimizer/layer_norm_fusion.cc @@ -61,6 +61,14 @@ static bool CheckAxesOnReduceMean(std::vector& axes_values, int64_t ran return true; } +// The fused ops broadcast the reduced result back over the reduced axes, so the sub-graph is only +// equivalent to a layer norm when ReduceMean keeps the reduced dims. +static bool KeepDimsOnReduceMean(const Node& reduce_mean_node) { + const onnxruntime::NodeAttributes& attributes = reduce_mean_node.GetAttributes(); + auto it = attributes.find("keepdims"); + return it == attributes.end() || it->second.i() != 0; +} + static std::vector GetAxesFromReduceMeanNode(Node& reduce_mean_node, const Graph& graph) { const onnxruntime::NodeAttributes& attributes = reduce_mean_node.GetAttributes(); std::vector axes_values; @@ -263,6 +271,7 @@ Status LayerNormFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, !graph_utils::IsSupportedProvider(reduce_mean_node, GetCompatibleExecutionProviders()) || (reduce_mean_node.GetOutputEdgesCount() != 1 && reduce_mean_node.GetOutputEdgesCount() != 2) || graph.NodeProducesGraphOutput(reduce_mean_node) || + !KeepDimsOnReduceMean(reduce_mean_node) || !IsSupportedDataType(reduce_mean_node, 1)) { continue; } @@ -426,6 +435,7 @@ Status LayerNormFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, if (!graph_utils::IsSupportedOptypeVersionAndDomain(reduce_mean2_node, "ReduceMean", {1, 11, 13, 18}) || reduce_mean2_node.GetExecutionProviderType() != reduce_mean_node.GetExecutionProviderType() || !optimizer_utils::CheckOutputEdges(graph, reduce_mean2_node, 1) || + !KeepDimsOnReduceMean(reduce_mean2_node) || !IsSupportedDataType(reduce_mean2_node, 1) || reduce_mean2_node.GetInputEdgesCount() == 0) { continue; @@ -667,8 +677,8 @@ Status SimplifiedLayerNormFusion::ApplyImpl(Graph& graph, bool& modified, int gr Node& reduce_mean_node = *graph.GetNode(p_reduce_mean->Index()); if (!graph_utils::IsSupportedOptypeVersionAndDomain(reduce_mean_node, "ReduceMean", {1, 11, 13, 18}) || reduce_mean_node.GetExecutionProviderType() != pow_node.GetExecutionProviderType() || - !optimizer_utils::CheckOutputEdges(graph, reduce_mean_node, 1) || !IsSupportedDataType(reduce_mean_node, 1) || - reduce_mean_node.GetInputEdgesCount() == 0) { + !optimizer_utils::CheckOutputEdges(graph, reduce_mean_node, 1) || !KeepDimsOnReduceMean(reduce_mean_node) || + !IsSupportedDataType(reduce_mean_node, 1) || reduce_mean_node.GetInputEdgesCount() == 0) { continue; } nodes_to_remove.push_back(reduce_mean_node); diff --git a/onnxruntime/test/optimizer/graph_transform_test_layernorm.cc b/onnxruntime/test/optimizer/graph_transform_test_layernorm.cc index 62db1844782d0..979e5a047d990 100644 --- a/onnxruntime/test/optimizer/graph_transform_test_layernorm.cc +++ b/onnxruntime/test/optimizer/graph_transform_test_layernorm.cc @@ -466,6 +466,93 @@ TEST_F(GraphTransformationTests, LayerNormWithCastFusionTest_9) { TransformerLevel::Level2, 1, nullptr, post_graph_checker)); } +// ReduceMean with keepdims=0 drops the reduced dim, so the sub-graph does not compute a layer norm. +TEST_F(GraphTransformationTests, LayerNormFusion_ReduceMeanNoKeepDims_NoFusion) { + auto build_test_case = [](ModelTestBuilder& builder) { + auto* input = builder.MakeInput({{4, 4}}); + auto* pow_exponent = builder.MakeInitializer({}, {2.0f}); + auto* epsilon = builder.MakeInitializer({}, {1e-5f}); + auto* scale = builder.MakeInitializer({4}, {1.0f, 1.0f, 1.0f, 1.0f}); + auto* bias = builder.MakeInitializer({4}, {0.0f, 0.0f, 0.0f, 0.0f}); + + auto* mean_out = builder.MakeIntermediate(); + auto* sub_out = builder.MakeIntermediate(); + auto* pow_out = builder.MakeIntermediate(); + auto* reduce_mean_out = builder.MakeIntermediate(); + auto* add_out = builder.MakeIntermediate(); + auto* sqrt_out = builder.MakeIntermediate(); + auto* div_out = builder.MakeIntermediate(); + auto* mul_out = builder.MakeIntermediate(); + auto* output = builder.MakeOutput(); + + Node& mean_node = builder.AddNode("ReduceMean", {input}, {mean_out}); + mean_node.AddAttribute("axes", std::vector{-1}); + mean_node.AddAttribute("keepdims", static_cast(0)); + builder.AddNode("Sub", {input, mean_out}, {sub_out}); + builder.AddNode("Pow", {sub_out, pow_exponent}, {pow_out}); + Node& reduce_mean_node = builder.AddNode("ReduceMean", {pow_out}, {reduce_mean_out}); + reduce_mean_node.AddAttribute("axes", std::vector{-1}); + reduce_mean_node.AddAttribute("keepdims", static_cast(0)); + builder.AddNode("Add", {reduce_mean_out, epsilon}, {add_out}); + builder.AddNode("Sqrt", {add_out}, {sqrt_out}); + builder.AddNode("Div", {sub_out, sqrt_out}, {div_out}); + builder.AddNode("Mul", {div_out, scale}, {mul_out}); + builder.AddNode("Add", {mul_out, bias}, {output}); + }; + + auto post_graph_checker = [](Graph& graph) { + const auto op_to_count = CountOpsInGraph(graph); + TEST_RETURN_IF_NOT(op_to_count.find("LayerNormalization") == op_to_count.end()); + const auto reduce_mean_it = op_to_count.find("ReduceMean"); + TEST_RETURN_IF_NOT(reduce_mean_it != op_to_count.end() && reduce_mean_it->second == 2); + return Status::OK(); + }; + + const InlinedHashSet no_limit_empty_ep_list = {}; + ASSERT_STATUS_OK(TestGraphTransformer( + build_test_case, 13, *logger_, + std::make_unique(no_limit_empty_ep_list, TransformerLevel::Level2), + TransformerLevel::Level2, 1, nullptr, post_graph_checker)); +} + +// Same as above for the simplified pattern, which is the one hit by the model in issue #30513. +TEST_F(GraphTransformationTests, SimplifiedLayerNormFusion_ReduceMeanNoKeepDims_NoFusion) { + auto build_test_case = [](ModelTestBuilder& builder) { + auto* input = builder.MakeInput({{4, 4}}); + auto* pow_exponent = builder.MakeInitializer({}, {2.0f}); + auto* epsilon = builder.MakeInitializer({}, {1e-5f}); + auto* scale = builder.MakeInitializer({4}, {1.0f, 1.0f, 1.0f, 1.0f}); + + auto* pow_out = builder.MakeIntermediate(); + auto* reduce_mean_out = builder.MakeIntermediate(); + auto* add_out = builder.MakeIntermediate(); + auto* sqrt_out = builder.MakeIntermediate(); + auto* div_out = builder.MakeIntermediate(); + auto* output = builder.MakeOutput(); + + builder.AddNode("Pow", {input, pow_exponent}, {pow_out}); + Node& reduce_mean_node = builder.AddNode("ReduceMean", {pow_out}, {reduce_mean_out}); + reduce_mean_node.AddAttribute("axes", std::vector{-1}); + reduce_mean_node.AddAttribute("keepdims", static_cast(0)); + builder.AddNode("Add", {reduce_mean_out, epsilon}, {add_out}); + builder.AddNode("Sqrt", {add_out}, {sqrt_out}); + builder.AddNode("Div", {input, sqrt_out}, {div_out}); + builder.AddNode("Mul", {div_out, scale}, {output}); + }; + + auto post_graph_checker = [](Graph& graph) { + const auto op_to_count = CountOpsInGraph(graph); + TEST_RETURN_IF_NOT(op_to_count.find("SimplifiedLayerNormalization") == op_to_count.end()); + const auto reduce_mean_it = op_to_count.find("ReduceMean"); + TEST_RETURN_IF_NOT(reduce_mean_it != op_to_count.end() && reduce_mean_it->second == 1); + return Status::OK(); + }; + + ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, 17, *logger_, + std::make_unique(), + TransformerLevel::Level2, 1, nullptr, post_graph_checker)); +} + TEST_F(GraphTransformationTests, SimplifiedLayerNormFusionTest) { constexpr const ORTCHAR_T* model_uri = MODEL_FOLDER "fusion/layer_norm_t5.onnx"; std::shared_ptr p_model;