From 58f1d15d19006464546c73ac6fbed95ff5c90b0a Mon Sep 17 00:00:00 2001 From: guyang3532 <62738430+guyang3532@users.noreply.github.com> Date: Fri, 27 Oct 2023 23:50:18 +0800 Subject: [PATCH] Replace Transpose with Replace if they are equivalent (#18096) ### Description Transpose is equivalent to a Reshape if: empty dimensions can change place, not empty dimensions must be in the same order in the permuted tenosr. Example: Shape=(1,1,1024,4096) -> perm=(2,0,3,1). This pr adds a graph transformer which replaces Transpose with Reshape if they are equivalent. Because Transpose need memory copy while Reshape needn't, this replacement can save overhead for memory copy. --- .../testdata/transform/transpose_graph_gen.py | 41 +++++++++++ .../transpose_to_reshape_invalid.onnx | Bin 0 -> 235 bytes .../transform/transpose_to_reshape_valid.onnx | Bin 0 -> 235 bytes .../core/optimizer/graph_transformer_utils.cc | 2 + .../core/optimizer/transpose_replacement..cc | 68 ++++++++++++++++++ .../core/optimizer/transpose_replacement.h | 38 ++++++++++ .../test/optimizer/graph_transform_test.cc | 41 +++++++++++ 7 files changed, 190 insertions(+) create mode 100644 onnxruntime/test/testdata/transform/transpose_graph_gen.py create mode 100644 onnxruntime/test/testdata/transform/transpose_to_reshape_invalid.onnx create mode 100644 onnxruntime/test/testdata/transform/transpose_to_reshape_valid.onnx create mode 100644 orttraining/orttraining/core/optimizer/transpose_replacement..cc create mode 100644 orttraining/orttraining/core/optimizer/transpose_replacement.h diff --git a/onnxruntime/test/testdata/transform/transpose_graph_gen.py b/onnxruntime/test/testdata/transform/transpose_graph_gen.py new file mode 100644 index 0000000000..14f2994a19 --- /dev/null +++ b/onnxruntime/test/testdata/transform/transpose_graph_gen.py @@ -0,0 +1,41 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +import onnx +from onnx import TensorProto, helper + + +def GenerateModel(model_name, valid): # noqa: N802 + nodes = [ + helper.make_node("Transpose", ["input_0"], ["transposed_input_0"], perm=[2, 1, 3, 0]), + helper.make_node("Add", ["transposed_input_0", "input_1"], ["output"]), + ] + + if valid: + inputs = [ + helper.make_tensor_value_info("input_0", TensorProto.FLOAT, [1, 1, 3, 3]), + helper.make_tensor_value_info("input_1", TensorProto.FLOAT, [3, 1, 3, 1]), + ] + outputs = [helper.make_tensor_value_info("output", TensorProto.FLOAT, [3, 1, 3, 1])] + else: + inputs = [ + helper.make_tensor_value_info("input_0", TensorProto.FLOAT, [1, 2, 3, 3]), + helper.make_tensor_value_info("input_1", TensorProto.FLOAT, [3, 2, 3, 1]), + ] + outputs = [helper.make_tensor_value_info("output", TensorProto.FLOAT, [3, 2, 3, 1])] + + graph = helper.make_graph( + nodes, + "TransposeAndAdd", # name + inputs, + outputs, + [], + ) + + model = helper.make_model(graph) + onnx.save(model, model_name) + + +GenerateModel("transpose_to_reshape_valid.onnx", True) +GenerateModel("transpose_to_reshape_invalid.onnx", False) diff --git a/onnxruntime/test/testdata/transform/transpose_to_reshape_invalid.onnx b/onnxruntime/test/testdata/transform/transpose_to_reshape_invalid.onnx new file mode 100644 index 0000000000000000000000000000000000000000..a09b13fc184a8fbf2f494fd89fd79483157e0940 GIT binary patch literal 235 zcmdCbAZ1`WNr4M$3oaE-Oaj6Hml!l} literal 0 HcmV?d00001 diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc index e5c65b2a96..57d76577f1 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc @@ -63,6 +63,7 @@ #include "orttraining/core/optimizer/scaled_sum_fusion.h" #include "orttraining/core/optimizer/shape_optimizer.h" #include "orttraining/core/optimizer/transformer_layer_recompute.h" +#include "orttraining/core/optimizer/transpose_replacement.h" #include "core/optimizer/compute_optimizer/upstream_gather.h" #include "core/optimizer/compute_optimizer/upstream_reshape.h" #include "core/optimizer/pre_shape_node_elimination.h" @@ -203,6 +204,7 @@ std::vector> GeneratePreTrainingTransformers( std::make_unique(optimizer_utils::GenerateRuleBasedTransformerName(level), compatible_eps); ORT_THROW_IF_ERROR(rule_transformer->Register(std::make_unique())); + ORT_THROW_IF_ERROR(rule_transformer->Register(std::make_unique())); } break; case TransformerLevel::Level3: { diff --git a/orttraining/orttraining/core/optimizer/transpose_replacement..cc b/orttraining/orttraining/core/optimizer/transpose_replacement..cc new file mode 100644 index 0000000000..48e9c4d6e6 --- /dev/null +++ b/orttraining/orttraining/core/optimizer/transpose_replacement..cc @@ -0,0 +1,68 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "orttraining/core/optimizer/transpose_replacement.h" + +#include "core/common/logging/logging.h" +#include "core/optimizer/rewrite_rule.h" +#include "core/optimizer/utils.h" +#include "core/graph/graph.h" +#include "core/graph/graph_utils.h" + +namespace onnxruntime { + +Status TransposeReplacement::Apply(Graph& graph, + Node& transpose_node, + RewriteRuleEffect& rule_effect, + const logging::Logger& logger) const { + auto& transpose_inputs = transpose_node.MutableInputDefs(); + auto& transpose_outputs = transpose_node.MutableOutputDefs(); + NodeArg* input = transpose_inputs[0]; + auto input_shape = input->Shape(); + if (!input_shape) { + LOG_DEBUG_INFO(logger, "Exit TransposeReplacement optimization for input shape is None."); + return Status::OK(); + } + auto perm = graph_utils::onnx_repeated_values::RetrieveValues(transpose_node.GetAttributes().at("perm")); + InlinedVector new_shape; + new_shape.reserve(perm.size()); + int64_t last_permuted_axis = 0; + for (int i = 0; i < static_cast(perm.size()); ++i) { + if (!input_shape->dim(static_cast(perm[i])).has_dim_value()) { + LOG_DEBUG_INFO(logger, "Exit TransposeReplacement optimization for not supporting symbolic shape."); + return Status::OK(); + } + new_shape.push_back(input_shape->dim(static_cast(perm[i])).dim_value()); + if (input_shape->dim(static_cast(perm[i])).dim_value() == 1) + continue; + if (perm[i] < last_permuted_axis) { + LOG_DEBUG_INFO(logger, "Exit TransposeReplacement optimization for not supporting shape."); + return Status::OK(); + } + last_permuted_axis = perm[i]; + } + + transpose_inputs.push_back( + optimizer::compute_optimizer::CreateInitializerFromVector(graph, + {static_cast(new_shape.size())}, + new_shape, + graph.GenerateNodeArgName("transpose_reshape_shape"))); + + Node& transpose_reshape_node = graph.AddNode(graph.GenerateNodeName("Transpose_Reshape"), + "Reshape", + "Transpose replaced Reshape", + transpose_inputs, + transpose_outputs, + nullptr, + kOnnxDomain); + transpose_reshape_node.SetExecutionProviderType(transpose_node.GetExecutionProviderType()); + graph_utils::FinalizeNodeFusion(graph, transpose_reshape_node, transpose_node); + rule_effect = RewriteRuleEffect::kRemovedCurrentNode; + return Status::OK(); +} + +bool TransposeReplacement::SatisfyCondition(const Graph&, const Node&, const logging::Logger&) const { + return true; +} + +} // namespace onnxruntime diff --git a/orttraining/orttraining/core/optimizer/transpose_replacement.h b/orttraining/orttraining/core/optimizer/transpose_replacement.h new file mode 100644 index 0000000000..c38e402339 --- /dev/null +++ b/orttraining/orttraining/core/optimizer/transpose_replacement.h @@ -0,0 +1,38 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/optimizer/rewrite_rule.h" +#include "core/optimizer/compute_optimizer/shared_utils.h" + +namespace onnxruntime { + +/** +@Class TransposeReplacement + +Transpose is equivalent to a Reshape if: + empty dimensions (which dim_value=1) can change place, not empty dimensions must be in + the same order in the permuted tenosr. + Example: Shape=(1,1,1024,4096) -> perm=(2,0,3,1). + +This Rewrite rule replaces Transpose which meets the requirments with Reshape. +Because Transpose need memory copy while Reshape needn't, this replacement can save overhead for memory copy. + +It is attempted to be triggered only on nodes with op type "Transpose". +*/ +class TransposeReplacement : public RewriteRule { + public: + TransposeReplacement() noexcept : RewriteRule("TransposeReplacement") {} + + std::vector TargetOpTypes() const noexcept override { + return {"Transpose"}; + } + + private: + bool SatisfyCondition(const Graph& graph, const Node& node, const logging::Logger& logger) const override; + + Status Apply(Graph& graph, Node& node, RewriteRuleEffect& rule_effect, const logging::Logger& logger) const override; +}; + +} // namespace onnxruntime diff --git a/orttraining/orttraining/test/optimizer/graph_transform_test.cc b/orttraining/orttraining/test/optimizer/graph_transform_test.cc index 94ca87b2ac..20b9354d85 100644 --- a/orttraining/orttraining/test/optimizer/graph_transform_test.cc +++ b/orttraining/orttraining/test/optimizer/graph_transform_test.cc @@ -18,6 +18,7 @@ #include "orttraining/core/optimizer/concat_replacement.h" #include "orttraining/core/optimizer/batchnorm_replacement.h" #include "orttraining/core/optimizer/localized_recompute.h" +#include "orttraining/core/optimizer/transpose_replacement.h" #include "test/optimizer/graph_transform_test_builder.h" #include "test/optimizer/graph_transform_test_fixture.h" #include "test/util/include/default_providers.h" @@ -551,6 +552,46 @@ TEST_F(GraphTransformationTests, ConcatReplacement) { ASSERT_EQ(op_to_count["com.microsoft.ConcatTraining"], 1); } +TEST_F(GraphTransformationTests, TransposeReplacement) { + { + auto model_uri = MODEL_FOLDER "transpose_to_reshape_valid.onnx"; + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model, nullptr, *logger_).IsOK()); + Graph& graph = p_model->MainGraph(); + + auto rule_transformer_L1 = std::make_unique("TransposeReplacement"); + ASSERT_STATUS_OK(rule_transformer_L1->Register(std::make_unique())); + onnxruntime::GraphTransformerManager graph_transformation_mgr{1}; + ASSERT_STATUS_OK(graph_transformation_mgr.Register(std::move(rule_transformer_L1), TransformerLevel::Level1)); + + ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); + + std::map op_to_count = CountOpsInGraph(graph); + + ASSERT_EQ(op_to_count["Transpose"], 0); + ASSERT_EQ(op_to_count["Reshape"], 1); + } + + { + auto model_uri = MODEL_FOLDER "transpose_to_reshape_invalid.onnx"; + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model, nullptr, *logger_).IsOK()); + Graph& graph = p_model->MainGraph(); + + auto rule_transformer_L1 = std::make_unique("TransposeReplacement"); + ASSERT_STATUS_OK(rule_transformer_L1->Register(std::make_unique())); + onnxruntime::GraphTransformerManager graph_transformation_mgr{1}; + ASSERT_STATUS_OK(graph_transformation_mgr.Register(std::move(rule_transformer_L1), TransformerLevel::Level1)); + + ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); + + std::map op_to_count = CountOpsInGraph(graph); + + ASSERT_EQ(op_to_count["Transpose"], 1); + ASSERT_EQ(op_to_count["Reshape"], 0); + } +} + TEST_F(GraphTransformationTests, MegatronMLPPartitionRank0) { auto model_uri = MODEL_FOLDER "model_parallel/mlp_megatron_basic_test.onnx"; std::shared_ptr p_model;