mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
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:
parent
18a3675bf7
commit
751aa8d31a
5 changed files with 155 additions and 169 deletions
|
|
@ -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()) {
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
};
|
||||
|
||||
/**
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue