已合并
MaxPool3dGrad支持CALCULATED #8211
MaxPool3dGrad支持CALCULATED #8211
已合并
sikaiwei创建于 26 天前
3 个文件变更+31-27
@@ -78,7 +78,7 @@
78 <tr>78 <tr>
79 <td>pads</td>79 <td>pads</td>
80 <td>属性</td>80 <td>属性</td>
81- <td>当pad模式为CALCULATED时,输入的orig_x的shape进行左右上下前后的扩展,填补的数据为-inf,暂不生效。</td>81+ <td>当pad模式为CALCULATED时,输入的orig_x的shape进行左右上下前后的扩展,填补的数据为-inf。</td>
82 <td>-</td>82 <td>-</td>
83 <td>-</td>83 <td>-</td>
84 </tr>84 </tr>
@@ -42,6 +42,7 @@ constexpr size_t C_DIM_OFFSET = 4;
42constexpr size_t D_DIM_OFFSET = 1;42constexpr size_t D_DIM_OFFSET = 1;
43constexpr size_t H_DIM_OFFSET = 2;43constexpr size_t H_DIM_OFFSET = 2;
44constexpr size_t W_DIM_OFFSET = 3;44constexpr size_t W_DIM_OFFSET = 3;
45+constexpr size_t PADS_SIZE = 6;
45 46 
46static const gert::Shape& EnsureNotScalar(const gert::Shape& inShape)47static const gert::Shape& EnsureNotScalar(const gert::Shape& inShape)
47{48{
@@ -167,6 +168,17 @@ ge::graphStatus MaxPool3DGradSimtTiling::GetShapeAttrsInfo()
167 int64_t wPadNeed = std::max(168 int64_t wPadNeed = std::max(
168 int64_t{0}, int64_t((gradShape.GetDim(wDimPos) - 1) * strideW + kSizeW - inputShape.GetDim(wDimPos)));169 int64_t{0}, int64_t((gradShape.GetDim(wDimPos) - 1) * strideW + kSizeW - inputShape.GetDim(wDimPos)));
169 padsW = wPadNeed / DIGIT_TWO;170 padsW = wPadNeed / DIGIT_TWO;
171+ } else if (padModeStr == "CALCULATED") {
172+ auto padsList = runtimeAttrs->GetListInt(PADDING_POS);
173+ OPS_CHECK_NULL_WITH_CONTEXT(context_, padsList);
174+ if (padsList->GetSize() != PADS_SIZE) {
175+ OP_LOGE_FOR_INVALID_LISTSIZE(opName_, "Length of pads", std::to_string(padsList->GetSize()).c_str(), "6");
176+ return ge::GRAPH_FAILED;
177+ }
178+ auto padsVector = padsList->GetData();
179+ padsD = padsVector[0];
G
Gguankarl18 天前

取值之前先判断长度

likedislike
sikaiwei
17 天前 评论:
180+ padsH = padsVector[2];
181+ padsW = padsVector[4];
atomgit-bot
atomgit-botatomgit-bot26 天前

🟠 High Priority

第 171 行直接链式调用 GetListInt(PADDING_POS)->GetData(),未对 GetListInt 的返回值进行空指针检查。如果 PADDING_POS 对应的属性不存在或为空,GetListInt 可能返回 nullptr,后续 ->GetData() 会导致空指针解引用崩溃。参考同一子目录下的 max_pool_3d_tiling_common.cpp 第 324-325 行,同类代码均先保存 GetListInt 返回值再通过 OPS_CHECK_NULL_WITH_CONTEXT 做空检查。

建议:将 GetListInt 返回值保存到临时变量,增加空指针检查后再使用。参考 max_pool_3d_tiling_common.cpp 的同模式代码。

改动建议
181
+ } else if (padModeStr == "CALCULATED") {
182
+ auto padding = runtimeAttrs->GetListInt(PADDING_POS);
183
+ OPS_CHECK_NULL_WITH_CONTEXT(context_, padding);
184
+ auto padsVector = padding->GetData();
185
+ padsD = padsVector[0];
186
+ padsH = padsVector[2];
181
- padsW = padsVector[4];
187
+ padsW = padsVector[4];
应用建议
likedislike
170 }182 }
171 if (padsD * DOUB > kSizeD || padsH * DOUB > kSizeH || padsW * DOUB > kSizeW) {183 if (padsD * DOUB > kSizeD || padsH * DOUB > kSizeH || padsW * DOUB > kSizeW) {
172 OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(184 OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(
@@ -68,34 +68,26 @@ ge::graphStatus CheckPadInfoForMaxPool3DGrad(const gert::InferShapeContext* cont
68 ge::GRAPH_SUCCESS),68 ge::GRAPH_SUCCESS),
69 OP_LOGE(context, "Cannot get platform info!"), return ge::GRAPH_FAILED);69 OP_LOGE(context, "Cannot get platform info!"), return ge::GRAPH_FAILED);
70 OP_LOGD(context, "soc version is %s", platform_info.str_info.short_soc_version.c_str());70 OP_LOGD(context, "soc version is %s", platform_info.str_info.short_soc_version.c_str());
71- if (platform_info.str_info.short_soc_version == "Ascend950") {71+ if (padding != "SAME" && padding != "VALID" && padding != "CALCULATED") {
72- if (padding != "SAME" && padding != "VALID") {72+ OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "padMode", padding.c_str(),
73- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "padMode", padding.c_str(),73+ "only support SAME, VALID or CALCULATED padding mode");
74- "only support SAME or VALID padding mode on Ascend950");74+ return ge::GRAPH_FAILED;
75+ }
76+ if (padding == "CALCULATED") {
77+ auto pads = attrs->GetAttrPointer<gert::ContinuousVector>(PADS_ATTR_INDEX);
78+ OP_CHECK_NULL_WITH_CONTEXT(context, pads);
79+ if (pads->GetSize() != PADS_SIZE) {
80+ OP_LOGE_FOR_INVALID_LISTSIZE(opName_, "Length of pads", std::to_string(pads->GetSize()).c_str(), "6");
75 return ge::GRAPH_FAILED;81 return ge::GRAPH_FAILED;
76 }82 }
77- } else {
78- if (padding != "SAME" && padding != "VALID" && padding != "CALCULATED") {
79- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, "padMode", padding.c_str(),
80- "only support SAME, VALID or CALCULATED padding mode");
81- return ge::GRAPH_FAILED;
82- }
83- if (padding == "CALCULATED") {
84- auto pads = attrs->GetAttrPointer<gert::ContinuousVector>(PADS_ATTR_INDEX);
85- OP_CHECK_NULL_WITH_CONTEXT(context, pads);
86- if (pads->GetSize() != PADS_SIZE) {
87- OP_LOGE_FOR_INVALID_LISTSIZE(opName_, "Length of pads", std::to_string(pads->GetSize()).c_str(), "6");
88- return ge::GRAPH_FAILED;
89- }
90 83 
91- auto pads_data = static_cast<const int64_t*>(pads->GetData());84+ auto pads_data = static_cast<const int64_t*>(pads->GetData());
92- for (uint32_t i = 0; i < static_cast<uint32_t>(pads->GetSize()); i++) {85+ for (uint32_t i = 0; i < static_cast<uint32_t>(pads->GetSize()); i++) {
93- if (pads_data[i] < 0) {86+ if (pads_data[i] < 0) {
94- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(87+ OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName_, (std::string("pads[") + std::to_string(i) + "]").c_str(),
95- opName_, (std::string("pads[") + std::to_string(i) + "]").c_str(),88+ std::to_string(pads_data[i]).c_str(),
96- std::to_string(pads_data[i]).c_str(), "pads value should be >= 0");89+ "pads value should be >= 0");
97- return ge::GRAPH_FAILED;90+ return ge::GRAPH_FAILED;
98- }
99 }91 }
100 }92 }
101 }93 }
@@ -206,4 +198,4 @@ static ge::graphStatus InferDataTypeForMaxPool3DGrad(gert::InferDataTypeContext*
206}198}
207 199 
208IMPL_OP_INFERSHAPE(MaxPool3DGrad).InferShape(InferShapeForMaxPool3DGrad).InferDataType(InferDataTypeForMaxPool3DGrad);200IMPL_OP_INFERSHAPE(MaxPool3DGrad).InferShape(InferShapeForMaxPool3DGrad).InferDataType(InferDataTypeForMaxPool3DGrad);
209-} // namespace ops201+} // namespace ops