diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index 321d67a45b..22a4980e96 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -683,6 +683,18 @@ const Node* GetInputNode(const Node& node, int arg_index) { return &(edge->GetNode()); } +inline std::string ToString(const std::vector& versions) { + std::ostringstream output; + if (!versions.empty()) { + // Convert all but the last element to avoid a trailing ";" + std::copy(versions.begin(), versions.end() - 1, + std::ostream_iterator(output, ";")); + // Now add the last element with no delimiter + output << versions.back(); + } + return output.str(); +} + bool FindPath(const Node& node, bool is_input_edge, const std::vector& edges_to_match, std::vector& result, const logging::Logger& logger) { result.clear(); result.reserve(edges_to_match.size()); @@ -690,12 +702,22 @@ bool FindPath(const Node& node, bool is_input_edge, const std::vectorInputEdgesBegin() : current_node->OutputEdgesBegin(); auto edges_end = is_input_edge ? current_node->InputEdgesEnd() : current_node->OutputEdgesEnd(); for (auto it = edges_begin; it != edges_end; ++it) { - - if (edge.dst_arg_index == it->GetDstArgIndex() && edge.src_arg_index == it->GetSrcArgIndex() && edge.op_type == it->GetNode().OpType() && MatchesOpSinceVersion(it->GetNode(), edge.versions) && MatchesOpSetDomain(it->GetNode(), edge.domain)) { +#ifndef NDEBUG + LOGS(logger, VERBOSE) << "E:" << it->GetSrcArgIndex() << "," << it->GetDstArgIndex() + << "," << it->GetNode().OpType() << "," << it->GetNode().Domain() << "," << it->GetNode().Op()->SinceVersion(); +#endif + if (edge.dst_arg_index == it->GetDstArgIndex() && + edge.src_arg_index == it->GetSrcArgIndex() && + edge.op_type == it->GetNode().OpType() && + MatchesOpSinceVersion(it->GetNode(), edge.versions) && + MatchesOpSetDomain(it->GetNode(), edge.domain)) { // For output edge, there could be multiple edges matched. // This function will return failure in such case by design. if (nullptr != edge_found) { diff --git a/onnxruntime/core/optimizer/attention_fusion.cc b/onnxruntime/core/optimizer/attention_fusion.cc new file mode 100644 index 0000000000..7a57b5a171 --- /dev/null +++ b/onnxruntime/core/optimizer/attention_fusion.cc @@ -0,0 +1,711 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/graph/graph_utils.h" +#include "core/optimizer/initializer.h" +#include "core/optimizer/attention_fusion.h" +#include "core/optimizer/utils.h" +#include + +#define DEBUG_LOG(x) LOGS(logger, VERBOSE) << x + +using namespace ONNX_NAMESPACE; +using namespace onnxruntime::common; +namespace onnxruntime { + +static bool ValidateMatMulInitializer(const Graph& graph, const Node& matmul, int64_t hidden_size) { + const NodeArg& input_b = *(matmul.InputDefs()[1]); + if (!graph_utils::IsInitializer(graph, input_b.Name(), true)) { + return false; + } + + return optimizer_utils::ValidateShape(input_b, {hidden_size, hidden_size}); +} + +static bool ValidateAddBiasInitializer(const Graph& graph, const Node& add, int64_t hidden_size) { + const NodeArg& input_b = *(add.InputDefs()[1]); + if (!graph_utils::IsInitializer(graph, input_b.Name(), true)) { + return false; + } + + return optimizer_utils::ValidateShape(input_b, {hidden_size}); +} + +// Merge 1-D weights (q, k and v) by concanating them one by one. +template +void MergeWeights(const T* q, const T* k, const T* v, std::vector& result, int64_t element_count) { + for (int64_t i = 0; i < element_count; i++) { + result.push_back(*q); + q++; + } + + for (int64_t i = 0; i < element_count; i++) { + result.push_back(*k); + k++; + } + + for (int64_t i = 0; i < element_count; i++) { + result.push_back(*v); + v++; + } +} + +// Merge 2-D weights (q, k and v) by concanating them row by row. +template +void MergeMatMulWeights(const T* q_weight, const T* k_weight, const T* v_weight, std::vector& result, int64_t hidden_size) { + const T* q = q_weight; + const T* k = k_weight; + const T* v = v_weight; + for (int64_t i = 0; i < hidden_size; i++, q += hidden_size, k += hidden_size, v += hidden_size) { + MergeWeights(q, k, v, result, hidden_size); + } +} + +// Load q, k and v weights, and validate their data types. +static bool LoadQkvWeights( + Graph& graph, + const Node& q, const Node& k, const Node& v, + const ONNX_NAMESPACE::TensorProto*& q_tensor, + const ONNX_NAMESPACE::TensorProto*& k_tensor, + const ONNX_NAMESPACE::TensorProto*& v_tensor) { + + if (!graph.GetInitializedTensor(q.InputDefs()[1]->Name(), q_tensor)) { + return false; + } + + // Attention Op requires float or float16 weights. + const auto data_type = q_tensor->data_type(); + if (data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT && + data_type != ONNX_NAMESPACE::TensorProto_DataType_FLOAT16) { + return false; + } + + if (!graph.GetInitializedTensor(k.InputDefs()[1]->Name(), k_tensor) || + data_type != k_tensor->data_type()) { + return false; + } + + if (!graph.GetInitializedTensor(v.InputDefs()[1]->Name(), v_tensor) || + data_type != v_tensor->data_type()) { + return false; + } + + return true; +} + +// Merge the weights of Q, K and V inputs for MatMul or Add (bias) into one input. +static NodeArg& MergeQkvWeights(Graph& graph, int64_t hidden_size, + const ONNX_NAMESPACE::TensorProto* q_tensor, + const ONNX_NAMESPACE::TensorProto* k_tensor, + const ONNX_NAMESPACE::TensorProto* v_tensor, + bool is_matmul) { + assert(nullptr != q_tensor); + assert(nullptr != k_tensor); + assert(nullptr != v_tensor); + auto q_initializer = onnxruntime::make_unique(*q_tensor); + auto k_initializer = onnxruntime::make_unique(*k_tensor); + auto v_initializer = onnxruntime::make_unique(*v_tensor); + auto data_type = q_tensor->data_type(); + + ONNX_NAMESPACE::TensorProto initializer; + initializer.set_name(graph.GenerateNodeArgName(is_matmul ? "qkv_weights" : "qkv_bias")); + // Shape of weights for MatMul is (hidden_size, 3 * hidden_size) + // Shape of weights for Add bias is (3 * hidden_size) + if (is_matmul) { + initializer.add_dims(hidden_size); + } + initializer.add_dims(3 * hidden_size); + initializer.set_data_type(data_type); + const int64_t element_count = 3 * hidden_size * (is_matmul ? hidden_size : 1); + + if (data_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT) { + const float* q_weight = q_initializer->data(); + const float* k_weight = k_initializer->data(); + const float* v_weight = v_initializer->data(); + std::vector result; + result.reserve(element_count); + if (is_matmul) { + MergeMatMulWeights(q_weight, k_weight, v_weight, result, hidden_size); + } else { + MergeWeights(q_weight, k_weight, v_weight, result, hidden_size); + } + initializer.set_raw_data(result.data(), element_count * sizeof(float)); + } else { // data_type == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 + const MLFloat16* q_weight = q_initializer->data(); + const MLFloat16* k_weight = k_initializer->data(); + const MLFloat16* v_weight = v_initializer->data(); + std::vector result; + result.reserve(element_count); + if (is_matmul) { + MergeMatMulWeights(q_weight, k_weight, v_weight, result, hidden_size); + } else { + MergeWeights(q_weight, k_weight, v_weight, result, hidden_size); + } + initializer.set_raw_data(result.data(), element_count * sizeof(MLFloat16)); + } + + return graph_utils::AddInitializer(graph, initializer); +} + +// Add a Cast to convert Mask from int64 to int32. +static NodeArg& CastMaskToInt32(Graph& graph, NodeArg* mask_input, ProviderType provider_type) { + const TensorShapeProto* mask_shape = mask_input->Shape(); + TypeProto mask_int32; + mask_int32.mutable_tensor_type()->set_elem_type(TensorProto_DataType_INT32); + auto dim0 = mask_int32.mutable_tensor_type()->mutable_shape()->add_dim(); + *dim0 = mask_shape->dim(0); + auto dim1 = mask_int32.mutable_tensor_type()->mutable_shape()->add_dim(); + *dim1 = mask_shape->dim(1); + auto& cast32 = graph.GetOrCreateNodeArg(graph.GenerateNodeArgName("Mask_Int32"), &mask_int32); + + Node& node = graph.AddNode(graph.GenerateNodeName("MaskCast"), + "Cast", + "Cast mask from int64 to int32", + {mask_input}, + {&cast32}, + nullptr, + kOnnxDomain); + + // Add attribute: "to" = 6 + ONNX_NAMESPACE::AttributeProto to; + to.set_name("to"); + to.set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT); + to.set_i(static_cast(ONNX_NAMESPACE::TensorProto_DataType_INT32)); + node.AddAttribute("to", to); + + node.SetExecutionProviderType(provider_type); + return cast32; +} + +static NodeArg& AddMaskReduceSum(Graph& graph, NodeArg* reduce_sum_input, TypeProto& output_type, ProviderType provider_type) { + NodeArg& reduce_sum_output = graph.GetOrCreateNodeArg(graph.GenerateNodeArgName("MaskIndex_Int32"), &output_type); + + const std::vector input_defs{reduce_sum_input}; + const std::vector output_defs{&reduce_sum_output}; + Node& node = graph.AddNode( + graph.GenerateNodeName("MaskIndex"), + "ReduceSum", + "Count number of words", + input_defs, + output_defs, + {}, + kOnnxDomain); + + // Add attribute: "axes" = [1] + ONNX_NAMESPACE::AttributeProto axes; + axes.set_name("axes"); + axes.set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INTS); + axes.add_ints(1); + node.AddAttribute("axes", axes); + + // Add attribute: "keepdims" = 0 + ONNX_NAMESPACE::AttributeProto keepdims; + keepdims.set_name("keepdims"); + keepdims.set_type(ONNX_NAMESPACE::AttributeProto_AttributeType::AttributeProto_AttributeType_INT); + keepdims.set_i(static_cast(0)); + node.AddAttribute("keepdims", keepdims); + + node.SetExecutionProviderType(provider_type); + + return reduce_sum_output; +} + +static NodeArg* ProcessMask(Graph& graph, NodeArg* mask_input, ProviderType provider_type, const logging::Logger& logger) { + // Validate mask input shape (batch_size, sequence_length) and data type. + // Note that batch_size and sequence_length could be symbolic. + const TensorShapeProto* mask_shape = mask_input->Shape(); + if (mask_shape == nullptr || mask_shape->dim_size() != 2 || mask_input->Type() == nullptr) { + DEBUG_LOG("Mask shape is unknown or not 2D, or data type unknown"); + return nullptr; + } + + auto data_type = mask_input->TypeAsProto()->tensor_type().elem_type(); + if (data_type != ONNX_NAMESPACE::TensorProto_DataType_INT64 && + data_type != ONNX_NAMESPACE::TensorProto_DataType_INT32) { + DEBUG_LOG("Mask data type is not int32 or int64"); + return nullptr; + } + + NodeArg* reduce_sum_input = mask_input; + if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT64) { + NodeArg& cast_int32 = CastMaskToInt32(graph, mask_input, provider_type); + reduce_sum_input = &cast_int32; + } + + // Construct shape based on mask input shape. Note that batch_size could be symbolic. + TypeProto output_type; + output_type.mutable_tensor_type()->set_elem_type(TensorProto_DataType_INT32); + auto dim = output_type.mutable_tensor_type()->mutable_shape()->add_dim(); + *dim = mask_shape->dim(0); + + NodeArg& output = AddMaskReduceSum(graph, reduce_sum_input, output_type, provider_type); + return &output; +} + +static NodeArg* GetOrCreateMaskIndex( + Graph& graph, + NodeArg* mask_input, + std::map& mask_index_map, + ProviderType provider_type, + const logging::Logger& logger) { + // Lookup in map, and return the mask index if created. + auto search = mask_index_map.find(mask_input->Name()); + if (search != mask_index_map.end()) { + return search->second; + } + + NodeArg* output = ProcessMask(graph, mask_input, provider_type, logger); + if (nullptr == output) { + return nullptr; + } + + // Add it to map for lookup later. + mask_index_map.insert(std::pair(mask_input->Name(), output)); + return output; +} + +Status AttentionFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const { + GraphViewer graph_viewer(graph); + const auto& node_topology_list = graph_viewer.GetNodesInTopologicalOrder(); + + // A map from mask input arg name to mask index output. + std::map mask_index_map; + + int fused_count = 0; + for (auto node_index : node_topology_list) { + auto* p_node = graph.GetNode(node_index); + if (p_node == nullptr) + continue; // we removed the node as part of an earlier fusion + + Node& node = *p_node; + ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); + + if (node.GetOutputEdgesCount() == 4 && + graph_utils::IsSupportedOptypeVersionAndDomain(node, "LayerNormalization", {1}, kOnnxDomain) && + graph_utils::IsSupportedProvider(node, GetCompatibleExecutionProviders())) { + // Get hidden size from layer norm bias tensor shape. + const NodeArg& layer_norm_bias = *(node.InputDefs()[2]); + if (!optimizer_utils::IsShapeKnownOnAllDims(layer_norm_bias, 1)) { + DEBUG_LOG("shape of layer norm bias tensor not expected"); + continue; + } + int64_t hidden_size = layer_norm_bias.Shape()->dim(0).dim_value(); + + // Check that LayerNormalization has 4 children: 1 Add, 3 MatMul + const Node* add_node = nullptr; + int add_count = 0; + int matmul_count = 0; + for (auto it = node.OutputNodesBegin(); it != node.OutputNodesEnd(); ++it) { + if ((*it).OpType().compare("Add") == 0) { + add_count++; + add_node = &(*it); + } else if ((*it).OpType().compare("MatMul") == 0) { + matmul_count++; + } + } + + if (add_count != 1 || matmul_count != 3) { + DEBUG_LOG("Attention subgraph expects 1 Add and 3 MatMul as children of LayerNormalization."); + continue; + } + + if (AttentionFusion::FuseSubGraph(node, *add_node, graph, hidden_size, mask_index_map, logger)) { + fused_count++; + modified = true; + } + } + } + + if (fused_count > 0) { + LOGS(logger, INFO) << "Total fused Attention node count: " << fused_count; + } + + return Status::OK(); +} + +/** Fuse Attention SubGraph. +@remark add_after_layer_norm is the Add node in the bottom of sub-graph. + Abbreviatios: B is batch_size, S is sequence_length, W is hidden_size + N is number of attention heads, H is head size, and W=N*H + B and S could be symbolic. + Graph before Fusion (q_, k_, v_, qk_, qkv_ and mask_ prefix is added before Operator type): + [Input](BxSxW) + | + LayerNormalization + / | | \ [Weights](WxW) + / | | \ / + | q_MatMul k_MatMul v_MatMul [Bias](W) + | | | | / + | q_Add k_Add v_Add [Shape=0,0,N,H] + | | | | / + | q_Reshape k_Reshape v_Reshape [Mask] (BxS) + | | | | | + |q_Transpose k_Transpose v_Transpose mask_Unsqueeze(axes=1) + | (0,2,1,3) (0,2,3,1) (perm=0,2,1,3) | + | \ / | mask_Unsqueeze(axes=2) + | qk_MatMul | | + | | [B=2] | [A=1] mask_Cast(to=1) + | | / | \ / + | qk_Div | mask_Sub [A=1000] + | \ | \ / + | mask_Add <-------- /---------------------mask_Mul + | | / + | Softmax / + | \ / + | \ / + | qkv_MatMul + | | + | Transpose (perm=0,2,1,3) + | | + | Reshape---[shape=0,0,W] + | | + | MatMul----[Weights](WxW) + | | + | Add----[Bias](W) + +-------------------|---+ + | | + Add + +After Fusion: + LayerNormalization [Weights](Wx3W) Mask + | \ / [Bias](3W) | + | \ / / | + | Attention <------------ReduceSum + \ | + \ MatMul + \ | + \ Add + +------|---+ + | | + Add +*/ +bool AttentionFusion::FuseSubGraph(Node& layer_norm, const Node& add_after_layer_norm, Graph& graph, int64_t hidden_size, std::map& mask_index_map, const logging::Logger& logger) { + std::vector parent_path{ + {0, 0, "Add", {7}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 0, "Reshape", {5}, kOnnxDomain}, + {0, 0, "Transpose", {1}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 1, "Transpose", {1}, kOnnxDomain}, + {0, 0, "Reshape", {5}, kOnnxDomain}, + {0, 0, "Add", {7}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 0, "LayerNormalization", {1}, kOnnxDomain}}; + + std::vector edges; + if (!graph_utils::FindPath(add_after_layer_norm, true, parent_path, edges, logger)) { + DEBUG_LOG("Faild to find path v"); + return false; + } + + const Node& add = edges[0]->GetNode(); + const Node& matmul = edges[1]->GetNode(); + const Node& reshape = edges[2]->GetNode(); + const Node& transpose = edges[3]->GetNode(); + const Node& qkv_matmul = edges[4]->GetNode(); + const Node& v_transpose = edges[5]->GetNode(); + const Node& v_reshape = edges[6]->GetNode(); + const Node& v_add = edges[7]->GetNode(); + const Node& v_matmul = edges[8]->GetNode(); + const Node& v_root = edges[9]->GetNode(); + if (v_root.Index() != layer_norm.Index()) { + return false; + } + + if (add.GetOutputEdgesCount() != 1 || + matmul.GetOutputEdgesCount() != 1 || + reshape.GetOutputEdgesCount() != 1 || + transpose.GetOutputEdgesCount() != 1 || + qkv_matmul.GetOutputEdgesCount() != 1 || + v_transpose.GetOutputEdgesCount() != 1 || + v_reshape.GetOutputEdgesCount() != 1 || + v_add.GetOutputEdgesCount() != 1 || + v_matmul.GetOutputEdgesCount() != 1 || + v_root.GetOutputEdgesCount() != 4) { + DEBUG_LOG("Output edge count not expected for nodes in path v"); + return false; + } + + std::vector perm; + if (!(graph_utils::GetRepeatedNodeAttributeValues(transpose, "perm", perm) && perm.size() == 4 && perm[0] == 0 && perm[1] == 2 && perm[2] == 1 && perm[3] == 3)) { + DEBUG_LOG("Failed in match Transpose attribute perm. Expected: 0, 2, 1, 3"); + return false; + } + if (!(graph_utils::GetRepeatedNodeAttributeValues(v_transpose, "perm", perm) && perm.size() == 4 && perm[0] == 0 && perm[1] == 2 && perm[2] == 1 && perm[3] == 3)) { + DEBUG_LOG("Failed in match v_transpose attribute perm. Expected: 0, 2, 1, 3"); + return false; + } + + std::vector v_reshape_shape; + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(v_reshape.InputDefs()[1]), v_reshape_shape) || + v_reshape_shape.size() != 4 || + v_reshape_shape[2] <= 0 || + v_reshape_shape[3] <= 0 || + hidden_size != v_reshape_shape[2] * v_reshape_shape[3]) { + DEBUG_LOG("v_reshape initializer value is not expected"); + return false; + } + + const int64_t num_attention_head = v_reshape_shape[2]; + const int64_t attention_head_size = v_reshape_shape[3]; + + std::vector reshape_shape; + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(reshape.InputDefs()[1]), reshape_shape) || + reshape_shape.size() != 3 || + reshape_shape[2] != hidden_size) { + DEBUG_LOG("reshape initializer value is not expected"); + return false; + } + + // Validate the input shape of MatMul and Add according to hidden_size. + if (!(ValidateAddBiasInitializer(graph, add, hidden_size) && + ValidateMatMulInitializer(graph, matmul, hidden_size) && + ValidateAddBiasInitializer(graph, v_add, hidden_size) && + ValidateMatMulInitializer(graph, v_matmul, hidden_size))) { + DEBUG_LOG("Failed in match v_matmul and v_add input shape"); + return false; + } + + // path 2 to find mask + std::vector mask_path{ + {0, 0, "Softmax", {1, 11}, kOnnxDomain}, + {0, 0, "Add", {7}, kOnnxDomain}, + {0, 1, "Mul", {7}, kOnnxDomain}, + {0, 0, "Sub", {7}, kOnnxDomain}, + {0, 1, "Cast", {9}, kOnnxDomain}, + {0, 0, "Unsqueeze", {1, 11}, kOnnxDomain}, + {0, 0, "Unsqueeze", {1, 11}, kOnnxDomain}}; + + if (!graph_utils::FindPath(qkv_matmul, true, mask_path, edges, logger)) { + DEBUG_LOG("Failed to find path for mask"); + return false; + } + + const Node& softmax = edges[0]->GetNode(); + const Node& mask_add = edges[1]->GetNode(); + const Node& mask_mul = edges[2]->GetNode(); + const Node& mask_sub = edges[3]->GetNode(); + const Node& mask_cast = edges[4]->GetNode(); + const Node& mask_unsqueeze_2 = edges[5]->GetNode(); + const Node& mask_unsqueeze_1 = edges[6]->GetNode(); + + if (softmax.GetOutputEdgesCount() != 1 || + mask_add.GetOutputEdgesCount() != 1 || + mask_sub.GetOutputEdgesCount() != 1 || + mask_cast.GetOutputEdgesCount() != 1 || + mask_unsqueeze_2.GetOutputEdgesCount() != 1 || + mask_unsqueeze_1.GetOutputEdgesCount() != 1) { + DEBUG_LOG("Output edge count not expected for mask nodes"); + return false; + } + + if (!optimizer_utils::IsAttributeWithExpectedValue(softmax, "axis", 3)) { + DEBUG_LOG("Softmax attribute axis is expected to be 3"); + return false; + } + + std::vector axes; + if (!(graph_utils::GetRepeatedNodeAttributeValues(mask_unsqueeze_1, "axes", axes) && axes.size() == 1 && axes[0] == 1)) { + DEBUG_LOG("mask_unsqueeze_1 axes not matched. Expect: 1"); + return false; + } + + if (!(graph_utils::GetRepeatedNodeAttributeValues(mask_unsqueeze_2, "axes", axes) && axes.size() == 1 && axes[0] == 2)) { + DEBUG_LOG("mask_unsqueeze_2 axes not matched. Expect: 2"); + return false; + } + + if (!optimizer_utils::IsInitializerWithExpectedValue(graph, *(mask_sub.InputDefs()[0]), float(1), false)) { + DEBUG_LOG("mask_sub const input not matched"); + return false; + } + + if (!optimizer_utils::IsInitializerWithExpectedValue(graph, *(mask_mul.InputDefs()[1]), float(-10000), false)) { + DEBUG_LOG("mask_mul const input not matched"); + return false; + } + + // path to q + std::vector q_path{ + {0, 0, "Div", {7}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 0, "Transpose", {1}, kOnnxDomain}, + {0, 0, "Reshape", {5}, kOnnxDomain}, + {0, 0, "Add", {7}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 0, "LayerNormalization", {1}, kOnnxDomain}}; + + if (!graph_utils::FindPath(mask_add, true, q_path, edges, logger)) { + DEBUG_LOG("Failed to find path for q"); + return false; + } + + const Node& qk_div = edges[0]->GetNode(); + const Node& qk_matmul = edges[1]->GetNode(); + const Node& q_transpose = edges[2]->GetNode(); + const Node& q_reshape = edges[3]->GetNode(); + const Node& q_add = edges[4]->GetNode(); + const Node& q_matmul = edges[5]->GetNode(); + const Node& q_root = edges[6]->GetNode(); + if (q_root.Index() != layer_norm.Index()) { + DEBUG_LOG("q root should be layer normalization"); + return false; + } + + std::vector q_reshape_shape; + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(q_reshape.InputDefs()[1]), q_reshape_shape) || + q_reshape_shape.size() != 4 || + q_reshape_shape[2] != num_attention_head || + q_reshape_shape[3] != attention_head_size) { + DEBUG_LOG("q_reshape const not matched"); + return false; + } + + float expected_value = std::sqrt(static_cast(attention_head_size)); + if (!optimizer_utils::IsInitializerWithExpectedValue(graph, *(qk_div.InputDefs()[1]), expected_value, false)) { + DEBUG_LOG("qk_div const not matched."); + return false; + } + + if (!(graph_utils::GetRepeatedNodeAttributeValues(q_transpose, "perm", perm) && perm.size() == 4 && perm[0] == 0 && perm[1] == 2 && perm[2] == 1 && perm[3] == 3)) { + DEBUG_LOG("q_transpose perm attribute not matched"); + return false; + } + + if (!(ValidateAddBiasInitializer(graph, q_add, hidden_size) && + ValidateMatMulInitializer(graph, q_matmul, hidden_size))) { + DEBUG_LOG("q_matmul and q_add shape not matched"); + return false; + } + + // path to k + std::vector k_path{ + {0, 1, "Transpose", {1}, kOnnxDomain}, + {0, 0, "Reshape", {5}, kOnnxDomain}, + {0, 0, "Add", {7}, kOnnxDomain}, + {0, 0, "MatMul", {1, 9}, kOnnxDomain}, + {0, 0, "LayerNormalization", {1}, kOnnxDomain}}; + + if (!graph_utils::FindPath(qk_matmul, true, k_path, edges, logger)) { + DEBUG_LOG("Failed to find path for k"); + return false; + } + + const Node& k_transpose = edges[0]->GetNode(); + const Node& k_reshape = edges[1]->GetNode(); + const Node& k_add = edges[2]->GetNode(); + const Node& k_matmul = edges[3]->GetNode(); + const Node& k_root = edges[4]->GetNode(); + if (k_root.Index() != layer_norm.Index()) { + DEBUG_LOG("k root is not layer norm"); + return false; + } + + if (!(graph_utils::GetRepeatedNodeAttributeValues(k_transpose, "perm", perm) && perm.size() == 4 && perm[0] == 0 && perm[1] == 2 && perm[2] == 3 && perm[3] == 1)) { + DEBUG_LOG("k_transpose perm attribute not matched"); + return false; + } + + if (!(ValidateAddBiasInitializer(graph, k_add, hidden_size) && + ValidateMatMulInitializer(graph, k_matmul, hidden_size))) { + DEBUG_LOG("k_matmul and k_add shape not matched"); + return false; + } + + std::vector k_reshape_shape; + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(k_reshape.InputDefs()[1]), k_reshape_shape) || + k_reshape_shape.size() != 4 || + k_reshape_shape[2] != num_attention_head || + k_reshape_shape[3] != attention_head_size) { + DEBUG_LOG("k_reshape const not matched"); + return false; + } + + // Load q, k and v weights + const ONNX_NAMESPACE::TensorProto* q_weight_tensor = nullptr; + const ONNX_NAMESPACE::TensorProto* k_weight_tensor = nullptr; + const ONNX_NAMESPACE::TensorProto* v_weight_tensor = nullptr; + if (!LoadQkvWeights(graph, q_matmul, k_matmul, v_matmul, q_weight_tensor, k_weight_tensor, v_weight_tensor)) { + DEBUG_LOG("Failed to load Q, K and V weights, or data type is not float or float16."); + return false; + } + + const ONNX_NAMESPACE::TensorProto* q_bias_tensor = nullptr; + const ONNX_NAMESPACE::TensorProto* k_bias_tensor = nullptr; + const ONNX_NAMESPACE::TensorProto* v_bias_tensor = nullptr; + if (!LoadQkvWeights(graph, q_add, k_add, v_add, q_bias_tensor, k_bias_tensor, v_bias_tensor)) { + DEBUG_LOG("Failed to load Q, K and V bias tensors, or data type is not float or float16."); + return false; + } + + // Now everything is ready, we will start fusing subgraph. + NodeArg* mask_input = graph.GetNode(mask_unsqueeze_1.Index())->MutableInputDefs()[0]; + NodeArg* mask_index = GetOrCreateMaskIndex(graph, mask_input, mask_index_map, layer_norm.GetExecutionProviderType(), logger); + if (nullptr == mask_index) { + DEBUG_LOG("Failed to create mask index"); + return false; + } + + // Merge Q, K and V weights + NodeArg& qkv_weights = MergeQkvWeights(graph, hidden_size, q_weight_tensor, k_weight_tensor, v_weight_tensor, true); + NodeArg& qkv_bias = MergeQkvWeights(graph, hidden_size, q_bias_tensor, k_bias_tensor, v_bias_tensor, false); + + // Create Attention Node. + const std::vector input_defs{layer_norm.MutableOutputDefs()[0], &qkv_weights, &qkv_bias, mask_index}; + const std::vector output_defs{graph.GetNode(reshape.Index())->MutableOutputDefs()[0]}; + Node& attention_node = graph.AddNode( + graph.GenerateNodeName("Attention"), + "Attention", + "Fused Attention subgraphs ", + input_defs, + output_defs, + nullptr, + kMSDomain); + attention_node.AddAttribute("num_heads", num_attention_head); + + // Assign provider to this new node. + attention_node.SetExecutionProviderType(layer_norm.GetExecutionProviderType()); + + // Remove nodes that are not used anymore. + std::vector nodes_to_remove{ + reshape.Index(), + transpose.Index(), + qkv_matmul.Index(), + v_transpose.Index(), + v_reshape.Index(), + v_add.Index(), + v_matmul.Index(), + softmax.Index(), + mask_add.Index(), + qk_div.Index(), + qk_matmul.Index(), + q_transpose.Index(), + q_reshape.Index(), + q_add.Index(), + q_matmul.Index(), + k_transpose.Index(), + k_reshape.Index(), + k_add.Index(), + k_matmul.Index()}; + + // When the last Attention node is fused. Original mask processing nodes can be removed safely. + if (mask_mul.GetOutputEdgesCount() == 1) { + nodes_to_remove.push_back(mask_mul.Index()); + nodes_to_remove.push_back(mask_sub.Index()); + nodes_to_remove.push_back(mask_cast.Index()); + nodes_to_remove.push_back(mask_unsqueeze_2.Index()); + nodes_to_remove.push_back(mask_unsqueeze_1.Index()); + } + + for (const auto& node_index : nodes_to_remove) { + Node* node = graph.GetNode(node_index); + graph_utils::RemoveNodeOutputEdges(graph, *node); + graph.RemoveNode(node->Index()); + } + + DEBUG_LOG("Fused an attention node."); + + return true; +} + +} // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/attention_fusion.h b/onnxruntime/core/optimizer/attention_fusion.h new file mode 100644 index 0000000000..1dff1e7ecf --- /dev/null +++ b/onnxruntime/core/optimizer/attention_fusion.h @@ -0,0 +1,25 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/optimizer/graph_transformer.h" + +namespace onnxruntime { + +/** +@Class AttentionFusion +Rewrite graph fusing attention subgraph to a single Attention node. +*/ +class AttentionFusion : public GraphTransformer { + public: + AttentionFusion(const std::unordered_set& compatible_execution_providers = {}) noexcept + : GraphTransformer("AttentionFusion", compatible_execution_providers) {} + + Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; + +private: + static bool FuseSubGraph(Node& layer_norm, const Node& add_after_layer_norm, Graph& graph, int64_t hidden_size, std::map& mask_index_map, const logging::Logger& logger); +}; + +} // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index b7e6351dd5..8845752d2c 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -23,6 +23,7 @@ #include "core/optimizer/layer_norm_fusion.h" #include "core/optimizer/skip_layer_norm_fusion.h" #include "core/optimizer/reshape_fusion.h" +#include "core/optimizer/attention_fusion.h" #include "core/mlas/inc/mlas.h" #include "core/session/inference_session.h" @@ -125,6 +126,7 @@ std::vector> GenerateTransformers(TransformerL std::unordered_set cpu_cuda_execution_providers = {onnxruntime::kCpuExecutionProvider, onnxruntime::kCudaExecutionProvider}; transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); + transformers.emplace_back(onnxruntime::make_unique(cpu_cuda_execution_providers)); std::unordered_set cuda_execution_providers = {onnxruntime::kCudaExecutionProvider}; transformers.emplace_back(onnxruntime::make_unique(cuda_execution_providers)); diff --git a/onnxruntime/core/optimizer/reshape_fusion.cc b/onnxruntime/core/optimizer/reshape_fusion.cc index b8f14a6a48..0b54e46745 100644 --- a/onnxruntime/core/optimizer/reshape_fusion.cc +++ b/onnxruntime/core/optimizer/reshape_fusion.cc @@ -10,32 +10,6 @@ using namespace ONNX_NAMESPACE; using namespace onnxruntime::common; namespace onnxruntime { -// Get values of integer tensor from initializer, and append them to a vector. -static bool LoadIntegerTensor(const Graph& graph, const NodeArg& input_arg, std::vector& data) { - const ONNX_NAMESPACE::TensorProto* tensor_proto = nullptr; - if (!graph.GetInitializedTensor(input_arg.Name(), tensor_proto)) { - return false; - } - - auto init_const = onnxruntime::make_unique(*tensor_proto); - const auto data_type = tensor_proto->data_type(); - if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT64) { - const int64_t* val = init_const->data(); - data.reserve(data.size() + init_const->size()); - data.insert(data.end(), val, val + init_const->size()); - } else if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT32) { - const int32_t* val = init_const->data(); - data.reserve(data.size() + init_const->size()); - for (int64_t i = 0; i < init_const->size(); i++) { - data.push_back(static_cast(val[i])); - } - } else { - return false; - } - - return true; -} - Status ReshapeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const { GraphViewer graph_viewer(graph); const auto& node_topology_list = graph_viewer.GetNodesInTopologicalOrder(); @@ -173,12 +147,12 @@ bool ReshapeFusion::Fuse_Subgraph1(Node& reshape, Graph& graph, const logging::L // We do not check whether the initializer is constant. // Some model uses constant initializer and some does not. // Here we assume that no one will override the initializer using graph input. - if (!LoadIntegerTensor(graph, *(concat.InputDefs()[2]), shape_value)) { + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(concat.InputDefs()[2]), shape_value)) { return false; } if (concat_input_count > 3) { - if (!LoadIntegerTensor(graph, *(concat.InputDefs()[3]), shape_value)) { + if (!optimizer_utils::AppendTensorFromInitializer(graph, *(concat.InputDefs()[3]), shape_value)) { return false; } } diff --git a/onnxruntime/core/optimizer/utils.cc b/onnxruntime/core/optimizer/utils.cc index 1f6df8ca80..3b3ad343ad 100644 --- a/onnxruntime/core/optimizer/utils.cc +++ b/onnxruntime/core/optimizer/utils.cc @@ -16,9 +16,7 @@ namespace onnxruntime { namespace optimizer_utils { bool IsFloatingPointDataType(const ONNX_NAMESPACE::TensorProto& tensor_proto) { - return tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT - || tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 - || tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_DOUBLE; + return tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT || tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 || tensor_proto.data_type() == ONNX_NAMESPACE::TensorProto_DataType_DOUBLE; } inline bool IsScalar(const NodeArg& input_arg) { @@ -44,7 +42,7 @@ bool IsInitializerWithExpectedValue(const Graph& graph, const NodeArg& input_arg } else if (!graph.GetInitializedTensor(input_arg.Name(), tensor_proto)) { return false; } - + if (tensor_proto == nullptr) { return false; } @@ -110,5 +108,73 @@ bool IsInitializerWithExpectedValue(const Graph& graph, const NodeArg& input_arg return true; } +bool IsAttributeWithExpectedValue(const Node& node, const std::string& attr_name, int64_t expected_value) { + const auto* attr_proto = graph_utils::GetNodeAttribute(node, attr_name); + if ((nullptr != attr_proto) && attr_proto->has_i()) { + return attr_proto->i() == expected_value; + } + return false; +} + +bool AppendTensorFromInitializer(const Graph& graph, const NodeArg& input_arg, std::vector& data) { + const ONNX_NAMESPACE::TensorProto* tensor_proto = nullptr; + if (!graph.GetInitializedTensor(input_arg.Name(), tensor_proto)) { + return false; + } + + auto init_const = onnxruntime::make_unique(*tensor_proto); + const auto data_type = tensor_proto->data_type(); + if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT64) { + const int64_t* val = init_const->data(); + data.reserve(data.size() + init_const->size()); + data.insert(data.end(), val, val + init_const->size()); + } else if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT32) { + const int32_t* val = init_const->data(); + data.reserve(data.size() + init_const->size()); + for (int64_t i = 0; i < init_const->size(); i++) { + data.push_back(static_cast(val[i])); + } + } else { + return false; + } + + return true; +} + +bool ValidateShape(const NodeArg& node_arg, const std::initializer_list& expected_dim_values) { + auto shape = node_arg.Shape(); + if (shape == nullptr || static_cast(shape->dim_size()) != expected_dim_values.size()) { + return false; + } + + int index = 0; + for (auto& expected_dim_value : expected_dim_values) { + if (expected_dim_value > 0) { + auto dim = shape->dim(index); + if (!utils::HasDimValue(dim) || expected_dim_value != dim.dim_value()) { + return false; + } + } + ++index; + } + + return true; +} + +bool IsShapeKnownOnAllDims(const NodeArg& node_arg, int expected_dim_size) { + auto shape = node_arg.Shape(); + if (shape == nullptr || shape->dim_size() != expected_dim_size) { + return false; + } + + for (int i = 0; i < expected_dim_size; i++) { + if (!utils::HasDimValue(shape->dim(i))) { + return false; + } + } + + return true; +} + } // namespace optimizer_utils } // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/utils.h b/onnxruntime/core/optimizer/utils.h index 4de2fb119e..0fd434993f 100644 --- a/onnxruntime/core/optimizer/utils.h +++ b/onnxruntime/core/optimizer/utils.h @@ -30,5 +30,27 @@ bool IsInitializerWithExpectedValue(const onnxruntime::Graph& graph, const onnxr */ bool IsInitializerWithExpectedValue(const onnxruntime::Graph& graph, const onnxruntime::NodeArg& input_arg, int64_t expected_value, bool is_constant); +/** Check whether an attribute of node has specified integer value. +@param expected_value is the expected value of the initializer. +*/ +bool IsAttributeWithExpectedValue(const Node& node, const std::string& attr_name, int64_t expected_value); + +/** Get values of an integer tensor from initializer, and append them to a vector. +@remarks only support int32 and int64 tensor. This function does not clear vector before appending. +*/ +bool AppendTensorFromInitializer(const Graph& graph, const NodeArg& input_arg, std::vector& data); + +/** Check Shape of node input or output. +@remarks when expected dim value > 0, the dim is expected to known and match the dim value. + when dim value <= 0, we do not check this dim. +*/ +bool ValidateShape(const NodeArg& node_arg, const std::initializer_list& expected_dim_values); + +/** Check check whether each dimension is known for shape of node_arg +@returns false when shape is nullptr, or total dimension is not same as expected_dim_size length, + or any dim is unknown (without dim value). +*/ +bool IsShapeKnownOnAllDims(const NodeArg& node_arg, int expected_dim_size); + } // namespace optimizer_utils } // namespace onnxruntime diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index e66b653f8b..9ef55cab71 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -29,6 +29,8 @@ #include "core/optimizer/slice_elimination.h" #include "core/optimizer/unsqueeze_elimination.h" #include "core/optimizer/reshape_fusion.h" +#include "core/optimizer/attention_fusion.h" +#include "core/optimizer/utils.h" #include "core/platform/env.h" #include "core/util/math.h" #include "test/capturing_sink.h" @@ -875,11 +877,11 @@ TEST(GraphTransformationTests, ReshapeFusionOneConstTest) { ASSERT_TRUE(ret.IsOK()); std::map op_to_count = CountOpsInGraph(graph); - ASSERT_TRUE(op_to_count["Shape"] == 0); - ASSERT_TRUE(op_to_count["Gather"] == 0); - ASSERT_TRUE(op_to_count["Unsqueeze"] == 0); - ASSERT_TRUE(op_to_count["Concat"] == 0); - ASSERT_TRUE(op_to_count["Reshape"] == 1); + ASSERT_EQ(op_to_count["Shape"], 0); + ASSERT_EQ(op_to_count["Gather"], 0); + ASSERT_EQ(op_to_count["Unsqueeze"], 0); + ASSERT_EQ(op_to_count["Concat"], 0); + ASSERT_EQ(op_to_count["Reshape"], 1); for (const Node& node : graph.Nodes()) { if (node.OpType() == "Reshape") { @@ -898,6 +900,167 @@ TEST(GraphTransformationTests, ReshapeFusionOneConstTest) { } } +static void ValidateAttention(Graph& graph) { + // Validate the merged weights (initializer) input for Attention node. + for (const Node& node : graph.Nodes()) { + if (node.OpType() == "Attention") { + int64_t expected_heads = 2; + ASSERT_TRUE(optimizer_utils::IsAttributeWithExpectedValue(node, "num_heads", expected_heads)); + + const ONNX_NAMESPACE::TensorProto* tensor_proto = graph_utils::GetConstantInitializer(graph, node.InputDefs()[1]->Name()); + ASSERT_TRUE(tensor_proto != nullptr); + EXPECT_EQ(tensor_proto->data_type(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + + auto initializer = onnxruntime::make_unique(*tensor_proto); + EXPECT_EQ(initializer->size(), 192); + + // Validate two rows (2x24 items) for sanity check. + std::vector expected_value = { + -0.10791015625, + -0.04193115234375, + 0.09051513671875, + 0.025787353515625, + -0.11572265625, + -0.126953125, + -0.043304443359375, + -0.02984619140625, + 0.022125244140625, + -0.017730712890625, + -0.03265380859375, + -0.05108642578125, + 0.0423583984375, + 0.112060546875, + 0.080810546875, + 0.09375, + -0.03643798828125, + 0.02862548828125, + 0.039764404296875, + 0.06097412109375, + -0.002288818359375, + -0.10797119140625, + -0.01171875, + 0.041717529296875, + + 0.033538818359375, + -0.05755615234375, + -0.04986572265625, + -0.01558685302734375, + -0.0352783203125, + 0.03546142578125, + 0.05218505859375, + 0.005565643310546875, + -0.043182373046875, + -0.05010986328125, + -0.063720703125, + -0.00824737548828125, + 0.1492919921875, + 0.048431396484375, + -0.0482177734375, + -0.1123046875, + 0.032196044921875, + 0.0135650634765625, + 0.020233154296875, + -0.05084228515625, + -0.011260986328125, + -0.1241455078125, + -0.0101165771484375, + -0.00490570068359375}; + + const float* data = initializer->data(); + for (size_t i = 0; i < expected_value.size(); i++) { + EXPECT_EQ(data[i], static_cast(expected_value[i])); + } + + tensor_proto = graph_utils::GetConstantInitializer(graph, node.InputDefs()[2]->Name()); + ASSERT_TRUE(tensor_proto != nullptr); + EXPECT_EQ(tensor_proto->data_type(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + + auto initializer2 = onnxruntime::make_unique(*tensor_proto); + EXPECT_EQ(initializer2->size(), 24); + + std::vector expected_value2 = { + -0.23681640625, + -0.16552734375, + 0.2191162109375, + -0.1756591796875, + -0.03460693359375, + -0.05316162109375, + -0.336181640625, + -0.253662109375, + 0.0246734619140625, + 0.011993408203125, + 0.0178375244140625, + 0.00998687744140625, + 0.0255126953125, + 0.076416015625, + -0.040771484375, + 0.0107879638671875, + -0.005893707275390625, + -0.00916290283203125, + 0.04541015625, + 0.0159454345703125, + -0.0029163360595703125, + -0.03472900390625, + 0.0535888671875, + 0.0091094970703125}; + + const float* data2 = initializer2->data(); + for (size_t i = 0; i < expected_value2.size(); i++) { + EXPECT_EQ(data2[i], static_cast(expected_value2[i])); + } + + } + } +} + +// Test Attention Fusion with int32 mask +TEST(GraphTransformationTests, AttentionFusionInt32Test) { + auto model_uri = MODEL_FOLDER "fusion/attention_int32_mask.onnx"; + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model, nullptr, DefaultLoggingManager().DefaultLogger()).IsOK()); + Graph& graph = p_model->MainGraph(); + + onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; + graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level2); + auto ret = graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level2, DefaultLoggingManager().DefaultLogger()); + ASSERT_TRUE(ret.IsOK()); + + std::map op_to_count = CountOpsInGraph(graph); + EXPECT_EQ(op_to_count["MatMul"], 1); + EXPECT_EQ(op_to_count["Add"], 2); + EXPECT_EQ(op_to_count["Transpose"], 0); + EXPECT_EQ(op_to_count["Reshape"], 0); + EXPECT_EQ(op_to_count["Cast"], 0); + EXPECT_EQ(op_to_count["ReduceSum"], 1); + EXPECT_EQ(op_to_count["Attention"], 1); + + ValidateAttention(graph); +} + +// Test Attention Fusion with int64 mask and symbolic batch dimension +TEST(GraphTransformationTests, AttentionFusionInt64Test) { + auto model_uri = MODEL_FOLDER "fusion/attention_symbolic_batch.onnx"; + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model, nullptr, DefaultLoggingManager().DefaultLogger()).IsOK()); + Graph& graph = p_model->MainGraph(); + + onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; + graph_transformation_mgr.Register(onnxruntime::make_unique(), TransformerLevel::Level2); + auto ret = graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level2, DefaultLoggingManager().DefaultLogger()); + ASSERT_TRUE(ret.IsOK()); + + std::map op_to_count = CountOpsInGraph(graph); + EXPECT_EQ(op_to_count["MatMul"], 1); + EXPECT_EQ(op_to_count["Add"], 2); + EXPECT_EQ(op_to_count["Transpose"], 0); + EXPECT_EQ(op_to_count["Reshape"], 0); + EXPECT_EQ(op_to_count["Cast"], 1); // Cast for int64 mask to int32 + EXPECT_EQ(op_to_count["ReduceSum"], 1); + EXPECT_EQ(op_to_count["Attention"], 1); + + ValidateAttention(graph); +} + #ifndef DISABLE_CONTRIB_OPS TEST(GraphTransformationTests, GeluFusionTest) { auto model_uri = MODEL_FOLDER "fusion/gelu.onnx"; diff --git a/onnxruntime/test/testdata/transform/fusion/attention_int32_mask.onnx b/onnxruntime/test/testdata/transform/fusion/attention_int32_mask.onnx new file mode 100644 index 0000000000..6241c329fb Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/attention_int32_mask.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/attention_symbolic_batch.onnx b/onnxruntime/test/testdata/transform/fusion/attention_symbolic_batch.onnx new file mode 100644 index 0000000000..2967eb93a8 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/attention_symbolic_batch.onnx differ