已合并
conv2d A16W8 InnerBatch LoadAL1 bug fix #7767
xinweiliu创建于 7月21日
conv2d A16W8 InnerBatch LoadAL1 bug fix #7767
已合并
共 2 个文件变更+25-25
| @@ -124,9 +124,10 @@ public: | |||
| 124 | { | 124 | { |
| 125 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { | 125 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { |
| 126 | if constexpr (Intf::formatFmap == ConvFormat::NCHW && Intf::c04Flag) { | 126 | if constexpr (Intf::formatFmap == ConvFormat::NCHW && Intf::c04Flag) { |
| 127 | - Load3DSetFMatrixCal(self_->ctx.innerBatch, | 127 | + Load3DSetFMatrixCal( |
| 128 | - AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, Intf::k0) / C04_CIN_SIZE, | 128 | + self_->ctx.innerBatch, |
| 129 | - padList); | 129 | + AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, Intf::k0FmapTail) / C04_CIN_SIZE, |
| 130 | + padList); | ||
| 130 | } else { | 131 | } else { |
| 131 | Load3DSetFMatrixCal(self_->ctx.innerBatch * hiLoadL1, self_->ctx.convTilingData->orgWi, padList); | 132 | Load3DSetFMatrixCal(self_->ctx.innerBatch * hiLoadL1, self_->ctx.convTilingData->orgWi, padList); |
| 132 | } | 133 | } |
| @@ -189,7 +190,7 @@ private: | |||
| 189 | intriParams.srcDValue = self_->ctx.convTilingData->orgHixWi; | 190 | intriParams.srcDValue = self_->ctx.convTilingData->orgHixWi; |
| 190 | intriParams.srcDnMatrixStride = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->orgHixWi; | 191 | intriParams.srcDnMatrixStride = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->orgHixWi; |
| 191 | intriParams.dstNzNStride = 1; | 192 | intriParams.dstNzNStride = 1; |
| 192 | - intriParams.dstNzMatrixStride = AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, Intf::k0); | 193 | + intriParams.dstNzMatrixStride = AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, Intf::k0FmapTail); |
| 193 | } | 194 | } |
| 194 | 195 | ||
| 195 | __aicore__ inline void SetDn2NzIntriParams(Dn2NzParams& intriParams, uint64_t kAL1Iter) | 196 | __aicore__ inline void SetDn2NzIntriParams(Dn2NzParams& intriParams, uint64_t kAL1Iter) |
| @@ -210,10 +211,10 @@ private: | |||
| 210 | 211 | ||
| 211 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { | 212 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { |
| 212 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; | 213 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; |
| 213 | - intriParams.dstNzMatrixStride = realHixWi * Intf::k0; | 214 | + intriParams.dstNzMatrixStride = realHixWi * Intf::k0FmapTail; |
| 214 | } else { | 215 | } else { |
| 215 | intriParams.dstNzC0Stride = realHixWi; | 216 | intriParams.dstNzC0Stride = realHixWi; |
| 216 | - intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0) * realHixWi; | 217 | + intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0FmapTail) * realHixWi; |
| 217 | } | 218 | } |
| 218 | } | 219 | } |
| 219 | 220 | ||
| @@ -227,7 +228,8 @@ private: | |||
| 227 | intriParams.nValue = realHixWi; | 228 | intriParams.nValue = realHixWi; |
| 228 | intriParams.srcNdMatrixStride = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->orgHixWi; | 229 | intriParams.srcNdMatrixStride = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->orgHixWi; |
| 229 | intriParams.dstNzNStride = 1; | 230 | intriParams.dstNzNStride = 1; |
| 230 | - intriParams.dstNzMatrixStride = AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, Intf::k0); | 231 | + intriParams.dstNzMatrixStride = AlignB(C04_CIN_SIZE * self_->ctx.convTilingData->orgHixWi, |
| 232 | + Intf::k0FmapTail); | ||
| 231 | } | 233 | } |
| 232 | intriParams.dValue = self_->ctx.convTilingData->orgCi; | 234 | intriParams.dValue = self_->ctx.convTilingData->orgCi; |
| 233 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi; | 235 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi; |
| @@ -251,10 +253,10 @@ private: | |||
| 251 | 253 | ||
| 252 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { | 254 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { |
| 253 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; | 255 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; |
| 254 | - intriParams.dstNzMatrixStride = realHixWi * Intf::k0; | 256 | + intriParams.dstNzMatrixStride = realHixWi * Intf::k0FmapTail; |
| 255 | } else { | 257 | } else { |
| 256 | intriParams.dstNzC0Stride = realHixWi; | 258 | intriParams.dstNzC0Stride = realHixWi; |
| 257 | - intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0) * realHixWi; | 259 | + intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0FmapTail) * realHixWi; |
| 258 | } | 260 | } |
| 259 | } | 261 | } |
| 260 | 262 | ||
| @@ -271,10 +273,10 @@ private: | |||
| 271 | 273 | ||
| 272 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { | 274 | if constexpr (Intf::ConvParam::innerBatch == static_cast<int8_t>(ConvInnerBatch::KERNEL_1X1_MULTI_BATCH)) { |
| 273 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; | 275 | intriParams.dstNzC0Stride = self_->ctx.innerBatch * realHixWi; |
| 274 | - intriParams.dstNzMatrixStride = realHixWi * Intf::k0; | 276 | + intriParams.dstNzMatrixStride = realHixWi * Intf::k0FmapTail; |
| 275 | } else { | 277 | } else { |
| 276 | intriParams.dstNzC0Stride = realHixWi; | 278 | intriParams.dstNzC0Stride = realHixWi; |
| 277 | - intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0) * realHixWi; | 279 | + intriParams.dstNzMatrixStride = AlignB(al1Ci, Intf::k0FmapTail) * realHixWi; |
| 278 | } | 280 | } |
| 279 | } | 281 | } |
| 280 | 282 | ||
| @@ -461,4 +463,4 @@ private: | |||
| 461 | 463 | ||
| 462 | }; // namespace Conv2dFunc | 464 | }; // namespace Conv2dFunc |
| 463 | 465 | ||
| 464 | -#endif // CONV2D_V2_INSTR_BASE_IMPL_H | 466 | +#endif // CONV2D_V2_INSTR_BASE_IMPL_H |
| @@ -113,7 +113,7 @@ public: | |||
| 113 | intriParams.dnNum = this->hiLoadL1; | 113 | intriParams.dnNum = this->hiLoadL1; |
| 114 | intriParams.nValue = this->wiLoadL1; | 114 | intriParams.nValue = this->wiLoadL1; |
| 115 | intriParams.srcDnMatrixStride = self_->ctx.convTilingData->orgWi; | 115 | intriParams.srcDnMatrixStride = self_->ctx.convTilingData->orgWi; |
| 116 | - intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0; | 116 | + intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0FmapTail; |
| 117 | } | 117 | } |
| 118 | 118 | ||
| 119 | if constexpr (Intf::groupOptPreloadFlag) { | 119 | if constexpr (Intf::groupOptPreloadFlag) { |
| @@ -181,7 +181,7 @@ public: | |||
| 181 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi; | 181 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi; |
| 182 | intriParams.dstNzC0Stride = this->hiLoadL1 * this->wiLoadL1; | 182 | intriParams.dstNzC0Stride = this->hiLoadL1 * this->wiLoadL1; |
| 183 | intriParams.dstNzNStride = 1; | 183 | intriParams.dstNzNStride = 1; |
| 184 | - intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0; | 184 | + intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0FmapTail; |
| 185 | } | 185 | } |
| 186 | 186 | ||
| 187 | __aicore__ inline void SetNd2NzIntriParamsInputHWNC(Nd2NzParams& intriParams, uint64_t kAL1Iter) | 187 | __aicore__ inline void SetNd2NzIntriParamsInputHWNC(Nd2NzParams& intriParams, uint64_t kAL1Iter) |
| @@ -195,7 +195,7 @@ public: | |||
| 195 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->batch; | 195 | intriParams.srcDValue = self_->ctx.convTilingData->orgCi * self_->ctx.convTilingData->batch; |
| 196 | intriParams.dstNzC0Stride = this->hiLoadL1 * this->wiLoadL1; | 196 | intriParams.dstNzC0Stride = this->hiLoadL1 * this->wiLoadL1; |
| 197 | intriParams.dstNzNStride = 1; | 197 | intriParams.dstNzNStride = 1; |
| 198 | - intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0; | 198 | + intriParams.dstNzMatrixStride = this->wiLoadL1 * Intf::k0FmapTail; |
| 199 | } | 199 | } |
| 200 | 200 | ||
| 201 | __aicore__ inline void SetNd2NzIntriParamsC04InputHWNC(Nd2NzParams& intriParams) | 201 | __aicore__ inline void SetNd2NzIntriParamsC04InputHWNC(Nd2NzParams& intriParams) |
| @@ -550,12 +550,10 @@ private: | |||
| 550 | uint64_t step = CeilDiv(this->hiLoadL1, hiLoadPerStep); | 550 | uint64_t step = CeilDiv(this->hiLoadL1, hiLoadPerStep); |
| 551 | uint64_t hiLoadTail = this->hiLoadL1 % hiLoadPerStep; | 551 | uint64_t hiLoadTail = this->hiLoadL1 % hiLoadPerStep; |
| 552 | if (hiLoadTail == 0) { | 552 | if (hiLoadTail == 0) { |
| 553 | - SetNd2NzIntriParamsC04(intriParams, step, | 553 | + SetNd2NzIntriParamsC04(intriParams, step, hiLoadPerStep * self_->ctx.convTilingData->orgWi); |
| 554 | - hiLoadPerStep * self_->ctx.convTilingData->orgWi); | ||
| 555 | DataCopy<typename Intf::FmapT, true>(self_->ctx.al1, self_->ctx.agm[aL1GmOffset], intriParams); | 554 | DataCopy<typename Intf::FmapT, true>(self_->ctx.al1, self_->ctx.agm[aL1GmOffset], intriParams); |
| 556 | } else { | 555 | } else { |
| 557 | - SetNd2NzIntriParamsC04(intriParams, step - 1, | 556 | + SetNd2NzIntriParamsC04(intriParams, step - 1, hiLoadPerStep * self_->ctx.convTilingData->orgWi); |
| 558 | - hiLoadPerStep * self_->ctx.convTilingData->orgWi); | ||
| 559 | DataCopy<typename Intf::FmapT, true>(self_->ctx.al1, self_->ctx.agm[aL1GmOffset], intriParams); | 557 | DataCopy<typename Intf::FmapT, true>(self_->ctx.al1, self_->ctx.agm[aL1GmOffset], intriParams); |
| 560 | 558 | ||
| 561 | // hiLoadTail | 559 | // hiLoadTail |
| @@ -563,8 +561,8 @@ private: | |||
| 563 | uint64_t aL1Offset = offset * C04_CIN_SIZE; | 561 | uint64_t aL1Offset = offset * C04_CIN_SIZE; |
| 564 | aL1GmOffset += offset * self_->ctx.convTilingData->orgCi; | 562 | aL1GmOffset += offset * self_->ctx.convTilingData->orgCi; |
| 565 | SetNd2NzIntriParamsC04(intriParams, 1, hiLoadTail * self_->ctx.convTilingData->orgWi); | 563 | SetNd2NzIntriParamsC04(intriParams, 1, hiLoadTail * self_->ctx.convTilingData->orgWi); |
| 566 | - DataCopy<typename Intf::FmapT, true>( | 564 | + DataCopy<typename Intf::FmapT, true>(self_->ctx.al1[aL1Offset], self_->ctx.agm[aL1GmOffset], |
| 567 | - self_->ctx.al1[aL1Offset], self_->ctx.agm[aL1GmOffset], intriParams); | 565 | + intriParams); |
| 568 | } | 566 | } |
| 569 | } else { | 567 | } else { |
| 570 | SetNd2NzIntriParamsC04(intriParams, 1, aL1Mi); | 568 | SetNd2NzIntriParamsC04(intriParams, 1, aL1Mi); |
| @@ -572,8 +570,8 @@ private: | |||
| 572 | } | 570 | } |
| 573 | } | 571 | } |
| 574 | 572 | ||
| 575 | - __aicore__ inline void SetNd2NzIntriParamsC04(Nd2NzParams &intriParams, uint64_t ndNum, uint64_t nValue) | 573 | + __aicore__ inline void SetNd2NzIntriParamsC04(Nd2NzParams& intriParams, uint64_t ndNum, uint64_t nValue) |
| 576 | - { | 574 | + { |
| 577 | intriParams.ndNum = ndNum; | 575 | intriParams.ndNum = ndNum; |
| 578 | intriParams.nValue = nValue; | 576 | intriParams.nValue = nValue; |
| 579 | intriParams.dValue = self_->ctx.convTilingData->singleCoreCi; | 577 | intriParams.dValue = self_->ctx.convTilingData->singleCoreCi; |
| @@ -582,7 +580,7 @@ private: | |||
| 582 | intriParams.dstNzNStride = 1; | 580 | intriParams.dstNzNStride = 1; |
| 583 | intriParams.dstNzMatrixStride = nValue * C04_CIN_SIZE; | 581 | intriParams.dstNzMatrixStride = nValue * C04_CIN_SIZE; |
| 584 | } | 582 | } |
| 585 | - __aicore__ inline void SetNd2NzIntriParams(Nd2NzParams &intriParams, uint64_t kAL1Iter) | 583 | + __aicore__ inline void SetNd2NzIntriParams(Nd2NzParams& intriParams, uint64_t kAL1Iter) |
| 586 | { | 584 | { |
| 587 | uint64_t aL1Mi = this->hiLoadL1 * self_->ctx.convTilingData->orgWi; | 585 | uint64_t aL1Mi = this->hiLoadL1 * self_->ctx.convTilingData->orgWi; |
| 588 | intriParams.ndNum = 1; | 586 | intriParams.ndNum = 1; |
| @@ -630,4 +628,4 @@ private: | |||
| 630 | 628 | ||
| 631 | }; // namespace Conv2dFunc | 629 | }; // namespace Conv2dFunc |
| 632 | 630 | ||
| 633 | -#endif // CONV2D_V2_INSTR_IMPL_H | 631 | +#endif // CONV2D_V2_INSTR_IMPL_H |