Skip to content
Merged
26 changes: 13 additions & 13 deletions onnxruntime/core/optimizer/attention_fusion.cc
Original file line number Diff line number Diff line change
Expand Up @@ -349,22 +349,22 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,

const Node* sequence_transpose = graph_utils::GetInputNode(qkv_matmul, 0);
if (sequence_transpose == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*sequence_transpose, "Transpose", {1, 13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*sequence_transpose, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain) ||
!HasExpectedPerm(*sequence_transpose, {0, 2, 1}) ||
!optimizer_utils::CheckOutputEdges(graph, *sequence_transpose, 1)) {
return false;
}

const Node* input_reshape = graph_utils::GetInputNode(*sequence_transpose, 0);
if (input_reshape == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*input_reshape, "Reshape", {5, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*input_reshape, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25}, kOnnxDomain) ||
!optimizer_utils::CheckOutputEdges(graph, *input_reshape, 1)) {
return fail("missing input Reshape before sequence transpose");
}

Node* qkv_reshape = GetOnlyChildByOutputIndex(graph, qkv_matmul, 0, "Reshape");
if (qkv_reshape == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*qkv_reshape, "Reshape", {5, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*qkv_reshape, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25}, kOnnxDomain) ||
!optimizer_utils::CheckOutputEdges(graph, *qkv_reshape, 1)) {
return fail("qkv Reshape after MatMul not matched");
}
Expand All @@ -379,9 +379,9 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,
Node* k_squeeze = GetOnlyChildByOutputIndex(graph, *split, 1, "Squeeze");
Node* v_transpose = GetOnlyChildByOutputIndex(graph, *split, 2, "Transpose");
if (q_transpose == nullptr || k_squeeze == nullptr || v_transpose == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*q_transpose, "Transpose", {1, 13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*k_squeeze, "Squeeze", {13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*v_transpose, "Transpose", {1, 13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*q_transpose, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*k_squeeze, "Squeeze", {13, 21, 23, 24, 25}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*v_transpose, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain) ||
!HasExpectedPerm(*q_transpose, {2, 0, 3, 1, 4}) ||
!HasExpectedPerm(*v_transpose, {2, 0, 3, 1, 4}) ||
!HasExpectedAxesInput(graph, *k_squeeze, {2})) {
Expand All @@ -391,8 +391,8 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,
Node* q_squeeze = GetOnlyChildByOutputIndex(graph, *q_transpose, 0, "Squeeze");
Node* v_squeeze = GetOnlyChildByOutputIndex(graph, *v_transpose, 0, "Squeeze");
if (q_squeeze == nullptr || v_squeeze == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*q_squeeze, "Squeeze", {13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*v_squeeze, "Squeeze", {13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*q_squeeze, "Squeeze", {13, 21, 23, 24, 25}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*v_squeeze, "Squeeze", {13, 21, 23, 24, 25}, kOnnxDomain) ||
!HasExpectedAxesInput(graph, *q_squeeze, {0}) ||
!HasExpectedAxesInput(graph, *v_squeeze, {0})) {
return fail("q/v squeeze pattern not matched");
Expand All @@ -402,7 +402,7 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,
Node* k_transpose = GetOnlyChildByOutputIndex(graph, *k_squeeze, 0, "Transpose");
if (q_scale_mul == nullptr || k_transpose == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*q_scale_mul, "Mul", {7, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*k_transpose, "Transpose", {1, 13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*k_transpose, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain) ||
!HasExpectedPerm(*k_transpose, {0, 2, 3, 1})) {
return fail("q scale Mul or k Transpose(0,2,3,1) not matched");
}
Expand Down Expand Up @@ -460,15 +460,15 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,

Node* transpose_3 = GetOnlyChildByOutputIndex(graph, *qkv_matmul_1, 0, "Transpose");
if (transpose_3 == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*transpose_3, "Transpose", {1, 13}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*transpose_3, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain) ||
!HasExpectedPerm(*transpose_3, {0, 2, 1, 3}) ||
!optimizer_utils::CheckOutputEdges(graph, *transpose_3, 1)) {
return fail("output Transpose(0,2,1,3) not matched");
}

Node* reshape_2 = GetOnlyChildByOutputIndex(graph, *transpose_3, 0, "Reshape");
if (reshape_2 == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_2, "Reshape", {5, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_2, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25}, kOnnxDomain) ||
!optimizer_utils::CheckOutputEdges(graph, *reshape_2, 1)) {
return fail("output Reshape not matched");
}
Expand Down Expand Up @@ -497,7 +497,7 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,
if (proj_gemm == nullptr) {
proj_gemm_input_reshape = GetOnlyChildByOutputIndex(graph, *reshape_2, 0, "Reshape");
if (proj_gemm_input_reshape == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*proj_gemm_input_reshape, "Reshape", {5, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*proj_gemm_input_reshape, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25}, kOnnxDomain) ||
!optimizer_utils::CheckOutputEdges(graph, *proj_gemm_input_reshape, 1)) {
Comment thread
yuslepukhin marked this conversation as resolved.
return fail("projection MatMul/Gemm not matched");
}
Expand All @@ -511,7 +511,7 @@ static bool TryFuseMobileClipMHA(Node& qkv_matmul,

proj_gemm_output_reshape = GetOnlyChildByOutputIndex(graph, *proj_gemm, 0, "Reshape");
if (proj_gemm_output_reshape == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*proj_gemm_output_reshape, "Reshape", {5, 13, 14}, kOnnxDomain) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*proj_gemm_output_reshape, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25}, kOnnxDomain) ||
!optimizer_utils::CheckOutputEdges(graph, *proj_gemm_output_reshape, 1)) {
return fail("normalized projection Gemm output Reshape not matched");
}
Expand Down
2 changes: 1 addition & 1 deletion onnxruntime/core/optimizer/attention_fusion_helper.h
Original file line number Diff line number Diff line change
Expand Up @@ -1447,7 +1447,7 @@ bool FuseGptAttention(Node& layer_norm, Graph& graph, int64_t hidden_size, std::
return false;
}

if (graph_utils::IsSupportedOptypeVersionAndDomain(*k_concat, "Transpose", {1, 13, 21}, kOnnxDomain)) {
if (graph_utils::IsSupportedOptypeVersionAndDomain(*k_concat, "Transpose", {1, 13, 21, 23, 24, 25}, kOnnxDomain)) {
Comment thread
yuslepukhin marked this conversation as resolved.
transpose_optimized_pattern = true;
DEBUG_LOG("Using transpose optimized pattern");
opt_k_transpose = k_concat;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,7 @@ bool MatchPreNormReshapeChain(Graph& graph,

Node* reshape_outer = graph.GetMutableProducerNode(consumer_input->Name());
if (reshape_outer == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_outer, "Reshape", {5, 13, 14, 19, 21, 23})) {
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_outer, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25})) {
return false;
}
if (reshape_outer->GetOutputEdgesCount() != 1) {
Expand Down Expand Up @@ -174,7 +174,7 @@ bool MatchPreNormReshapeChain(Graph& graph,
}
Node* reshape_inner = graph.GetMutableProducerNode(sln->InputDefs()[0]->Name());
if (reshape_inner == nullptr ||
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_inner, "Reshape", {5, 13, 14, 19, 21, 23})) {
!graph_utils::IsSupportedOptypeVersionAndDomain(*reshape_inner, "Reshape", {5, 13, 14, 19, 21, 23, 24, 25})) {
return false;
}
if (reshape_inner->GetOutputEdgesCount() != 1) {
Expand Down
8 changes: 6 additions & 2 deletions onnxruntime/core/optimizer/reshape_fusion.cc
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,10 @@ namespace onnxruntime {
bool GetAxesFromUnsqueezeNode(const Graph& graph, const Node& unsqueeze, InlinedVector<int64_t>& axes) {
if (graph_utils::MatchesOpSinceVersion(unsqueeze, {1, 11})) {
return graph_utils::GetRepeatedNodeAttributeValues(unsqueeze, "axes", axes);
} else if (graph_utils::MatchesOpSinceVersion(unsqueeze, {13})) {
}

// Opset 13+ moved axes from attribute to input[1].
if (unsqueeze.InputDefs().size() > 1) {
const NodeArg* axes_node_arg = unsqueeze.InputDefs()[1];
return optimizer_utils::AppendTensorFromInitializer(graph, *axes_node_arg, axes, true);
Comment thread
yuslepukhin marked this conversation as resolved.
}
Comment thread
yuslepukhin marked this conversation as resolved.
Expand Down Expand Up @@ -169,7 +172,8 @@ bool ReshapeFusion::Match_One_Element_Output_Subgraph_1(Graph& graph, const Node
const Node& gather = edges[1]->GetNode();
const Node& shape = edges[2]->GetNode();

if (graph_utils::MatchesOpSinceVersion(shape, {15})) {
// Opset 15+ added start/end attributes to Shape. Reject partial-shape queries.
if (shape.SinceVersion() >= 15) {
const ONNX_NAMESPACE::AttributeProto* start_attr = graph_utils::GetNodeAttribute(shape, "start");
const ONNX_NAMESPACE::AttributeProto* end_attr = graph_utils::GetNodeAttribute(shape, "end");
if (!((!start_attr || static_cast<int>(start_attr->i()) == 0) && (!end_attr))) {
Expand Down
74 changes: 58 additions & 16 deletions onnxruntime/test/optimizer/graph_transform_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -6883,6 +6883,21 @@ TEST_F(GraphTransformationTests, AttentionFusionMobileClipMhaTest) {
std::make_unique<AttentionFusion>());
}

TEST_F(GraphTransformationTests, AttentionFusionMobileClipMhaOpset25Test) {
auto build_test_case = [](ModelTestBuilder& builder) {
BuildMobileClipAttentionTestCase(builder, MobileClipProjectionType::MatMulAdd);
};

TransformerTester(build_test_case,
CheckMobileClipAttentionFusedSession,
TransformerLevel::Level1,
TransformerLevel::Level2,
25,
1e-3,
0.0,
std::make_unique<AttentionFusion>());
}

TEST_F(GraphTransformationTests, AttentionFusionMobileClipMhaProjectionGemmTest) {
auto build_test_case = [](ModelTestBuilder& builder) {
BuildMobileClipAttentionTestCase(builder, MobileClipProjectionType::GemmWithReshapes);
Expand Down Expand Up @@ -8219,8 +8234,7 @@ TEST_F(GraphTransformationTests, ReshapeFusionOpsetTest) {
return Status::OK();
};

const std::vector<int> opsets{11, 12, 13, 14, 15, 18};
bool shape_test_for_opset15 = false;
const std::vector<int> opsets{11, 12, 13, 14, 15, 18, 19, 21, 23, 24, 25};

for (auto& opset : opsets) {
auto build_test_case = [&](ModelTestBuilder& builder) {
Expand All @@ -8245,14 +8259,7 @@ TEST_F(GraphTransformationTests, ReshapeFusionOpsetTest) {

builder.AddNode("Add", {input_arg0, input_arg1}, {add_out});
if (opset_version >= 15) {
if (shape_test_for_opset15) {
auto& shape_1 = builder.AddNode("Shape", {add_out}, {shape_out});
shape_1.AddAttribute("start", (int64_t)1);
shape_1.AddAttribute("end", (int64_t)2);
} else {
builder.AddNode("Shape", {add_out}, {shape_out}).AddAttribute("start", (int64_t)0);
shape_test_for_opset15 = true;
}
builder.AddNode("Shape", {add_out}, {shape_out}).AddAttribute("start", (int64_t)0);
} else {
builder.AddNode("Shape", {add_out}, {shape_out});
}
Expand All @@ -8271,13 +8278,48 @@ TEST_F(GraphTransformationTests, ReshapeFusionOpsetTest) {
builder.AddNode("Reshape", {add_out, concattraining1_out}, {out});
};

// Test that the fusion fires for every opset.
std::unique_ptr<GraphTransformer> transformer = std::make_unique<ReshapeFusion>();
if (opset >= 15 && shape_test_for_opset15) {
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset, *logger_, std::move(transformer), TransformerLevel::Level1, 1,
pre_graph_checker, pre_graph_checker));
} else {
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset, *logger_, std::move(transformer), TransformerLevel::Level1, 1,
pre_graph_checker, post_graph_checker));
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset, *logger_, std::move(transformer), TransformerLevel::Level1, 1,
pre_graph_checker, post_graph_checker));

// For opset >= 15, also test that partial Shape (start=1, end=2) prevents fusion.
if (opset >= 15) {
auto build_partial_shape_case = [&](ModelTestBuilder& builder) {
auto* input_arg0 = builder.MakeInput<float>({{batch_size, seq_lenth, hidden_size}});
auto* input_arg1 = builder.MakeInput<float>({{hidden_size}});
auto* scalar_int_0 = builder.MakeInitializer<int64_t>({}, {0});
auto* scalar_int_1 = builder.MakeInitializer<int64_t>({}, {1});
auto* single_value_1d_int_0 = builder.MakeInitializer<int64_t>({1}, {0});
auto* single_value_1d_int_16 = builder.MakeInitializer<int64_t>({1}, {16});
auto* single_value_1d_int_64 = builder.MakeInitializer<int64_t>({1}, {64});
auto* add_out = builder.MakeIntermediate();
auto* shape_out = builder.MakeIntermediate();
auto* gather_out_0 = builder.MakeIntermediate();
auto* gather_out_1 = builder.MakeIntermediate();
auto* unsqueeze_out_0 = builder.MakeIntermediate();
auto* unsqueeze_out_1 = builder.MakeIntermediate();
auto* concattraining1_out = builder.MakeIntermediate();
auto* concattraining1_length = builder.MakeIntermediate();
auto* out = builder.MakeOutput();

builder.AddNode("Add", {input_arg0, input_arg1}, {add_out});
auto& shape_1 = builder.AddNode("Shape", {add_out}, {shape_out});
shape_1.AddAttribute("start", (int64_t)1);
shape_1.AddAttribute("end", (int64_t)2);
builder.AddNode("Gather", {shape_out, scalar_int_0}, {gather_out_0});
builder.AddNode("Gather", {shape_out, scalar_int_1}, {gather_out_1});
builder.AddNode("Unsqueeze", {gather_out_0, single_value_1d_int_0}, {unsqueeze_out_0});
builder.AddNode("Unsqueeze", {gather_out_1, single_value_1d_int_0}, {unsqueeze_out_1});
builder.AddNode("ConcatTraining", {unsqueeze_out_0, unsqueeze_out_1, single_value_1d_int_16, single_value_1d_int_64},
{concattraining1_out, concattraining1_length}, "com.microsoft")
.AddAttribute("axis", static_cast<int64_t>(0));
builder.AddNode("Reshape", {add_out, concattraining1_out}, {out});
};

std::unique_ptr<GraphTransformer> transformer_no_fuse = std::make_unique<ReshapeFusion>();
ASSERT_STATUS_OK(TestGraphTransformer(build_partial_shape_case, opset, *logger_, std::move(transformer_no_fuse),
TransformerLevel::Level1, 1, pre_graph_checker, pre_graph_checker));
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -258,6 +258,13 @@ TEST_F(GraphTransformationTests, GroupQueryAttentionPreNormFusionFusesQwenPatter
TransformerLevel::Level2, /*steps=*/1, nullptr, CheckFusedGraph));
}

TEST_F(GraphTransformationTests, GroupQueryAttentionPreNormFusionFusesQwenPatternOpset25) {
auto build = [](ModelTestBuilder& builder) { BuildQwenQkPostNormPattern(builder, BuildOptions{}); };
ASSERT_STATUS_OK(TestGraphTransformer(
build, /*opset_version=*/25, *logger_, MakeWebGpuTransformer(),
TransformerLevel::Level2, /*steps=*/1, nullptr, CheckFusedGraph));
}

TEST_F(GraphTransformationTests, GroupQueryAttentionPreNormFusionMatchesUnfusedWebGpuResults) {
if (!DefaultWebGpuExecutionProvider()) {
GTEST_SKIP() << "WebGPU EP unavailable in this build.";
Expand Down
Loading