已合并
perf: softmax_grad算子FMA指令融合优化 #7040
zhuzixian-lr创建于 7月4日
perf: softmax_grad算子FMA指令融合优化 #7040
已合并
zhuzixian-lr创建于 7月4日
5 个文件变更+25-25
@@ -178,8 +178,8 @@ __aicore__ inline void SoftmaxGradAR<T>::NormComputeSmallR(const int64_t aSize)
178 Duplicate(reg2, reg2, pFull);178 Duplicate(reg2, reg2, pFull);
179 179 
180 Mul(reg1, reg0, reg1, pMask);180 Mul(reg1, reg0, reg1, pMask);
181- Mul(reg0, reg0, reg2, pMask);181+ Neg(reg0, reg0, pMask);
182- Sub(reg1, reg1, reg0, pMask);182+ MulAddDst(reg1, reg0, reg2, pMask);
183 183 
184 StoreTensorForDtypeTOut(dst, reg1, pMask, i * rAligned);184 StoreTensorForDtypeTOut(dst, reg1, pMask, i * rAligned);
185 }185 }
@@ -214,13 +214,13 @@ __aicore__ inline void SoftmaxGradAR<T>::NormComputeSmallR(const int64_t aSize)
214 Duplicate(reg2, reg2, pFull);214 Duplicate(reg2, reg2, pFull);
215 215 
216 Mul(reg1, reg0, reg1, pFull);216 Mul(reg1, reg0, reg1, pFull);
217- Mul(reg0, reg0, reg2, pFull);217+ Neg(reg0, reg0, pFull);
218- Sub(reg1, reg1, reg0, pFull);218+ MulAddDst(reg1, reg0, reg2, pFull);
219 StoreTensorForDtypeTOut(dst, reg1, pFull, i * rAligned);219 StoreTensorForDtypeTOut(dst, reg1, pFull, i * rAligned);
220 220 
221 Mul(reg1_1, reg0_1, reg1_1, pMask);221 Mul(reg1_1, reg0_1, reg1_1, pMask);
222- Mul(reg0_1, reg0_1, reg2, pMask);222+ Neg(reg0_1, reg0_1, pMask);
223- Sub(reg1_1, reg1_1, reg0_1, pMask);223+ MulAddDst(reg1_1, reg0_1, reg2, pMask);
224 StoreTensorForDtypeTOut(dst, reg1_1, pMask, i * rAligned + VL_FP32);224 StoreTensorForDtypeTOut(dst, reg1_1, pMask, i * rAligned + VL_FP32);
225 }225 }
226 }226 }
@@ -377,8 +377,8 @@ __aicore__ inline void SoftmaxGradAR<T>::NormComputePost(const LocalTensor<T>& d
377 LoadTensorForDtypeTIn(x0, reg0, maskOri, offset);377 LoadTensorForDtypeTIn(x0, reg0, maskOri, offset);
378 LoadTensorForDtypeTIn(x1, reg1, maskOri, offset);378 LoadTensorForDtypeTIn(x1, reg1, maskOri, offset);
379 Mul(reg1, reg0, reg1, maskOri);379 Mul(reg1, reg0, reg1, maskOri);
380- Mul(reg0, reg0, reg2, maskOri);380+ Neg(reg0, reg0, maskOri);
381- Sub(reg1, reg1, reg0, maskOri);381+ MulAddDst(reg1, reg0, reg2, maskOri);
382 StoreTensorForDtypeTOut(dst, reg1, maskOri, offset);382 StoreTensorForDtypeTOut(dst, reg1, maskOri, offset);
383 }383 }
384 }384 }
@@ -415,8 +415,8 @@ __aicore__ inline void SoftmaxGradAR<T>::NormComputePost(const LocalTensor<T>& d
415 LoadTensorForDtypeTIn(x1, reg1, maskOri, offset);415 LoadTensorForDtypeTIn(x1, reg1, maskOri, offset);
416 416 
417 Mul(reg1, reg0, reg1, maskOri);417 Mul(reg1, reg0, reg1, maskOri);
418- Mul(reg0, reg0, reg2, maskOri);418+ Neg(reg0, reg0, maskOri);
419- Sub(reg1, reg1, reg0, maskOri);419+ MulAddDst(reg1, reg0, reg2, maskOri);
420 StoreTensorForDtypeTOut(dst, reg1, maskOri, offset);420 StoreTensorForDtypeTOut(dst, reg1, maskOri, offset);
421 }421 }
422 }422 }
@@ -283,9 +283,9 @@ __aicore__ inline void SoftmaxGradArRecompute<T>::CalcOutVF(uint32_t ubFactor)
283 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, offset);283 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, offset);
284 284 
285 Mul(x1Reg, x0Reg, x1Reg, pregMask);285 Mul(x1Reg, x0Reg, x1Reg, pregMask);
286- Mul(x0Reg, x0Reg, sumReg, pregMask);286+ Neg(x0Reg, x0Reg, pregMask);
287- Sub(x0Reg, x1Reg, x0Reg, pregMask);287+ MulAddDst(x1Reg, x0Reg, sumReg, pregMask);
288- StoreTensorForDtypeTOut(yLocal, x0Reg, pregMask, offset);288+ StoreTensorForDtypeTOut(yLocal, x1Reg, pregMask, offset);
289 }289 }
290 }290 }
291 291 
@@ -181,14 +181,14 @@ private:
181 LoadTensorForDtypeT(x0Local, x0Reg, pregMask, xOffset);181 LoadTensorForDtypeT(x0Local, x0Reg, pregMask, xOffset);
182 182 
183 DataCopy(x1Reg, tmpAddr2 + xOffset);183 DataCopy(x1Reg, tmpAddr2 + xOffset);
184- Mul(x0Reg, x0Reg, sumReg, pregMask);184+ Neg(x0Reg, x0Reg, pregMask);
185- Sub(x0Reg, x1Reg, x0Reg, pregMask);185+ MulAddDst(x1Reg, x0Reg, sumReg, pregMask);
186 186 
187 if constexpr (xToFp32_) {187 if constexpr (xToFp32_) {
188- MicroAPI::DataCopy(tmpAddrTy + xOffset, x0Reg, pregMask);188+ MicroAPI::DataCopy(tmpAddrTy + xOffset, x1Reg, pregMask);
189 } else { // fp16、bf16189 } else { // fp16、bf16
190 RegTensor<T> xFp16;190 RegTensor<T> xFp16;
191- MicroAPI::Cast<T, float, castTraitFp32ToFp16>(xFp16, x0Reg, pregMask);191+ MicroAPI::Cast<T, float, castTraitFp32ToFp16>(xFp16, x1Reg, pregMask);
192 MicroAPI::DataCopy<T, MicroAPI::StoreDist::DIST_PACK_B32>(tmpAddrTy + xOffset, xFp16, pregMask);192 MicroAPI::DataCopy<T, MicroAPI::StoreDist::DIST_PACK_B32>(tmpAddrTy + xOffset, xFp16, pregMask);
193 }193 }
194 }194 }
@@ -177,14 +177,14 @@ private:
177 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, xOffset);177 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, xOffset);
178 178 
179 Mul(x1Reg, x0Reg, x1Reg, pregMask);179 Mul(x1Reg, x0Reg, x1Reg, pregMask);
180- Mul(x0Reg, x0Reg, sumReg, pregMask);180+ Neg(x0Reg, x0Reg, pregMask);
181- Sub(x0Reg, x1Reg, x0Reg, pregMask);181+ MulAddDst(x1Reg, x0Reg, sumReg, pregMask);
182 182 
183 if constexpr (IsSameType<T, float>::value) {183 if constexpr (IsSameType<T, float>::value) {
184- DataCopy(((__local_mem__ float*)yLocal) + xOffset, x0Reg, pregMask);184+ DataCopy(((__local_mem__ float*)yLocal) + xOffset, x1Reg, pregMask);
185 } else { // fp16、bf16185 } else { // fp16、bf16
186 RegTensor<T> xFp16;186 RegTensor<T> xFp16;
187- Cast<T, float, castTraitFp32ToFp16>(xFp16, x0Reg, pregMask);187+ Cast<T, float, castTraitFp32ToFp16>(xFp16, x1Reg, pregMask);
188 DataCopy<T, StoreDist::DIST_PACK_B32>(((__local_mem__ T*)yLocal) + xOffset, xFp16, pregMask);188 DataCopy<T, StoreDist::DIST_PACK_B32>(((__local_mem__ T*)yLocal) + xOffset, xFp16, pregMask);
189 }189 }
190 }190 }
@@ -279,15 +279,15 @@ private:
279 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, xOffset);279 LoadTensorForDtypeT(x1Local, x1Reg, pregMask, xOffset);
280 280 
281 Mul(x1Reg, x0Reg, x1Reg, pregMask);281 Mul(x1Reg, x0Reg, x1Reg, pregMask);
282- Mul(x0Reg, x0Reg, sumReg, pregMask);282+ Neg(x0Reg, x0Reg, pregMask);
283- Sub(x0Reg, x1Reg, x0Reg, pregMask);283+ MulAddDst(x1Reg, x0Reg, sumReg, pregMask);
284 284 
285 // copy out285 // copy out
286 if constexpr (IsSameType<T, float>::value) {286 if constexpr (IsSameType<T, float>::value) {
287- DataCopy(((__local_mem__ float*)yLocal) + xOffset, x0Reg, pregMask);287+ DataCopy(((__local_mem__ float*)yLocal) + xOffset, x1Reg, pregMask);
288 } else { // fp16、bf16288 } else { // fp16、bf16
289 RegTensor<T> xFp16;289 RegTensor<T> xFp16;
290- Cast<T, float, castTraitFp32ToFp16>(xFp16, x0Reg, pregMask);290+ Cast<T, float, castTraitFp32ToFp16>(xFp16, x1Reg, pregMask);
291 DataCopy<T, StoreDist::DIST_PACK_B32>(((__local_mem__ T*)yLocal) + xOffset, xFp16, pregMask);291 DataCopy<T, StoreDist::DIST_PACK_B32>(((__local_mem__ T*)yLocal) + xOffset, xFp16, pregMask);
292 }292 }
293 }293 }