已合并
MaxPool3dGrad支持CALCULATED #8211
sikaiwei创建于 26 天前
MaxPool3dGrad支持CALCULATED #8211
已合并
共 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; | |||||||||||||||||||
| 42 | constexpr size_t D_DIM_OFFSET = 1; | 42 | constexpr size_t D_DIM_OFFSET = 1; | ||||||||||||||||
| 43 | constexpr size_t H_DIM_OFFSET = 2; | 43 | constexpr size_t H_DIM_OFFSET = 2; | ||||||||||||||||
| 44 | constexpr size_t W_DIM_OFFSET = 3; | 44 | constexpr size_t W_DIM_OFFSET = 3; | ||||||||||||||||
| 45 | +constexpr size_t PADS_SIZE = 6; | ||||||||||||||||||
| 45 | 46 | ||||||||||||||||||
| 46 | static const gert::Shape& EnsureNotScalar(const gert::Shape& inShape) | 47 | static 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 | |||||||||||||||||||
| 180 | + padsH = padsVector[2]; | ||||||||||||||||||
| 181 | + padsW = padsVector[4]; | ||||||||||||||||||
🟠 High Priority 第 171 行直接链式调用 建议:将 改动建议
![]() ![]() | |||||||||||||||||||
| 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 | ||
| 208 | IMPL_OP_INFERSHAPE(MaxPool3DGrad).InferShape(InferShapeForMaxPool3DGrad).InferDataType(InferDataTypeForMaxPool3DGrad); | 200 | IMPL_OP_INFERSHAPE(MaxPool3DGrad).InferShape(InferShapeForMaxPool3DGrad).InferDataType(InferDataTypeForMaxPool3DGrad); |
| 209 | -} // namespace ops | 201 | +} // namespace ops |


取值之前先判断长度