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.
This commit is contained in:
guyang3532 2023-10-27 23:50:18 +08:00 committed by GitHub
parent b5f242e978
commit 58f1d15d19
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
7 changed files with 190 additions and 0 deletions

View file

@ -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)

Binary file not shown.

View file

@ -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<std::unique_ptr<GraphTransformer>> GeneratePreTrainingTransformers(
std::make_unique<RuleBasedGraphTransformer>(optimizer_utils::GenerateRuleBasedTransformerName(level),
compatible_eps);
ORT_THROW_IF_ERROR(rule_transformer->Register(std::make_unique<ConcatReplacement>()));
ORT_THROW_IF_ERROR(rule_transformer->Register(std::make_unique<TransposeReplacement>()));
} break;
case TransformerLevel::Level3: {

View file

@ -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<int64_t>(transpose_node.GetAttributes().at("perm"));
InlinedVector<int64_t> new_shape;
new_shape.reserve(perm.size());
int64_t last_permuted_axis = 0;
for (int i = 0; i < static_cast<int>(perm.size()); ++i) {
if (!input_shape->dim(static_cast<int>(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<int>(perm[i])).dim_value());
if (input_shape->dim(static_cast<int>(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<int64_t>(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

View file

@ -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<std::string> 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

View file

@ -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<Model> 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<RuleBasedGraphTransformer>("TransposeReplacement");
ASSERT_STATUS_OK(rule_transformer_L1->Register(std::make_unique<TransposeReplacement>()));
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<std::string, int> 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<Model> 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<RuleBasedGraphTransformer>("TransposeReplacement");
ASSERT_STATUS_OK(rule_transformer_L1->Register(std::make_unique<TransposeReplacement>()));
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<std::string, int> 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<Model> p_model;