mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
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:
parent
b5f242e978
commit
58f1d15d19
7 changed files with 190 additions and 0 deletions
41
onnxruntime/test/testdata/transform/transpose_graph_gen.py
vendored
Normal file
41
onnxruntime/test/testdata/transform/transpose_graph_gen.py
vendored
Normal 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)
|
||||
BIN
onnxruntime/test/testdata/transform/transpose_to_reshape_invalid.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/transpose_to_reshape_invalid.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/transpose_to_reshape_valid.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/transpose_to_reshape_valid.onnx
vendored
Normal file
Binary file not shown.
|
|
@ -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: {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Reference in a new issue