已合并
修改cleancode以及补充池化类算子ut #4888
SimonZzz创建于 5月15日
修改cleancode以及补充池化类算子ut #4888
已合并
SimonZzz创建于 5月15日
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+#include <gtest/gtest.h>
17+#include <iostream>
18+#include <vector>
19+ 
20+#include "infershape_case_executor.h"
21+#include "kernel_run_context_facker.h"
22+#include "register/op_impl_registry.h"
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+#include "exe_graph/runtime/storage_format.h"
17+#include "exe_graph/runtime/storage_shape.h"
18+#include "gtest/gtest.h"
19+#include "kernel_run_context_facker.h"
20+#include "log/log.h"
21+#include "register/op_impl_registry.h"
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+#include <iostream>
12+#include <sstream>
13+#include "exe_graph/runtime/storage_format.h"
14+#include "exe_graph/runtime/storage_shape.h"
15+#include <gtest/gtest.h>
16+#include "kernel_run_context_facker.h"
17+#include "register/op_impl_registry.h"
18+#include "log/log.h"
19+#include "platform/platform_info.h"
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 
73static ge::graphStatus HandleUnknownShape(73static 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} // namespace400} // 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}