已合并
修改cleancode以及补充池化类算子ut #4888
SimonZzz创建于 5月15日
修改cleancode以及补充池化类算子ut #4888
已合并
共 11 个文件变更+2200-3
| @@ -62,7 +62,7 @@ static ge::graphStatus InferShape4AdaptiveAvgPool2d(gert::InferShapeContext* con | |||
| 62 | } | 62 | } |
| 63 | const int64_t* output_size = static_cast<const int64_t*>(output_size_ptr->GetData()); | 63 | const int64_t* output_size = static_cast<const int64_t*>(output_size_ptr->GetData()); |
| 64 | for (int i = 0; i < output_size_num; i++) { | 64 | for (int i = 0; i < output_size_num; i++) { |
| 65 | - y_shape->AppendDim((int64_t)output_size[i]); | 65 | + y_shape->AppendDim(static_cast<int64_t>(output_size[i])); |
| 66 | } | 66 | } |
| 67 | OP_LOGD(context->GetNodeName(), "runtime2.0 AdaptiveAvgPool2d infershape run success."); | 67 | OP_LOGD(context->GetNodeName(), "runtime2.0 AdaptiveAvgPool2d infershape run success."); |
| 68 | return ge::GRAPH_SUCCESS; | 68 | return ge::GRAPH_SUCCESS; |
| @@ -0,0 +1,252 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | +/*! | ||
| 12 | + * \file test_adaptive_avg_pool2d_infershape.cpp | ||
| 13 | + * \brief InferShape测试 - AdaptiveAvgPool2d | ||
| 14 | + */ | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | +using namespace std; | ||
| 25 | +using namespace ge; | ||
| 26 | + | ||
| 27 | +class AdaptiveAvgPool2dInfershape : public testing::Test { | ||
| 28 | +protected: | ||
| 29 | + static void SetUpTestCase() | ||
| 30 | + { | ||
| 31 | + std::cout << "AdaptiveAvgPool2dInfershape SetUp" << std::endl; | ||
| 32 | + } | ||
| 33 | + | ||
| 34 | + static void TearDownTestCase() | ||
| 35 | + { | ||
| 36 | + std::cout << "AdaptiveAvgPool2dInfershape TearDown" << std::endl; | ||
| 37 | + } | ||
| 38 | +}; | ||
| 39 | + | ||
| 40 | +// 正常用例: 4D 输入 + 合法 output_size | ||
| 41 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_4d_basic) | ||
| 42 | +{ | ||
| 43 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 44 | + "AdaptiveAvgPool2d", | ||
| 45 | + { | ||
| 46 | + {{{2, 3, 224, 224}, {2, 3, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 47 | + }, | ||
| 48 | + { | ||
| 49 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 50 | + }, | ||
| 51 | + { | ||
| 52 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7})}, | ||
| 53 | + }); | ||
| 54 | + std::vector<std::vector<int64_t>> expectOutputShape = {{2, 3, 7, 7}}; | ||
| 55 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 56 | +} | ||
| 57 | + | ||
| 58 | +// 正常用例: 3D 输入 + 合法 output_size | ||
| 59 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_3d_basic) | ||
| 60 | +{ | ||
| 61 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 62 | + "AdaptiveAvgPool2d", | ||
| 63 | + { | ||
| 64 | + {{{3, 224, 224}, {3, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 65 | + }, | ||
| 66 | + { | ||
| 67 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 68 | + }, | ||
| 69 | + { | ||
| 70 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1})}, | ||
| 71 | + }); | ||
| 72 | + std::vector<std::vector<int64_t>> expectOutputShape = {{3, 1, 1}}; | ||
| 73 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 74 | +} | ||
| 75 | + | ||
| 76 | +// 正常用例: FP16 数据类型 | ||
| 77 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_fp16) | ||
| 78 | +{ | ||
| 79 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 80 | + "AdaptiveAvgPool2d", | ||
| 81 | + { | ||
| 82 | + {{{1, 64, 7, 7}, {1, 64, 7, 7}}, ge::DT_FLOAT16, ge::FORMAT_ND}, | ||
| 83 | + }, | ||
| 84 | + { | ||
| 85 | + {{{}, {}}, ge::DT_FLOAT16, ge::FORMAT_ND}, | ||
| 86 | + }, | ||
| 87 | + { | ||
| 88 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1})}, | ||
| 89 | + }); | ||
| 90 | + std::vector<std::vector<int64_t>> expectOutputShape = {{1, 64, 1, 1}}; | ||
| 91 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 92 | +} | ||
| 93 | + | ||
| 94 | +// 正常用例: BF16 数据类型 | ||
| 95 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_bf16) | ||
| 96 | +{ | ||
| 97 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 98 | + "AdaptiveAvgPool2d", | ||
| 99 | + { | ||
| 100 | + {{{1, 32, 14, 14}, {1, 32, 14, 14}}, ge::DT_BF16, ge::FORMAT_ND}, | ||
| 101 | + }, | ||
| 102 | + { | ||
| 103 | + {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND}, | ||
| 104 | + }, | ||
| 105 | + { | ||
| 106 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7})}, | ||
| 107 | + }); | ||
| 108 | + std::vector<std::vector<int64_t>> expectOutputShape = {{1, 32, 7, 7}}; | ||
| 109 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape); | ||
| 110 | +} | ||
| 111 | + | ||
| 112 | +// 异常用例: output_size 为空 (size=0, !=2) | ||
| 113 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_invalid_output_size_empty) | ||
| 114 | +{ | ||
| 115 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 116 | + "AdaptiveAvgPool2d", | ||
| 117 | + { | ||
| 118 | + {{{2, 3, 224, 224}, {2, 3, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 119 | + }, | ||
| 120 | + { | ||
| 121 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 122 | + }, | ||
| 123 | + { | ||
| 124 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({})}, | ||
| 125 | + }); | ||
| 126 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 127 | +} | ||
| 128 | + | ||
| 129 | +// 异常用例: output_size 元素数量为 3 (不为 2) | ||
| 130 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_invalid_output_size_3_elements) | ||
| 131 | +{ | ||
| 132 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 133 | + "AdaptiveAvgPool2d", | ||
| 134 | + { | ||
| 135 | + {{{2, 3, 224, 224}, {2, 3, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 136 | + }, | ||
| 137 | + { | ||
| 138 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 139 | + }, | ||
| 140 | + { | ||
| 141 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7, 7})}, | ||
| 142 | + }); | ||
| 143 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 144 | +} | ||
| 145 | + | ||
| 146 | +// 异常用例: output_size 元素数量为 1 (不为 2) | ||
| 147 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_invalid_output_size_1_element) | ||
| 148 | +{ | ||
| 149 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 150 | + "AdaptiveAvgPool2d", | ||
| 151 | + { | ||
| 152 | + {{{2, 3, 224, 224}, {2, 3, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 153 | + }, | ||
| 154 | + { | ||
| 155 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 156 | + }, | ||
| 157 | + { | ||
| 158 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7})}, | ||
| 159 | + }); | ||
| 160 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 161 | +} | ||
| 162 | + | ||
| 163 | +// 异常用例: 输入为 2D (不支持,仅支持 3D/4D) | ||
| 164 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_invalid_input_dims_2d) | ||
| 165 | +{ | ||
| 166 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 167 | + "AdaptiveAvgPool2d", | ||
| 168 | + { | ||
| 169 | + {{{224, 224}, {224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 170 | + }, | ||
| 171 | + { | ||
| 172 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 173 | + }, | ||
| 174 | + { | ||
| 175 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7})}, | ||
| 176 | + }); | ||
| 177 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 178 | +} | ||
| 179 | + | ||
| 180 | +// 异常用例: 输入为 5D (不支持,仅支持 3D/4D) | ||
| 181 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_invalid_input_dims_5d) | ||
| 182 | +{ | ||
| 183 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 184 | + "AdaptiveAvgPool2d", | ||
| 185 | + { | ||
| 186 | + {{{2, 3, 4, 224, 224}, {2, 3, 4, 224, 224}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 187 | + }, | ||
| 188 | + { | ||
| 189 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 190 | + }, | ||
| 191 | + { | ||
| 192 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7})}, | ||
| 193 | + }); | ||
| 194 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED); | ||
| 195 | +} | ||
| 196 | + | ||
| 197 | +// ======================== 异常用例: Unknown Rank ======================== | ||
| 198 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infershape_unknown_rank) | ||
| 199 | +{ | ||
| 200 | + gert::InfershapeContextPara infershapeContextPara( | ||
| 201 | + "AdaptiveAvgPool2d", | ||
| 202 | + { | ||
| 203 | + {{{-2}, {-2}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 204 | + }, | ||
| 205 | + { | ||
| 206 | + {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND}, | ||
| 207 | + }, | ||
| 208 | + { | ||
| 209 | + {"output_size", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({7, 7})}, | ||
| 210 | + }); | ||
| 211 | + ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS); | ||
| 212 | +} | ||
| 213 | + | ||
| 214 | +// ======================== InferDataType ======================== | ||
| 215 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infer_dtype_fp16) | ||
| 216 | +{ | ||
| 217 | + ge::DataType inputDtype = ge::DT_FLOAT16; | ||
| 218 | + ge::DataType outputDtype = ge::DT_UNDEFINED; | ||
| 219 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 220 | + .NodeIoNum(1, 1) | ||
| 221 | + .IrInstanceNum({1}) | ||
| 222 | + .NodeInputTd(0, inputDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 223 | + .NodeOutputTd(0, inputDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 224 | + .InputDataTypes({&inputDtype}) | ||
| 225 | + .OutputDataTypes({&outputDtype}) | ||
| 226 | + .Build(); | ||
| 227 | + | ||
| 228 | + auto inferDtypeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("AdaptiveAvgPool2d")->infer_datatype; | ||
| 229 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 230 | + ASSERT_NE(context, nullptr); | ||
| 231 | + ASSERT_EQ(inferDtypeFunc(context), ge::GRAPH_SUCCESS); | ||
| 232 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT16); | ||
| 233 | +} | ||
| 234 | + | ||
| 235 | +TEST_F(AdaptiveAvgPool2dInfershape, test_infer_dtype_bfloat16) | ||
| 236 | +{ | ||
| 237 | + ge::DataType inputDtype = ge::DT_BF16; | ||
| 238 | + ge::DataType outputDtype = ge::DT_UNDEFINED; | ||
| 239 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 240 | + .NodeIoNum(1, 1) | ||
| 241 | + .IrInstanceNum({1}) | ||
| 242 | + .NodeInputTd(0, inputDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 243 | + .NodeOutputTd(0, inputDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 244 | + .InputDataTypes({&inputDtype}) | ||
| 245 | + .OutputDataTypes({&outputDtype}) | ||
| 246 | + .Build(); | ||
| 247 | + | ||
| 248 | + auto inferDtypeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("AdaptiveAvgPool2d")->infer_datatype; | ||
| 249 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 250 | + ASSERT_EQ(inferDtypeFunc(context), ge::GRAPH_SUCCESS); | ||
| 251 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_BF16); | ||
| 252 | +} | ||
| @@ -0,0 +1,487 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | +/*! | ||
| 12 | + * \file test_avg_pool_v2_grad_infershape.cpp | ||
| 13 | + * \brief | ||
| 14 | + */ | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | +static constexpr size_t INPUT_NUM = 2; | ||
| 24 | +static constexpr size_t OUTPUT_NUM = 1; | ||
| 25 | + | ||
| 26 | +class AvgPoolV2GradInfershape : public testing::Test { | ||
| 27 | +protected: | ||
| 28 | + static void SetUpTestCase() | ||
| 29 | + { | ||
| 30 | + std::cout << "AvgPoolV2GradInfershape SetUp" << std::endl; | ||
| 31 | + } | ||
| 32 | + | ||
| 33 | + static void TearDownTestCase() | ||
| 34 | + { | ||
| 35 | + std::cout << "AvgPoolV2GradInfershape TearDown" << std::endl; | ||
| 36 | + } | ||
| 37 | +}; | ||
| 38 | + | ||
| 39 | +static void ExecuteAvgPoolV2GradInfershapeTest( | ||
| 40 | + const std::vector<int32_t>& inputShapeData, | ||
| 41 | + const gert::StorageShape& gradShape, | ||
| 42 | + ge::Format gradFormat, | ||
| 43 | + ge::DataType gradDtype, | ||
| 44 | + const std::string& dataFormat, | ||
| 45 | + const std::string& paddingMode, | ||
| 46 | + const std::vector<int64_t>& ksize, | ||
| 47 | + const std::vector<int64_t>& strides, | ||
| 48 | + const std::vector<int64_t>& pads, | ||
| 49 | + bool globalPooling, | ||
| 50 | + bool ceilMode, | ||
| 51 | + bool exclusive, | ||
| 52 | + int64_t divisorOverride, | ||
| 53 | + ge::graphStatus expectResult, | ||
| 54 | + const std::string& expectOutputStr) | ||
| 55 | +{ | ||
| 56 | + int64_t input0Size = static_cast<int64_t>(inputShapeData.size()); | ||
| 57 | + gert::StorageShape input0Shape = {{input0Size}, {input0Size}}; | ||
| 58 | + gert::StorageShape outputShape = {{}, {}}; | ||
| 59 | + | ||
| 60 | + size_t totalSize = 0; | ||
| 61 | + auto tensorHolder = | ||
| 62 | + gert::Tensor::CreateFollowing(input0Size, ge::DT_INT32, totalSize); | ||
| 63 | + auto tensor = reinterpret_cast<gert::Tensor*>(tensorHolder.get()); | ||
| 64 | + tensor->MutableStorageShape().AppendDim(input0Size); | ||
| 65 | + tensor->MutableOriginShape().AppendDim(input0Size); | ||
| 66 | + tensor->SetOriginFormat(ge::FORMAT_ND); | ||
| 67 | + tensor->SetStorageFormat(ge::FORMAT_ND); | ||
| 68 | + (void)memcpy_s( | ||
| 69 | + tensor->GetData<uint8_t>(), totalSize - sizeof(gert::Tensor), inputShapeData.data(), | ||
| 70 | + inputShapeData.size() * sizeof(int32_t)); | ||
| 71 | + | ||
| 72 | + auto holder = gert::InferShapeContextFaker() | ||
| 73 | + .NodeIoNum(INPUT_NUM, OUTPUT_NUM) | ||
| 74 | + .IrInstanceNum({1, 1}) | ||
| 75 | + .InputShapes({tensor, const_cast<gert::StorageShape*>(&gradShape)}) | ||
| 76 | + .OutputShapes({&outputShape}) | ||
| 77 | + .NodeAttrs( | ||
| 78 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>(ksize)}, | ||
| 79 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>(strides)}, | ||
| 80 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>(paddingMode)}, | ||
| 81 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>(pads)}, | ||
| 82 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>(dataFormat)}, | ||
| 83 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(globalPooling)}, | ||
| 84 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(ceilMode)}, | ||
| 85 | + {"exclusive", Ops::NN::AnyValue::CreateFrom<bool>(exclusive)}, | ||
| 86 | + {"divisor_override", Ops::NN::AnyValue::CreateFrom<int64_t>(divisorOverride)}}) | ||
| 87 | + .NodeInputTd(0, ge::DT_INT32, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 88 | + .NodeInputTd(1, gradDtype, gradFormat, gradFormat) | ||
| 89 | + .NodeOutputTd(0, gradDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 90 | + .Build(); | ||
| 91 | + | ||
| 92 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("AvgPoolV2Grad")->infer_shape; | ||
| 93 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), expectResult); | ||
| 94 | + if (expectResult == ge::GRAPH_SUCCESS) { | ||
| 95 | + auto output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 96 | + ASSERT_EQ(Ops::Base::ToString(*output), expectOutputStr); | ||
| 97 | + } | ||
| 98 | +} | ||
| 99 | + | ||
| 100 | +// ======================== Success Cases ======================== | ||
| 101 | + | ||
| 102 | +TEST_F(AvgPoolV2GradInfershape, nchw_4d_basic) | ||
| 103 | +{ | ||
| 104 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 105 | + {2, 3, 8, 8}, | ||
| 106 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 107 | + ge::FORMAT_NCHW, | ||
| 108 | + ge::DT_FLOAT16, | ||
| 109 | + "NCHW", | ||
| 110 | + "CALCULATED", | ||
| 111 | + {1, 1, 2, 2}, | ||
| 112 | + {1, 1, 2, 2}, | ||
| 113 | + {0, 0, 0, 0}, | ||
| 114 | + false, false, true, 0, | ||
| 115 | + ge::GRAPH_SUCCESS, | ||
| 116 | + "[2, 3, 8, 8]"); | ||
| 117 | +} | ||
| 118 | + | ||
| 119 | +TEST_F(AvgPoolV2GradInfershape, nchw_4d_valid) | ||
| 120 | +{ | ||
| 121 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 122 | + {2, 128, 14, 14}, | ||
| 123 | + {{2, 128, 7, 7}, {2, 128, 7, 7}}, | ||
| 124 | + ge::FORMAT_NCHW, | ||
| 125 | + ge::DT_FLOAT, | ||
| 126 | + "NCHW", | ||
| 127 | + "VALID", | ||
| 128 | + {1, 1, 2, 2}, | ||
| 129 | + {1, 1, 2, 2}, | ||
| 130 | + {0, 0, 0, 0}, | ||
| 131 | + false, false, true, 0, | ||
| 132 | + ge::GRAPH_SUCCESS, | ||
| 133 | + "[2, 128, 14, 14]"); | ||
| 134 | +} | ||
| 135 | + | ||
| 136 | +TEST_F(AvgPoolV2GradInfershape, nchw_4d_same) | ||
| 137 | +{ | ||
| 138 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 139 | + {4, 64, 32, 32}, | ||
| 140 | + {{4, 64, 16, 16}, {4, 64, 16, 16}}, | ||
| 141 | + ge::FORMAT_NCHW, | ||
| 142 | + ge::DT_FLOAT, | ||
| 143 | + "NCHW", | ||
| 144 | + "SAME", | ||
| 145 | + {1, 1, 3, 3}, | ||
| 146 | + {1, 1, 2, 2}, | ||
| 147 | + {0, 0, 0, 0}, | ||
| 148 | + true, false, false, 0, | ||
| 149 | + ge::GRAPH_SUCCESS, | ||
| 150 | + "[4, 64, 32, 32]"); | ||
| 151 | +} | ||
| 152 | + | ||
| 153 | +TEST_F(AvgPoolV2GradInfershape, nd_format_4d) | ||
| 154 | +{ | ||
| 155 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 156 | + {2, 3, 8, 8}, | ||
| 157 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 158 | + ge::FORMAT_ND, | ||
| 159 | + ge::DT_FLOAT16, | ||
| 160 | + "NCHW", | ||
| 161 | + "CALCULATED", | ||
| 162 | + {1, 1, 2, 2}, | ||
| 163 | + {1, 1, 2, 2}, | ||
| 164 | + {0, 0, 0, 0}, | ||
| 165 | + false, false, true, 0, | ||
| 166 | + ge::GRAPH_SUCCESS, | ||
| 167 | + "[2, 3, 8, 8]"); | ||
| 168 | +} | ||
| 169 | + | ||
| 170 | +TEST_F(AvgPoolV2GradInfershape, nhwc_4d_basic) | ||
| 171 | +{ | ||
| 172 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 173 | + {2, 8, 8, 3}, | ||
| 174 | + {{2, 4, 4, 3}, {2, 4, 4, 3}}, | ||
| 175 | + ge::FORMAT_NHWC, | ||
| 176 | + ge::DT_FLOAT, | ||
| 177 | + "NHWC", | ||
| 178 | + "CALCULATED", | ||
| 179 | + {1, 2, 2, 1}, | ||
| 180 | + {1, 2, 2, 1}, | ||
| 181 | + {0, 0, 0, 0}, | ||
| 182 | + false, false, false, 0, | ||
| 183 | + ge::GRAPH_SUCCESS, | ||
| 184 | + "[2, 8, 8, 3]"); | ||
| 185 | +} | ||
| 186 | + | ||
| 187 | +TEST_F(AvgPoolV2GradInfershape, nd_3d_input) | ||
| 188 | +{ | ||
| 189 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 190 | + {3, 8, 8}, | ||
| 191 | + {{3, 8, 8}, {3, 8, 8}}, | ||
| 192 | + ge::FORMAT_ND, | ||
| 193 | + ge::DT_FLOAT16, | ||
| 194 | + "NCHW", | ||
| 195 | + "CALCULATED", | ||
| 196 | + {1, 1, 2, 2}, | ||
| 197 | + {1, 1, 2, 2}, | ||
| 198 | + {0, 0, 0, 0}, | ||
| 199 | + false, false, true, 0, | ||
| 200 | + ge::GRAPH_SUCCESS, | ||
| 201 | + "[3, 8, 8]"); | ||
| 202 | +} | ||
| 203 | + | ||
| 204 | +// ======================== Failure Cases ======================== | ||
| 205 | + | ||
| 206 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_padding_mode) | ||
| 207 | +{ | ||
| 208 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 209 | + {2, 3, 8, 8}, | ||
| 210 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 211 | + ge::FORMAT_NCHW, | ||
| 212 | + ge::DT_FLOAT16, | ||
| 213 | + "NCHW", | ||
| 214 | + "INVALID", | ||
| 215 | + {1, 1, 2, 2}, | ||
| 216 | + {1, 1, 2, 2}, | ||
| 217 | + {0, 0, 0, 0}, | ||
| 218 | + false, false, true, 0, | ||
| 219 | + ge::GRAPH_FAILED, | ||
| 220 | + ""); | ||
| 221 | +} | ||
| 222 | + | ||
| 223 | +TEST_F(AvgPoolV2GradInfershape, fail_nchw_ksize0_not_one) | ||
| 224 | +{ | ||
| 225 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 226 | + {2, 3, 8, 8}, | ||
| 227 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 228 | + ge::FORMAT_NCHW, | ||
| 229 | + ge::DT_FLOAT16, | ||
| 230 | + "NCHW", | ||
| 231 | + "CALCULATED", | ||
| 232 | + {2, 1, 2, 2}, | ||
| 233 | + {1, 1, 2, 2}, | ||
| 234 | + {0, 0, 0, 0}, | ||
| 235 | + false, false, true, 0, | ||
| 236 | + ge::GRAPH_FAILED, | ||
| 237 | + ""); | ||
| 238 | +} | ||
| 239 | + | ||
| 240 | +TEST_F(AvgPoolV2GradInfershape, fail_nchw_ksize1_not_one) | ||
| 241 | +{ | ||
| 242 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 243 | + {2, 3, 8, 8}, | ||
| 244 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 245 | + ge::FORMAT_NCHW, | ||
| 246 | + ge::DT_FLOAT16, | ||
| 247 | + "NCHW", | ||
| 248 | + "CALCULATED", | ||
| 249 | + {1, 2, 2, 2}, | ||
| 250 | + {1, 1, 2, 2}, | ||
| 251 | + {0, 0, 0, 0}, | ||
| 252 | + false, false, true, 0, | ||
| 253 | + ge::GRAPH_FAILED, | ||
| 254 | + ""); | ||
| 255 | +} | ||
| 256 | + | ||
| 257 | +TEST_F(AvgPoolV2GradInfershape, fail_nchw_strides0_not_one) | ||
| 258 | +{ | ||
| 259 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 260 | + {2, 3, 8, 8}, | ||
| 261 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 262 | + ge::FORMAT_NCHW, | ||
| 263 | + ge::DT_FLOAT16, | ||
| 264 | + "NCHW", | ||
| 265 | + "CALCULATED", | ||
| 266 | + {1, 1, 2, 2}, | ||
| 267 | + {2, 1, 2, 2}, | ||
| 268 | + {0, 0, 0, 0}, | ||
| 269 | + false, false, true, 0, | ||
| 270 | + ge::GRAPH_FAILED, | ||
| 271 | + ""); | ||
| 272 | +} | ||
| 273 | + | ||
| 274 | +TEST_F(AvgPoolV2GradInfershape, fail_nchw_strides1_not_one) | ||
| 275 | +{ | ||
| 276 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 277 | + {2, 3, 8, 8}, | ||
| 278 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 279 | + ge::FORMAT_NCHW, | ||
| 280 | + ge::DT_FLOAT16, | ||
| 281 | + "NCHW", | ||
| 282 | + "CALCULATED", | ||
| 283 | + {1, 1, 2, 2}, | ||
| 284 | + {1, 3, 2, 2}, | ||
| 285 | + {0, 0, 0, 0}, | ||
| 286 | + false, false, true, 0, | ||
| 287 | + ge::GRAPH_FAILED, | ||
| 288 | + ""); | ||
| 289 | +} | ||
| 290 | + | ||
| 291 | +TEST_F(AvgPoolV2GradInfershape, fail_nhwc_ksize3_not_one) | ||
| 292 | +{ | ||
| 293 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 294 | + {2, 8, 8, 3}, | ||
| 295 | + {{2, 4, 4, 3}, {2, 4, 4, 3}}, | ||
| 296 | + ge::FORMAT_NHWC, | ||
| 297 | + ge::DT_FLOAT, | ||
| 298 | + "NHWC", | ||
| 299 | + "CALCULATED", | ||
| 300 | + {1, 2, 2, 2}, | ||
| 301 | + {1, 2, 2, 1}, | ||
| 302 | + {0, 0, 0, 0}, | ||
| 303 | + false, false, false, 0, | ||
| 304 | + ge::GRAPH_FAILED, | ||
| 305 | + ""); | ||
| 306 | +} | ||
| 307 | + | ||
| 308 | +TEST_F(AvgPoolV2GradInfershape, fail_nhwc_strides3_not_one) | ||
| 309 | +{ | ||
| 310 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 311 | + {2, 8, 8, 3}, | ||
| 312 | + {{2, 4, 4, 3}, {2, 4, 4, 3}}, | ||
| 313 | + ge::FORMAT_NHWC, | ||
| 314 | + ge::DT_FLOAT, | ||
| 315 | + "NHWC", | ||
| 316 | + "CALCULATED", | ||
| 317 | + {1, 2, 2, 1}, | ||
| 318 | + {1, 2, 2, 3}, | ||
| 319 | + {0, 0, 0, 0}, | ||
| 320 | + false, false, false, 0, | ||
| 321 | + ge::GRAPH_FAILED, | ||
| 322 | + ""); | ||
| 323 | +} | ||
| 324 | + | ||
| 325 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_ksize_length) | ||
| 326 | +{ | ||
| 327 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 328 | + {2, 3, 8, 8}, | ||
| 329 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 330 | + ge::FORMAT_NCHW, | ||
| 331 | + ge::DT_FLOAT16, | ||
| 332 | + "NCHW", | ||
| 333 | + "CALCULATED", | ||
| 334 | + {2, 2}, | ||
| 335 | + {1, 1, 2, 2}, | ||
| 336 | + {0, 0, 0, 0}, | ||
| 337 | + false, false, true, 0, | ||
| 338 | + ge::GRAPH_FAILED, | ||
| 339 | + ""); | ||
| 340 | +} | ||
| 341 | + | ||
| 342 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_strides_length) | ||
| 343 | +{ | ||
| 344 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 345 | + {2, 3, 8, 8}, | ||
| 346 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 347 | + ge::FORMAT_NCHW, | ||
| 348 | + ge::DT_FLOAT16, | ||
| 349 | + "NCHW", | ||
| 350 | + "CALCULATED", | ||
| 351 | + {1, 1, 2, 2}, | ||
| 352 | + {2, 2}, | ||
| 353 | + {0, 0, 0, 0}, | ||
| 354 | + false, false, true, 0, | ||
| 355 | + ge::GRAPH_FAILED, | ||
| 356 | + ""); | ||
| 357 | +} | ||
| 358 | + | ||
| 359 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_dim_count_5d) | ||
| 360 | +{ | ||
| 361 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 362 | + {2, 3, 8, 8, 4}, | ||
| 363 | + {{2, 3, 8, 8, 4}, {2, 3, 8, 8, 4}}, | ||
| 364 | + ge::FORMAT_ND, | ||
| 365 | + ge::DT_FLOAT, | ||
| 366 | + "NCHW", | ||
| 367 | + "CALCULATED", | ||
| 368 | + {1, 1, 2, 2, 2}, | ||
| 369 | + {1, 1, 2, 2, 2}, | ||
| 370 | + {0, 0, 0, 0}, | ||
| 371 | + false, false, true, 0, | ||
| 372 | + ge::GRAPH_FAILED, | ||
| 373 | + ""); | ||
| 374 | +} | ||
| 375 | + | ||
| 376 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_dim_count_2d) | ||
| 377 | +{ | ||
| 378 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 379 | + {8, 8}, | ||
| 380 | + {{8, 8}, {8, 8}}, | ||
| 381 | + ge::FORMAT_ND, | ||
| 382 | + ge::DT_FLOAT, | ||
| 383 | + "NCHW", | ||
| 384 | + "CALCULATED", | ||
| 385 | + {1, 1, 2, 2}, | ||
| 386 | + {1, 1, 2, 2}, | ||
| 387 | + {0, 0, 0, 0}, | ||
| 388 | + false, false, true, 0, | ||
| 389 | + ge::GRAPH_FAILED, | ||
| 390 | + ""); | ||
| 391 | +} | ||
| 392 | + | ||
| 393 | +TEST_F(AvgPoolV2GradInfershape, fail_invalid_grad_format) | ||
| 394 | +{ | ||
| 395 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 396 | + {2, 3, 8, 8}, | ||
| 397 | + {{2, 3, 4, 4}, {2, 3, 4, 4}}, | ||
| 398 | + ge::FORMAT_FRACTAL_NZ, | ||
| 399 | + ge::DT_FLOAT16, | ||
| 400 | + "NCHW", | ||
| 401 | + "CALCULATED", | ||
| 402 | + {1, 1, 2, 2}, | ||
| 403 | + {1, 1, 2, 2}, | ||
| 404 | + {0, 0, 0, 0}, | ||
| 405 | + false, false, true, 0, | ||
| 406 | + ge::GRAPH_FAILED, | ||
| 407 | + ""); | ||
| 408 | +} | ||
| 409 | + | ||
| 410 | +// ======================== NHWC Additional Failure Cases ======================== | ||
| 411 | + | ||
| 412 | +TEST_F(AvgPoolV2GradInfershape, fail_nhwc_ksize0_not_one) | ||
| 413 | +{ | ||
| 414 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 415 | + {2, 8, 8, 3}, | ||
| 416 | + {{2, 4, 4, 3}, {2, 4, 4, 3}}, | ||
| 417 | + ge::FORMAT_NHWC, | ||
| 418 | + ge::DT_FLOAT, | ||
| 419 | + "NHWC", | ||
| 420 | + "CALCULATED", | ||
| 421 | + {2, 2, 2, 1}, | ||
| 422 | + {1, 2, 2, 1}, | ||
| 423 | + {0, 0, 0, 0}, | ||
| 424 | + false, false, false, 0, | ||
| 425 | + ge::GRAPH_FAILED, | ||
| 426 | + ""); | ||
| 427 | +} | ||
| 428 | + | ||
| 429 | +TEST_F(AvgPoolV2GradInfershape, fail_nhwc_strides0_not_one) | ||
| 430 | +{ | ||
| 431 | + ExecuteAvgPoolV2GradInfershapeTest( | ||
| 432 | + {2, 8, 8, 3}, | ||
| 433 | + {{2, 4, 4, 3}, {2, 4, 4, 3}}, | ||
| 434 | + ge::FORMAT_NHWC, | ||
| 435 | + ge::DT_FLOAT, | ||
| 436 | + "NHWC", | ||
| 437 | + "CALCULATED", | ||
| 438 | + {1, 2, 2, 1}, | ||
| 439 | + {3, 2, 2, 1}, | ||
| 440 | + {0, 0, 0, 0}, | ||
| 441 | + false, false, false, 0, | ||
| 442 | + ge::GRAPH_FAILED, | ||
| 443 | + ""); | ||
| 444 | +} | ||
| 445 | + | ||
| 446 | +// ======================== InferDataType ======================== | ||
| 447 | + | ||
| 448 | +TEST_F(AvgPoolV2GradInfershape, infer_dtype_basic) | ||
| 449 | +{ | ||
| 450 | + ge::DataType gradDtype = ge::DT_FLOAT16; | ||
| 451 | + ge::DataType outputDtype = ge::DT_UNDEFINED; | ||
| 452 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 453 | + .NodeIoNum(INPUT_NUM, OUTPUT_NUM) | ||
| 454 | + .IrInstanceNum({1, 1}) | ||
| 455 | + .NodeInputTd(0, ge::DT_INT32, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 456 | + .NodeInputTd(1, gradDtype, ge::FORMAT_NCHW, ge::FORMAT_NCHW) | ||
| 457 | + .NodeOutputTd(0, gradDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 458 | + .InputDataTypes({nullptr, &gradDtype}) | ||
| 459 | + .OutputDataTypes({&outputDtype}) | ||
| 460 | + .Build(); | ||
| 461 | + | ||
| 462 | + auto inferDtypeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("AvgPoolV2Grad")->infer_datatype; | ||
| 463 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 464 | + ASSERT_NE(context, nullptr); | ||
| 465 | + ASSERT_EQ(inferDtypeFunc(context), ge::GRAPH_SUCCESS); | ||
| 466 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT16); | ||
| 467 | +} | ||
| 468 | + | ||
| 469 | +TEST_F(AvgPoolV2GradInfershape, infer_dtype_float32) | ||
| 470 | +{ | ||
| 471 | + ge::DataType gradDtype = ge::DT_FLOAT; | ||
| 472 | + ge::DataType outputDtype = ge::DT_UNDEFINED; | ||
| 473 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 474 | + .NodeIoNum(INPUT_NUM, OUTPUT_NUM) | ||
| 475 | + .IrInstanceNum({1, 1}) | ||
| 476 | + .NodeInputTd(0, ge::DT_INT32, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 477 | + .NodeInputTd(1, gradDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 478 | + .NodeOutputTd(0, gradDtype, ge::FORMAT_ND, ge::FORMAT_ND) | ||
| 479 | + .InputDataTypes({nullptr, &gradDtype}) | ||
| 480 | + .OutputDataTypes({&outputDtype}) | ||
| 481 | + .Build(); | ||
| 482 | + | ||
| 483 | + auto inferDtypeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("AvgPoolV2Grad")->infer_datatype; | ||
| 484 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 485 | + ASSERT_EQ(inferDtypeFunc(context), ge::GRAPH_SUCCESS); | ||
| 486 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT); | ||
| 487 | +} | ||
| @@ -57,4 +57,383 @@ TEST_F(MaxPool3DUT, InferShapeOk) { | |||
| 57 | 57 | ||
| 58 | EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | 58 | EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); |
| 59 | } | 59 | } |
| 60 | + | ||
| 61 | +// ======================== VALID Padding ======================== | ||
| 62 | + | ||
| 63 | +TEST_F(MaxPool3DUT, InferShapeValid) { | ||
| 64 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 65 | + | ||
| 66 | + gert::StorageShape x_shape = {{2, 2, 8, 8, 7}, {2, 2, 8, 8, 7}}; | ||
| 67 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 68 | + std::vector<void *> output_shapes_ref(1); | ||
| 69 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 70 | + | ||
| 71 | + auto holder = gert::InferShapeContextFaker() | ||
| 72 | + .NodeIoNum(1, 1) | ||
| 73 | + .IrInstanceNum({1}) | ||
| 74 | + .InputShapes({&x_shape}) | ||
| 75 | + .OutputShapes(output_shapes_ref) | ||
| 76 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 77 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 78 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 79 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 80 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 81 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 82 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 83 | + .Build(); | ||
| 84 | + | ||
| 85 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 86 | +} | ||
| 87 | + | ||
| 88 | +// ======================== CALCULATED Padding ======================== | ||
| 89 | + | ||
| 90 | +TEST_F(MaxPool3DUT, InferShapeCalculated) { | ||
| 91 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 92 | + | ||
| 93 | + gert::StorageShape x_shape = {{2, 2, 8, 8, 7}, {2, 2, 8, 8, 7}}; | ||
| 94 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 95 | + std::vector<void *> output_shapes_ref(1); | ||
| 96 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 97 | + | ||
| 98 | + auto holder = gert::InferShapeContextFaker() | ||
| 99 | + .NodeIoNum(1, 1) | ||
| 100 | + .IrInstanceNum({1}) | ||
| 101 | + .InputShapes({&x_shape}) | ||
| 102 | + .OutputShapes(output_shapes_ref) | ||
| 103 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 104 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 105 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 106 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 107 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 108 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 109 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 110 | + .Build(); | ||
| 111 | + | ||
| 112 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 113 | +} | ||
| 114 | + | ||
| 115 | +TEST_F(MaxPool3DUT, InferShapeCalculatedCeilMode) { | ||
| 116 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 117 | + | ||
| 118 | + gert::StorageShape x_shape = {{2, 2, 7, 7, 7}, {2, 2, 7, 7, 7}}; | ||
| 119 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 120 | + std::vector<void *> output_shapes_ref(1); | ||
| 121 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 122 | + | ||
| 123 | + auto holder = gert::InferShapeContextFaker() | ||
| 124 | + .NodeIoNum(1, 1) | ||
| 125 | + .IrInstanceNum({1}) | ||
| 126 | + .InputShapes({&x_shape}) | ||
| 127 | + .OutputShapes(output_shapes_ref) | ||
| 128 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 3, 3, 3})}, | ||
| 129 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 130 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 131 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 132 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 133 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(true)}, | ||
| 134 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 135 | + .Build(); | ||
| 136 | + | ||
| 137 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 138 | +} | ||
| 139 | + | ||
| 140 | +// ======================== NDHWC Format ======================== | ||
| 141 | + | ||
| 142 | +TEST_F(MaxPool3DUT, InferShapeNdhwc) { | ||
| 143 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 144 | + | ||
| 145 | + gert::StorageShape x_shape = {{2, 8, 8, 7, 2}, {2, 8, 8, 7, 2}}; | ||
| 146 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 147 | + std::vector<void *> output_shapes_ref(1); | ||
| 148 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 149 | + | ||
| 150 | + auto holder = gert::InferShapeContextFaker() | ||
| 151 | + .NodeIoNum(1, 1) | ||
| 152 | + .IrInstanceNum({1}) | ||
| 153 | + .InputShapes({&x_shape}) | ||
| 154 | + .OutputShapes(output_shapes_ref) | ||
| 155 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 156 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 157 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 158 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 159 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 160 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 161 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 162 | + .Build(); | ||
| 163 | + | ||
| 164 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 165 | +} | ||
| 166 | + | ||
| 167 | +// ======================== Unknown Shape / Unknown Rank ======================== | ||
| 168 | + | ||
| 169 | +TEST_F(MaxPool3DUT, InferShapeUnknownRank) { | ||
| 170 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 171 | + | ||
| 172 | + gert::StorageShape x_shape = {{-2}, {}}; | ||
| 173 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 174 | + std::vector<void *> output_shapes_ref(1); | ||
| 175 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 176 | + | ||
| 177 | + auto holder = gert::InferShapeContextFaker() | ||
| 178 | + .NodeIoNum(1, 1) | ||
| 179 | + .IrInstanceNum({1}) | ||
| 180 | + .InputShapes({&x_shape}) | ||
| 181 | + .OutputShapes(output_shapes_ref) | ||
| 182 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 183 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 184 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 185 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 186 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 187 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 188 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 189 | + .Build(); | ||
| 190 | + | ||
| 191 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 192 | +} | ||
| 193 | + | ||
| 194 | +TEST_F(MaxPool3DUT, InferShapeUnknownShape) { | ||
| 195 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 196 | + | ||
| 197 | + gert::StorageShape x_shape = {{2, -1, -1, 8, 7}, {}}; | ||
| 198 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 199 | + std::vector<void *> output_shapes_ref(1); | ||
| 200 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 201 | + | ||
| 202 | + auto holder = gert::InferShapeContextFaker() | ||
| 203 | + .NodeIoNum(1, 1) | ||
| 204 | + .IrInstanceNum({1}) | ||
| 205 | + .InputShapes({&x_shape}) | ||
| 206 | + .OutputShapes(output_shapes_ref) | ||
| 207 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 208 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 209 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 210 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 211 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 212 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 213 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 214 | + .Build(); | ||
| 215 | + | ||
| 216 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 217 | +} | ||
| 218 | + | ||
| 219 | +// ======================== Failure Cases ======================== | ||
| 220 | + | ||
| 221 | +TEST_F(MaxPool3DUT, InferShapeFailInvalidPadding) { | ||
| 222 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 223 | + | ||
| 224 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 225 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 226 | + std::vector<void *> output_shapes_ref(1); | ||
| 227 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 228 | + | ||
| 229 | + auto holder = gert::InferShapeContextFaker() | ||
| 230 | + .NodeIoNum(1, 1) | ||
| 231 | + .IrInstanceNum({1}) | ||
| 232 | + .InputShapes({&x_shape}) | ||
| 233 | + .OutputShapes(output_shapes_ref) | ||
| 234 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 235 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 236 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("INVALID")}, | ||
| 237 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 238 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 239 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 240 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 241 | + .Build(); | ||
| 242 | + | ||
| 243 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 244 | +} | ||
| 245 | + | ||
| 246 | +TEST_F(MaxPool3DUT, InferShapeFailValidKsizeLength) { | ||
| 247 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 248 | + | ||
| 249 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 250 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 251 | + std::vector<void *> output_shapes_ref(1); | ||
| 252 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 253 | + | ||
| 254 | + auto holder = gert::InferShapeContextFaker() | ||
| 255 | + .NodeIoNum(1, 1) | ||
| 256 | + .IrInstanceNum({1}) | ||
| 257 | + .InputShapes({&x_shape}) | ||
| 258 | + .OutputShapes(output_shapes_ref) | ||
| 259 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2})}, | ||
| 260 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 261 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 262 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 263 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 264 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 265 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 266 | + .Build(); | ||
| 267 | + | ||
| 268 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 269 | +} | ||
| 270 | + | ||
| 271 | +TEST_F(MaxPool3DUT, InferShapeFailValidStridesLength) { | ||
| 272 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 273 | + | ||
| 274 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 275 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 276 | + std::vector<void *> output_shapes_ref(1); | ||
| 277 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 278 | + | ||
| 279 | + auto holder = gert::InferShapeContextFaker() | ||
| 280 | + .NodeIoNum(1, 1) | ||
| 281 | + .IrInstanceNum({1}) | ||
| 282 | + .InputShapes({&x_shape}) | ||
| 283 | + .OutputShapes(output_shapes_ref) | ||
| 284 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 285 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2})}, | ||
| 286 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 287 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 288 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 289 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 290 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 291 | + .Build(); | ||
| 292 | + | ||
| 293 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 294 | +} | ||
| 295 | + | ||
| 296 | +TEST_F(MaxPool3DUT, InferShapeFailStridesZero) { | ||
| 297 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 298 | + | ||
| 299 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 300 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 301 | + std::vector<void *> output_shapes_ref(1); | ||
| 302 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 303 | + | ||
| 304 | + auto holder = gert::InferShapeContextFaker() | ||
| 305 | + .NodeIoNum(1, 1) | ||
| 306 | + .IrInstanceNum({1}) | ||
| 307 | + .InputShapes({&x_shape}) | ||
| 308 | + .OutputShapes(output_shapes_ref) | ||
| 309 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 310 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 0, 2, 2})}, | ||
| 311 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 312 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 313 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 314 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 315 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 316 | + .Build(); | ||
| 317 | + | ||
| 318 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 319 | +} | ||
| 320 | + | ||
| 321 | +TEST_F(MaxPool3DUT, InferShapeFailCalculatedPadsLength) { | ||
| 322 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 323 | + | ||
| 324 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 325 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 326 | + std::vector<void *> output_shapes_ref(1); | ||
| 327 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 328 | + | ||
| 329 | + auto holder = gert::InferShapeContextFaker() | ||
| 330 | + .NodeIoNum(1, 1) | ||
| 331 | + .IrInstanceNum({1}) | ||
| 332 | + .InputShapes({&x_shape}) | ||
| 333 | + .OutputShapes(output_shapes_ref) | ||
| 334 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 335 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 336 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 337 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 338 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 339 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 340 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 341 | + .Build(); | ||
| 342 | + | ||
| 343 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 344 | +} | ||
| 345 | + | ||
| 346 | +TEST_F(MaxPool3DUT, InferShapeFailCalculatedStridesZero) { | ||
| 347 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 348 | + | ||
| 349 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 350 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 351 | + std::vector<void *> output_shapes_ref(1); | ||
| 352 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 353 | + | ||
| 354 | + auto holder = gert::InferShapeContextFaker() | ||
| 355 | + .NodeIoNum(1, 1) | ||
| 356 | + .IrInstanceNum({1}) | ||
| 357 | + .InputShapes({&x_shape}) | ||
| 358 | + .OutputShapes(output_shapes_ref) | ||
| 359 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 360 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 0, 2, 2})}, | ||
| 361 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 362 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 363 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 364 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 365 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 366 | + .Build(); | ||
| 367 | + | ||
| 368 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 369 | +} | ||
| 370 | + | ||
| 371 | +TEST_F(MaxPool3DUT, InferShapeFailSAMEStridesZero) { | ||
| 372 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_shape; | ||
| 373 | + | ||
| 374 | + gert::StorageShape x_shape = {{2, 2, 5, 2, 7}, {2, 2, 5, 2, 7}}; | ||
| 375 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 376 | + std::vector<void *> output_shapes_ref(1); | ||
| 377 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 378 | + | ||
| 379 | + auto holder = gert::InferShapeContextFaker() | ||
| 380 | + .NodeIoNum(1, 1) | ||
| 381 | + .IrInstanceNum({1}) | ||
| 382 | + .InputShapes({&x_shape}) | ||
| 383 | + .OutputShapes(output_shapes_ref) | ||
| 384 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2, 2})}, | ||
| 385 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 0, 2, 2})}, | ||
| 386 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 387 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 388 | + {"dilation", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1, 1})}, | ||
| 389 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 390 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCDHW")}}) | ||
| 391 | + .Build(); | ||
| 392 | + | ||
| 393 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 394 | +} | ||
| 395 | + | ||
| 396 | +// ======================== InferDataType ======================== | ||
| 397 | + | ||
| 398 | +TEST_F(MaxPool3DUT, InferDataTypeFP16) { | ||
| 399 | + auto infer_dtype_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_datatype; | ||
| 400 | + ASSERT_NE(infer_dtype_func, nullptr); | ||
| 401 | + | ||
| 402 | + ge::DataType input_dtype = ge::DT_FLOAT16; | ||
| 403 | + ge::DataType output_dtype = ge::DT_UNDEFINED; | ||
| 404 | + | ||
| 405 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 406 | + .NodeIoNum(1, 1) | ||
| 407 | + .IrInstanceNum({1}) | ||
| 408 | + .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 409 | + .NodeOutputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 410 | + .InputDataTypes({&input_dtype}) | ||
| 411 | + .OutputDataTypes({&output_dtype}) | ||
| 412 | + .Build(); | ||
| 413 | + | ||
| 414 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 415 | + ASSERT_EQ(infer_dtype_func(context), ge::GRAPH_SUCCESS); | ||
| 416 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT16); | ||
| 417 | +} | ||
| 418 | + | ||
| 419 | +TEST_F(MaxPool3DUT, InferDataTypeBF16) { | ||
| 420 | + auto infer_dtype_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3D")->infer_datatype; | ||
| 421 | + | ||
| 422 | + ge::DataType input_dtype = ge::DT_BF16; | ||
| 423 | + ge::DataType output_dtype = ge::DT_UNDEFINED; | ||
| 424 | + | ||
| 425 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 426 | + .NodeIoNum(1, 1) | ||
| 427 | + .IrInstanceNum({1}) | ||
| 428 | + .NodeInputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 429 | + .NodeOutputTd(0, ge::DT_BF16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 430 | + .InputDataTypes({&input_dtype}) | ||
| 431 | + .OutputDataTypes({&output_dtype}) | ||
| 432 | + .Build(); | ||
| 433 | + | ||
| 434 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 435 | + ASSERT_EQ(infer_dtype_func(context), ge::GRAPH_SUCCESS); | ||
| 436 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_BF16); | ||
| 437 | +} | ||
| 438 | + | ||
| 60 | } | 439 | } |
| @@ -0,0 +1,18 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +message(STATUS "=== Debug: start ops.pooling.max_pool3d_grad.tests.CMakeLists.txt ") | ||
| 12 | +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 13 | +message(STATUS "=== Debug: CURRENT_DIRS =${CURRENT_DIRS} ") | ||
| 14 | +foreach(SUB_DIR ${CURRENT_DIRS}) | ||
| 15 | + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") | ||
| 16 | + add_subdirectory(${SUB_DIR}) | ||
| 17 | + endif() | ||
| 18 | +endforeach() | ||
| @@ -0,0 +1,17 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 12 | +message(STATUS "=== Debug: CURRENT_DIRS =${CURRENT_DIRS} ") | ||
| 13 | +foreach(SUB_DIR ${CURRENT_DIRS}) | ||
| 14 | + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") | ||
| 15 | + add_subdirectory(${SUB_DIR}) | ||
| 16 | + endif() | ||
| 17 | +endforeach() | ||
| @@ -0,0 +1,15 @@ | |||
| 1 | +# ---------------------------------------------------------------------------- | ||
| 2 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | +# ---------------------------------------------------------------------------- | ||
| 10 | + | ||
| 11 | +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) | ||
| 12 | +if(UT_TEST_ALL OR OP_HOST_UT) | ||
| 13 | + add_modules_ut_sources(HOSTNAME ${OP_TILING_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 14 | + add_modules_ut_sources(HOSTNAME ${OP_INFERSHAPE_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) | ||
| 15 | +endif() | ||
| @@ -0,0 +1,503 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | +namespace { | ||
| 22 | +template <typename T> | ||
| 23 | +std::string Shape2String(const T& shape) | ||
| 24 | +{ | ||
| 25 | + std::ostringstream oss; | ||
| 26 | + oss << "["; | ||
| 27 | + if (shape.GetDimNum() > 0) { | ||
| 28 | + for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) { | ||
| 29 | + oss << shape.GetDim(i) << ", "; | ||
| 30 | + } | ||
| 31 | + oss << shape.GetDim(shape.GetDimNum() - 1); | ||
| 32 | + } | ||
| 33 | + oss << "]"; | ||
| 34 | + return oss.str(); | ||
| 35 | +} | ||
| 36 | + | ||
| 37 | +class MaxPool3DGradInfer : public testing::Test { | ||
| 38 | +protected: | ||
| 39 | + static void SetUpTestCase() | ||
| 40 | + { | ||
| 41 | + std::cout << "MaxPool3DGradInferTest SetUp" << std::endl; | ||
| 42 | + fe::PlatformInfo platformInfo; | ||
| 43 | + platformInfo.soc_info.ai_core_cnt = 64; | ||
| 44 | + platformInfo.str_info.short_soc_version = "Ascend950"; | ||
| 45 | + fe::PlatformInfoManager::Instance().platform_info_map_["Ascend950"] = platformInfo; | ||
| 46 | + fe::OptionalInfo optiCompilationInfo; | ||
| 47 | + optiCompilationInfo.soc_version = "Ascend950"; | ||
| 48 | + fe::PlatformInfoManager::Instance().SetOptionalCompilationInfo(optiCompilationInfo); | ||
| 49 | + } | ||
| 50 | + | ||
| 51 | + static void TearDownTestCase() | ||
| 52 | + { | ||
| 53 | + std::cout << "MaxPool3DGradInferTest TearDown" << std::endl; | ||
| 54 | + } | ||
| 55 | +}; | ||
| 56 | + | ||
| 57 | +// ======================== Success Cases ======================== | ||
| 58 | + | ||
| 59 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fp32_valid) | ||
| 60 | +{ | ||
| 61 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 62 | + | ||
| 63 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 64 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 65 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 66 | + gert::StorageShape yShape = {{}, {}}; | ||
| 67 | + | ||
| 68 | + auto holder = gert::InferShapeContextFaker() | ||
| 69 | + .NodeIoNum(3, 1) | ||
| 70 | + .IrInstanceNum({1, 1, 1}) | ||
| 71 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 72 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 73 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 74 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 75 | + .NodeAttrs( | ||
| 76 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 77 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 78 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 79 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 80 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 81 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 82 | + .OutputShapes({&yShape}) | ||
| 83 | + .Build(); | ||
| 84 | + | ||
| 85 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 86 | + gert::Shape* output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 87 | + ASSERT_EQ(Shape2String(*output), "[2, 4, 8, 8, 3]"); | ||
| 88 | +} | ||
| 89 | + | ||
| 90 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fp16_same) | ||
| 91 | +{ | ||
| 92 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 93 | + | ||
| 94 | + gert::StorageShape xShape = {{1, 16, 7, 7, 8}, {}}; | ||
| 95 | + gert::StorageShape yInShape = {{1, 16, 4, 4, 8}, {}}; | ||
| 96 | + gert::StorageShape gradShape = {{1, 16, 4, 4, 8}, {}}; | ||
| 97 | + gert::StorageShape yShape = {{}, {}}; | ||
| 98 | + | ||
| 99 | + auto holder = gert::InferShapeContextFaker() | ||
| 100 | + .NodeIoNum(3, 1) | ||
| 101 | + .IrInstanceNum({1, 1, 1}) | ||
| 102 | + .NodeInputTd(0, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 103 | + .NodeInputTd(1, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 104 | + .NodeInputTd(2, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 105 | + .NodeOutputTd(0, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 106 | + .NodeAttrs( | ||
| 107 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 108 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 109 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 110 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 111 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 112 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 113 | + .OutputShapes({&yShape}) | ||
| 114 | + .Build(); | ||
| 115 | + | ||
| 116 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 117 | + gert::Shape* output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 118 | + ASSERT_EQ(Shape2String(*output), "[1, 16, 7, 7, 8]"); | ||
| 119 | +} | ||
| 120 | + | ||
| 121 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_bf16) | ||
| 122 | +{ | ||
| 123 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 124 | + | ||
| 125 | + gert::StorageShape xShape = {{1, 32, 4, 4, 4}, {}}; | ||
| 126 | + gert::StorageShape yInShape = {{1, 32, 2, 2, 4}, {}}; | ||
| 127 | + gert::StorageShape gradShape = {{1, 32, 2, 2, 4}, {}}; | ||
| 128 | + gert::StorageShape yShape = {{}, {}}; | ||
| 129 | + | ||
| 130 | + auto holder = gert::InferShapeContextFaker() | ||
| 131 | + .NodeIoNum(3, 1) | ||
| 132 | + .IrInstanceNum({1, 1, 1}) | ||
| 133 | + .NodeInputTd(0, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 134 | + .NodeInputTd(1, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 135 | + .NodeInputTd(2, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 136 | + .NodeOutputTd(0, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 137 | + .NodeAttrs( | ||
| 138 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 139 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 140 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 141 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 142 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 143 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 144 | + .OutputShapes({&yShape}) | ||
| 145 | + .Build(); | ||
| 146 | + | ||
| 147 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 148 | + gert::Shape* output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 149 | + ASSERT_EQ(Shape2String(*output), "[1, 32, 4, 4, 4]"); | ||
| 150 | +} | ||
| 151 | + | ||
| 152 | +// ======================== Unknown Shape / Unknown Rank ======================== | ||
| 153 | + | ||
| 154 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_unknown_shape) | ||
| 155 | +{ | ||
| 156 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 157 | + | ||
| 158 | + gert::StorageShape xShape = {{2, -1, -1, 8, 3}, {}}; | ||
| 159 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 160 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 161 | + gert::StorageShape yShape = {{}, {}}; | ||
| 162 | + | ||
| 163 | + auto holder = gert::InferShapeContextFaker() | ||
| 164 | + .NodeIoNum(3, 1) | ||
| 165 | + .IrInstanceNum({1, 1, 1}) | ||
| 166 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 167 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 168 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 169 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 170 | + .NodeAttrs( | ||
| 171 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 172 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 173 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 174 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 175 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 176 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 177 | + .OutputShapes({&yShape}) | ||
| 178 | + .Build(); | ||
| 179 | + | ||
| 180 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 181 | + gert::Shape* output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 182 | + ASSERT_EQ(Shape2String(*output), "[-1, -1, -1, -1, -1]"); | ||
| 183 | +} | ||
| 184 | + | ||
| 185 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_unknown_rank) | ||
| 186 | +{ | ||
| 187 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 188 | + | ||
| 189 | + gert::StorageShape xShape = {{-2}, {}}; | ||
| 190 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 191 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 192 | + gert::StorageShape yShape = {{}, {}}; | ||
| 193 | + | ||
| 194 | + auto holder = gert::InferShapeContextFaker() | ||
| 195 | + .NodeIoNum(3, 1) | ||
| 196 | + .IrInstanceNum({1, 1, 1}) | ||
| 197 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 198 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 199 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 200 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 201 | + .NodeAttrs( | ||
| 202 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 203 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 204 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 205 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 206 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 207 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 208 | + .OutputShapes({&yShape}) | ||
| 209 | + .Build(); | ||
| 210 | + | ||
| 211 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 212 | +} | ||
| 213 | + | ||
| 214 | +// ======================== Failure Cases ======================== | ||
| 215 | + | ||
| 216 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fail_invalid_ksize_length) | ||
| 217 | +{ | ||
| 218 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 219 | + | ||
| 220 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 221 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 222 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 223 | + gert::StorageShape yShape = {{}, {}}; | ||
| 224 | + | ||
| 225 | + auto holder = gert::InferShapeContextFaker() | ||
| 226 | + .NodeIoNum(3, 1) | ||
| 227 | + .IrInstanceNum({1, 1, 1}) | ||
| 228 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 229 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 230 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 231 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 232 | + .NodeAttrs( | ||
| 233 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({2, 2, 2})}, | ||
| 234 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 235 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 236 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 237 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 238 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 239 | + .OutputShapes({&yShape}) | ||
| 240 | + .Build(); | ||
| 241 | + | ||
| 242 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 243 | +} | ||
| 244 | + | ||
| 245 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fail_ksize_zero) | ||
| 246 | +{ | ||
| 247 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 248 | + | ||
| 249 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 250 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 251 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 252 | + gert::StorageShape yShape = {{}, {}}; | ||
| 253 | + | ||
| 254 | + auto holder = gert::InferShapeContextFaker() | ||
| 255 | + .NodeIoNum(3, 1) | ||
| 256 | + .IrInstanceNum({1, 1, 1}) | ||
| 257 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 258 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 259 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 260 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 261 | + .NodeAttrs( | ||
| 262 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 0, 2, 2, 1})}, | ||
| 263 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 264 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 265 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 266 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 267 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 268 | + .OutputShapes({&yShape}) | ||
| 269 | + .Build(); | ||
| 270 | + | ||
| 271 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 272 | +} | ||
| 273 | + | ||
| 274 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fail_invalid_strides_length) | ||
| 275 | +{ | ||
| 276 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 277 | + | ||
| 278 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 279 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 280 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 281 | + gert::StorageShape yShape = {{}, {}}; | ||
| 282 | + | ||
| 283 | + auto holder = gert::InferShapeContextFaker() | ||
| 284 | + .NodeIoNum(3, 1) | ||
| 285 | + .IrInstanceNum({1, 1, 1}) | ||
| 286 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 287 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 288 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 289 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 290 | + .NodeAttrs( | ||
| 291 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 292 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({2, 2, 2})}, | ||
| 293 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 294 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 295 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 296 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 297 | + .OutputShapes({&yShape}) | ||
| 298 | + .Build(); | ||
| 299 | + | ||
| 300 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 301 | +} | ||
| 302 | + | ||
| 303 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fail_strides_zero) | ||
| 304 | +{ | ||
| 305 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 306 | + | ||
| 307 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 308 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 309 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 310 | + gert::StorageShape yShape = {{}, {}}; | ||
| 311 | + | ||
| 312 | + auto holder = gert::InferShapeContextFaker() | ||
| 313 | + .NodeIoNum(3, 1) | ||
| 314 | + .IrInstanceNum({1, 1, 1}) | ||
| 315 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 316 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 317 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 318 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 319 | + .NodeAttrs( | ||
| 320 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 321 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 0, 2, 2, 1})}, | ||
| 322 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 323 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 324 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 325 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 326 | + .OutputShapes({&yShape}) | ||
| 327 | + .Build(); | ||
| 328 | + | ||
| 329 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 330 | +} | ||
| 331 | + | ||
| 332 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_infershape_fail_invalid_padding) | ||
| 333 | +{ | ||
| 334 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 335 | + | ||
| 336 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 337 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 338 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 339 | + gert::StorageShape yShape = {{}, {}}; | ||
| 340 | + | ||
| 341 | + auto holder = gert::InferShapeContextFaker() | ||
| 342 | + .NodeIoNum(3, 1) | ||
| 343 | + .IrInstanceNum({1, 1, 1}) | ||
| 344 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 345 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 346 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 347 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 348 | + .NodeAttrs( | ||
| 349 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 350 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 351 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("INVALID")}, | ||
| 352 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 353 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 354 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 355 | + .OutputShapes({&yShape}) | ||
| 356 | + .Build(); | ||
| 357 | + | ||
| 358 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 359 | +} | ||
| 360 | + | ||
| 361 | +// ======================== InferDataType ======================== | ||
| 362 | + | ||
| 363 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_inferdatatype_fp16) | ||
| 364 | +{ | ||
| 365 | + auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_datatype; | ||
| 366 | + | ||
| 367 | + if (data_type_func != nullptr) { | ||
| 368 | + ge::DataType input_ref = ge::DT_FLOAT16; | ||
| 369 | + ge::DataType output_ref = ge::DT_FLOAT16; | ||
| 370 | + auto context_holder = gert::InferDataTypeContextFaker() | ||
| 371 | + .NodeIoNum(3, 1) | ||
| 372 | + .IrInstanceNum({1, 1, 1}) | ||
| 373 | + .NodeInputTd(0, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 374 | + .NodeInputTd(1, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 375 | + .NodeInputTd(2, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 376 | + .NodeOutputTd(0, ge::DT_FLOAT16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 377 | + .InputDataTypes({&input_ref, &input_ref, &input_ref}) | ||
| 378 | + .OutputDataTypes({&output_ref}) | ||
| 379 | + .NodeAttrs( | ||
| 380 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 381 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 382 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 383 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 384 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 385 | + .Build(); | ||
| 386 | + auto context = context_holder.GetContext<gert::InferDataTypeContext>(); | ||
| 387 | + EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS); | ||
| 388 | + ASSERT_NE(context, nullptr); | ||
| 389 | + EXPECT_EQ(context->GetOutputDataType(0), output_ref); | ||
| 390 | + } | ||
| 391 | +} | ||
| 392 | + | ||
| 393 | +TEST_F(MaxPool3DGradInfer, max_pool3d_grad_inferdatatype_bf16) | ||
| 394 | +{ | ||
| 395 | + auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_datatype; | ||
| 396 | + | ||
| 397 | + if (data_type_func != nullptr) { | ||
| 398 | + ge::DataType input_ref = ge::DT_BF16; | ||
| 399 | + ge::DataType output_ref = ge::DT_BF16; | ||
| 400 | + auto context_holder = gert::InferDataTypeContextFaker() | ||
| 401 | + .NodeIoNum(3, 1) | ||
| 402 | + .IrInstanceNum({1, 1, 1}) | ||
| 403 | + .NodeInputTd(0, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 404 | + .NodeInputTd(1, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 405 | + .NodeInputTd(2, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 406 | + .NodeOutputTd(0, ge::DT_BF16, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 407 | + .InputDataTypes({&input_ref, &input_ref, &input_ref}) | ||
| 408 | + .OutputDataTypes({&output_ref}) | ||
| 409 | + .NodeAttrs( | ||
| 410 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 411 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 412 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 413 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0, 0, 0})}, | ||
| 414 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 415 | + .Build(); | ||
| 416 | + auto context = context_holder.GetContext<gert::InferDataTypeContext>(); | ||
| 417 | + EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS); | ||
| 418 | + ASSERT_NE(context, nullptr); | ||
| 419 | + EXPECT_EQ(context->GetOutputDataType(0), output_ref); | ||
| 420 | + } | ||
| 421 | +} | ||
| 422 | + | ||
| 423 | +// ======================== Non-Ascend950 Platform Tests ======================== | ||
| 424 | + | ||
| 425 | +class MaxPool3DGradInferNon950 : public testing::Test { | ||
| 426 | +protected: | ||
| 427 | + static void SetUpTestCase() | ||
| 428 | + { | ||
| 429 | + std::cout << "MaxPool3DGradInferNon950Test SetUp" << std::endl; | ||
| 430 | + fe::PlatformInfo platformInfo; | ||
| 431 | + platformInfo.soc_info.ai_core_cnt = 64; | ||
| 432 | + platformInfo.str_info.short_soc_version = "Ascend910B"; | ||
| 433 | + fe::PlatformInfoManager::Instance().platform_info_map_["Ascend910B"] = platformInfo; | ||
| 434 | + fe::OptionalInfo optiCompilationInfo; | ||
| 435 | + optiCompilationInfo.soc_version = "Ascend910B"; | ||
| 436 | + fe::PlatformInfoManager::Instance().SetOptionalCompilationInfo(optiCompilationInfo); | ||
| 437 | + } | ||
| 438 | + | ||
| 439 | + static void TearDownTestCase() | ||
| 440 | + { | ||
| 441 | + std::cout << "MaxPool3DGradInferNon950Test TearDown" << std::endl; | ||
| 442 | + } | ||
| 443 | +}; | ||
| 444 | + | ||
| 445 | +TEST_F(MaxPool3DGradInferNon950, max_pool3d_grad_infershape_calculated_success) | ||
| 446 | +{ | ||
| 447 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 448 | + | ||
| 449 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 450 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 451 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 452 | + gert::StorageShape yShape = {{}, {}}; | ||
| 453 | + | ||
| 454 | + auto holder = gert::InferShapeContextFaker() | ||
| 455 | + .NodeIoNum(3, 1) | ||
| 456 | + .IrInstanceNum({1, 1, 1}) | ||
| 457 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 458 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 459 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 460 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 461 | + .NodeAttrs( | ||
| 462 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 463 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 464 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 465 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 1, 0, 1, 0, 0})}, | ||
| 466 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 467 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 468 | + .OutputShapes({&yShape}) | ||
| 469 | + .Build(); | ||
| 470 | + | ||
| 471 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 472 | +} | ||
| 473 | + | ||
| 474 | +TEST_F(MaxPool3DGradInferNon950, max_pool3d_grad_infershape_fail_negative_pads) | ||
| 475 | +{ | ||
| 476 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPool3DGrad")->infer_shape; | ||
| 477 | + | ||
| 478 | + gert::StorageShape xShape = {{2, 4, 8, 8, 3}, {}}; | ||
| 479 | + gert::StorageShape yInShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 480 | + gert::StorageShape gradShape = {{2, 3, 4, 4, 3}, {}}; | ||
| 481 | + gert::StorageShape yShape = {{}, {}}; | ||
| 482 | + | ||
| 483 | + auto holder = gert::InferShapeContextFaker() | ||
| 484 | + .NodeIoNum(3, 1) | ||
| 485 | + .IrInstanceNum({1, 1, 1}) | ||
| 486 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 487 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 488 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 489 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NDHWC, ge::Format::FORMAT_RESERVED) | ||
| 490 | + .NodeAttrs( | ||
| 491 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 492 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 2, 1})}, | ||
| 493 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 494 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, -1, 0, 1, 0, 0})}, | ||
| 495 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NDHWC")}}) | ||
| 496 | + .InputShapes({&xShape, &yInShape, &gradShape}) | ||
| 497 | + .OutputShapes({&yShape}) | ||
| 498 | + .Build(); | ||
| 499 | + | ||
| 500 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 501 | +} | ||
| 502 | + | ||
| 503 | +} // namespace | ||
| @@ -44,7 +44,7 @@ static inline void SetAllUnknownDim(const int64_t rank, gert::Shape* output_shap | |||
| 44 | } | 44 | } |
| 45 | } | 45 | } |
| 46 | 46 | ||
| 47 | -static ge::graphStatus CheckAttrsValid(gert::InferShapeContext* context) | 47 | +static ge::graphStatus CheckAttrsValid(const gert::InferShapeContext* context) |
| 48 | { | 48 | { |
| 49 | auto attrs = context->GetAttrs(); | 49 | auto attrs = context->GetAttrs(); |
| 50 | OP_CHECK_NULL_WITH_CONTEXT(context, attrs); | 50 | OP_CHECK_NULL_WITH_CONTEXT(context, attrs); |
| @@ -71,7 +71,7 @@ static ge::graphStatus CheckAttrsValid(gert::InferShapeContext* context) | |||
| 71 | } | 71 | } |
| 72 | 72 | ||
| 73 | static ge::graphStatus HandleUnknownShape( | 73 | static ge::graphStatus HandleUnknownShape( |
| 74 | - gert::InferShapeContext* context, const gert::Shape* x1Shape, const gert::Shape* x2Shape, | 74 | + const gert::InferShapeContext* context, const gert::Shape* x1Shape, const gert::Shape* x2Shape, |
| 75 | const gert::Shape* gradShape, gert::Shape* yShape) | 75 | const gert::Shape* gradShape, gert::Shape* yShape) |
| 76 | { | 76 | { |
| 77 | size_t x1DimNum = x1Shape->GetDimNum(); | 77 | size_t x1DimNum = x1Shape->GetDimNum(); |
| @@ -224,4 +224,177 @@ TEST_F(MaxPoolGradInfer, max_pool_grad_inferdatatype_fp16) | |||
| 224 | } | 224 | } |
| 225 | } | 225 | } |
| 226 | 226 | ||
| 227 | +// ======================== Unknown Shape ======================== | ||
| 228 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_unknown_shape) | ||
| 229 | +{ | ||
| 230 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 231 | + | ||
| 232 | + gert::StorageShape x1Shape = {{1, 1, -1, -1}, {}}; | ||
| 233 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 234 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 235 | + gert::StorageShape yShape = {{}, {}}; | ||
| 236 | + | ||
| 237 | + auto holder = gert::InferShapeContextFaker() | ||
| 238 | + .NodeIoNum(3, 1) | ||
| 239 | + .IrInstanceNum({1, 1, 1}) | ||
| 240 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 241 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 242 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 243 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 244 | + .NodeAttrs( | ||
| 245 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 246 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 247 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 248 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 249 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 250 | + .OutputShapes({&yShape}) | ||
| 251 | + .Build(); | ||
| 252 | + | ||
| 253 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 254 | + gert::Shape* output = holder.GetContext<gert::InferShapeContext>()->GetOutputShape(0); | ||
| 255 | + ASSERT_EQ(Shape2String(*output), "[-1, -1, -1, -1]"); | ||
| 256 | +} | ||
| 257 | + | ||
| 258 | +// ======================== Unknown Rank ======================== | ||
| 259 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_unknown_rank) | ||
| 260 | +{ | ||
| 261 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 262 | + | ||
| 263 | + gert::StorageShape x1Shape = {{-2}, {}}; | ||
| 264 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 265 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 266 | + gert::StorageShape yShape = {{}, {}}; | ||
| 267 | + | ||
| 268 | + auto holder = gert::InferShapeContextFaker() | ||
| 269 | + .NodeIoNum(3, 1) | ||
| 270 | + .IrInstanceNum({1, 1, 1}) | ||
| 271 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 272 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 273 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 274 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 275 | + .NodeAttrs( | ||
| 276 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 277 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 278 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 279 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 280 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 281 | + .OutputShapes({&yShape}) | ||
| 282 | + .Build(); | ||
| 283 | + | ||
| 284 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 285 | +} | ||
| 286 | + | ||
| 287 | +// ======================== Failure Cases ======================== | ||
| 288 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_fail_invalid_format) | ||
| 289 | +{ | ||
| 290 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 291 | + | ||
| 292 | + gert::StorageShape x1Shape = {{1, 1, 4, 4}, {}}; | ||
| 293 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 294 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 295 | + gert::StorageShape yShape = {{}, {}}; | ||
| 296 | + | ||
| 297 | + auto holder = gert::InferShapeContextFaker() | ||
| 298 | + .NodeIoNum(3, 1) | ||
| 299 | + .IrInstanceNum({1, 1, 1}) | ||
| 300 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_FRACTAL_NZ, ge::Format::FORMAT_RESERVED) | ||
| 301 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_FRACTAL_NZ, ge::Format::FORMAT_RESERVED) | ||
| 302 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_FRACTAL_NZ, ge::Format::FORMAT_RESERVED) | ||
| 303 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 304 | + .NodeAttrs( | ||
| 305 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 306 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 307 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 308 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 309 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 310 | + .OutputShapes({&yShape}) | ||
| 311 | + .Build(); | ||
| 312 | + | ||
| 313 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 314 | +} | ||
| 315 | + | ||
| 316 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_fail_invalid_ksize_length) | ||
| 317 | +{ | ||
| 318 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 319 | + | ||
| 320 | + gert::StorageShape x1Shape = {{1, 1, 4, 4}, {}}; | ||
| 321 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 322 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 323 | + gert::StorageShape yShape = {{}, {}}; | ||
| 324 | + | ||
| 325 | + auto holder = gert::InferShapeContextFaker() | ||
| 326 | + .NodeIoNum(3, 1) | ||
| 327 | + .IrInstanceNum({1, 1, 1}) | ||
| 328 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 329 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 330 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 331 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 332 | + .NodeAttrs( | ||
| 333 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({2, 2})}, | ||
| 334 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 335 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 336 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 337 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 338 | + .OutputShapes({&yShape}) | ||
| 339 | + .Build(); | ||
| 340 | + | ||
| 341 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 342 | +} | ||
| 343 | + | ||
| 344 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_fail_invalid_strides_length) | ||
| 345 | +{ | ||
| 346 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 347 | + | ||
| 348 | + gert::StorageShape x1Shape = {{1, 1, 4, 4}, {}}; | ||
| 349 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 350 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 351 | + gert::StorageShape yShape = {{}, {}}; | ||
| 352 | + | ||
| 353 | + auto holder = gert::InferShapeContextFaker() | ||
| 354 | + .NodeIoNum(3, 1) | ||
| 355 | + .IrInstanceNum({1, 1, 1}) | ||
| 356 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 357 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 358 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 359 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 360 | + .NodeAttrs( | ||
| 361 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 362 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({2, 2, 1})}, | ||
| 363 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 364 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 365 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 366 | + .OutputShapes({&yShape}) | ||
| 367 | + .Build(); | ||
| 368 | + | ||
| 369 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 370 | +} | ||
| 371 | + | ||
| 372 | +TEST_F(MaxPoolGradInfer, max_pool_grad_infershape_fail_invalid_padding) | ||
| 373 | +{ | ||
| 374 | + auto inferShapeFunc = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolGrad")->infer_shape; | ||
| 375 | + | ||
| 376 | + gert::StorageShape x1Shape = {{1, 1, 4, 4}, {}}; | ||
| 377 | + gert::StorageShape x2Shape = {{1, 1, 2, 2}, {}}; | ||
| 378 | + gert::StorageShape gradShape = {{1, 1, 2, 2}, {}}; | ||
| 379 | + gert::StorageShape yShape = {{}, {}}; | ||
| 380 | + | ||
| 381 | + auto holder = gert::InferShapeContextFaker() | ||
| 382 | + .NodeIoNum(3, 1) | ||
| 383 | + .IrInstanceNum({1, 1, 1}) | ||
| 384 | + .NodeInputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 385 | + .NodeInputTd(1, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 386 | + .NodeInputTd(2, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 387 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::Format::FORMAT_NCHW, ge::Format::FORMAT_RESERVED) | ||
| 388 | + .NodeAttrs( | ||
| 389 | + {{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 390 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 391 | + {"padding", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 392 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}}) | ||
| 393 | + .InputShapes({&x1Shape, &x2Shape, &gradShape}) | ||
| 394 | + .OutputShapes({&yShape}) | ||
| 395 | + .Build(); | ||
| 396 | + | ||
| 397 | + ASSERT_EQ(inferShapeFunc(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 398 | +} | ||
| 399 | + | ||
| 227 | } // namespace | 400 | } // namespace |
| @@ -57,4 +57,357 @@ TEST_F(MaxPoolV3UT, InferShapeOk) { | |||
| 57 | 57 | ||
| 58 | EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | 58 | EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); |
| 59 | } | 59 | } |
| 60 | + | ||
| 61 | +// ======================== Padding Modes ======================== | ||
| 62 | + | ||
| 63 | +TEST_F(MaxPoolV3UT, InferShapeValid) { | ||
| 64 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 65 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 66 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 67 | + std::vector<void *> output_shapes_ref(1); | ||
| 68 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 69 | + | ||
| 70 | + auto holder = gert::InferShapeContextFaker() | ||
| 71 | + .NodeIoNum(1, 1) | ||
| 72 | + .IrInstanceNum({1}) | ||
| 73 | + .InputShapes({&x_shape}) | ||
| 74 | + .OutputShapes(output_shapes_ref) | ||
| 75 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 76 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 77 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 78 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 79 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 80 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 81 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 82 | + .Build(); | ||
| 83 | + | ||
| 84 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 85 | +} | ||
| 86 | + | ||
| 87 | +TEST_F(MaxPoolV3UT, InferShapeSame) { | ||
| 88 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 89 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 90 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 91 | + std::vector<void *> output_shapes_ref(1); | ||
| 92 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 93 | + | ||
| 94 | + auto holder = gert::InferShapeContextFaker() | ||
| 95 | + .NodeIoNum(1, 1) | ||
| 96 | + .IrInstanceNum({1}) | ||
| 97 | + .InputShapes({&x_shape}) | ||
| 98 | + .OutputShapes(output_shapes_ref) | ||
| 99 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 100 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 101 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("SAME")}, | ||
| 102 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 103 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 104 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 105 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 106 | + .Build(); | ||
| 107 | + | ||
| 108 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 109 | +} | ||
| 110 | + | ||
| 111 | +TEST_F(MaxPoolV3UT, InferShapeCalculatedCeilMode) { | ||
| 112 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 113 | + gert::StorageShape x_shape = {{1, 3, 7, 7}, {1, 3, 7, 7}}; | ||
| 114 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 115 | + std::vector<void *> output_shapes_ref(1); | ||
| 116 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 117 | + | ||
| 118 | + auto holder = gert::InferShapeContextFaker() | ||
| 119 | + .NodeIoNum(1, 1) | ||
| 120 | + .IrInstanceNum({1}) | ||
| 121 | + .InputShapes({&x_shape}) | ||
| 122 | + .OutputShapes(output_shapes_ref) | ||
| 123 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 3, 3})}, | ||
| 124 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 125 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 126 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 127 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 128 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 129 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(true)}}) | ||
| 130 | + .Build(); | ||
| 131 | + | ||
| 132 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 133 | +} | ||
| 134 | + | ||
| 135 | +// ======================== Global Pooling ======================== | ||
| 136 | + | ||
| 137 | +TEST_F(MaxPoolV3UT, InferShapeGlobalPooling) { | ||
| 138 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 139 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 140 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 141 | + std::vector<void *> output_shapes_ref(1); | ||
| 142 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 143 | + | ||
| 144 | + auto holder = gert::InferShapeContextFaker() | ||
| 145 | + .NodeIoNum(1, 1) | ||
| 146 | + .IrInstanceNum({1}) | ||
| 147 | + .InputShapes({&x_shape}) | ||
| 148 | + .OutputShapes(output_shapes_ref) | ||
| 149 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 150 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 151 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 152 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 153 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 154 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(true)}, | ||
| 155 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 156 | + .Build(); | ||
| 157 | + | ||
| 158 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 159 | +} | ||
| 160 | + | ||
| 161 | +// ======================== Unknown Rank ======================== | ||
| 162 | + | ||
| 163 | +TEST_F(MaxPoolV3UT, InferShapeUnknownRank) { | ||
| 164 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 165 | + gert::StorageShape x_shape = {{-2}, {}}; | ||
| 166 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 167 | + std::vector<void *> output_shapes_ref(1); | ||
| 168 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 169 | + | ||
| 170 | + auto holder = gert::InferShapeContextFaker() | ||
| 171 | + .NodeIoNum(1, 1) | ||
| 172 | + .IrInstanceNum({1}) | ||
| 173 | + .InputShapes({&x_shape}) | ||
| 174 | + .OutputShapes(output_shapes_ref) | ||
| 175 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 176 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 177 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 178 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 179 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 180 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 181 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 182 | + .Build(); | ||
| 183 | + | ||
| 184 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 185 | +} | ||
| 186 | + | ||
| 187 | +// ======================== NHWC Format ======================== | ||
| 188 | + | ||
| 189 | +TEST_F(MaxPoolV3UT, InferShapeNhwc) { | ||
| 190 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 191 | + gert::StorageShape x_shape = {{1, 8, 8, 3}, {1, 8, 8, 3}}; | ||
| 192 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 193 | + std::vector<void *> output_shapes_ref(1); | ||
| 194 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 195 | + | ||
| 196 | + auto holder = gert::InferShapeContextFaker() | ||
| 197 | + .NodeIoNum(1, 1) | ||
| 198 | + .IrInstanceNum({1}) | ||
| 199 | + .InputShapes({&x_shape}) | ||
| 200 | + .OutputShapes(output_shapes_ref) | ||
| 201 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 202 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 2, 2, 1})}, | ||
| 203 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("VALID")}, | ||
| 204 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 205 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NHWC")}, | ||
| 206 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 207 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 208 | + .Build(); | ||
| 209 | + | ||
| 210 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 211 | +} | ||
| 212 | + | ||
| 213 | +// ======================== Failure Cases ======================== | ||
| 214 | + | ||
| 215 | +TEST_F(MaxPoolV3UT, InferShapeFailInvalidPadding) { | ||
| 216 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 217 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 218 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 219 | + std::vector<void *> output_shapes_ref(1); | ||
| 220 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 221 | + | ||
| 222 | + auto holder = gert::InferShapeContextFaker() | ||
| 223 | + .NodeIoNum(1, 1) | ||
| 224 | + .IrInstanceNum({1}) | ||
| 225 | + .InputShapes({&x_shape}) | ||
| 226 | + .OutputShapes(output_shapes_ref) | ||
| 227 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 228 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 229 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("INVALID")}, | ||
| 230 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 231 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 232 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 233 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 234 | + .Build(); | ||
| 235 | + | ||
| 236 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 237 | +} | ||
| 238 | + | ||
| 239 | +TEST_F(MaxPoolV3UT, InferShapeFailKsizeLength) { | ||
| 240 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 241 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 242 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 243 | + std::vector<void *> output_shapes_ref(1); | ||
| 244 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 245 | + | ||
| 246 | + auto holder = gert::InferShapeContextFaker() | ||
| 247 | + .NodeIoNum(1, 1) | ||
| 248 | + .IrInstanceNum({1}) | ||
| 249 | + .InputShapes({&x_shape}) | ||
| 250 | + .OutputShapes(output_shapes_ref) | ||
| 251 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2})}, | ||
| 252 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 253 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 254 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 255 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 256 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 257 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 258 | + .Build(); | ||
| 259 | + | ||
| 260 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 261 | +} | ||
| 262 | + | ||
| 263 | +TEST_F(MaxPoolV3UT, InferShapeFailStridesLength) { | ||
| 264 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 265 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 266 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 267 | + std::vector<void *> output_shapes_ref(1); | ||
| 268 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 269 | + | ||
| 270 | + auto holder = gert::InferShapeContextFaker() | ||
| 271 | + .NodeIoNum(1, 1) | ||
| 272 | + .IrInstanceNum({1}) | ||
| 273 | + .InputShapes({&x_shape}) | ||
| 274 | + .OutputShapes(output_shapes_ref) | ||
| 275 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 276 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2})}, | ||
| 277 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 278 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 279 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 280 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 281 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 282 | + .Build(); | ||
| 283 | + | ||
| 284 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 285 | +} | ||
| 286 | + | ||
| 287 | +TEST_F(MaxPoolV3UT, InferShapeFailStridesZero) { | ||
| 288 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 289 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 290 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 291 | + std::vector<void *> output_shapes_ref(1); | ||
| 292 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 293 | + | ||
| 294 | + auto holder = gert::InferShapeContextFaker() | ||
| 295 | + .NodeIoNum(1, 1) | ||
| 296 | + .IrInstanceNum({1}) | ||
| 297 | + .InputShapes({&x_shape}) | ||
| 298 | + .OutputShapes(output_shapes_ref) | ||
| 299 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 300 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 0, 2})}, | ||
| 301 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 302 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0, 0})}, | ||
| 303 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 304 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 305 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 306 | + .Build(); | ||
| 307 | + | ||
| 308 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 309 | +} | ||
| 310 | + | ||
| 311 | +TEST_F(MaxPoolV3UT, InferShapeFailPadsLength) { | ||
| 312 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 313 | + gert::StorageShape x_shape = {{1, 3, 8, 8}, {1, 3, 8, 8}}; | ||
| 314 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 315 | + std::vector<void *> output_shapes_ref(1); | ||
| 316 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 317 | + | ||
| 318 | + auto holder = gert::InferShapeContextFaker() | ||
| 319 | + .NodeIoNum(1, 1) | ||
| 320 | + .IrInstanceNum({1}) | ||
| 321 | + .InputShapes({&x_shape}) | ||
| 322 | + .OutputShapes(output_shapes_ref) | ||
| 323 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 324 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 2, 2})}, | ||
| 325 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 326 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({0, 0, 0})}, | ||
| 327 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 328 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 329 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(false)}}) | ||
| 330 | + .Build(); | ||
| 331 | + | ||
| 332 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_FAILED); | ||
| 333 | +} | ||
| 334 | + | ||
| 335 | +// ======================== divRtn edge case ======================== | ||
| 336 | + | ||
| 337 | +TEST_F(MaxPoolV3UT, InferShapeCalculatedWithPadDivRtn) { | ||
| 338 | + auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_shape; | ||
| 339 | + gert::StorageShape x_shape = {{1, 3, 5, 5}, {1, 3, 5, 5}}; | ||
| 340 | + std::vector<gert::StorageShape> output_shapes(1); | ||
| 341 | + std::vector<void *> output_shapes_ref(1); | ||
| 342 | + output_shapes_ref[0] = &output_shapes[0]; | ||
| 343 | + | ||
| 344 | + auto holder = gert::InferShapeContextFaker() | ||
| 345 | + .NodeIoNum(1, 1) | ||
| 346 | + .IrInstanceNum({1}) | ||
| 347 | + .InputShapes({&x_shape}) | ||
| 348 | + .OutputShapes(output_shapes_ref) | ||
| 349 | + .NodeAttrs({{"ksize", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 4, 4})}, | ||
| 350 | + {"strides", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 3, 3})}, | ||
| 351 | + {"padding_mode", Ops::NN::AnyValue::CreateFrom<std::string>("CALCULATED")}, | ||
| 352 | + {"pads", Ops::NN::AnyValue::CreateFrom<std::vector<int64_t>>({1, 1, 1, 1})}, | ||
| 353 | + {"data_format", Ops::NN::AnyValue::CreateFrom<std::string>("NCHW")}, | ||
| 354 | + {"global_pooling", Ops::NN::AnyValue::CreateFrom<bool>(false)}, | ||
| 355 | + {"ceil_mode", Ops::NN::AnyValue::CreateFrom<bool>(true)}}) | ||
| 356 | + .Build(); | ||
| 357 | + | ||
| 358 | + EXPECT_EQ(infer_shape_func(holder.GetContext<gert::InferShapeContext>()), ge::GRAPH_SUCCESS); | ||
| 359 | +} | ||
| 360 | + | ||
| 361 | +} // namespace | ||
| 362 | + | ||
| 363 | +// ======================== InferDataType ======================== | ||
| 364 | +class MaxPoolV3InferDataType : public testing::Test { | ||
| 365 | +protected: | ||
| 366 | + static void SetUpTestCase() { | ||
| 367 | + std::cout << "MaxPoolV3InferDataType SetUp" << std::endl; | ||
| 368 | + } | ||
| 369 | + static void TearDownTestCase() { | ||
| 370 | + std::cout << "MaxPoolV3InferDataType TearDown" << std::endl; | ||
| 371 | + } | ||
| 372 | +}; | ||
| 373 | + | ||
| 374 | +TEST_F(MaxPoolV3InferDataType, InferDataTypeFP16) { | ||
| 375 | + auto infer_dtype_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_datatype; | ||
| 376 | + ASSERT_NE(infer_dtype_func, nullptr); | ||
| 377 | + | ||
| 378 | + ge::DataType input_dtype = ge::DT_FLOAT16; | ||
| 379 | + ge::DataType output_dtype = ge::DT_UNDEFINED; | ||
| 380 | + | ||
| 381 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 382 | + .NodeIoNum(1, 1) | ||
| 383 | + .IrInstanceNum({1}) | ||
| 384 | + .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 385 | + .NodeOutputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 386 | + .InputDataTypes({&input_dtype}) | ||
| 387 | + .OutputDataTypes({&output_dtype}) | ||
| 388 | + .Build(); | ||
| 389 | + | ||
| 390 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 391 | + ASSERT_EQ(infer_dtype_func(context), ge::GRAPH_SUCCESS); | ||
| 392 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT16); | ||
| 393 | +} | ||
| 394 | + | ||
| 395 | +TEST_F(MaxPoolV3InferDataType, InferDataTypeFP32) { | ||
| 396 | + auto infer_dtype_func = gert::OpImplRegistry::GetInstance().GetOpImpl("MaxPoolV3")->infer_datatype; | ||
| 397 | + | ||
| 398 | + ge::DataType input_dtype = ge::DT_FLOAT; | ||
| 399 | + ge::DataType output_dtype = ge::DT_UNDEFINED; | ||
| 400 | + | ||
| 401 | + auto holder = gert::InferDataTypeContextFaker() | ||
| 402 | + .NodeIoNum(1, 1) | ||
| 403 | + .IrInstanceNum({1}) | ||
| 404 | + .NodeInputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 405 | + .NodeOutputTd(0, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_RESERVED) | ||
| 406 | + .InputDataTypes({&input_dtype}) | ||
| 407 | + .OutputDataTypes({&output_dtype}) | ||
| 408 | + .Build(); | ||
| 409 | + | ||
| 410 | + auto context = holder.GetContext<gert::InferDataTypeContext>(); | ||
| 411 | + ASSERT_EQ(infer_dtype_func(context), ge::GRAPH_SUCCESS); | ||
| 412 | + EXPECT_EQ(context->GetOutputDataType(0), ge::DT_FLOAT); | ||
| 60 | } | 413 | } |