fix axis of layernorm for UpstreamReshape (#18425)

Similar to https://github.com/microsoft/onnxruntime/pull/17255
update axis for Layernormalization when Reshape upstream it.
This commit is contained in:
guyang3532 2023-11-16 16:29:00 +08:00 committed by GitHub
parent 18a3675bf7
commit 751aa8d31a
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 155 additions and 169 deletions

View file

@ -338,8 +338,8 @@ std::optional<SliceInfo> IsSupportedGather(Graph& graph, Node& node,
auto axis = static_cast<int>(node.GetAttributes().at("axis").i());
axis = axis < 0 ? axis + data_rank : axis;
size_t dim_size = static_cast<size_t>(indices_shape->dim_size());
bool is_single_value_1d_tensor = dim_size != 0 && (dim_size == 1 && utils::HasDimValue(indices_shape->dim(0)) &&
indices_shape->dim(0).dim_value() == 1);
bool is_single_value_1d_tensor = dim_size == 1 && utils::HasDimValue(indices_shape->dim(0)) &&
indices_shape->dim(0).dim_value() == 1;
if (dim_size != 0 && !is_single_value_1d_tensor) {
if (dim_size == 1 && utils::HasDimValue(data_shape->dim(axis)) &&
data_shape->dim(axis).dim_value() > indices_shape->dim(0).dim_value()) {

View file

@ -3,6 +3,7 @@
#ifdef ENABLE_TRAINING
#include <onnx/defs/attr_proto_util.h>
#include "core/optimizer/utils.h"
#include "core/optimizer/compute_optimizer/upstream_reshape_actors.h"
@ -282,6 +283,23 @@ bool LayerNormalizationReshapeActor::PreCheck(
return propagate_input_indices.size() > 0;
}
bool LayerNormalizationReshapeActor::PostProcess(
Graph& /* graph */, Node& current_node, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) {
auto axis = static_cast<int64_t>(current_node.GetAttributes().at("axis").i());
// When Reshape(from 3D to 2D, with the first two dimensions be merged) upstream a LayerNormalization,
// The axis attribute of LayerNormalization should be decreased by 1 if it is greater than 1.
if (axis > 1) {
auto new_axis = axis - 1;
auto& attributes = current_node.GetMutableAttributes();
attributes["axis"] = ONNX_NAMESPACE::MakeAttribute("axis", static_cast<int64_t>(new_axis));
}
return true;
}
template class SimplePointwiseReshapeActor<true>;
template class SimplePointwiseReshapeActor<false>;

View file

@ -111,13 +111,11 @@ class UpStreamReshapeOperatorActorBase : public UpStreamOperatorActorBase {
* So far, we don't have requirements to override PostProcess function.
*/
bool PostProcess(Graph& /* graph */, Node& /* current_node */, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) {
return true;
}
virtual bool PostProcess(Graph& /* graph */, Node& /* current_node */, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) = 0;
};
// The inputs are broad-cast-able. The outputs should have the same shape (fully broadcasted shape)
@ -133,6 +131,14 @@ class SimplePointwiseReshapeActor : public UpStreamReshapeOperatorActorBase {
std::vector<int>& propagate_input_indices,
std::unordered_map<int, std::vector<DimCompare>>& all_input_cmp_rets,
std::function<void(Node& node)>& shape_update_func) override;
bool PostProcess(Graph& /* graph */, Node& /* current_node */, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) override {
return true;
}
};
class MatMulReshapeActor : public UpStreamReshapeOperatorActorBase {
@ -145,6 +151,14 @@ class MatMulReshapeActor : public UpStreamReshapeOperatorActorBase {
std::vector<int>& propagate_input_indices,
std::unordered_map<int, std::vector<DimCompare>>& all_input_cmp_rets,
std::function<void(Node& node)>& shape_update_func) override;
bool PostProcess(Graph& /* graph */, Node& /* current_node */, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) override {
return true;
}
};
class LayerNormalizationReshapeActor : public UpStreamReshapeOperatorActorBase {
@ -157,6 +171,12 @@ class LayerNormalizationReshapeActor : public UpStreamReshapeOperatorActorBase {
std::vector<int>& propagate_input_indices,
std::unordered_map<int, std::vector<DimCompare>>& all_input_cmp_rets,
std::function<void(Node& node)>& shape_update_func) override;
bool PostProcess(Graph& /* graph */, Node& current_node, const ReshapeInfo& /* info_without_node */,
const logging::Logger& /* logger */,
std::vector<int>& /* propagate_input_indices */,
const std::unordered_map<int, std::vector<DimCompare>>& /* all_input_cmp_rets */,
const std::unordered_map<int, ReshapeInfo>& /* new_reshape_infos */) override;
};
/**

View file

@ -847,7 +847,7 @@ Test graph includes multiple equivalent subgraphs as below.
Add an Identity node because currently, we don't allow Gather generates graph output.
*/
TEST(ComputeOptimizerTests, GatherLayerNormalization) {
std::vector<std::tuple<int, int64_t, int64_t, bool>> test_config_pairs{
std::vector<std::tuple<bool, int64_t, int64_t, bool>> test_config_pairs{
// {
// is_scalar_slice,
// ln_axis_before_propagation,
@ -929,13 +929,6 @@ TEST(ComputeOptimizerTests, GatherLayerNormalization) {
const ONNX_NAMESPACE::TensorShapeProto* slice_out_shape = producer_node->OutputDefs()[0]->Shape();
TEST_RETURN_IF_NOT(slice_out_shape != nullptr);
auto& attrs = node.GetAttributes();
TEST_RETURN_IF_NOT(attrs.find("axis") != attrs.end());
auto& axis_attr = attrs.at("axis");
auto axis_value = (int)axis_attr.i();
TEST_RETURN_IF_NOT(axis_value == ln_axis_after);
if (is_scalar_slice) {
TEST_RETURN_IF_NOT(slice_out_shape->dim_size() == 2);
TEST_RETURN_IF_NOT(utils::HasDimValue(slice_out_shape->dim(0)) &&
@ -951,10 +944,15 @@ TEST(ComputeOptimizerTests, GatherLayerNormalization) {
TEST_RETURN_IF_NOT(utils::HasDimValue(slice_out_shape->dim(2)) &&
slice_out_shape->dim(2).dim_value() == 256);
}
} else {
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
auto& attrs = node.GetAttributes();
TEST_RETURN_IF_NOT(attrs.find("axis") != attrs.end());
auto& axis_attr = attrs.at("axis");
auto axis_value = (int)axis_attr.i();
TEST_RETURN_IF_NOT(axis_value == ln_axis_after);
}
}
@ -2841,165 +2839,110 @@ Test graph include multiple equivalent subgraphs as below.
Add an Identity node because currently we don't allow Reshape generate graph output.
*/
TEST(ComputeOptimizerTests, ReshapeLayerNormalization_PropagationOnOneBranch) {
const logging::Logger* logger = &logging::LoggingManager::DefaultLogger();
auto pre_graph_checker = [](Graph& graph) -> Status {
auto op_count_pre = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_pre.size() == 3U);
TEST_RETURN_IF_NOT(op_count_pre["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Identity"] == 1);
return Status::OK();
TEST(ComputeOptimizerTests, ReshapeLayerNormalization) {
std::vector<std::tuple<int64_t, int64_t, bool>> test_config_pairs{
// {
// ln_axis_before_propagation,
// expected_ln_axis_after_propagation,
// expected to propagate
// }
{0, 0, false},
{1, 1, false},
{2, 1, true},
{-3, -3, false},
{-2, -2, false},
{-1, -1, true},
};
auto post_graph_checker = [](Graph& graph) {
auto op_count_post = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_post.size() == 3U);
TEST_RETURN_IF_NOT(op_count_post["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Identity"] == 1);
for (auto p : test_config_pairs) {
int64_t ln_axis_before = std::get<0>(p);
int64_t ln_axis_after = std::get<1>(p);
bool expected_to_propagate = std::get<2>(p);
for (Node& node : graph.Nodes()) {
if (node.OpType() == "LayerNormalization") {
const auto& input_defs = node.InputDefs();
{
auto producer_node = graph.GetProducerNode(input_defs[0]->Name());
TEST_RETURN_IF_NOT(producer_node != nullptr);
TEST_RETURN_IF_NOT(producer_node->OpType() == "Reshape");
InlinedVector<int64_t> values;
constexpr bool require_constant = true;
NodeArg* initializer_node_arg = graph.GetNodeArg(producer_node->InputDefs()[1]->Name());
TEST_RETURN_IF_NOT(optimizer_utils::AppendTensorFromInitializer(graph, *initializer_node_arg, values, require_constant));
TEST_RETURN_IF_NOT(values.size() == 2);
TEST_RETURN_IF_NOT(values[0] == -1);
TEST_RETURN_IF_NOT(values[1] == 1024);
}
{
auto producer_node = graph.GetProducerNode(input_defs[1]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
{
auto producer_node = graph.GetProducerNode(input_defs[2]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
}
}
return Status::OK();
};
std::vector<int> fist_dim_values = {-1, 128};
for (auto first_dim_value : fist_dim_values) {
auto build_test_case = [&first_dim_value](ModelTestBuilder& builder) {
auto* input1_arg = builder.MakeInput<float>({{4, 32, 1024}});
auto* input2_arg = builder.MakeInput<float>({{1024}});
auto* input3_arg = builder.MakeInput<float>({{1024}});
auto* ln_out = builder.MakeIntermediate();
builder.AddNode("LayerNormalization", {input1_arg, input2_arg, input3_arg}, {ln_out})
.AddAttribute("axis", static_cast<int64_t>(-1));
auto* shape_initializer = builder.MakeInitializer<int64_t>({2}, {first_dim_value, 1024});
auto* reshape_out = builder.MakeIntermediate();
builder.AddNode("Reshape", {ln_out, shape_initializer}, {reshape_out});
auto* identity_out = builder.MakeOutput();
builder.AddNode("Identity", {reshape_out}, {identity_out});
const logging::Logger* logger = &logging::LoggingManager::DefaultLogger();
auto pre_graph_checker = [](Graph& graph) -> Status {
auto op_count_pre = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_pre.size() == 3U);
TEST_RETURN_IF_NOT(op_count_pre["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Identity"] == 1);
return Status::OK();
};
const std::vector<int> opsets{12, 13, 14};
for (auto& opset_version : opsets) {
std::unique_ptr<GraphTransformer> transformer = std::make_unique<UpStreamReshapeGraphTransformer>();
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset_version, *logger, std::move(transformer),
TransformerLevel::Level1,
1, pre_graph_checker, post_graph_checker));
}
}
}
auto post_graph_checker = [ln_axis_after, expected_to_propagate](Graph& graph) {
auto op_count_post = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_post.size() == 3U);
TEST_RETURN_IF_NOT(op_count_post["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Identity"] == 1);
/*
Test graph include multiple equivalent subgraphs as below.
graph input [4, 32, 1024] (float) graph input [1024] (float) graph input [1024] (float)
| | /
\_____________ _______/ __________________________/
\ / /
LayerNormalization
|
Reshape
|
Identity
|
graph out [128, 1024] (float)
for (Node& node : graph.Nodes()) {
if (node.OpType() == "LayerNormalization") {
const auto& input_defs = node.InputDefs();
Add an Identity node because currently we don't allow Reshape generate graph output.
*/
TEST(ComputeOptimizerTests, ReshapeLayerNormalization_NoPropagation) {
const logging::Logger* logger = &logging::LoggingManager::DefaultLogger();
auto pre_graph_checker = [](Graph& graph) -> Status {
auto op_count_pre = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_pre.size() == 3U);
TEST_RETURN_IF_NOT(op_count_pre["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_pre["Identity"] == 1);
return Status::OK();
};
if (expected_to_propagate) {
auto producer_node = graph.GetProducerNode(input_defs[0]->Name());
TEST_RETURN_IF_NOT(producer_node != nullptr);
TEST_RETURN_IF_NOT(producer_node->OpType() == "Reshape");
auto post_graph_checker = [](Graph& graph) {
auto op_count_post = CountOpsInGraph(graph);
TEST_RETURN_IF_NOT(op_count_post.size() == 3U);
TEST_RETURN_IF_NOT(op_count_post["LayerNormalization"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Reshape"] == 1);
TEST_RETURN_IF_NOT(op_count_post["Identity"] == 1);
InlinedVector<int64_t> values;
constexpr bool require_constant = true;
NodeArg* initializer_node_arg = graph.GetNodeArg(producer_node->InputDefs()[1]->Name());
TEST_RETURN_IF_NOT(optimizer_utils::AppendTensorFromInitializer(graph, *initializer_node_arg, values, require_constant));
TEST_RETURN_IF_NOT(values.size() == 2);
TEST_RETURN_IF_NOT(values[0] == -1);
TEST_RETURN_IF_NOT(values[1] == 1024);
} else {
auto producer_node = graph.GetProducerNode(input_defs[0]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
for (Node& node : graph.Nodes()) {
if (node.OpType() == "LayerNormalization") {
const auto& input_defs = node.InputDefs();
{
auto producer_node = graph.GetProducerNode(input_defs[1]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
{
auto producer_node = graph.GetProducerNode(input_defs[0]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
{
auto producer_node = graph.GetProducerNode(input_defs[2]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
{
auto producer_node = graph.GetProducerNode(input_defs[1]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
}
auto& attrs = node.GetAttributes();
TEST_RETURN_IF_NOT(attrs.find("axis") != attrs.end());
{
auto producer_node = graph.GetProducerNode(input_defs[2]->Name());
TEST_RETURN_IF_NOT(producer_node == nullptr);
auto& axis_attr = attrs.at("axis");
auto axis_value = (int)axis_attr.i();
TEST_RETURN_IF_NOT(axis_value == ln_axis_after);
}
}
}
return Status::OK();
};
std::vector<int> fist_dim_values = {-1, 128};
for (auto first_dim_value : fist_dim_values) {
auto build_test_case = [&first_dim_value](ModelTestBuilder& builder) {
auto* input1_arg = builder.MakeInput<float>({{4, 32, 1024}});
auto* input2_arg = builder.MakeInput<float>({{1024}});
auto* input3_arg = builder.MakeInput<float>({{1024}});
auto* ln_out = builder.MakeIntermediate();
builder.AddNode("LayerNormalization", {input1_arg, input2_arg, input3_arg}, {ln_out})
.AddAttribute("axis", static_cast<int64_t>(1));
auto* shape_initializer = builder.MakeInitializer<int64_t>({2}, {first_dim_value, 1024});
auto* reshape_out = builder.MakeIntermediate();
builder.AddNode("Reshape", {ln_out, shape_initializer}, {reshape_out});
auto* identity_out = builder.MakeOutput();
builder.AddNode("Identity", {reshape_out}, {identity_out});
return Status::OK();
};
const std::vector<int> opsets{12, 13, 14};
for (auto& opset_version : opsets) {
std::unique_ptr<GraphTransformer> transformer = std::make_unique<UpStreamReshapeGraphTransformer>();
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset_version, *logger, std::move(transformer),
TransformerLevel::Level1,
1, pre_graph_checker, post_graph_checker));
std::vector<int> fist_dim_values = {-1, 128};
for (auto first_dim_value : fist_dim_values) {
auto build_test_case = [ln_axis_before, &first_dim_value](ModelTestBuilder& builder) {
auto* input1_arg = builder.MakeInput<float>({{4, 32, 1024}});
auto* input2_arg = builder.MakeInput<float>({{1024}});
auto* input3_arg = builder.MakeInput<float>({{1024}});
auto* ln_out = builder.MakeIntermediate();
builder.AddNode("LayerNormalization", {input1_arg, input2_arg, input3_arg}, {ln_out})
.AddAttribute("axis", ln_axis_before);
auto* shape_initializer = builder.MakeInitializer<int64_t>({2}, {first_dim_value, 1024});
auto* reshape_out = builder.MakeIntermediate();
builder.AddNode("Reshape", {ln_out, shape_initializer}, {reshape_out});
auto* identity_out = builder.MakeOutput();
builder.AddNode("Identity", {reshape_out}, {identity_out});
};
const std::vector<int> opsets{12, 13, 14};
for (auto& opset_version : opsets) {
std::unique_ptr<GraphTransformer> transformer = std::make_unique<UpStreamReshapeGraphTransformer>();
ASSERT_STATUS_OK(TestGraphTransformer(build_test_case, opset_version, *logger, std::move(transformer),
TransformerLevel::Level1,
1, pre_graph_checker, post_graph_checker));
}
}
}
}

View file

@ -5761,6 +5761,7 @@ def test_runtime_inspector_label_and_embed_sparsity_detection(embed_is_sparse, l
("MatMul", 1),
("Dropout", 0),
("LayerNormalization", 0),
("LayerNormalization", 1),
("Cast", 0),
("BiasGelu", 0),
("Gelu", 0),
@ -5773,12 +5774,18 @@ def test_ops_for_padding_elimination(test_cases):
test_op = test_cases[0]
case = test_cases[1]
vocab_size, hidden_size = 50265, 768
batch_size, max_seq_length = 8, 128
class ToyModel(torch.nn.Module):
def __init__(self, vocab_size, hidden_size, pad_token_id):
super().__init__()
self.word_embeddings = nn.Embedding(vocab_size, hidden_size, padding_idx=pad_token_id)
if test_op == "LayerNormalization":
self.LayerNorm = nn.LayerNorm(hidden_size, eps=1e-05)
if case == 0:
self.LayerNorm = nn.LayerNorm(hidden_size, eps=1e-05)
else:
self.LayerNorm = nn.LayerNorm([max_seq_length, hidden_size], eps=1e-05)
self.hidden_size = hidden_size
# test test_elementwise op for padding elimination
@ -5889,8 +5896,6 @@ def test_ops_for_padding_elimination(test_cases):
batched_inputs.append(torch.cat((input_id, padding)))
return torch.stack(batched_inputs)
vocab_size, hidden_size = 50265, 768
batch_size, max_seq_length = 8, 128
device = "cuda"
model = ORTModule(ToyModel(vocab_size, hidden_size, 1).to(device))
x = generate_inputs(batch_size, max_seq_length, vocab_size)
@ -5908,7 +5913,7 @@ def test_ops_for_padding_elimination(test_cases):
assert len([node.op_type for node in training_model.graph.node if node.op_type == "FlattenAndUnpad"]) == 3
else:
assert len([node.op_type for node in training_model.graph.node if node.op_type == "FlattenAndUnpad"]) == 2
gathergrad_node = next(node for node in training_model.graph.node if node.op_type == "PadAndUnflatten")
recover_pad_node = next(node for node in training_model.graph.node if node.op_type == "PadAndUnflatten")
def find_input_node_type(model, arg):
result = []
@ -5917,14 +5922,14 @@ def test_ops_for_padding_elimination(test_cases):
result.append(node)
return result[0].op_type if len(result) == 1 else None
gathergrad_input_optypes = [find_input_node_type(training_model, arg) for arg in gathergrad_node.input]
recover_pad_input_optypes = [find_input_node_type(training_model, arg) for arg in recover_pad_node.input]
if test_op == "Add" or test_op == "Mul" or test_op == "Sub":
assert test_op in gathergrad_input_optypes
assert test_op in recover_pad_input_optypes
else:
if case == 0:
assert test_op in gathergrad_input_optypes
assert test_op in recover_pad_input_optypes
else:
assert "ATen" in gathergrad_input_optypes
assert "ATen" in recover_pad_input_optypes
del os.environ["ORTMODULE_ENABLE_EMBEDDING_SPARSE_OPTIMIZER"]