已合并
conv2d A16W8 InnerBatch LoadAL1 bug fix #7767
conv2d A16W8 InnerBatch LoadAL1 bug fix #7767
已合并
xinweiliu创建于 7月21日
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 Conv2dFunc464}; // namespace Conv2dFunc
463 465 
464-#endif // CONV2D_V2_INSTR_BASE_IMPL_H466+#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 // hiLoadTail559 // 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 Conv2dFunc629}; // namespace Conv2dFunc
632 630 
633-#endif // CONV2D_V2_INSTR_IMPL_H631+#endif // CONV2D_V2_INSTR_IMPL_H