diff --git a/onnxruntime/core/optimizer/transpose_optimizer/transpose_optimizer.cc b/onnxruntime/core/optimizer/transpose_optimizer/transpose_optimizer.cc index b117fa5526d3b..6cdb63319a5bf 100644 --- a/onnxruntime/core/optimizer/transpose_optimizer/transpose_optimizer.cc +++ b/onnxruntime/core/optimizer/transpose_optimizer/transpose_optimizer.cc @@ -1979,11 +1979,18 @@ OptimizeResult OptimizeImpl(OptimizerCtx& ctx) { continue; } - auto consumers = ctx.graph.GetValueConsumers(transpose_node.Outputs()[0]); - bool is_part_of_qdq_group = std::find_if(consumers->nodes.cbegin(), consumers->nodes.cend(), + // Check if Transpose node is the only consumer of dq node + auto consumers_of_dq_node = ctx.graph.GetValueConsumers(dq_node->Outputs()[0]); + if (!consumers_of_dq_node->comprehensive || consumers_of_dq_node->nodes.size() > 1) { + continue; + } + + auto consumers_of_transpose_node = ctx.graph.GetValueConsumers(transpose_node.Outputs()[0]); + bool is_part_of_qdq_group = std::find_if(consumers_of_transpose_node->nodes.cbegin(), + consumers_of_transpose_node->nodes.cend(), [](const std::unique_ptr& node) { return node->OpType() == "QuantizeLinear"; - }) != consumers->nodes.cend(); + }) != consumers_of_transpose_node->nodes.cend(); if (is_part_of_qdq_group) { continue; } diff --git a/onnxruntime/test/optimizer/qdq_test_utils.cc b/onnxruntime/test/optimizer/qdq_test_utils.cc index ab392d9640708..a18d699560632 100644 --- a/onnxruntime/test/optimizer/qdq_test_utils.cc +++ b/onnxruntime/test/optimizer/qdq_test_utils.cc @@ -171,5 +171,15 @@ GetQDQTestCaseFn BuildQDQMatMulTestCase(const std::vector& input1_shape }; } +std::vector GetNodeOpTypesInTopologicalOrder(const Graph& graph) { + std::vector op_types{}; + GraphViewer graph_viewer{graph}; + const auto& ordering = graph_viewer.GetNodesInTopologicalOrder(); + for (const auto node_idx : ordering) { + op_types.push_back(graph.GetNode(node_idx)->OpType()); + } + return op_types; +} + } // namespace test } // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/test/optimizer/qdq_test_utils.h b/onnxruntime/test/optimizer/qdq_test_utils.h index 2a8d2d9a06a75..3d9f5c7271a9c 100644 --- a/onnxruntime/test/optimizer/qdq_test_utils.h +++ b/onnxruntime/test/optimizer/qdq_test_utils.h @@ -3,6 +3,9 @@ #pragma once +#include +#include + #include "graph_transform_test_builder.h" #include "core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.h" @@ -359,5 +362,7 @@ GetQDQTestCaseFn BuildQDQGemmTestCase(const std::vector& input1_shape, }; } +std::vector GetNodeOpTypesInTopologicalOrder(const Graph& graph); + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/optimizer/qdq_transformer_test.cc b/onnxruntime/test/optimizer/qdq_transformer_test.cc index 787cbc0d7021d..79bb64b1223f9 100644 --- a/onnxruntime/test/optimizer/qdq_transformer_test.cc +++ b/onnxruntime/test/optimizer/qdq_transformer_test.cc @@ -36,16 +36,6 @@ namespace onnxruntime { namespace test { -static std::vector GetNodeOpTypesInTopologicalOrder(const Graph& graph) { - std::vector op_types{}; - GraphViewer graph_viewer{graph}; - const auto& ordering = graph_viewer.GetNodesInTopologicalOrder(); - for (const auto node_idx : ordering) { - op_types.push_back(graph.GetNode(node_idx)->OpType()); - } - return op_types; -} - #if !defined(DISABLE_CONTRIB_OPS) template diff --git a/onnxruntime/test/optimizer/transpose_optimizer_test.cc b/onnxruntime/test/optimizer/transpose_optimizer_test.cc index 176d1541583fd..68e2be1b34cf7 100644 --- a/onnxruntime/test/optimizer/transpose_optimizer_test.cc +++ b/onnxruntime/test/optimizer/transpose_optimizer_test.cc @@ -8,6 +8,7 @@ #include "graph_transform_test_builder.h" #include "core/graph/graph.h" +#include "qdq_test_utils.h" #include "test/test_environment.h" #include "test/util/include/asserts.h" @@ -3591,6 +3592,42 @@ TEST(TransposeOptimizerTests, TestDequantizeLinearNoAxis) { /*opset_version*/ 10); } +TEST(TransposeOptimizerTests, TestDequantizeLinearTransposePropagation) { + auto build_test_case_1 = [&](ModelTestBuilder& builder) { + auto* input0_arg = MakeInput(builder, {{2, -1, 6, 3}}, {2, 4, 6, 3}, 0, 5); + auto* input1_arg = MakeInput(builder, {std::vector{}}, std::vector{}, {2.3f}); + auto* input2_arg = MakeInput(builder, {std::vector{}}, std::vector{}, {10}); + auto* dequantizelinear_1_out_0 = builder.MakeIntermediate(); + auto* transpose_1_out_0 = builder.MakeOutput(); + auto* transpose_2_out_0 = builder.MakeOutput(); + + builder.AddNode("DequantizeLinear", {input0_arg, input1_arg, input2_arg}, {dequantizelinear_1_out_0}); + + auto& transpose_1 = builder.AddNode("Transpose", {dequantizelinear_1_out_0}, {transpose_1_out_0}); + transpose_1.AddAttribute("perm", std::vector{0, 3, 1, 2}); + + auto& transpose_2 = builder.AddNode("Transpose", {dequantizelinear_1_out_0}, {transpose_2_out_0}); + transpose_2.AddAttribute("perm", std::vector{0, 2, 3, 1}); + }; + + auto check_graph = [&](InferenceSessionWrapper& session) { + std::vector expected_op_types_in_order{ + "DequantizeLinear", + "Transpose", + "Transpose"}; + + const auto op_types_in_order = GetNodeOpTypesInTopologicalOrder(session.GetGraph()); + EXPECT_EQ(op_types_in_order, expected_op_types_in_order); + }; + + + TransformerTester(build_test_case_1, + check_graph, + TransformerLevel::Default, + TransformerLevel::Level1, + /*opset_version*/ 10); +} + TEST(TransposeOptimizerTests, TestCast) { auto build_test_case_1 = [&](ModelTestBuilder& builder) { auto* input0_arg = MakeInput(builder, {{-1, 4, -1, 5}}, {2, 4, 6, 5}, -1, 5);