From 5dce9be4f9d78e76afd52376332550a34d922f7d Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Mon, 25 Nov 2019 14:46:37 -0800 Subject: [PATCH] Add Attention Fusion Transformer (#2445) Add Attention Fusion Transformer to fuse multi-head self attention subgraph to one node for optimizing Bert model inference performance. It supports BERT model exported from PyTorch. It fuses about 20 nodes into one Attention node, and could significantly improve the inference speed of BERT model. Support symbolic (first dimension for batch size) in input shape. --- onnxruntime/core/graph/graph_utils.cc | 28 +- .../core/optimizer/attention_fusion.cc | 711 ++++++++++++++++++ onnxruntime/core/optimizer/attention_fusion.h | 25 + .../core/optimizer/graph_transformer_utils.cc | 2 + onnxruntime/core/optimizer/reshape_fusion.cc | 30 +- onnxruntime/core/optimizer/utils.cc | 74 +- onnxruntime/core/optimizer/utils.h | 22 + .../test/optimizer/graph_transform_test.cc | 173 ++++- .../fusion/attention_int32_mask.onnx | Bin 0 -> 4029 bytes .../fusion/attention_symbolic_batch.onnx | Bin 0 -> 4044 bytes 10 files changed, 1025 insertions(+), 40 deletions(-) create mode 100644 onnxruntime/core/optimizer/attention_fusion.cc create mode 100644 onnxruntime/core/optimizer/attention_fusion.h create mode 100644 onnxruntime/test/testdata/transform/fusion/attention_int32_mask.onnx create mode 100644 onnxruntime/test/testdata/transform/fusion/attention_symbolic_batch.onnx 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 0000000000000000000000000000000000000000..6241c329fb723eb6572cb15157beb85f52d66a59 GIT binary patch literal 4029 zcmbtXeQaCR6+ixZsWNF$4bo?<1i9a??`(1y$ z*iI^<&eHX}_nvcp=bYa?=lWUL797u%%TB(WE99-hVmY^vJCWgw{V(oh8XL;x7njPG z9Q5ijapbOi>5(PJIpJ86rN_bncOAKs`7_6zlIf{>hFFY28x~Tpg>324W)@y z^&IpQRZ>-tpIrxN1%s!cio)TgIgYM|eVDZnwH`#>(zQz?nxA~Di>Lxo1EMYphnD6! zq8|2P#`ZC9TxU#GV^+~AqE>&4^YZS4iS9jbAr%{j$#)X6_k$U4R7 zd}bNnv8a0{(ULQNGz!*>%h9=9rnHV~h{7$|EL_Kpun)6!*OIGbNUm$yG9=gaEb|%? zQ?}yzVQ0XxA-5d^P#$nBuRB&<$BJdBu4CkEs9Vdqb(Z}(8(P~nAhot@85|mH#{=Dx zC_R)}bfln9iW%beJDt|jx^1cCNa3N*k{OJPp(Dk_9Y>0pd}*;z;w6oZaKPP&j4}VB zQ(Q2~^qL;i$LkdC40+I<3@BOccPFpA<2Kr$jRG4b-fyF<+wc&TIM*?HLzFFSd`yVk zR90cC_GPD7*tnFbuidwMU@FD%rN+3uBo24VR)z6;@kF?W#711wv$d5?cwjt6ZJ?&q^#i*JMw)@X z*WAAMNw@&J4C)qEZghvCtxIh3WBRYKRL12S%{uv#)7J>F#!bnoMAo<|TU6f2L3uyu zL7{DF;R5DLy478`HBWoG_qxRPwqsbBbF$f7{=rgomft|P7m5qfuR6I0A1ZfRn(eu? zV*kCQi`TdqnfcDI3BRWFaB!o5-5(t6Uh%r_m5$^rx60bB-E%wd?%wwkf`4>qKW==l zKRw*ixA2QrDirx`y#ZhD#h)C&m;Zl!r58UwfZzLjwYB)~8{|jr#Wx1<)t>lwb@{nQ z@PQv6k+6}T>EJ*g_W*>u0OO;cFz6+7gMDzM^UmO8EnXmzs|3&JP}--e%0XcmCj7I$ zsU8%})>iMPwwb2BndanXn(@svje`OqVg8AHUcd`MI7H0@DQYgKsKUoBDWV(kd@Mzc zS8-XxdCaF1t=l2s-9Rub{jm+zlRDE}fOka)ciz^iR6YT#1o4cooUWMKX z_?U*xcBExM_6J~1C&?Vg^V?X%XXPct_6qRoz}tss05O@c_rl-ftJE+M^TH~bA?$NM zdDv&Mmr9U*8gs1aNyy%U4DmNT^>oh#8nHSwH{P#P^)$wBF`mJA13ZY+)I5dfOSSan zUZZ`)(_fGP$gdsqt`~%vtC>isVNzAH&{VNwR-DMU4{l zpGr^&2S=$Sg*+Iah3v1;d%mr2PvB60$Ys;icnAK@!|olBF0E3P)`)%$ay4My4WGv` zc7iAsm!NkQFFfpIRyDXA)5gX4z%$cWM9V~ zaxrm0O+S2Ji+nD?_No=4zau7x6YsvU_VHrAnIQ99z_|i>d;u|i2Ks-)I0wuI;)-Ga z5cbgnp&=B_#5CErL4Io;*vRqXD%mv@2ruSkU^T$)M<^f^7b&no0h<@5sCgDe`+LL< z{wq~r7BJ3&>)nuJ-2iv?cM<>26{^06Lgx)|6#jky*-5N}kfF%!0ysMf-QOiqq^r;Y z-WkM-VyW(c4qKyU7PjAnjvsS7N!1Z>SWZ%C40(JJIqiV{3y9~@7BAh`>3&|WhAz&a z8HS%_=#Rq}YKTn46aj`0ao>+T&mreKpc@AV|3$2Oz-4un$if-qaRz%=$o?AY16)*x zfinY}EObzB_8-B=F2wa`taoDlGUOHThni^q0(u|7mjitGsXkq!D)O|^jFl*@vF7(VPnEaXWOs1g>8L&yTI8FYm8T@v literal 0 HcmV?d00001 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 0000000000000000000000000000000000000000..2967eb93a89384822a4bb7806a47bb604e8b9cd5 GIT binary patch literal 4044 zcmbtXdu&^06+eD-zW$b^iA&Sqjn%tt)zu8o_u(7g#*XjYY@jOLfOZ;eZ1+0#?P`|T zneCLO6Pnje3>EPdWgP<50#)5qgy0VlP3&M08V?DfsS*MNlO@D!Wq%+c@dpU;`~0}E zokWX@rO)?!obx;9{JwLJuZ7zJa=VPVvD z)Js%JRb5_o2cQ)U9)cPajxWt|bXD|V)!^BAxIdeP>9{U>Fk3e*IYPSRn3hdla!k)Mw=OYd zJEotu2Mp_S+b{s-0mJgTVbw9Ln7Zm1M$Edpy_j2f*&nkWYnujStZiBbhdSGNvYQg6 z)0stE3V5Wb9vuF#-CCLtE|qLa^mUd@XZ#p4QZ%^pL@|>uEfz{VrLiIUosIAq^Df%O z1%nK?;W9itPhn1%2hGWVlGT25^13;$6C2b~prge5b(D1-?xGUsGJ3CzvWbm%aFMIZ zDpb|J=oAwhTd8{5b$bA+QWRfmlB?hI#A zRYVV_>@F_L%TMEoYS7P9mcyX{dCuO)&s*N2438I1l@~H+Tlz84*IBlEORhA58*nr3 zqZ-4P!f#O393Nu<)to%zET}D;lf49T%adEN(79uUt4~WcgG>ENo?{R>aVa=M)Qqi?R?4ZD+FkxQ*tU1H9BRB%6m8{uLmt4 zv^7mMU@oMY-E~{?f~(utC3djw!@``M&F1n?l_Ime13gqIE<`?I=bk!U?vyk;u(e=+ zU($km9E{9-=huX{DNPJ)JYY`-2CFxCN4KRTIMd0pc4_x?=l$J&KOuNWeTUKU-QM(2 zOWwpUTB%Uv-FgGQ+>1XsfG_|5_)0H+d;q`q_iB6oKQYLU+KaCb;Hy3U@9XljNAQ3j z50S8up6S3q9rpl;U4ZdXmk4^v*q|R6>AW-8S(7J7c$eTB^`$+ksvH%Dpu#`vo9a=) zXl?aw#x|4GHj|v(OftTiq<&N&1k69-hXuS4h#_hoNl|k-MHN2oPZ8aN=kqCQ%woQf zN^iFm!pAy**6})puGT1YzD|`37_3glHH@(}YMiL2w^uUZi*5lQxL9$M@qs!uKAxb; zr;}6}!u+KbYX0Xs(JM)!be*hxf@q{h%@Z(vrAE~QNost(M)Yce8XEL&PLcIbxb$zF z{h&sT_tvTTg(Nj+@cc%Cs$;-hsgd!XBvt=9MTVzF*3A_P{SNwnfzG2-MAsm94mPHt zGmNlwoc#t^(@8SM@%$3juvxhb-(Ce?9e9WE^us3u`fk{JVU-#>d|p^3BZPf!ClCEB z_EHJ5E?|x|eF|r<;|%^cUG;QN0}Y!Unj3G`sd@q9ml&VN_zrjoPE+$GJTKSMx7tSg zh^Idw0FY%}hhsND^4oO;vPPB!0zjzcPg3LQIvKN&c@DvkVm^kwN0MZHe~KC<$bU9L zAzU1#k`!`bydP(Ogxt$*d20ff`Wio*rpBAFcO80ff^=z>se!^OX`O3iPBya00Uryz%L{&$?s00$S^_%hBui9N(( z;DVZ7*nT(SxdGi>D@1>WPc|psd1LM4#rR%=j4uLb2jch;e0m?`|AcV`m<{+9#r`qu zqXt4gB+bM$Sr6j;fjY1ey;WzlN zRDoH*xCE~E;~eV-xU;?r|M#v?^>0XYUI3?H@0&O~k97cNNOG$H&dx*jS4kx4DrA6n z5xyc>syiXW)~K0p%GNz z+cU5!6xXP72sLA&I2GjMFl0XpuD=DIpI=Mg+FzaU;~n_WFL2GE zCxre}-G@-#!1`;T`rXY_YHqp z9uu86t-v0(TW|%&m@lwnE>k{vnz;oJhVa;3T^TvBi;W4sz;K_;+&<>K&us~kc8~i1 E2crS4j{pDw literal 0 HcmV?d00001