From 0b1bbd7ccc15d780c6aaf676944557b26c5d10b1 Mon Sep 17 00:00:00 2001 From: xiaoh Date: Mon, 28 Sep 2026 02:39:26 -0500 Subject: [PATCH 1/3] Only push a transpose through a shared output when it can cancel. The cost check treated any output that leads to a transpose as a benefit, so a shared QDQ value was pushed even when another branch could not cancel the permutation. Count that benefit only for a single consumer or a path that reaches the inverse permutation. Co-authored-by: Cursor --- .../onnx_transpose_optimization.cc | 72 +++++++++++++++++-- .../optimizer/transpose_optimizer_test.cc | 63 ++++++++++++++++ 2 files changed, 130 insertions(+), 5 deletions(-) diff --git a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc index 421451a0c1f68..fb0cf7012a53f 100755 --- a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc +++ b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc @@ -14,6 +14,7 @@ #include #include +#include "core/common/inlined_containers.h" #include "core/common/make_string.h" #include "core/graph/constants.h" @@ -2971,12 +2972,69 @@ static const HandlerInfo* GetHandler(api::NodeRef& node, const HandlerMap& exten return nullptr; } -static int CalculateCost(const api::GraphRef& graph, const api::NodeRef& node, +static bool HasPathToCancelingTranspose(OptimizerCtx& ctx, std::string_view output, + const std::vector& perm, + const std::unordered_set& outputs_leading_to_transpose) { + const auto perm_inv = InvertPerm(perm); + onnxruntime::InlinedVector pending{std::string(output)}; + onnxruntime::InlinedHashSet visited; + while (!pending.empty()) { + auto value = std::move(pending.back()); + pending.pop_back(); + if (!visited.insert(value).second || outputs_leading_to_transpose.count(value) == 0) { + continue; + } + + auto consumers = ctx.graph.GetValueConsumers(value); + for (auto& consumer : consumers->nodes) { + if (consumer->IsOp("Transpose")) { + auto downstream_perm = GetPermAttrIfValid(*consumer); + if (downstream_perm && *downstream_perm == perm_inv) { + return true; + } + // A non-canceling transpose ends this chain. + continue; + } + + const auto* handler = GetHandler(*consumer, ctx.extended_handlers); + if (handler == nullptr || !handler->transposes_outputs) { + continue; + } + + // These handlers can change the permutation when changing rank. Do not compare the + // original permutation against a transpose beyond them. + if (handler->handler_fn == squeeze_handler.handler_fn || + handler->handler_fn == unsqueeze_handler.handler_fn || + handler->handler_fn == gather_handler.handler_fn || + ((handler->handler_fn == reduce_op_handler.handler_fn || + handler->handler_fn == arg_min_max_handler.handler_fn) && + consumer->GetAttributeIntDefault("keepdims", 1) == 0)) { + continue; + } + + const auto inputs = consumer->Inputs(); + const auto input_indices = handler->transposible_inputs_fn(ctx, *consumer); + if (std::none_of(input_indices.begin(), input_indices.end(), + [&](size_t index) { return inputs[index] == value; })) { + continue; + } + + for (auto consumer_output : consumer->Outputs()) { + pending.emplace_back(consumer_output); + } + } + } + + return false; +} + +static int CalculateCost(OptimizerCtx& ctx, const api::NodeRef& node, const std::vector& perm, const std::unordered_set& outputs_leading_to_transpose, const HandlerInfo& info, const std::vector& input_indices, const HandlerMap& extended_handlers) { + const auto& graph = ctx.graph; // We require the input cost (number of transposes before the op) and the total cost to strictly decrease. // Strict decrease of the input cost ensures the optimization is stable, since the total cost decrease is just an // estimate (the transpose after the op may or may not cancel with a subsequent transpose). We don't want @@ -2993,7 +3051,11 @@ static int CalculateCost(const api::GraphRef& graph, const api::NodeRef& node, for (auto out : outputs) { out_cost = std::max(out_cost, EstimateValueRank(graph, out)); if (outputs_leading_to_transpose.find(std::string(out)) != outputs_leading_to_transpose.end()) { - has_output_leading_to_transpose = true; + auto consumers = graph.GetValueConsumers(out); + if (consumers->nodes.size() <= 1 || + HasPathToCancelingTranspose(ctx, out, perm, outputs_leading_to_transpose)) { + has_output_leading_to_transpose = true; + } } } @@ -3006,7 +3068,7 @@ static int CalculateCost(const api::GraphRef& graph, const api::NodeRef& node, } // Default cost check. Returns `true` if pushing the Transpose through the node is considered to be beneficial. -static bool DefaultCostCheck(const api::GraphRef& graph, const api::NodeRef& node, +static bool DefaultCostCheck(OptimizerCtx& ctx, const api::NodeRef& node, const std::vector& perm, const std::unordered_set& outputs_leading_to_transpose, const HandlerInfo& info, @@ -3016,7 +3078,7 @@ static bool DefaultCostCheck(const api::GraphRef& graph, const api::NodeRef& nod return true; } - int cost = CalculateCost(graph, node, perm, outputs_leading_to_transpose, info, transposable_input_indices, + int cost = CalculateCost(ctx, node, perm, outputs_leading_to_transpose, info, transposable_input_indices, extended_handlers); return cost < 0; } @@ -3043,7 +3105,7 @@ bool ProcessTranspose(OptimizerCtx& ctx, api::NodeRef& transpose, api::NodeRef& } if (cost == CostCheckResult::kFallThrough) { - cost = DefaultCostCheck(ctx.graph, node, perm, outputs_leading_to_transpose, *info, input_indices, + cost = DefaultCostCheck(ctx, node, perm, outputs_leading_to_transpose, *info, input_indices, ctx.extended_handlers) ? CostCheckResult::kPushTranspose : CostCheckResult::kStop; diff --git a/onnxruntime/test/optimizer/transpose_optimizer_test.cc b/onnxruntime/test/optimizer/transpose_optimizer_test.cc index 84f6e7fca8460..2e848b7f95f99 100644 --- a/onnxruntime/test/optimizer/transpose_optimizer_test.cc +++ b/onnxruntime/test/optimizer/transpose_optimizer_test.cc @@ -107,6 +107,69 @@ int EstimateTransposeCost(const Graph& graph) { return cost; } +TEST(TransposeOptimizerTests, SharedQDQOutputRequiresCancelingTranspose) { + for (bool branched : {false, true}) { + for (bool cancel : {false, true}) { + SCOPED_TRACE(MakeString("branched=", branched, ", cancel=", cancel)); + std::string q1_name; + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input = builder.MakeInput({2, 3, 4}, 0.0f, 1.0f); + auto add_qdq = [&](NodeArg* data) { + auto* q = builder.MakeIntermediate(); + auto* dq = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(data, 0.01f, 0, q); + builder.AddDequantizeLinearNode(q, 0.01f, 0, dq); + return dq; + }; + auto* dq0 = add_qdq(input); + auto* t1 = builder.MakeIntermediate(); + builder.AddNode("Transpose", {dq0}, {t1}).AddAttribute("perm", std::vector{1, 2, 0}); + auto* q1 = builder.MakeIntermediate(); + q1_name = builder.AddQuantizeLinearNode(t1, 0.01f, 0, q1).Name(); + + auto add_branch = [&](bool with_transpose) { + auto* dq = builder.MakeIntermediate(); + builder.AddDequantizeLinearNode(q1, 0.01f, 0, dq); + auto* starts = builder.MakeInitializer({1}, {0}); + auto* ends = builder.MakeInitializer({1}, {2}); + auto* axes = builder.MakeInitializer({1}, {0}); + auto* slice = with_transpose ? builder.MakeIntermediate() : builder.MakeOutput(); + builder.AddNode("Slice", {dq, starts, ends, axes}, {slice}); + if (with_transpose) { + auto* dq_b = add_qdq(slice); + auto* output = builder.MakeOutput(); + builder.AddNode("Transpose", {dq_b}, {output}) + .AddAttribute("perm", cancel ? std::vector{2, 0, 1} : std::vector{1, 2, 0}); + } + }; + add_branch(true); + if (branched) { + add_branch(false); + add_branch(false); + } + }; + + auto check_graph = [&](InferenceSessionWrapper& session) { + auto& graph = session.GetGraph(); + auto counts = CountOpsInGraph(graph); + EXPECT_EQ(counts["Transpose"], (branched ? 1 : 0) + (cancel ? 0 : 1)); + bool q1_has_transpose_input = false; + for (const auto& node : graph.Nodes()) { + if (node.Name() == q1_name) { + const auto* producer = graph.GetProducerNode(node.InputDefs()[0]->Name()); + q1_has_transpose_input = producer != nullptr && producer->OpType() == "Transpose"; + } + } + EXPECT_EQ(q1_has_transpose_input, branched && !cancel); + }; + + TransformerTester(build_test_case, check_graph, TransformerLevel::Default, TransformerLevel::Level1, + 15, 0.0, 0.0, + std::make_unique(TestCPUExecutionProvider()->CreatePreferredAllocators()[0])); + } + } +} + TEST(TransposeOptimizerTests, TestSplit) { auto build_test_case_1 = [&](ModelTestBuilder& builder) { auto* input0_arg = builder.MakeInput({4, 6, 10}, 0.0, 1.0); From 34d4a9c42ffd8a002b672617167fadaa5471dfdb Mon Sep 17 00:00:00 2001 From: xiaoh Date: Mon, 28 Sep 2026 05:48:23 -0500 Subject: [PATCH 2/3] Do not treat a graph output as a canceling transpose consumer. GetValueConsumers only lists node consumers, so a graph output or subgraph input was treated as a single consumer and the shared transpose was still pushed. Co-authored-by: Cursor --- .../onnx_transpose_optimization.cc | 7 ++- .../optimizer/transpose_optimizer_test.cc | 50 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc index fb0cf7012a53f..2da62a5e6d772 100755 --- a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc +++ b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc @@ -3051,9 +3051,12 @@ static int CalculateCost(OptimizerCtx& ctx, const api::NodeRef& node, for (auto out : outputs) { out_cost = std::max(out_cost, EstimateValueRank(graph, out)); if (outputs_leading_to_transpose.find(std::string(out)) != outputs_leading_to_transpose.end()) { + // `nodes` lists only node consumers. A graph output or subgraph input sets `comprehensive` to false + // and cannot cancel the permutation, so it blocks the benefit the same way a second branch does. auto consumers = graph.GetValueConsumers(out); - if (consumers->nodes.size() <= 1 || - HasPathToCancelingTranspose(ctx, out, perm, outputs_leading_to_transpose)) { + if (consumers->comprehensive && + (consumers->nodes.size() <= 1 || + HasPathToCancelingTranspose(ctx, out, perm, outputs_leading_to_transpose))) { has_output_leading_to_transpose = true; } } diff --git a/onnxruntime/test/optimizer/transpose_optimizer_test.cc b/onnxruntime/test/optimizer/transpose_optimizer_test.cc index 2e848b7f95f99..28be73a1d35e9 100644 --- a/onnxruntime/test/optimizer/transpose_optimizer_test.cc +++ b/onnxruntime/test/optimizer/transpose_optimizer_test.cc @@ -170,6 +170,56 @@ TEST(TransposeOptimizerTests, SharedQDQOutputRequiresCancelingTranspose) { } } +// The quantized value is a graph output and also feeds one branch that could cancel the permutation. +// Pushing would leave that transpose on the graph output, so the transpose stays in front of the QuantizeLinear. +TEST(TransposeOptimizerTests, GraphOutputBlocksSharedTransposePush) { + std::string q1_name; + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input = builder.MakeInput({2, 3, 4}, 0.0f, 1.0f); + auto add_qdq = [&](NodeArg* data) { + auto* q = builder.MakeIntermediate(); + auto* dq = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(data, 0.01f, 0, q); + builder.AddDequantizeLinearNode(q, 0.01f, 0, dq); + return dq; + }; + auto* dq0 = add_qdq(input); + auto* t1 = builder.MakeIntermediate(); + builder.AddNode("Transpose", {dq0}, {t1}).AddAttribute("perm", std::vector{1, 2, 0}); + auto* q1 = builder.MakeOutput(); + q1_name = builder.AddQuantizeLinearNode(t1, 0.01f, 0, q1).Name(); + + auto* dq = builder.MakeIntermediate(); + builder.AddDequantizeLinearNode(q1, 0.01f, 0, dq); + auto* starts = builder.MakeInitializer({1}, {0}); + auto* ends = builder.MakeInitializer({1}, {2}); + auto* axes = builder.MakeInitializer({1}, {0}); + auto* slice = builder.MakeIntermediate(); + builder.AddNode("Slice", {dq, starts, ends, axes}, {slice}); + auto* dq_b = add_qdq(slice); + auto* output = builder.MakeOutput(); + builder.AddNode("Transpose", {dq_b}, {output}).AddAttribute("perm", std::vector{2, 0, 1}); + }; + + auto check_graph = [&](InferenceSessionWrapper& session) { + auto& graph = session.GetGraph(); + auto counts = CountOpsInGraph(graph); + EXPECT_EQ(counts["Transpose"], 2); + bool q1_has_transpose_input = false; + for (const auto& node : graph.Nodes()) { + if (node.Name() == q1_name) { + const auto* producer = graph.GetProducerNode(node.InputDefs()[0]->Name()); + q1_has_transpose_input = producer != nullptr && producer->OpType() == "Transpose"; + } + } + EXPECT_TRUE(q1_has_transpose_input); + }; + + TransformerTester(build_test_case, check_graph, TransformerLevel::Default, TransformerLevel::Level1, + 15, 0.0, 0.0, + std::make_unique(TestCPUExecutionProvider()->CreatePreferredAllocators()[0])); +} + TEST(TransposeOptimizerTests, TestSplit) { auto build_test_case_1 = [&](ModelTestBuilder& builder) { auto* input0_arg = builder.MakeInput({4, 6, 10}, 0.0, 1.0); From 599b99c74343b4a7c9161bfdbc2dce63d807da5f Mon Sep 17 00:00:00 2001 From: xiaoh Date: Mon, 28 Sep 2026 08:33:09 -0500 Subject: [PATCH 3/3] Stop the cancel walk at a graph output or subgraph input. The cost check already ignored a non-comprehensive consumer on the node output. A later value on the path could still be a graph output and was credited when a node consumer reached the inverse permutation. Co-authored-by: Cursor --- .../onnx_transpose_optimization.cc | 5 ++ .../optimizer/transpose_optimizer_test.cc | 56 +++++++++++++++++++ 2 files changed, 61 insertions(+) diff --git a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc index 2da62a5e6d772..946f9aa0e993a 100755 --- a/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc +++ b/onnxruntime/core/optimizer/transpose_optimization/onnx_transpose_optimization.cc @@ -2986,6 +2986,11 @@ static bool HasPathToCancelingTranspose(OptimizerCtx& ctx, std::string_view outp } auto consumers = ctx.graph.GetValueConsumers(value); + // A graph output or subgraph input is another use that cannot cancel. Do not credit a + // transpose reached through that value. + if (!consumers->comprehensive) { + continue; + } for (auto& consumer : consumers->nodes) { if (consumer->IsOp("Transpose")) { auto downstream_perm = GetPermAttrIfValid(*consumer); diff --git a/onnxruntime/test/optimizer/transpose_optimizer_test.cc b/onnxruntime/test/optimizer/transpose_optimizer_test.cc index 28be73a1d35e9..109c020101243 100644 --- a/onnxruntime/test/optimizer/transpose_optimizer_test.cc +++ b/onnxruntime/test/optimizer/transpose_optimizer_test.cc @@ -220,6 +220,62 @@ TEST(TransposeOptimizerTests, GraphOutputBlocksSharedTransposePush) { std::make_unique(TestCPUExecutionProvider()->CreatePreferredAllocators()[0])); } +// q1 has two node consumers, so the cancel walk runs. One branch reaches an inverse transpose only by +// passing through a value that is also a graph output. That hidden use cannot cancel, so the transpose +// stays in front of the QuantizeLinear. +TEST(TransposeOptimizerTests, IntermediateGraphOutputBlocksCancelingPath) { + std::string q1_name; + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input = builder.MakeInput({2, 3, 4}, 0.0f, 1.0f); + auto add_qdq = [&](NodeArg* data) { + auto* q = builder.MakeIntermediate(); + auto* dq = builder.MakeIntermediate(); + builder.AddQuantizeLinearNode(data, 0.01f, 0, q); + builder.AddDequantizeLinearNode(q, 0.01f, 0, dq); + return dq; + }; + auto* dq0 = add_qdq(input); + auto* t1 = builder.MakeIntermediate(); + builder.AddNode("Transpose", {dq0}, {t1}).AddAttribute("perm", std::vector{1, 2, 0}); + auto* q1 = builder.MakeIntermediate(); + q1_name = builder.AddQuantizeLinearNode(t1, 0.01f, 0, q1).Name(); + + auto add_slice = [&](NodeArg* data, NodeArg* slice_out) { + auto* dq = builder.MakeIntermediate(); + builder.AddDequantizeLinearNode(data, 0.01f, 0, dq); + auto* starts = builder.MakeInitializer({1}, {0}); + auto* ends = builder.MakeInitializer({1}, {2}); + auto* axes = builder.MakeInitializer({1}, {0}); + builder.AddNode("Slice", {dq, starts, ends, axes}, {slice_out}); + }; + add_slice(q1, builder.MakeOutput()); + + auto* mid = builder.MakeOutput(); + add_slice(q1, mid); + auto* dq_b = add_qdq(mid); + auto* output = builder.MakeOutput(); + builder.AddNode("Transpose", {dq_b}, {output}).AddAttribute("perm", std::vector{2, 0, 1}); + }; + + auto check_graph = [&](InferenceSessionWrapper& session) { + auto& graph = session.GetGraph(); + auto counts = CountOpsInGraph(graph); + EXPECT_EQ(counts["Transpose"], 2); + bool q1_has_transpose_input = false; + for (const auto& node : graph.Nodes()) { + if (node.Name() == q1_name) { + const auto* producer = graph.GetProducerNode(node.InputDefs()[0]->Name()); + q1_has_transpose_input = producer != nullptr && producer->OpType() == "Transpose"; + } + } + EXPECT_TRUE(q1_has_transpose_input); + }; + + TransformerTester(build_test_case, check_graph, TransformerLevel::Default, TransformerLevel::Level1, + 15, 0.0, 0.0, + std::make_unique(TestCPUExecutionProvider()->CreatePreferredAllocators()[0])); +} + TEST(TransposeOptimizerTests, TestSplit) { auto build_test_case_1 = [&](ModelTestBuilder& builder) { auto* input0_arg = builder.MakeInput({4, 6, 10}, 0.0, 1.0);