Skip to content

Commit e5e9174

Browse files
authored
[QNN EP] Add RMSNorm Op support in QNN EP (microsoft#26853)
### Description - Add standalone RMSNorm op translation in QNN EP - Add unit tests ### Motivation and Context - This fixes the CPU fallback of ONNX RMSNormalization operator when running inferences using QNN EP
1 parent cfd5667 commit e5e9174

8 files changed

Lines changed: 398 additions & 0 deletions

File tree

onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.cc

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -868,6 +868,21 @@ bool ScatterElementsNodeGroupSelector::Check(const GraphViewer& graph_viewer, co
868868
return true;
869869
}
870870

871+
bool RMSNormalizationNodeGroupSelector::Check(const GraphViewer& graph_viewer, const Node& node,
872+
const Node* redundant_clip_node,
873+
const std::vector<const Node*>& dq_nodes,
874+
const std::vector<const Node*>& q_nodes) const {
875+
if (!CheckQDQNodes(graph_viewer, node, redundant_clip_node, dq_nodes, q_nodes)) {
876+
return false;
877+
}
878+
879+
int32_t dt_input = dq_nodes[0]->InputDefs()[0]->TypeAsProto()->tensor_type().elem_type();
880+
int32_t dt_output = q_nodes[0]->OutputDefs()[0]->TypeAsProto()->tensor_type().elem_type();
881+
882+
// input and output need to be the same type.
883+
return (dt_input == dt_output);
884+
}
885+
871886
} // namespace QDQ
872887
} // namespace onnxruntime
873888

onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -331,6 +331,15 @@ class ScatterElementsNodeGroupSelector : public NodeGroupSelector {
331331
const std::vector<const Node*>& q_nodes) const override;
332332
};
333333

334+
// Input: DQ nodes for input, scale
335+
// Output: Q node for output
336+
class RMSNormalizationNodeGroupSelector : public NodeGroupSelector {
337+
private:
338+
bool Check(const GraphViewer& graph_viewer, const Node& node, const Node* redundant_clip_node,
339+
const std::vector<const Node*>& dq_nodes,
340+
const std::vector<const Node*>& q_nodes) const override;
341+
};
342+
334343
/*
335344
* NodeSelector instances for use in the QDQ::SelectorActionTransformer.
336345
*/

onnxruntime/core/optimizer/qdq_transformer/selectors_actions/shared/utils.cc

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -161,6 +161,10 @@ static const OpVersionsAndSelector::OpVersionsMap GetScatterElementsOpVersionsMa
161161
return {{"ScatterElements", {}}};
162162
}
163163

164+
static const OpVersionsAndSelector::OpVersionsMap GetRMSNormalizationOpVersionsMap() {
165+
return {{"RMSNormalization", {}}};
166+
}
167+
164168
/* Selector rules registration related */
165169
void RegisterMiscSelectors(Selectors& qdq_selectors) {
166170
/* register selectors for miscellaneous ops */
@@ -310,6 +314,13 @@ void RegisterScatterElementsSelector(Selectors& qdq_selectors) {
310314
std::move(selector));
311315
}
312316

317+
void RegisterRMSNormalizationSelector(Selectors& qdq_selectors) {
318+
/* register selector for RMSNormalization op */
319+
std::unique_ptr<NodeGroupSelector> selector = std::make_unique<RMSNormalizationNodeGroupSelector>();
320+
qdq_selectors.RegisterSelector(GetRMSNormalizationOpVersionsMap(),
321+
std::move(selector));
322+
}
323+
313324
void SelectorManager::CreateSelectors() {
314325
RegisterMiscSelectors(qdq_selectors_);
315326
RegisterDropDQSelectors(qdq_selectors_);
@@ -332,6 +343,7 @@ void SelectorManager::CreateSelectors() {
332343
RegisterTopKSelector(qdq_selectors_);
333344
RegisterCumSumSelector(qdq_selectors_);
334345
RegisterScatterElementsSelector(qdq_selectors_);
346+
RegisterRMSNormalizationSelector(qdq_selectors_);
335347
}
336348

337349
void SelectorManager::InitializeSelectorsMap() {

onnxruntime/core/providers/qnn/builder/op_builder_factory.cc

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -160,6 +160,10 @@ OpBuilderRegistrations::OpBuilderRegistrations() {
160160
CreateLayerNormOpBuilder("LayerNormalization", *this);
161161
}
162162

163+
{
164+
CreateRMSNormOpBuilder("RMSNormalization", *this);
165+
}
166+
163167
{
164168
CreateLRNOpBuilder("LRN", *this);
165169
}

onnxruntime/core/providers/qnn/builder/op_builder_factory.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -91,6 +91,8 @@ void CreateBatchNormOpBuilder(const std::string& op_type, OpBuilderRegistrations
9191

9292
void CreateLayerNormOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations);
9393

94+
void CreateRMSNormOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations);
95+
9496
void CreateLRNOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations);
9597

9698
void CreateTransposeOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations);

onnxruntime/core/providers/qnn/builder/opbuilder/base_op_builder.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -224,6 +224,7 @@ class BaseOpBuilder : public IOpBuilder {
224224
{"InstanceNormalization", QNN_OP_INSTANCE_NORM},
225225
{"BatchNormalization", QNN_OP_BATCHNORM},
226226
{"LayerNormalization", QNN_OP_LAYER_NORM},
227+
{"RMSNormalization", QNN_OP_RMS_NORM},
227228

228229
{"LRN", QNN_OP_LRN},
229230

Lines changed: 180 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,180 @@
1+
// Copyright (c) Qualcomm. All rights reserved.
2+
// Licensed under the MIT License.
3+
4+
#include <cassert>
5+
6+
#include "core/providers/qnn/builder/opbuilder/base_op_builder.h"
7+
#include "core/providers/qnn/builder/qnn_utils.h"
8+
#include "core/providers/qnn/builder/qnn_model_wrapper.h"
9+
#include "core/providers/qnn/builder/op_builder_factory.h"
10+
11+
namespace onnxruntime {
12+
namespace qnn {
13+
14+
class RMSNormOpBuilder : public BaseOpBuilder {
15+
public:
16+
RMSNormOpBuilder() : BaseOpBuilder("RMSNormOpBuilder") {}
17+
ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(RMSNormOpBuilder);
18+
19+
Status IsOpSupported(QnnModelWrapper& qnn_model_wrapper,
20+
const NodeUnit& node_unit,
21+
const logging::Logger& logger) const override final ORT_MUST_USE_RESULT;
22+
23+
protected:
24+
Status ProcessInputs(QnnModelWrapper& qnn_model_wrapper,
25+
const NodeUnit& node_unit,
26+
const logging::Logger& logger,
27+
std::vector<std::string>& input_names,
28+
bool do_op_validation) const override ORT_MUST_USE_RESULT;
29+
Status ProcessAttributesAndOutputs(QnnModelWrapper& qnn_model_wrapper,
30+
const NodeUnit& node_unit,
31+
std::vector<std::string>&& input_names,
32+
const logging::Logger& logger,
33+
bool do_op_validation) const override ORT_MUST_USE_RESULT;
34+
};
35+
36+
Status RMSNormOpBuilder::IsOpSupported(QnnModelWrapper& qnn_model_wrapper,
37+
const NodeUnit& node_unit,
38+
const logging::Logger& logger) const {
39+
const auto& inputs = node_unit.Inputs();
40+
const auto& outputs = node_unit.Outputs();
41+
42+
// Validate scale input is present
43+
constexpr size_t SCALE_IDX = 1;
44+
const bool has_scale_input = inputs.size() > SCALE_IDX && inputs[SCALE_IDX].node_arg.Exists();
45+
ORT_RETURN_IF_NOT(has_scale_input, "QNN EP requires scale input for RMSNorm operator");
46+
47+
// Validate input and output rank constraints
48+
std::vector<uint32_t> input_shape;
49+
ORT_RETURN_IF_NOT(qnn_model_wrapper.GetOnnxShape(inputs[0].node_arg, input_shape), "Cannot get shape of input 0");
50+
const size_t input_rank = input_shape.size();
51+
ORT_RETURN_IF(input_rank > 4, "QNN RMSNorm only supports input rank <= 4");
52+
53+
std::vector<uint32_t> output_shape;
54+
ORT_RETURN_IF_NOT(qnn_model_wrapper.GetOnnxShape(outputs[0].node_arg, output_shape), "Cannot get shape of output 0");
55+
const size_t output_rank = output_shape.size();
56+
ORT_RETURN_IF(output_rank > 4, "QNN RMSNorm only supports output rank <= 4");
57+
58+
// Additional constraints for NPU backend
59+
bool is_npu_backend = IsNpuBackend(qnn_model_wrapper.GetQnnBackendType());
60+
if (is_npu_backend) {
61+
int32_t axis = -1;
62+
Qnn_Scalar_t axis_qnn_scalar = QNN_SCALAR_INIT;
63+
ORT_RETURN_IF_ERROR(ProcessAxisAttribute(qnn_model_wrapper, node_unit, axis_qnn_scalar, axis));
64+
ORT_RETURN_IF(static_cast<size_t>(axis) != input_rank - 1,
65+
"QNN RMSNorm for NPU backend only supports axis with last input dimension");
66+
}
67+
68+
return AddToModelBuilder(qnn_model_wrapper, node_unit, logger, true);
69+
}
70+
71+
Status RMSNormOpBuilder::ProcessInputs(QnnModelWrapper& qnn_model_wrapper,
72+
const NodeUnit& node_unit,
73+
const logging::Logger& logger,
74+
std::vector<std::string>& input_names,
75+
bool do_op_validation) const {
76+
ORT_UNUSED_PARAMETER(do_op_validation);
77+
78+
const auto& inputs = node_unit.Inputs();
79+
constexpr size_t X_IDX = 0;
80+
constexpr size_t SCALE_IDX = 1;
81+
82+
ORT_RETURN_IF_ERROR(ProcessInput(qnn_model_wrapper, inputs[X_IDX], logger, input_names));
83+
ORT_RETURN_IF_ERROR(ProcessInput(qnn_model_wrapper, inputs[SCALE_IDX], logger, input_names));
84+
85+
// Create dummy beta tensor for NPU backend
86+
bool is_npu_backend = IsNpuBackend(qnn_model_wrapper.GetQnnBackendType());
87+
if (is_npu_backend) {
88+
TensorInfo scale_info = {};
89+
ORT_RETURN_IF_ERROR(qnn_model_wrapper.GetTensorInfo(inputs[SCALE_IDX], scale_info));
90+
91+
std::vector<uint32_t> beta_shape = scale_info.shape;
92+
93+
// Match beta datatype to scale for float types, use UFIXED_POINT_8 for INT types
94+
Qnn_DataType_t beta_data_type = QNN_DATATYPE_UFIXED_POINT_8;
95+
if (scale_info.qnn_data_type == QNN_DATATYPE_FLOAT_32 ||
96+
scale_info.qnn_data_type == QNN_DATATYPE_FLOAT_16) {
97+
beta_data_type = scale_info.qnn_data_type;
98+
}
99+
100+
// Use appropriate quantization parameters for zero values
101+
QnnQuantParamsWrapper beta_quant_param;
102+
if (scale_info.quant_param.IsQuantized()) {
103+
float quant_scale = 1.0f;
104+
int32_t zero_point = 0;
105+
beta_quant_param = QnnQuantParamsWrapper(quant_scale, zero_point);
106+
}
107+
108+
const size_t beta_size_in_bytes = utils::GetQnnTensorDataSizeInBytes(beta_shape, beta_data_type);
109+
std::vector<uint8_t> beta_data(beta_size_in_bytes, 0);
110+
const std::string beta_tensor_name = node_unit.Name() + "_beta_dummy";
111+
QnnTensorWrapper beta_tensor_wrapper(beta_tensor_name,
112+
QNN_TENSOR_TYPE_STATIC,
113+
beta_data_type,
114+
std::move(beta_quant_param),
115+
std::move(beta_shape),
116+
std::move(beta_data));
117+
118+
ORT_RETURN_IF_NOT(qnn_model_wrapper.AddTensorWrapper(std::move(beta_tensor_wrapper)),
119+
"Failed to add dummy beta tensor for QNN RMSNorm node.");
120+
input_names.push_back(beta_tensor_name);
121+
}
122+
123+
return Status::OK();
124+
}
125+
126+
Status RMSNormOpBuilder::ProcessAttributesAndOutputs(QnnModelWrapper& qnn_model_wrapper,
127+
const NodeUnit& node_unit,
128+
std::vector<std::string>&& input_names,
129+
const logging::Logger& logger,
130+
bool do_op_validation) const {
131+
NodeAttrHelper node_helper(node_unit);
132+
std::vector<std::string> param_tensor_names;
133+
134+
// Process epsilon attribute
135+
const float epsilon = node_helper.Get("epsilon", 1e-05f);
136+
Qnn_Scalar_t epsilon_param = QNN_SCALAR_INIT;
137+
epsilon_param.dataType = QNN_DATATYPE_FLOAT_32;
138+
epsilon_param.floatValue = epsilon;
139+
QnnParamWrapper epsilon_param_wrapper(node_unit.Index(),
140+
node_unit.Name(),
141+
QNN_OP_RMS_NORM_PARAM_EPSILON,
142+
epsilon_param);
143+
param_tensor_names.push_back(epsilon_param_wrapper.GetParamTensorName());
144+
qnn_model_wrapper.AddParamWrapper(std::move(epsilon_param_wrapper));
145+
146+
// Process axis attribute and create axes parameter
147+
std::vector<uint32_t> input_shape;
148+
ORT_RETURN_IF_NOT(qnn_model_wrapper.GetOnnxShape(node_unit.Inputs()[0].node_arg, input_shape), "Cannot get shape of Input 0");
149+
const size_t input_rank = input_shape.size();
150+
int32_t axis = -1;
151+
Qnn_Scalar_t axis_qnn_scalar = QNN_SCALAR_INIT;
152+
ORT_RETURN_IF_ERROR(ProcessAxisAttribute(qnn_model_wrapper, node_unit, axis_qnn_scalar, axis));
153+
size_t axes_rank = input_rank - static_cast<size_t>(axis);
154+
std::vector<uint32_t> axes(axes_rank, 0);
155+
std::vector<uint32_t> axes_shape{SafeInt<uint32_t>(axes_rank)};
156+
axes[0] = static_cast<uint32_t>(axis);
157+
for (size_t i = 1; i < axes.size(); ++i) {
158+
axes[i] = axes[i - 1] + 1;
159+
}
160+
161+
QnnParamWrapper axes_param(node_unit.Index(), node_unit.Name(), QNN_OP_RMS_NORM_PARAM_AXES,
162+
std::move(axes_shape), std::move(axes));
163+
param_tensor_names.push_back(axes_param.GetParamTensorName());
164+
qnn_model_wrapper.AddParamWrapper(std::move(axes_param));
165+
166+
ORT_RETURN_IF_ERROR(ProcessOutputs(qnn_model_wrapper, node_unit,
167+
std::move(input_names),
168+
std::move(param_tensor_names),
169+
logger,
170+
do_op_validation,
171+
GetQnnOpType(node_unit.OpType())));
172+
return Status::OK();
173+
}
174+
175+
void CreateRMSNormOpBuilder(const std::string& op_type, OpBuilderRegistrations& op_registrations) {
176+
op_registrations.AddOpBuilder(op_type, std::make_unique<RMSNormOpBuilder>());
177+
}
178+
179+
} // namespace qnn
180+
} // namespace onnxruntime

0 commit comments

Comments
 (0)