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
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
#include <vector>

#include <gsl/gsl>
#include "core/common/inlined_containers.h"
#include "core/common/make_string.h"
#include "core/graph/constants.h"

Expand Down Expand Up @@ -2971,12 +2972,74 @@ 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<int64_t>& perm,
const std::unordered_set<std::string>& outputs_leading_to_transpose) {
const auto perm_inv = InvertPerm(perm);
onnxruntime::InlinedVector<std::string> pending{std::string(output)};
onnxruntime::InlinedHashSet<std::string> 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) {
Comment thread
xiaohanAMD marked this conversation as resolved.
continue;
}

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);
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;
Comment thread
xiaohanAMD marked this conversation as resolved.
}

// 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) &&
Comment thread
xiaohanAMD marked this conversation as resolved.
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<int64_t>& perm,
const std::unordered_set<std::string>& outputs_leading_to_transpose,
const HandlerInfo& info,
const std::vector<size_t>& 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
Expand All @@ -2993,7 +3056,14 @@ 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;
// `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->comprehensive &&
(consumers->nodes.size() <= 1 ||
HasPathToCancelingTranspose(ctx, out, perm, outputs_leading_to_transpose))) {
has_output_leading_to_transpose = true;
}
}
}

Expand All @@ -3006,7 +3076,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<int64_t>& perm,
const std::unordered_set<std::string>& outputs_leading_to_transpose,
const HandlerInfo& info,
Expand All @@ -3016,7 +3086,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;
}
Expand All @@ -3043,7 +3113,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;
Expand Down
169 changes: 169 additions & 0 deletions onnxruntime/test/optimizer/transpose_optimizer_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,175 @@ 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<float>({2, 3, 4}, 0.0f, 1.0f);
auto add_qdq = [&](NodeArg* data) {
auto* q = builder.MakeIntermediate();
auto* dq = builder.MakeIntermediate();
builder.AddQuantizeLinearNode<uint8_t>(data, 0.01f, 0, q);
builder.AddDequantizeLinearNode<uint8_t>(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<int64_t>{1, 2, 0});
auto* q1 = builder.MakeIntermediate();
q1_name = builder.AddQuantizeLinearNode<uint8_t>(t1, 0.01f, 0, q1).Name();

auto add_branch = [&](bool with_transpose) {
auto* dq = builder.MakeIntermediate();
builder.AddDequantizeLinearNode<uint8_t>(q1, 0.01f, 0, dq);
auto* starts = builder.MakeInitializer<int64_t>({1}, {0});
auto* ends = builder.MakeInitializer<int64_t>({1}, {2});
auto* axes = builder.MakeInitializer<int64_t>({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<int64_t>{2, 0, 1} : std::vector<int64_t>{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<TransposeOptimizer>(TestCPUExecutionProvider()->CreatePreferredAllocators()[0]));
}
}
}

// 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<float>({2, 3, 4}, 0.0f, 1.0f);
auto add_qdq = [&](NodeArg* data) {
auto* q = builder.MakeIntermediate();
auto* dq = builder.MakeIntermediate();
builder.AddQuantizeLinearNode<uint8_t>(data, 0.01f, 0, q);
builder.AddDequantizeLinearNode<uint8_t>(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<int64_t>{1, 2, 0});
auto* q1 = builder.MakeOutput();
q1_name = builder.AddQuantizeLinearNode<uint8_t>(t1, 0.01f, 0, q1).Name();

auto* dq = builder.MakeIntermediate();
builder.AddDequantizeLinearNode<uint8_t>(q1, 0.01f, 0, dq);
auto* starts = builder.MakeInitializer<int64_t>({1}, {0});
auto* ends = builder.MakeInitializer<int64_t>({1}, {2});
auto* axes = builder.MakeInitializer<int64_t>({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<int64_t>{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<TransposeOptimizer>(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<float>({2, 3, 4}, 0.0f, 1.0f);
auto add_qdq = [&](NodeArg* data) {
auto* q = builder.MakeIntermediate();
auto* dq = builder.MakeIntermediate();
builder.AddQuantizeLinearNode<uint8_t>(data, 0.01f, 0, q);
builder.AddDequantizeLinearNode<uint8_t>(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<int64_t>{1, 2, 0});
auto* q1 = builder.MakeIntermediate();
q1_name = builder.AddQuantizeLinearNode<uint8_t>(t1, 0.01f, 0, q1).Name();

auto add_slice = [&](NodeArg* data, NodeArg* slice_out) {
auto* dq = builder.MakeIntermediate();
builder.AddDequantizeLinearNode<uint8_t>(data, 0.01f, 0, dq);
auto* starts = builder.MakeInitializer<int64_t>({1}, {0});
auto* ends = builder.MakeInitializer<int64_t>({1}, {2});
auto* axes = builder.MakeInitializer<int64_t>({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<int64_t>{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<TransposeOptimizer>(TestCPUExecutionProvider()->CreatePreferredAllocators()[0]));
}

TEST(TransposeOptimizerTests, TestSplit) {
auto build_test_case_1 = [&](ModelTestBuilder& builder) {
auto* input0_arg = builder.MakeInput<float>({4, 6, 10}, 0.0, 1.0);
Expand Down
Loading