From d0f94ba7f1df0f40acaabc1cf57cd7efdf63bd1b Mon Sep 17 00:00:00 2001 From: Yuduo Date: Mon, 19 May 2025 11:36:26 -0700 Subject: [PATCH 1/5] [QNN-EP] Fuse pre-scale (multiply) into Softmax op --- cmake/onnxruntime_unittests.cmake | 1 + .../core/optimizer/bias_softmax_fusion.cc | 2 +- .../builder/qnn_node_group/qnn_node_group.cc | 2 + .../qnn_node_group/scale_softmax_fusion.cc | 210 ++++++++++++++++++ .../qnn_node_group/scale_softmax_fusion.h | 50 +++++ .../scale_softmax_fusion_test.cc | 145 ++++++++++++ 6 files changed, 409 insertions(+), 1 deletion(-) create mode 100644 onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc create mode 100644 onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h create mode 100644 onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc diff --git a/cmake/onnxruntime_unittests.cmake b/cmake/onnxruntime_unittests.cmake index b31fdd4ea1ee1..601795de58f5c 100644 --- a/cmake/onnxruntime_unittests.cmake +++ b/cmake/onnxruntime_unittests.cmake @@ -724,6 +724,7 @@ endif() # or reduced op builds. if(onnxruntime_USE_QNN AND NOT onnxruntime_MINIMAL_BUILD AND NOT onnxruntime_REDUCED_OPS_BUILD) list(APPEND onnxruntime_test_framework_src_patterns ${TEST_SRC_DIR}/providers/qnn/*) + list(APPEND onnxruntime_test_framework_src_patterns ${TEST_SRC_DIR}/providers/qnn/qnn_node_group/*) list(APPEND onnxruntime_test_framework_libs onnxruntime_providers_qnn) list(APPEND onnxruntime_test_providers_dependencies onnxruntime_providers_qnn) if(NOT onnxruntime_BUILD_QNN_EP_STATIC_LIB) diff --git a/onnxruntime/core/optimizer/bias_softmax_fusion.cc b/onnxruntime/core/optimizer/bias_softmax_fusion.cc index bcbb70ba8fac5..2bbc70db16cde 100644 --- a/onnxruntime/core/optimizer/bias_softmax_fusion.cc +++ b/onnxruntime/core/optimizer/bias_softmax_fusion.cc @@ -135,7 +135,7 @@ bool TrySelectInputAndBiasWithAlignment(Node& add_node, Node& softmax_node, Node new_axis = (int)HandleNegativeAxis(axis, rank); // The axis attribute for Softmax in OpSet-11 and OpSet-13 are different. - // Details in function documentatin. + // Details in function documentation. if (is_since_opset_13 && new_axis != rank - 1) return false; int singlebatch_rank = rank - new_axis; diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/qnn_node_group.cc b/onnxruntime/core/providers/qnn/builder/qnn_node_group/qnn_node_group.cc index 20b37a2fb2b22..839079e6c1a8e 100644 --- a/onnxruntime/core/providers/qnn/builder/qnn_node_group/qnn_node_group.cc +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/qnn_node_group.cc @@ -15,6 +15,7 @@ #include "core/providers/qnn/builder/qnn_node_group/hardsigmoid_mul_fusion.h" #include "core/providers/qnn/builder/qnn_node_group/qnn_node_group.h" #include "core/providers/qnn/builder/qnn_node_group/reshape_gemm_fusion.h" +#include "core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h" #include "core/providers/qnn/builder/qnn_utils.h" #include "core/providers/qnn/ort_api.h" @@ -90,6 +91,7 @@ static std::unique_ptr TryQnnFusions( {"DequantizeLinear", DQQFusion::TryFusion}, {"HardSigmoid", HardSigmoidMulFusion::TryFusion}, {"Gemm", ReshapeGemmFusion::TryFusion}, + {"Mul", ScaleSoftmaxFusion::TryFusion}, }; // For now, all fusions involve standalone node units (i.e., no wrapping DQ/Q nodes). diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc new file mode 100644 index 0000000000000..62bd86f3d287f --- /dev/null +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc @@ -0,0 +1,210 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h" + +#include +#include +#include +#include +#include +#include +#include + +#include "core/providers/qnn/builder/qnn_utils.h" +#include "core/providers/qnn/builder/op_builder_factory.h" +#include "core/providers/qnn/builder/qnn_node_group/utils.h" +#include "core/providers/qnn/builder/qnn_model_wrapper.h" +#include "core/providers/qnn/builder/opbuilder/base_op_builder.h" + +namespace onnxruntime { +namespace qnn { +namespace { + +constexpr char kOpMul[] = "Mul"; +constexpr char kOpSoftmax[] = "Softmax"; + +/// @brief Get the index of the scalar input in the mul node +/// @param mul Multiply node unit +/// @return The index of the scalar input (0 or 1) if found, otherwise std::nullopt +std::optional GetMulScalarInputIndex(const NodeUnit* mul) { + const NodeArg* mul_x = mul->GetNode().InputDefs()[0]; + const NodeArg* mul_y = mul->GetNode().InputDefs()[1]; + auto mul_x_shape = utils::GetTensorProtoShape(*mul_x->Shape()); + auto mul_y_shape = utils::GetTensorProtoShape(*mul_y->Shape()); + bool is_x_scalar = mul_x_shape.NumDimensions() == 0; + bool is_y_scalar = mul_y_shape.NumDimensions() == 0; + if (!is_x_scalar && !is_y_scalar) { + return std::nullopt; + } + return is_y_scalar ? 1 : 0; +} + +/// @brief Get the axis for softmax +/// @param node_units The node units containing the softmax and mul nodes +/// @return The axis for softmax +std::optional GetPositiveSoftmaxAxis(const NodeUnit* mul, const NodeUnit* softmax) { + NodeAttrHelper softmax_attr_helper(softmax->GetNode()); + std::optional param_axis = softmax_attr_helper.GetInt64(QNN_OP_SOFTMAX_PARAM_AXIS); + if (!param_axis.has_value()) { + return std::nullopt; + } + int64_t axis_value = param_axis.value(); + if (axis_value < 0) { + size_t input_scale_index = GetMulScalarInputIndex(mul).value(); + size_t input_other_index = 1 - input_scale_index; + int rank = mul->GetNode().InputDefs()[input_other_index]->Shape()->dim_size(); + axis_value += static_cast(rank); + } + return static_cast(axis_value); +} + +/// @brief Identify scalar input from mul node if present +/// @param mul Multiply node unit +/// @return The scalar input float value if found, otherwise std::nullopt +std::optional ExtractScalarValueFromMul(const GraphViewer& graph_viewer, const NodeUnit* mul) { + std::optional input_scale_index = GetMulScalarInputIndex(mul); + if (!input_scale_index.has_value()) { + return std::nullopt; + } + const NodeArg* scalar_arg = mul->GetNode().InputDefs()[input_scale_index.value()]; + if (!graph_viewer.IsConstantInitializer(scalar_arg->Name(), true)) { + return std::nullopt; + } + const auto* scalar_tensor = graph_viewer.GetConstantInitializer(scalar_arg->Name()); + if (!scalar_tensor) { + return std::nullopt; + } + if (scalar_tensor->data_type() != ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + return std::nullopt; + } + const auto& raw_data = scalar_tensor->raw_data(); + if (raw_data.size() != sizeof(float) || reinterpret_cast(raw_data.data()) % alignof(float) != 0) { + return std::nullopt; + } + return *reinterpret_cast(raw_data.data()); +} + +/// @brief Create or validate the QNN node +/// @param qnn_model_wrapper QNN model wrapper +/// @param node_units The node units containing the softmax and mul nodes +/// @param validate Whether to validate the QNN node +/// @return Status +Status CreateOrValidateOnQnn( + QnnModelWrapper* qnn_model_wrapper, + std::array node_units, + bool validate) { + const NodeUnit* mul = node_units[0]; + const NodeUnit* softmax = node_units[1]; + ORT_RETURN_IF_NOT(mul->OpType() == kOpMul, + "Expected scale node to be of type Mul, got ", mul->OpType()); + ORT_RETURN_IF_NOT(softmax->OpType() == kOpSoftmax, + "Expected softmax node to be of type Softmax, got ", softmax->OpType()); + size_t input_scale_index = GetMulScalarInputIndex(mul).value(); + size_t input_other_index = 1 - input_scale_index; + const NodeUnitIODef& mul_input_other = mul->Inputs()[input_other_index]; + const NodeUnitIODef& softmax_output = softmax->Outputs()[0]; + + std::vector param_tensor_names; + { // axis + std::optional axis = GetPositiveSoftmaxAxis(mul, softmax); + if (axis.has_value()) { + Qnn_Scalar_t axis_scalar = QNN_SCALAR_INIT; + axis_scalar.dataType = QNN_DATATYPE_UINT_32; + axis_scalar.uint32Value = axis.value(); + QnnParamWrapper param_wrapper(softmax->Index(), + softmax->Name(), + QNN_OP_SOFTMAX_PARAM_AXIS, + axis_scalar); + ORT_RETURN_IF_NOT(qnn_model_wrapper->AddParamWrapper(std::move(param_wrapper)), "Failed to add param"); + param_tensor_names.push_back(param_wrapper.GetParamTensorName()); + } + } + { // beta + NodeAttrHelper softmax_attr_helper(softmax->GetNode()); + std::optional beta = softmax_attr_helper.GetFloat(QNN_OP_SOFTMAX_PARAM_BETA); + float scale = ExtractScalarValueFromMul(qnn_model_wrapper->GetGraphViewer(), mul).value_or(1.0f); + Qnn_Scalar_t beta_scalar = QNN_SCALAR_INIT; + beta_scalar.dataType = QNN_DATATYPE_FLOAT_32; + beta_scalar.floatValue = scale * beta.value_or(1.0f); + QnnParamWrapper param_wrapper(softmax->Index(), + softmax->Name(), + QNN_OP_SOFTMAX_PARAM_BETA, + beta_scalar); + ORT_RETURN_IF_NOT(qnn_model_wrapper->AddParamWrapper(std::move(param_wrapper)), "Failed to add param"); + param_tensor_names.push_back(param_wrapper.GetParamTensorName()); + } + + QnnTensorWrapper fused_softmax_input; + QnnTensorWrapper fused_softmax_output; + ORT_RETURN_IF_ERROR(qnn_model_wrapper->MakeTensorWrapper(mul_input_other, fused_softmax_input)); + ORT_RETURN_IF_ERROR(qnn_model_wrapper->MakeTensorWrapper(softmax_output, fused_softmax_output)); + + if (validate) { + ORT_RETURN_IF_ERROR(qnn_model_wrapper->ValidateQnnNode(softmax->Name(), + QNN_OP_PACKAGE_NAME_QTI_AISW, + QNN_OP_SOFTMAX, + {fused_softmax_input.GetQnnTensor()}, + {fused_softmax_output.GetQnnTensor()}, + {})); + } else { + ORT_RETURN_IF_NOT(qnn_model_wrapper->AddTensorWrapper(std::move(fused_softmax_input)), "Failed to add input"); + ORT_RETURN_IF_NOT(qnn_model_wrapper->AddTensorWrapper(std::move(fused_softmax_output)), "Failed to add output"); + ORT_RETURN_IF_NOT(qnn_model_wrapper->CreateQnnNode(softmax->Name(), + QNN_OP_PACKAGE_NAME_QTI_AISW, + QNN_OP_SOFTMAX, + {mul_input_other.node_arg.Name()}, + {softmax_output.node_arg.Name()}, + std::move(param_tensor_names), + validate), + "Failed to add fused " + std::string(kOpSoftmax) + " node."); + } + return Status::OK(); +} + +} // namespace + +std::unique_ptr ScaleSoftmaxFusion::TryFusion( + QnnModelWrapper& qnn_model_wrapper, + const NodeUnit& mul_node_unit, + const std::unordered_map& node_to_node_unit, + const std::unordered_map& node_unit_to_qnn_node_group, + [[maybe_unused]] const logging::Logger& logger) { + if (mul_node_unit.OpType() != kOpMul || mul_node_unit.UnitType() != NodeUnit::Type::SingleNode) { + return nullptr; + } + // Check if the mul node has a scalar input that can fold into the softmax's beta + const GraphViewer& graph_viewer = qnn_model_wrapper.GetGraphViewer(); + std::optional scalar = ExtractScalarValueFromMul(graph_viewer, &mul_node_unit); + if (!scalar.has_value()) { + return nullptr; + } + + // Mul node must have a single Softmax node as child + const std::array child_op_types{kOpSoftmax}; + const NodeUnit* softmax = GetOnlyChildOfType(graph_viewer, mul_node_unit, child_op_types, + node_to_node_unit, node_unit_to_qnn_node_group); + if (softmax == nullptr) { + return nullptr; + } + + std::array node_units{&mul_node_unit, softmax}; + if (CreateOrValidateOnQnn(&qnn_model_wrapper, node_units, /*validate=*/true) != Status::OK()) { + return nullptr; + } + + return std::make_unique(node_units); +} + +Status ScaleSoftmaxFusion::IsSupported( + QnnModelWrapper& qnn_model_wrapper, [[maybe_unused]] const logging::Logger& logger) const { + return CreateOrValidateOnQnn(&qnn_model_wrapper, node_units_, /*validate=*/true); +} + +Status ScaleSoftmaxFusion::AddToModelBuilder( + QnnModelWrapper& qnn_model_wrapper, [[maybe_unused]] const logging::Logger& logger) const { + return CreateOrValidateOnQnn(&qnn_model_wrapper, node_units_, /*validate=*/false); +} + +} // namespace qnn +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h new file mode 100644 index 0000000000000..726694a20ddfd --- /dev/null +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h @@ -0,0 +1,50 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include +#include +#include +#include + +#include "core/providers/qnn/builder/qnn_node_group/qnn_node_group.h" +#include "core/providers/qnn/ort_api.h" + +namespace onnxruntime { +namespace qnn { + +class QnnModelWrapper; + +/// +/// Represents a fusion of pattern: Softmax(Mul(x, scalar_scale)) => QnnSoftmax(x, beta=scalar_scale) +/// +class ScaleSoftmaxFusion : public IQnnNodeGroup { + public: + explicit ScaleSoftmaxFusion(const std::array& node_units) : node_units_(node_units) {} + ORT_DISALLOW_COPY_AND_ASSIGNMENT(ScaleSoftmaxFusion); + + Status IsSupported(QnnModelWrapper& qnn_model_wrapper, const logging::Logger& logger) const override; + Status AddToModelBuilder(QnnModelWrapper& qnn_model_wrapper, const logging::Logger& logger) const override; + gsl::span GetNodeUnits() const override { return node_units_; } + const NodeUnit* GetTargetNodeUnit() const override { return node_units_[1]; } + std::string_view Type() const override { return "ScaleSoftmaxFusion"; } + + /// + /// Traverses graph to check if the given starting NodeUnit is part of a valid Softmax -> Mul sequence. + /// If so, returns a IQnnNodeGroup that contains the Softmax and Mul NodeUnits. + /// + static std::unique_ptr TryFusion( + QnnModelWrapper& qnn_model_wrapper, + const NodeUnit& mul_node_unit, + const std::unordered_map& node_to_node_unit, + const std::unordered_map& node_unit_to_qnn_node_group, + const logging::Logger& logger); + + private: + std::array node_units_; +}; + +} // namespace qnn +} // namespace onnxruntime diff --git a/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc new file mode 100644 index 0000000000000..82fd1347f128d --- /dev/null +++ b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc @@ -0,0 +1,145 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#if !defined(ORT_MINIMAL_BUILD) + +#include "core/graph/graph.h" +#include "core/graph/node_attr_utils.h" + +#include "test/optimizer/qdq_test_utils.h" +#include "test/providers/qnn/qnn_test_utils.h" +#include "gtest/gtest.h" + +namespace onnxruntime { +namespace test { + +namespace { + +GetTestModelFn BuildTestCaseScalar( + const TestInputDef& input_def, + float scale_value, + bool use_constant, + bool reverse_input_order, + std::optional softmax_axis = std::nullopt) { + return [&](ModelTestBuilder& builder) -> void { + NodeArg* input = MakeTestInput(builder, input_def); + NodeArg* scale{nullptr}; + if (use_constant) { + onnx::TensorProto scale_value_proto; + scale_value_proto.set_data_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + utils::SetRawDataInTensorProto(scale_value_proto, reinterpret_cast(&scale_value), sizeof(float)); + scale = builder.MakeIntermediate(); + builder.AddNode("Constant", {}, {scale}).AddAttribute("value", scale_value_proto); + } else { + scale = builder.MakeScalarInitializer(scale_value); + } + NodeArg* intermediate = builder.MakeIntermediate(); + auto mul_inputs = reverse_input_order ? std::vector{scale, input} : std::vector{input, scale}; + builder.AddNode("Mul", mul_inputs, {intermediate}); + Node& softmax = builder.AddNode("Softmax", {intermediate}, {builder.MakeOutput()}); + if (softmax_axis.has_value()) { + softmax.AddAttribute("axis", softmax_axis.value()); + } + }; +} + +GetTestModelFn BuildTestCaseNoScalar(const TestInputDef& input_def1, const TestInputDef& input_def2) { + return [&input_def1, input_def2](ModelTestBuilder& builder) -> void { + NodeArg* input = MakeTestInput(builder, input_def1); + NodeArg* scale = MakeTestInput(builder, input_def2); + NodeArg* intermediate = builder.MakeIntermediate(); + builder.AddNode("Mul", {input, scale}, {intermediate}); + builder.AddNode("Softmax", {intermediate}, {builder.MakeOutput()}); + }; +} + +ProviderOptions GetProviderOptions() { + ProviderOptions provider_options; + provider_options["backend_type"] = "htp"; + provider_options["offload_graph_io_quantization"] = "0"; + return provider_options; +} + +} // namespace + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarInitializer) { + ProviderOptions provider_options = GetProviderOptions(); + + auto input_def = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseScalar(input_def, 0.125f, /*use_constant=*/false, /*reverse_input_order=*/false), + provider_options, + /*opset_version=*/13, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); + + +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarConstant) { + ProviderOptions provider_options = GetProviderOptions(); + + auto input_def = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseScalar(input_def, 0.375f, /*use_constant=*/true, /*reverse_input_order=*/false), + provider_options, + /*opset_version=*/14, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarInitializerReversed) { + ProviderOptions provider_options = GetProviderOptions(); + auto input_def = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseScalar(input_def, 0.375f, /*use_constant=*/false, /*reverse_input_order=*/true), + provider_options, + /*opset_version=*/15, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarConstantReversed) { + ProviderOptions provider_options = GetProviderOptions(); + auto input_def = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseScalar(input_def, 0.125f, /*use_constant=*/true, /*reverse_input _order=*/true), + provider_options, + /*opset_version=*/16, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionSoftmaxNegativeAxis) { + ProviderOptions provider_options = GetProviderOptions(); + auto input_def = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseScalar(input_def, 0.125f, + /*use_constant=*/true, /*reverse_input_order=*/true, /*softmax_axis=*/-1), + provider_options, + /*opset_version=*/22, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionSkipNoScalar4d) { + ProviderOptions provider_options = GetProviderOptions(); + auto input_def1 = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + auto input_def2 = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseNoScalar(input_def1, input_def2), + provider_options, + /*opset_version=*/13, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionSkipNoScalar1d) { + ProviderOptions provider_options = GetProviderOptions(); + auto input_def1 = TestInputDef({1, 3, 5, 5}, false, -0.5f, 0.5f); + auto input_def2 = TestInputDef({1}, false, -0.5f, 0.5f); + RunQnnModelTest(BuildTestCaseNoScalar(input_def1, input_def2), + provider_options, + /*opset_version=*/13, + /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, + /*fp32_abs_err=*/1e-2f); +} + +} // namespace test +} // namespace onnxruntime + +#endif // !defined(ORT_MINIMAL_BUILD) From 8372bfe8e70cc1ff8342cbee3e64fecbfb3429f1 Mon Sep 17 00:00:00 2001 From: Yuduo Date: Mon, 19 May 2025 13:53:06 -0700 Subject: [PATCH 2/5] Address review feedback --- .../qnn_node_group/scale_softmax_fusion.cc | 26 ++++++++++++------- .../qnn_node_group/scale_softmax_fusion.h | 8 ++++-- .../scale_softmax_fusion_test.cc | 2 -- 3 files changed, 22 insertions(+), 14 deletions(-) diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc index 62bd86f3d287f..ccf7732542403 100644 --- a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc @@ -9,7 +9,8 @@ #include #include #include -#include +#include +#include #include "core/providers/qnn/builder/qnn_utils.h" #include "core/providers/qnn/builder/op_builder_factory.h" @@ -37,11 +38,12 @@ std::optional GetMulScalarInputIndex(const NodeUnit* mul) { if (!is_x_scalar && !is_y_scalar) { return std::nullopt; } - return is_y_scalar ? 1 : 0; + return is_y_scalar ? 1U : 0U; } /// @brief Get the axis for softmax -/// @param node_units The node units containing the softmax and mul nodes +/// @param mul Multiply node unit +/// @param softmax Softmax node unit /// @return The axis for softmax std::optional GetPositiveSoftmaxAxis(const NodeUnit* mul, const NodeUnit* softmax) { NodeAttrHelper softmax_attr_helper(softmax->GetNode()); @@ -52,7 +54,7 @@ std::optional GetPositiveSoftmaxAxis(const NodeUnit* mul, const NodeUn int64_t axis_value = param_axis.value(); if (axis_value < 0) { size_t input_scale_index = GetMulScalarInputIndex(mul).value(); - size_t input_other_index = 1 - input_scale_index; + size_t input_other_index = 1U - input_scale_index; int rank = mul->GetNode().InputDefs()[input_other_index]->Shape()->dim_size(); axis_value += static_cast(rank); } @@ -92,7 +94,7 @@ std::optional ExtractScalarValueFromMul(const GraphViewer& graph_viewer, /// @return Status Status CreateOrValidateOnQnn( QnnModelWrapper* qnn_model_wrapper, - std::array node_units, + gsl::span node_units, bool validate) { const NodeUnit* mul = node_units[0]; const NodeUnit* softmax = node_units[1]; @@ -101,7 +103,7 @@ Status CreateOrValidateOnQnn( ORT_RETURN_IF_NOT(softmax->OpType() == kOpSoftmax, "Expected softmax node to be of type Softmax, got ", softmax->OpType()); size_t input_scale_index = GetMulScalarInputIndex(mul).value(); - size_t input_other_index = 1 - input_scale_index; + size_t input_other_index = 1U - input_scale_index; const NodeUnitIODef& mul_input_other = mul->Inputs()[input_other_index]; const NodeUnitIODef& softmax_output = softmax->Outputs()[0]; @@ -188,22 +190,26 @@ std::unique_ptr ScaleSoftmaxFusion::TryFusion( return nullptr; } - std::array node_units{&mul_node_unit, softmax}; + std::array node_unit_array{&mul_node_unit, softmax}; + auto node_units = gsl::make_span(node_unit_array.data(), 2); if (CreateOrValidateOnQnn(&qnn_model_wrapper, node_units, /*validate=*/true) != Status::OK()) { return nullptr; } - return std::make_unique(node_units); } +gsl::span ScaleSoftmaxFusion::GetNodeUnits() const { + return gsl::span{node_units_.data(), node_units_.size()}; +} + Status ScaleSoftmaxFusion::IsSupported( QnnModelWrapper& qnn_model_wrapper, [[maybe_unused]] const logging::Logger& logger) const { - return CreateOrValidateOnQnn(&qnn_model_wrapper, node_units_, /*validate=*/true); + return CreateOrValidateOnQnn(&qnn_model_wrapper, GetNodeUnits(), /*validate=*/true); } Status ScaleSoftmaxFusion::AddToModelBuilder( QnnModelWrapper& qnn_model_wrapper, [[maybe_unused]] const logging::Logger& logger) const { - return CreateOrValidateOnQnn(&qnn_model_wrapper, node_units_, /*validate=*/false); + return CreateOrValidateOnQnn(&qnn_model_wrapper, GetNodeUnits(), /*validate=*/false); } } // namespace qnn diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h index 726694a20ddfd..66eb892e7a884 100644 --- a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h @@ -22,12 +22,16 @@ class QnnModelWrapper; /// class ScaleSoftmaxFusion : public IQnnNodeGroup { public: - explicit ScaleSoftmaxFusion(const std::array& node_units) : node_units_(node_units) {} + explicit ScaleSoftmaxFusion(gsl::span node_units) { + ORT_ENFORCE(node_units.size() == 2, "Pattern expect exactly 2 NodeUnits."); + node_units_[0] = node_units[0]; + node_units_[1] = node_units[1]; + } ORT_DISALLOW_COPY_AND_ASSIGNMENT(ScaleSoftmaxFusion); Status IsSupported(QnnModelWrapper& qnn_model_wrapper, const logging::Logger& logger) const override; Status AddToModelBuilder(QnnModelWrapper& qnn_model_wrapper, const logging::Logger& logger) const override; - gsl::span GetNodeUnits() const override { return node_units_; } + gsl::span GetNodeUnits() const override; const NodeUnit* GetTargetNodeUnit() const override { return node_units_[1]; } std::string_view Type() const override { return "ScaleSoftmaxFusion"; } diff --git a/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc index 82fd1347f128d..ef4400cc7c140 100644 --- a/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc +++ b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc @@ -71,8 +71,6 @@ TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarInitializer) { /*opset_version=*/13, /*expected_ep_assignment=*/ExpectedEPNodeAssignment::All, /*fp32_abs_err=*/1e-2f); - - } TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarConstant) { From 7129899ed7d5e9caa6be4d35085e002086cecb0d Mon Sep 17 00:00:00 2001 From: Yuduo Date: Tue, 20 May 2025 14:13:54 -0700 Subject: [PATCH 3/5] Address review feedback --- .../qnn/builder/qnn_node_group/scale_softmax_fusion.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc index ccf7732542403..479175adfe0e3 100644 --- a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc @@ -4,13 +4,13 @@ #include "core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.h" #include -#include #include #include #include #include #include #include +#include #include "core/providers/qnn/builder/qnn_utils.h" #include "core/providers/qnn/builder/op_builder_factory.h" From c7ad7df414ed11cb6768476baf74cd0897bf99f7 Mon Sep 17 00:00:00 2001 From: Yuduo Date: Thu, 22 May 2025 09:48:42 -0700 Subject: [PATCH 4/5] Fix CI runs on Windows tests --- .../providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc index ef4400cc7c140..aa8dc492a95c9 100644 --- a/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc +++ b/onnxruntime/test/providers/qnn/qnn_node_group/scale_softmax_fusion_test.cc @@ -62,6 +62,8 @@ ProviderOptions GetProviderOptions() { } // namespace +#if defined(__aarch64__) || defined(_M_ARM64) || defined(__linux__) + TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionScalarInitializer) { ProviderOptions provider_options = GetProviderOptions(); @@ -137,6 +139,8 @@ TEST_F(QnnHTPBackendTests, ScaleSoftmaxFusionSkipNoScalar1d) { /*fp32_abs_err=*/1e-2f); } +#endif // defined(__aarch64__) || defined(_M_ARM64) || defined(__linux__) + } // namespace test } // namespace onnxruntime From 09a403dd1aa3e721f08fe2e2349a64668c33de38 Mon Sep 17 00:00:00 2001 From: Yuduo Date: Fri, 23 May 2025 12:55:23 -0700 Subject: [PATCH 5/5] Fix CI failure --- .../qnn_node_group/scale_softmax_fusion.cc | 26 +++++++++++++------ 1 file changed, 18 insertions(+), 8 deletions(-) diff --git a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc index 479175adfe0e3..5c7091b3be3cc 100644 --- a/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc +++ b/onnxruntime/core/providers/qnn/builder/qnn_node_group/scale_softmax_fusion.cc @@ -29,16 +29,26 @@ constexpr char kOpSoftmax[] = "Softmax"; /// @param mul Multiply node unit /// @return The index of the scalar input (0 or 1) if found, otherwise std::nullopt std::optional GetMulScalarInputIndex(const NodeUnit* mul) { - const NodeArg* mul_x = mul->GetNode().InputDefs()[0]; const NodeArg* mul_y = mul->GetNode().InputDefs()[1]; - auto mul_x_shape = utils::GetTensorProtoShape(*mul_x->Shape()); - auto mul_y_shape = utils::GetTensorProtoShape(*mul_y->Shape()); - bool is_x_scalar = mul_x_shape.NumDimensions() == 0; - bool is_y_scalar = mul_y_shape.NumDimensions() == 0; - if (!is_x_scalar && !is_y_scalar) { - return std::nullopt; + const NodeArg* mul_x = mul->GetNode().InputDefs()[0]; + auto y_shape_proto = mul_y->Shape(); + auto x_shape_proto = mul_x->Shape(); + bool is_y_scalar = false; + if (y_shape_proto != nullptr) { + auto y_shape = utils::GetTensorProtoShape(*y_shape_proto); + is_y_scalar = y_shape.NumDimensions() == 0; + } + bool is_x_scalar = false; + if (x_shape_proto != nullptr) { + auto x_shape = utils::GetTensorProtoShape(*x_shape_proto); + is_x_scalar = x_shape.NumDimensions() == 0; + } + if (is_y_scalar) { + return 1U; + } else if (is_x_scalar) { + return 0U; } - return is_y_scalar ? 1U : 0U; + return std::nullopt; } /// @brief Get the axis for softmax