diff --git a/onnxruntime/core/optimizer/nchwc_transformer.cc b/onnxruntime/core/optimizer/nchwc_transformer.cc index a4a6716344..c3e39bb6e5 100644 --- a/onnxruntime/core/optimizer/nchwc_transformer.cc +++ b/onnxruntime/core/optimizer/nchwc_transformer.cc @@ -465,7 +465,11 @@ void NchwcTransformerImpl::TransformPool(Node& node) { const size_t nchwc_block_size = MlasNchwcGetBlockSize(); - auto* input_shape = input_defs[0]->Shape(); + const auto* input_type = input_defs[0]->TypeAsProto(); + if ((input_type == nullptr) || (input_type->tensor_type().elem_type() != TensorProto_DataType_FLOAT)) { + return; + } + const auto* input_shape = input_defs[0]->Shape(); if ((input_shape == nullptr) || (input_shape->dim_size() != 4)) { return; } @@ -927,9 +931,7 @@ void NchwcTransformerImpl::Transform(Node& node) { graph_utils::IsSupportedOptypeVersionAndDomain(node, "FusedConv", {1}, kMSDomain)) { TransformConv(node); } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "MaxPool", {1, 8, 10, 11, 12}) || - graph_utils::IsSupportedOptypeVersionAndDomain(node, "AveragePool", {1, 7, 10, 11}) || - graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalMaxPool", {1}) || - graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalAveragePool", {1})) { + graph_utils::IsSupportedOptypeVersionAndDomain(node, "AveragePool", {1, 7, 10, 11})) { TransformPool(node); } else if (node.GetInputEdgesCount() == 0 && node.InputDefs().size() != 0) { // The following transforms only run when the input edge count has already @@ -955,6 +957,10 @@ void NchwcTransformerImpl::Transform(Node& node) { } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "Upsample", {9, 13}) || graph_utils::IsSupportedOptypeVersionAndDomain(node, "Resize", {10, 11, 13})) { TransformResize(node); + } else if (graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalMaxPool", {1}) || + graph_utils::IsSupportedOptypeVersionAndDomain(node, "GlobalAveragePool", {1})) { + // Convert these pooling types only if the input is already in NCHWc format. + TransformPool(node); } } diff --git a/onnxruntime/test/optimizer/nchwc_optimizer_test.cc b/onnxruntime/test/optimizer/nchwc_optimizer_test.cc index 64ba6f5e2a..4b9f7def1a 100644 --- a/onnxruntime/test/optimizer/nchwc_optimizer_test.cc +++ b/onnxruntime/test/optimizer/nchwc_optimizer_test.cc @@ -44,6 +44,7 @@ struct NchwcTestHelper { NchwcTestHelper(Graph& graph) : graph_(graph), fill_value_(0), per_sample_tolerance_(0.0) { } + template NodeArg* MakeInput(const std::vector& shape, const ONNX_NAMESPACE::TypeProto& type_proto) { int64_t num_elements = 1; for (auto& dim : shape) { @@ -51,23 +52,24 @@ struct NchwcTestHelper { } OrtValue input_value; - CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), shape, - FillRandomData(static_cast(num_elements)), &input_value); + CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), shape, + FillRandomData(static_cast(num_elements)), &input_value); std::string name = graph_.GenerateNodeArgName("input"); feeds_.insert(std::make_pair(name, input_value)); return &graph_.GetOrCreateNodeArg(name, &type_proto); } + template NodeArg* MakeInput(const std::vector& shape) { ONNX_NAMESPACE::TypeProto type_proto; - type_proto.mutable_tensor_type()->set_elem_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + type_proto.mutable_tensor_type()->set_elem_type(utils::ToTensorProtoElementType()); for (auto& dim : shape) { type_proto.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(dim); } - return MakeInput(shape, type_proto); + return MakeInput(shape, type_proto); } NodeArg* MakeOutput() { @@ -101,7 +103,7 @@ struct NchwcTestHelper { NodeArg* MakeInitializer(const std::vector& shape) { int64_t num_elements = std::accumulate(shape.begin(), shape.end(), int64_t(1), std::multiplies{}); - return MakeInitializer(shape, FillRandomData(static_cast(num_elements))); + return MakeInitializer(shape, FillRandomData(static_cast(num_elements))); } NodeArg* Make1DInitializer(const std::vector& data) { @@ -161,14 +163,15 @@ struct NchwcTestHelper { return AddTransposeNode(input_arg, output_arg, {1, 0, 2, 3}); } - std::vector FillRandomData(size_t count) { + template + std::vector FillRandomData(size_t count) { constexpr int min_fill_value = -23; constexpr int max_fill_value = 23; - std::vector random_data; + std::vector random_data; random_data.resize(count); for (size_t n = 0; n < count; n++) { - random_data[n] = static_cast(fill_value_); + random_data[n] = static_cast(fill_value_); fill_value_++; if (fill_value_ == max_fill_value) { fill_value_ = min_fill_value; @@ -186,7 +189,7 @@ struct NchwcTestHelper { void NchwcOptimizerTester(const std::function& build_test_case, const std::function& check_nchwc_graph, - int opset_version = 11) { + int opset_version = 12) { // Ignore the test if NCHWc is not supported by the platform. if (MlasNchwcGetBlockSize() <= 1) { return; @@ -251,7 +254,7 @@ void NchwcOptimizerTester(const std::function& bu TEST(NchwcOptimizerTests, ConvNchw) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({16, 3, 112, 112}); + auto* input_arg = helper.MakeInput({16, 3, 112, 112}); auto* output_arg = helper.MakeOutput(); auto* conv_output_arg = output_arg; @@ -291,7 +294,7 @@ TEST(NchwcOptimizerTests, ConvNchw) { TEST(NchwcOptimizerTests, ConvNchwc) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({16, 64, 28, 28}); + auto* input_arg = helper.MakeInput({16, 64, 28, 28}); auto* output_arg = helper.MakeOutput(); auto* conv_output_arg = output_arg; @@ -329,7 +332,7 @@ TEST(NchwcOptimizerTests, ConvNchwc) { TEST(NchwcOptimizerTests, ConvNchwcGrouped) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({16, 48, 28, 28}); + auto* input_arg = helper.MakeInput({16, 48, 28, 28}); auto* output_arg = helper.MakeOutput(); auto* conv_output_arg = output_arg; @@ -364,7 +367,7 @@ TEST(NchwcOptimizerTests, ConvNchwcGrouped) { TEST(NchwcOptimizerTests, ConvDepthwise) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({16, 96, 28, 28}); + auto* input_arg = helper.MakeInput({16, 96, 28, 28}); auto* output_arg = helper.MakeOutput(); auto* conv_output_arg = output_arg; @@ -399,7 +402,7 @@ TEST(NchwcOptimizerTests, ConvDepthwise) { TEST(NchwcOptimizerTests, ConvPointwise) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({16, 64, 28, 42}); + auto* input_arg = helper.MakeInput({16, 64, 28, 42}); auto* output_arg = helper.MakeOutput(); auto* conv_output_arg = output_arg; @@ -432,7 +435,7 @@ TEST(NchwcOptimizerTests, ConvPointwise) { TEST(NchwcOptimizerTests, ConvMaxPool) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 34, 34}); + auto* input_arg = helper.MakeInput({1, 48, 34, 34}); auto* conv_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -456,7 +459,7 @@ TEST(NchwcOptimizerTests, ConvMaxPool) { TEST(NchwcOptimizerTests, ConvMaxPoolDilations) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 66, 77}); + auto* input_arg = helper.MakeInput({1, 48, 66, 77}); auto* conv_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -481,7 +484,7 @@ TEST(NchwcOptimizerTests, ConvMaxPoolDilations) { TEST(NchwcOptimizerTests, ConvAveragePool) { auto test_case = [&](bool count_include_pad) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 34, 34}); + auto* input_arg = helper.MakeInput({1, 48, 34, 34}); auto* conv_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -513,7 +516,7 @@ TEST(NchwcOptimizerTests, ConvAveragePool) { TEST(NchwcOptimizerTests, ConvGlobalPool) { auto test_case = [&](const std::string& op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 96, 54, 54}); + auto* input_arg = helper.MakeInput({1, 96, 54, 54}); auto* conv_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -543,7 +546,7 @@ TEST(NchwcOptimizerTests, ConvGlobalPool) { TEST(NchwcOptimizerTests, ConvAddFusion) { auto test_case = [&](const std::string& op_type, int opset_version, bool do_relu) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 32, 28, 28}); + auto* input_arg = helper.MakeInput({1, 32, 28, 28}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* conv2_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -575,7 +578,7 @@ TEST(NchwcOptimizerTests, ConvAddFusion) { // Verify that Add or Sum can be fused into a preceding NCHWc Conv node, // with an optional Relu node following. std::vector op_types{"Add", "Sum"}; - static const int opset_versions[] = {7, 10, 11}; + static const int opset_versions[] = {7, 10, 11, 12}; for (auto& op_type : op_types) { for (auto opset_version : opset_versions) { test_case(op_type, opset_version, false); @@ -586,7 +589,7 @@ TEST(NchwcOptimizerTests, ConvAddFusion) { TEST(NchwcOptimizerTests, ConvNoBiasAddFusion) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 32, 28, 28}); + auto* input_arg = helper.MakeInput({1, 32, 28, 28}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* conv2_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -612,7 +615,7 @@ TEST(NchwcOptimizerTests, ConvNoBiasAddFusion) { TEST(NchwcOptimizerTests, FusedConvAddFusion) { auto test_case = [&](bool do_relu1, bool do_relu2, int add_count) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 32, 28, 28}); + auto* input_arg = helper.MakeInput({1, 32, 28, 28}); auto* add1_input_arg = helper.MakeIntermediate(); auto* add2_input_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -659,7 +662,7 @@ TEST(NchwcOptimizerTests, FusedConvAddFusion) { TEST(NchwcOptimizerTests, ConvBinary) { auto test_case = [&](const std::string& op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 32, 23, 23}); + auto* input_arg = helper.MakeInput({1, 32, 23, 23}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* conv2_output_arg = helper.MakeIntermediate(); auto* relu1_output_arg = helper.MakeIntermediate(); @@ -697,7 +700,7 @@ TEST(NchwcOptimizerTests, ConvBinary) { TEST(NchwcOptimizerTests, ConvConcat) { auto test_case = [&](int axis, int channel_count, int reorder_output_count) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 17, 34}); + auto* input_arg = helper.MakeInput({1, 48, 17, 34}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* conv2_output_arg = helper.MakeIntermediate(); auto* conv3_output_arg = helper.MakeIntermediate(); @@ -733,7 +736,7 @@ TEST(NchwcOptimizerTests, ConvConcat) { TEST(NchwcOptimizerTests, ConvReuseWeightsOIHWBiBo) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 64, 7, 7}); + auto* input_arg = helper.MakeInput({1, 64, 7, 7}); auto* output1_arg = helper.MakeOutput(); auto* output2_arg = helper.MakeOutput(); auto* output3_arg = helper.MakeOutput(); @@ -774,10 +777,10 @@ TEST(NchwcOptimizerTests, ConvReuseWeightsOIHWBiBo) { TEST(NchwcOptimizerTests, ConvReuseWeightsOIHWBo) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input1_arg = helper.MakeInput({1, 64, 7, 7}); - auto* input2_arg = helper.MakeInput({1, 64, 7, 7}); - auto* input3_arg = helper.MakeInput({1, 1, 7, 7}); - auto* input4_arg = helper.MakeInput({1, 1, 7, 7}); + auto* input1_arg = helper.MakeInput({1, 64, 7, 7}); + auto* input2_arg = helper.MakeInput({1, 64, 7, 7}); + auto* input3_arg = helper.MakeInput({1, 1, 7, 7}); + auto* input4_arg = helper.MakeInput({1, 1, 7, 7}); auto* output1_arg = helper.MakeOutput(); auto* output2_arg = helper.MakeOutput(); auto* output3_arg = helper.MakeOutput(); @@ -831,7 +834,7 @@ TEST(NchwcOptimizerTests, ShapeInferencing) { type_proto.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_param("input_height"); type_proto.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_param("input_width"); - auto* input_arg = helper.MakeInput({1, 3, 50, 100}, type_proto); + auto* input_arg = helper.MakeInput({1, 3, 50, 100}, type_proto); auto* output_arg = helper.MakeOutput(); // With these padding and kernel arguments, the shape along each spatial @@ -887,7 +890,7 @@ TEST(NchwcOptimizerTests, ShapeInferencing2) { type_proto.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_param("input_height"); type_proto.mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_param("input_width"); - auto* input_arg = helper.MakeInput({1, 1, 49, 98}, type_proto); + auto* input_arg = helper.MakeInput({1, 1, 49, 98}, type_proto); auto* output_arg = helper.MakeOutput(); auto* conv1_output_arg = helper.MakeIntermediate(); @@ -925,7 +928,7 @@ TEST(NchwcOptimizerTests, ShapeInferencing2) { TEST(NchwcOptimizerTests, MixedOutputUsage) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({6, 5, 11, 11}); + auto* input_arg = helper.MakeInput({6, 5, 11, 11}); auto* output_arg = helper.MakeOutput(); auto* conv1_output_arg = helper.MakeIntermediate(); @@ -957,24 +960,24 @@ TEST(NchwcOptimizerTests, MixedOutputUsage) { TEST(NchwcOptimizerTests, TensorAlignment) { auto build_test_case = [&](NchwcTestHelper& helper) { // Input channel count must currently be a multiple of the NCHWc block size. - auto* input1_arg = helper.MakeInput({1, 60, 28, 42}); + auto* input1_arg = helper.MakeInput({1, 60, 28, 42}); auto* output1_arg = helper.MakeOutput(); helper.AddConvNode(input1_arg, output1_arg, {128, 60, 1, 1}); // Grouped input channel count must be a multiple of the NCHWc block size. - auto* input2_arg = helper.MakeInput({1, 48, 28, 42}); + auto* input2_arg = helper.MakeInput({1, 48, 28, 42}); auto* output2_arg = helper.MakeOutput(); auto& conv2_node = helper.AddConvNode(input2_arg, output2_arg, {128, 12, 3, 3}); conv2_node.AddAttribute("group", static_cast(4)); // Grouped output channel count must be a multiple of the NCHWc block size. - auto* input3_arg = helper.MakeInput({1, 64, 28, 42}); + auto* input3_arg = helper.MakeInput({1, 64, 28, 42}); auto* output3_arg = helper.MakeOutput(); auto& conv3_node = helper.AddConvNode(input3_arg, output3_arg, {48, 16, 3, 3}); conv3_node.AddAttribute("group", static_cast(4)); // Channel count must currently be a multiple of the NCHWc block size. - auto* input4_arg = helper.MakeInput({1, 60, 12, 12}); + auto* input4_arg = helper.MakeInput({1, 60, 12, 12}); auto* output4_arg = helper.MakeOutput(); auto& pool_node = helper.AddNode("MaxPool", {input4_arg}, {output4_arg}); pool_node.AddAttribute("kernel_shape", std::vector{2, 2}); @@ -996,7 +999,7 @@ TEST(NchwcOptimizerTests, TensorAlignment) { TEST(NchwcOptimizerTests, IntermediatesAsGraphOutputs) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 34, 34}); + auto* input_arg = helper.MakeInput({1, 48, 34, 34}); auto* conv_output_arg = helper.MakeOutput(); auto* output_arg = helper.MakeOutput(); @@ -1028,7 +1031,7 @@ TEST(NchwcOptimizerTests, IntermediatesAsGraphOutputs) { TEST(NchwcOptimizerTests, BatchNormalization) { auto test_case = [&](bool training_outputs) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 1, 23, 21}); + auto* input_arg = helper.MakeInput({1, 1, 23, 21}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* conv2_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -1104,7 +1107,7 @@ TEST(NchwcOptimizerTests, BatchNormalization) { TEST(NchwcOptimizerTests, ConvReorderOutputNhwc) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 64, 28, 32}); + auto* input_arg = helper.MakeInput({1, 64, 28, 32}); auto* conv_output_arg = helper.MakeIntermediate(); auto* nhwc_output_arg = helper.MakeOutput(); @@ -1126,7 +1129,7 @@ TEST(NchwcOptimizerTests, ConvReorderOutputNhwc) { TEST(NchwcOptimizerTests, ConvReorderOutputBoth) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({5, 64, 33, 37}); + auto* input_arg = helper.MakeInput({5, 64, 33, 37}); auto* conv_output_arg = helper.MakeIntermediate(); auto* nchw_output_arg = helper.MakeOutput(); auto* nhwc_output_arg = helper.MakeOutput(); @@ -1151,7 +1154,7 @@ TEST(NchwcOptimizerTests, ConvReorderOutputBoth) { TEST(NchwcOptimizerTests, ConvReorderOutputCnhw) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 64, 28, 32}); + auto* input_arg = helper.MakeInput({1, 64, 28, 32}); auto* conv_output_arg = helper.MakeIntermediate(); auto* nhwc_output_arg = helper.MakeOutput(); @@ -1174,7 +1177,7 @@ TEST(NchwcOptimizerTests, ConvReorderOutputCnhw) { TEST(NchwcOptimizerTests, Upsample) { auto test_case = [&](int opset_version, float scale_h, float scale_w) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({3, 16, 27, 15}); + auto* input_arg = helper.MakeInput({3, 16, 27, 15}); auto* conv_output_arg = helper.MakeIntermediate(); auto* output_arg = helper.MakeOutput(); @@ -1211,7 +1214,7 @@ TEST(NchwcOptimizerTests, Upsample) { // Verify that upsample nodes can be converted to the NCHWc format for // various versions of the operator. - static const int opset_versions[] = {9, 10, 11}; + static const int opset_versions[] = {9, 10, 11, 12}; for (auto opset_version : opset_versions) { test_case(opset_version, 1.f, 1.f); test_case(opset_version, 2.f, 2.f); @@ -1222,7 +1225,7 @@ TEST(NchwcOptimizerTests, Upsample) { TEST(NchwcOptimizerTests, Activation) { auto test_case = [&](const std::string& activation_op_type) { auto build_test_case = [&](NchwcTestHelper& helper) { - auto* input_arg = helper.MakeInput({1, 48, 11, 15}); + auto* input_arg = helper.MakeInput({1, 48, 11, 15}); auto* conv1_output_arg = helper.MakeIntermediate(); auto* activation_output_arg = helper.MakeIntermediate(); auto* mul_output_arg = helper.MakeIntermediate(); @@ -1254,6 +1257,32 @@ TEST(NchwcOptimizerTests, Activation) { } } +TEST(NchwcOptimizerTests, MaxPoolTypeCheck) { + auto build_test_case = [&](NchwcTestHelper& helper) { + auto add_pool_node = [&](NchwcTestHelper& helper, NodeArg* input_arg) { + auto* output_arg = helper.MakeOutput(); + auto& pool_node = helper.AddNode("MaxPool", {input_arg}, {output_arg}); + pool_node.AddAttribute("pads", std::vector{0, 0, 0, 0}); + pool_node.AddAttribute("kernel_shape", std::vector{3, 3}); + }; + + const std::vector input_shape{1, 32, 13, 13}; + add_pool_node(helper, helper.MakeInput(input_shape)); + add_pool_node(helper, helper.MakeInput(input_shape)); + }; + + auto check_nchwc_graph = [&](NchwcInferenceSession& session) { + auto op_to_count = session.CountOpsInGraph(); + EXPECT_EQ(op_to_count["MaxPool"], 1); + EXPECT_EQ(op_to_count["nchwc.MaxPool"], 1); + EXPECT_EQ(op_to_count["nchwc.ReorderInput"], 1); + EXPECT_EQ(op_to_count["nchwc.ReorderOutput"], 1); + }; + + // Verify that the optimizer checks the type of the MaxPool node. + NchwcOptimizerTester(build_test_case, check_nchwc_graph, 12); +} + #endif } // namespace test