Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 103 additions & 6 deletions onnxruntime/core/optimizer/group_query_attention_fusion.cc
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include <limits>

#include "core/optimizer/initializer.h"
#include "core/optimizer/group_query_attention_fusion.h"
#include "core/graph/graph_utils.h"
Expand Down Expand Up @@ -145,8 +147,9 @@ static std::vector<NodeArg*> MergeQkvWeightsForMatMulNBits(

TensorProto qkv_zp_initializer;

// We use 4 bit quantization, hence dividing by 2 since we need 1/2 of the bytes.
int64_t zp_elements_count = output_hidden_size * blocks / 2;
// We use 4 bit quantization, so each byte stores two block zero points per output row.
const int64_t zero_point_blocks = blocks / 2 + blocks % 2;
int64_t zp_elements_count = output_hidden_size * zero_point_blocks;

qkv_zp_initializer.set_name(graph.GenerateNodeArgName("qkv_zp"));
qkv_zp_initializer.add_dims(zp_elements_count);
Expand All @@ -156,7 +159,7 @@ static std::vector<NodeArg*> MergeQkvWeightsForMatMulNBits(
merged_qkv_zp.reserve(gsl::narrow<size_t>(zp_elements_count));

optimizer_utils::MergeMatMulWeightsByBlocks(q_zero_points_data, k_zero_points_data, v_zero_points_data,
merged_qkv_zp, q_hidden_size, kv_hidden_size, blocks / 2, 1);
merged_qkv_zp, q_hidden_size, kv_hidden_size, zero_point_blocks, 1);

utils::SetRawDataInTensorProto(qkv_zp_initializer, merged_qkv_zp.data(), zp_elements_count * sizeof(uint8_t));

Expand Down Expand Up @@ -252,6 +255,72 @@ static bool NodeArgExists(const NodeArg* node_arg) {
return node_arg != nullptr && node_arg->Exists();
}

static bool ProjectionTensorShapesMatch(bool quantization_used,
int64_t q_hidden_size,
int64_t kv_hidden_size,
const TensorProto& q_tensor,
const TensorProto& k_tensor,
const TensorProto& v_tensor) {
const int expected_rank = quantization_used ? 3 : 2;
if (q_tensor.dims_size() != expected_rank || k_tensor.dims_size() != expected_rank ||
v_tensor.dims_size() != expected_rank) {
return false;
}

if (quantization_used) {
return q_tensor.dims(0) == q_hidden_size && k_tensor.dims(0) == kv_hidden_size &&
v_tensor.dims(0) == kv_hidden_size && q_tensor.dims(1) == k_tensor.dims(1) &&
q_tensor.dims(1) == v_tensor.dims(1) && q_tensor.dims(2) == k_tensor.dims(2) &&
q_tensor.dims(2) == v_tensor.dims(2) && q_tensor.dims(1) > 0 && q_tensor.dims(2) > 0;
}

return q_tensor.dims(1) == q_hidden_size && k_tensor.dims(1) == kv_hidden_size &&
v_tensor.dims(1) == kv_hidden_size && q_tensor.dims(0) == k_tensor.dims(0) &&
q_tensor.dims(0) == v_tensor.dims(0);
}

static bool TensorShapeMatchesRowsAndBlocks(const TensorProto& tensor, int64_t rows, int64_t blocks) {
if (tensor.dims_size() == 2) {
return tensor.dims(0) == rows && tensor.dims(1) == blocks;
}

if (tensor.dims_size() != 1 || blocks <= 0 || tensor.dims(0) < 0 || tensor.dims(0) % blocks != 0) {
return false;
}

return tensor.dims(0) / blocks == rows;
}

static bool QuantizedAuxiliaryTensorShapesMatch(int64_t q_hidden_size,
int64_t kv_hidden_size,
int64_t blocks,
const TensorProto& q_scale_tensor,
const TensorProto& k_scale_tensor,
const TensorProto& v_scale_tensor,
const TensorProto* q_zero_point_tensor,
const TensorProto* k_zero_point_tensor,
const TensorProto* v_zero_point_tensor) {
const int scale_rank = q_scale_tensor.dims_size();
if (k_scale_tensor.dims_size() != scale_rank || v_scale_tensor.dims_size() != scale_rank ||
!TensorShapeMatchesRowsAndBlocks(q_scale_tensor, q_hidden_size, blocks) ||
!TensorShapeMatchesRowsAndBlocks(k_scale_tensor, kv_hidden_size, blocks) ||
!TensorShapeMatchesRowsAndBlocks(v_scale_tensor, kv_hidden_size, blocks)) {
return false;
}

if (q_zero_point_tensor == nullptr || k_zero_point_tensor == nullptr || v_zero_point_tensor == nullptr) {
return q_zero_point_tensor == nullptr && k_zero_point_tensor == nullptr && v_zero_point_tensor == nullptr;
}

const int64_t zero_point_blocks = blocks / 2 + blocks % 2;
const int zero_point_rank = q_zero_point_tensor->dims_size();
return k_zero_point_tensor->dims_size() == zero_point_rank &&
v_zero_point_tensor->dims_size() == zero_point_rank &&
TensorShapeMatchesRowsAndBlocks(*q_zero_point_tensor, q_hidden_size, zero_point_blocks) &&
TensorShapeMatchesRowsAndBlocks(*k_zero_point_tensor, kv_hidden_size, zero_point_blocks) &&
TensorShapeMatchesRowsAndBlocks(*v_zero_point_tensor, kv_hidden_size, zero_point_blocks);
}

struct RotaryEmbeddingArgs {
NodeArg* cos_cache_arg = nullptr;
NodeArg* sin_cache_arg = nullptr;
Expand Down Expand Up @@ -481,9 +550,37 @@ Status GroupQueryAttentionFusion::ApplyImpl(
int64_t head_size = past_key_values_key_arg->Shape()->dim(3).dim_value();
int64_t num_heads = node.GetAttributes().at("num_heads").i();
int64_t kv_num_heads = node.GetAttributes().at("kv_num_heads").i();
int64_t q_hidden_size = num_heads * head_size;
int64_t kv_hidden_size = kv_num_heads * head_size;
int64_t output_hidden_size = q_hidden_size + 2 * kv_hidden_size;
if (head_size <= 0 || num_heads <= 0 || kv_num_heads <= 0) {
DEBUG_LOG("GQA head attributes and cache head size must be positive");
continue;
}

constexpr int64_t max_int64 = std::numeric_limits<int64_t>::max();
if (num_heads > max_int64 / head_size || kv_num_heads > max_int64 / head_size) {
DEBUG_LOG("GQA hidden size calculation overflowed");
continue;
}
const int64_t q_hidden_size = num_heads * head_size;
const int64_t kv_hidden_size = kv_num_heads * head_size;
if (kv_hidden_size > (max_int64 - q_hidden_size) / 2) {
DEBUG_LOG("GQA output hidden size calculation overflowed");
continue;
}
const int64_t output_hidden_size = q_hidden_size + 2 * kv_hidden_size;

if (!ProjectionTensorShapesMatch(quantization_used, q_hidden_size, kv_hidden_size,
*q_proj_tensor, *k_proj_tensor, *v_proj_tensor)) {
DEBUG_LOG("GQA projection tensor shapes do not match the head attributes");
continue;
}

if (quantization_used &&
!QuantizedAuxiliaryTensorShapesMatch(q_hidden_size, kv_hidden_size, q_proj_tensor->dims(1),
*q_scale_tensor, *k_scale_tensor, *v_scale_tensor,
q_zero_points_tensor, k_zero_points_tensor, v_zero_points_tensor)) {
DEBUG_LOG("GQA quantized auxiliary tensor shapes do not match the projection tensors");
continue;
}

// Ensure the output shape has 3 dimensions [batch_size, sequence_length, hidden_size]
if (matmul_or_nbits_output_shape->dim_size() == 3) {
Expand Down
70 changes: 68 additions & 2 deletions onnxruntime/test/optimizer/graph_transform_test_layernorm.cc
Original file line number Diff line number Diff line change
Expand Up @@ -952,6 +952,32 @@ static void TestGQAFusion(const std::basic_string<ORTCHAR_T>& file_path, int mat
ASSERT_TRUE(op_to_count["com.microsoft.GroupQueryAttention"] == 1);
}

static void TestQuantizedGQAFusionRejectsInitializerShape(const std::string& initializer_name,
logging::Logger* logger) {
constexpr const ORTCHAR_T* model_uri = MODEL_FOLDER "fusion/gqa_fusion_quantized_simple.onnx";
std::shared_ptr<Model> model;
ASSERT_STATUS_OK(Model::Load(model_uri, model, nullptr, *logger));
Graph& graph = model->MainGraph();

const TensorProto* initializer = nullptr;
ASSERT_TRUE(graph.GetInitializedTensor(initializer_name, initializer));
TensorProto malformed_initializer = *initializer;
malformed_initializer.clear_dims();
malformed_initializer.add_dims(1);
graph.RemoveInitializedTensor(initializer_name);
graph.AddInitializedTensor(malformed_initializer);

GraphTransformerManager graph_transformation_mgr{3};
ASSERT_STATUS_OK(graph_transformation_mgr.Register(std::make_unique<GroupQueryAttentionFusion>(),
TransformerLevel::Level2));
ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level2, *logger));

const auto op_to_count = CountOpsInGraph(graph);
EXPECT_EQ(op_to_count.at("com.microsoft.MatMulNBits"), 3);
EXPECT_EQ(op_to_count.at("com.microsoft.RotaryEmbedding"), 2);
EXPECT_EQ(op_to_count.at("com.microsoft.GroupQueryAttention"), 1);
}

enum class RotaryEmbeddingDomain {
kOnnx,
kMS,
Expand All @@ -962,7 +988,8 @@ static void BuildRotaryEmbeddingGQAFusionGraph(ModelTestBuilder& builder,
bool include_position_ids,
int64_t q_interleaved = 0,
int64_t k_interleaved = 0,
int64_t rotary_embedding_dim = 0) {
int64_t rotary_embedding_dim = 0,
int64_t gqa_num_heads = 2) {
constexpr int64_t batch_size = 1;
constexpr int64_t sequence_length = 2;
constexpr int64_t input_hidden_size = 8;
Expand Down Expand Up @@ -1047,7 +1074,7 @@ static void BuildRotaryEmbeddingGQAFusionGraph(ModelTestBuilder& builder,
seqlens_k, total_sequence_length},
{gqa_output},
kMSDomain);
gqa.AddAttribute("num_heads", num_heads);
gqa.AddAttribute("num_heads", gqa_num_heads);
gqa.AddAttribute("kv_num_heads", kv_num_heads);
}

Expand All @@ -1071,6 +1098,26 @@ static Status CheckOnnxRotaryEmbeddingGQANotFused(Graph& graph) {
return Status::OK();
}

static Status CheckMsRotaryEmbeddingGQANotFused(Graph& graph) {
const auto op_to_count = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(OpCount(op_to_count, "com.microsoft.RotaryEmbedding") == 2);
TEST_RETURN_IF_NOT(OpCount(op_to_count, "MatMul") == 3);
TEST_RETURN_IF_NOT(OpCount(op_to_count, "com.microsoft.GroupQueryAttention") == 1);

for (const Node& node : graph.Nodes()) {
if (node.OpType() != "GroupQueryAttention") {
continue;
}

TEST_RETURN_IF_NOT(node.InputDefs().size() == 7);
const auto& attrs = node.GetAttributes();
const auto do_rotary_attr = attrs.find("do_rotary");
TEST_RETURN_IF_NOT(do_rotary_attr == attrs.end() || do_rotary_attr->second.i() == 0);
}

return Status::OK();
}

static Status CheckMsRotaryEmbeddingGQAFused(Graph& graph, int64_t expected_interleaved) {
const auto op_to_count = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(OpCount(op_to_count, "com.microsoft.RotaryEmbedding") == 0);
Expand Down Expand Up @@ -1355,6 +1402,25 @@ TEST_F(GraphTransformationTests, GroupQueryAttentionFusionMsRotaryEmbeddingForwa
check_fused_graph));
}

TEST_F(GraphTransformationTests, GroupQueryAttentionFusionSkipsMismatchedProjectionSizesTest) {
auto build_test_case = [](ModelTestBuilder& builder) {
BuildRotaryEmbeddingGQAFusionGraph(builder, RotaryEmbeddingDomain::kMS, true, 0, 0, 0, 3);
};

ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, 23, *logger_,
std::make_unique<GroupQueryAttentionFusion>(),
TransformerLevel::Level2, 3, nullptr,
CheckMsRotaryEmbeddingGQANotFused));
}

TEST_F(GraphTransformationTests, GroupQueryAttentionFusionSkipsMismatchedScaleShape) {
TestQuantizedGQAFusionRejectsInitializerShape("scales_q", logger_.get());
}

TEST_F(GraphTransformationTests, GroupQueryAttentionFusionSkipsMismatchedZeroPointShape) {
TestQuantizedGQAFusionRejectsInitializerShape("zero_points_q", logger_.get());
}

TEST_F(GraphTransformationTests, SkipLayerNormFusionWithCastTest) {
TestSkipLayerNormFusion(MODEL_FOLDER "fusion/skip_layer_norm_format1_with_cast.onnx", 0, 0, 1, 3, logger_.get());
TestSkipLayerNormFusion(MODEL_FOLDER "fusion/skip_layer_norm_format2_with_cast.onnx", 0, 0, 1, 3, logger_.get());
Expand Down
Loading