已合并
cleancode: norm 类算子代码去重 #7977
rk创建于 7月27日
cleancode: norm 类算子代码去重 #7977
已合并
共 7 个文件变更+125-204
| @@ -55,6 +55,7 @@ __aicore__ inline void LoadForHandleRemainV1(__local_mem__ T* mainAddr, __local_ | |||
| 55 | __local_mem__ float* xFp32MainAddr, __local_mem__ float* xFp32TailAddr, | 55 | __local_mem__ float* xFp32MainAddr, __local_mem__ float* xFp32TailAddr, |
| 56 | __local_mem__ T* mainAddr2, __local_mem__ T* tailAddr2) | 56 | __local_mem__ T* mainAddr2, __local_mem__ T* tailAddr2) |
| 57 | { | 57 | { |
| 58 | + RegTensor<float> mainA2, mainB2, tailA2, tailB2; | ||
| 58 | if constexpr (IsSameType<T, half>::value) { | 59 | if constexpr (IsSameType<T, half>::value) { |
| 59 | // x1 load and cast | 60 | // x1 load and cast |
| 60 | RegTensor<half> xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; | 61 | RegTensor<half> xFp16MainA, xFp16MainB, xFp16TailA, xFp16TailB; |
| @@ -72,24 +73,10 @@ __aicore__ inline void LoadForHandleRemainV1(__local_mem__ T* mainAddr, __local_ | |||
| 72 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16MainB2, mainAddr2 + offset2); | 73 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16MainB2, mainAddr2 + offset2); |
| 73 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16TailA2, tailAddr2 + offset1); | 74 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16TailA2, tailAddr2 + offset1); |
| 74 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16TailB2, tailAddr2 + offset2); | 75 | DataCopy<half, LoadDist::DIST_UNPACK_B16>(xFp16TailB2, tailAddr2 + offset2); |
| 75 | - RegTensor<float> mainA2, mainB2, tailA2, tailB2; | ||
| 76 | Cast<float, half, castTraitB162B32>(mainA2, xFp16MainA2, pregLoop); | 76 | Cast<float, half, castTraitB162B32>(mainA2, xFp16MainA2, pregLoop); |
| 77 | Cast<float, half, castTraitB162B32>(mainB2, xFp16MainB2, pregLoop); | 77 | Cast<float, half, castTraitB162B32>(mainB2, xFp16MainB2, pregLoop); |
| 78 | Cast<float, half, castTraitB162B32>(tailA2, xFp16TailA2, pregLoop); | 78 | Cast<float, half, castTraitB162B32>(tailA2, xFp16TailA2, pregLoop); |
| 79 | Cast<float, half, castTraitB162B32>(tailB2, xFp16TailB2, pregLoop); | 79 | Cast<float, half, castTraitB162B32>(tailB2, xFp16TailB2, pregLoop); |
| 80 | - // add x1 + x2 | ||
| 81 | - Add(mainA, mainA, mainA2, pregLoop); | ||
| 82 | - Add(mainB, mainB, mainB2, pregLoop); | ||
| 83 | - Add(tailA, tailA, tailA2, pregLoop); | ||
| 84 | - Add(tailB, tailB, tailB2, pregLoop); | ||
| 85 | - DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); | ||
| 86 | - DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); | ||
| 87 | - DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); | ||
| 88 | - DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); | ||
| 89 | - Mul(mainA, mainA, mainA, pregLoop); | ||
| 90 | - Mul(mainB, mainB, mainB, pregLoop); | ||
| 91 | - Mul(tailA, tailA, tailA, pregLoop); | ||
| 92 | - Mul(tailB, tailB, tailB, pregLoop); | ||
| 93 | } else if constexpr (IsSameType<T, bfloat16_t>::value) { | 80 | } else if constexpr (IsSameType<T, bfloat16_t>::value) { |
| 94 | // x1 load and cast | 81 | // x1 load and cast |
| 95 | RegTensor<bfloat16_t> xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; | 82 | RegTensor<bfloat16_t> xBFp16MainA, xBFp16MainB, xBFp16TailA, xBFp16TailB; |
| @@ -108,50 +95,34 @@ __aicore__ inline void LoadForHandleRemainV1(__local_mem__ T* mainAddr, __local_ | |||
| 108 | DataCopy<bfloat16_t, LoadDist::DIST_UNPACK_B16>(xBFp16TailA2, tailAddr2 + offset1); | 95 | DataCopy<bfloat16_t, LoadDist::DIST_UNPACK_B16>(xBFp16TailA2, tailAddr2 + offset1); |
| 109 | DataCopy<bfloat16_t, LoadDist::DIST_UNPACK_B16>(xBFp16TailB2, tailAddr2 + offset2); | 96 | DataCopy<bfloat16_t, LoadDist::DIST_UNPACK_B16>(xBFp16TailB2, tailAddr2 + offset2); |
| 110 | // x2 cast | 97 | // x2 cast |
| 111 | - RegTensor<float> mainA2, mainB2, tailA2, tailB2; | ||
| 112 | Cast<float, bfloat16_t, castTraitB162B32>(mainA2, xBFp16MainA2, pregLoop); | 98 | Cast<float, bfloat16_t, castTraitB162B32>(mainA2, xBFp16MainA2, pregLoop); |
| 113 | Cast<float, bfloat16_t, castTraitB162B32>(mainB2, xBFp16MainB2, pregLoop); | 99 | Cast<float, bfloat16_t, castTraitB162B32>(mainB2, xBFp16MainB2, pregLoop); |
| 114 | Cast<float, bfloat16_t, castTraitB162B32>(tailA2, xBFp16TailA2, pregLoop); | 100 | Cast<float, bfloat16_t, castTraitB162B32>(tailA2, xBFp16TailA2, pregLoop); |
| 115 | Cast<float, bfloat16_t, castTraitB162B32>(tailB2, xBFp16TailB2, pregLoop); | 101 | Cast<float, bfloat16_t, castTraitB162B32>(tailB2, xBFp16TailB2, pregLoop); |
| 116 | - // add x1 + x2 | ||
| 117 | - Add(mainA, mainA, mainA2, pregLoop); | ||
| 118 | - Add(mainB, mainB, mainB2, pregLoop); | ||
| 119 | - Add(tailA, tailA, tailA2, pregLoop); | ||
| 120 | - Add(tailB, tailB, tailB2, pregLoop); | ||
| 121 | - DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); | ||
| 122 | - DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); | ||
| 123 | - DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); | ||
| 124 | - DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); | ||
| 125 | - Mul(mainA, mainA, mainA, pregLoop); | ||
| 126 | - Mul(mainB, mainB, mainB, pregLoop); | ||
| 127 | - Mul(tailA, tailA, tailA, pregLoop); | ||
| 128 | - Mul(tailB, tailB, tailB, pregLoop); | ||
| 129 | } else { | 102 | } else { |
| 130 | DataCopy(mainA, mainAddr + offset1); | 103 | DataCopy(mainA, mainAddr + offset1); |
| 131 | DataCopy(mainB, mainAddr + offset2); | 104 | DataCopy(mainB, mainAddr + offset2); |
| 132 | DataCopy(tailA, tailAddr + offset1); | 105 | DataCopy(tailA, tailAddr + offset1); |
| 133 | DataCopy(tailB, tailAddr + offset2); | 106 | DataCopy(tailB, tailAddr + offset2); |
| 134 | // load x2 | 107 | // load x2 |
| 135 | - RegTensor<float> mainA2, mainB2, tailA2, tailB2; | ||
| 136 | DataCopy(mainA2, mainAddr2 + offset1); | 108 | DataCopy(mainA2, mainAddr2 + offset1); |
| 137 | DataCopy(mainB2, mainAddr2 + offset2); | 109 | DataCopy(mainB2, mainAddr2 + offset2); |
| 138 | DataCopy(tailA2, tailAddr2 + offset1); | 110 | DataCopy(tailA2, tailAddr2 + offset1); |
| 139 | DataCopy(tailB2, tailAddr2 + offset2); | 111 | DataCopy(tailB2, tailAddr2 + offset2); |
| 140 | - // add x1 + x2 | ||
| 141 | - Add(mainA, mainA, mainA2, pregLoop); | ||
| 142 | - Add(mainB, mainB, mainB2, pregLoop); | ||
| 143 | - Add(tailA, tailA, tailA2, pregLoop); | ||
| 144 | - Add(tailB, tailB, tailB2, pregLoop); | ||
| 145 | - DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); | ||
| 146 | - DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); | ||
| 147 | - DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); | ||
| 148 | - DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); | ||
| 149 | - // x * x | ||
| 150 | - Mul(mainA, mainA, mainA, pregLoop); | ||
| 151 | - Mul(mainB, mainB, mainB, pregLoop); | ||
| 152 | - Mul(tailA, tailA, tailA, pregLoop); | ||
| 153 | - Mul(tailB, tailB, tailB, pregLoop); | ||
| 154 | } | 112 | } |
| 113 | + // add x1 + x2 | ||
| 114 | + Add(mainA, mainA, mainA2, pregLoop); | ||
| 115 | + Add(mainB, mainB, mainB2, pregLoop); | ||
| 116 | + Add(tailA, tailA, tailA2, pregLoop); | ||
| 117 | + Add(tailB, tailB, tailB2, pregLoop); | ||
| 118 | + DataCopy(xFp32MainAddr + offset1, mainA, pregLoop); | ||
| 119 | + DataCopy(xFp32MainAddr + offset2, mainB, pregLoop); | ||
| 120 | + DataCopy(xFp32TailAddr + offset1, tailA, pregLoop); | ||
| 121 | + DataCopy(xFp32TailAddr + offset2, tailB, pregLoop); | ||
| 122 | + Mul(mainA, mainA, mainA, pregLoop); | ||
| 123 | + Mul(mainB, mainB, mainB, pregLoop); | ||
| 124 | + Mul(tailA, tailA, tailA, pregLoop); | ||
| 125 | + Mul(tailB, tailB, tailB, pregLoop); | ||
| 155 | } | 126 | } |
| 156 | 127 | ||
| 157 | template <typename T> | 128 | template <typename T> |
| @@ -166,6 +166,30 @@ __aicore__ inline void InitOptionalGmBuffers(GlobalTensor<T_SMOOTH_SCALE>& smoot | |||
| 166 | } | 166 | } |
| 167 | } | 167 | } |
| 168 | 168 | ||
| 169 | +template <bool HAS_BETA, bool HAS_SMOOTH_SCALE> | ||
| 170 | +__aicore__ inline void ComputeYAndAbsMaxVF(RegTensor<float>& xRegFp32, RegTensor<float>& yRegFp32, | ||
| 171 | + RegTensor<float>& rstdReg, RegTensor<float>& gammaRegFp32, | ||
| 172 | + RegTensor<float>& betaRegFp32, RegTensor<float>& smoothScaleRegFp32, | ||
| 173 | + RegTensor<float>& scaleReg, MaskReg& maskReg, MaskReg& maskRegFull, | ||
| 174 | + __local_mem__ float* yTmpAddr, uint16_t idx) | ||
| 175 | +{ | ||
| 176 | + Mul(xRegFp32, xRegFp32, rstdReg, maskReg); | ||
| 177 | + Mul(xRegFp32, xRegFp32, gammaRegFp32, maskReg); | ||
| 178 | + if constexpr (HAS_BETA) { | ||
| 179 | + Add(xRegFp32, xRegFp32, betaRegFp32, maskReg); | ||
| 180 | + } | ||
| 181 | + if constexpr (HAS_SMOOTH_SCALE) { | ||
| 182 | + Mul(yRegFp32, xRegFp32, smoothScaleRegFp32, maskReg); | ||
| 183 | + DataCopy<float>(yTmpAddr + idx * V_LENGTH, yRegFp32, maskReg); | ||
| 184 | + Abs(yRegFp32, yRegFp32, maskReg); // VF abs is zeroing mode | ||
| 185 | + Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 186 | + } else { | ||
| 187 | + DataCopy<float>(yTmpAddr + idx * V_LENGTH, xRegFp32, maskReg); | ||
| 188 | + Abs(yRegFp32, xRegFp32, maskReg); // VF abs is zeroing mode | ||
| 189 | + Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 190 | + } | ||
| 191 | +} | ||
| 192 | + | ||
| 169 | template <typename T_X, typename T_GAMMA, typename T_SMOOTH_SCALE = float, bool HAS_SMOOTH_SCALE = true, | 193 | template <typename T_X, typename T_GAMMA, typename T_SMOOTH_SCALE = float, bool HAS_SMOOTH_SCALE = true, |
| 170 | bool HAS_BETA = false, typename T_Y> | 194 | bool HAS_BETA = false, typename T_Y> |
| 171 | __aicore__ inline void ComputeYScale(LocalTensor<T_Y>& yLocal, LocalTensor<float>& scaleLocal, LocalTensor<T_X>& xLocal, | 195 | __aicore__ inline void ComputeYScale(LocalTensor<T_Y>& yLocal, LocalTensor<float>& scaleLocal, LocalTensor<T_X>& xLocal, |
| @@ -211,21 +235,9 @@ __aicore__ inline void ComputeYScale(LocalTensor<T_Y>& yLocal, LocalTensor<float | |||
| 211 | if constexpr (HAS_BETA) { | 235 | if constexpr (HAS_BETA) { |
| 212 | NormCommon::LoadCastRegVF(betaRegFp32, betaAddr, idx, maskReg); | 236 | NormCommon::LoadCastRegVF(betaRegFp32, betaAddr, idx, maskReg); |
| 213 | } | 237 | } |
| 214 | - Mul(xRegFp32, xRegFp32, rstdReg, maskReg); | 238 | + ComputeYAndAbsMaxVF<HAS_BETA, HAS_SMOOTH_SCALE>(xRegFp32, yRegFp32, rstdReg, gammaRegFp32, betaRegFp32, |
| 215 | - Mul(xRegFp32, xRegFp32, gammaRegFp32, maskReg); | 239 | + smoothScaleRegFp32, scaleReg, maskReg, maskRegFull, |
| 216 | - if constexpr (HAS_BETA) { | 240 | + yTmpAddr, idx); |
| 217 | - Add(xRegFp32, xRegFp32, betaRegFp32, maskReg); | ||
| 218 | - } | ||
| 219 | - if constexpr (HAS_SMOOTH_SCALE) { | ||
| 220 | - Mul(yRegFp32, xRegFp32, smoothScaleRegFp32, maskReg); | ||
| 221 | - DataCopy<float>(yTmpAddr + idx * V_LENGTH, yRegFp32, maskReg); | ||
| 222 | - Abs(yRegFp32, yRegFp32, maskReg); // VF abs is zeroing mode | ||
| 223 | - Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 224 | - } else { | ||
| 225 | - DataCopy<float>(yTmpAddr + idx * V_LENGTH, xRegFp32, maskReg); | ||
| 226 | - Abs(yRegFp32, xRegFp32, maskReg); // VF abs is zeroing mode | ||
| 227 | - Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 228 | - } | ||
| 229 | } | 241 | } |
| 230 | ReduceMax(scaleReg, scaleReg, maskRegFull); | 242 | ReduceMax(scaleReg, scaleReg, maskRegFull); |
| 231 | if constexpr (IsSameType<T_Y, int8_t>::value) { | 243 | if constexpr (IsSameType<T_Y, int8_t>::value) { |
| @@ -314,21 +326,9 @@ __aicore__ inline void ComputeReduceMax(LocalTensor<float>& scaleLocal, LocalTen | |||
| 314 | if constexpr (HAS_SMOOTH_SCALE) { | 326 | if constexpr (HAS_SMOOTH_SCALE) { |
| 315 | NormCommon::LoadCastRegVF(smoothScaleRegFp32, smoothScaleAddr, idx, maskReg); | 327 | NormCommon::LoadCastRegVF(smoothScaleRegFp32, smoothScaleAddr, idx, maskReg); |
| 316 | } | 328 | } |
| 317 | - Mul(xRegFp32, xRegFp32, rstdReg, maskReg); | 329 | + ComputeYAndAbsMaxVF<HAS_BETA, HAS_SMOOTH_SCALE>(xRegFp32, yRegFp32, rstdReg, gammaRegFp32, betaRegFp32, |
| 318 | - Mul(xRegFp32, xRegFp32, gammaRegFp32, maskReg); | 330 | + smoothScaleRegFp32, scaleReg, maskReg, maskRegFull, |
| 319 | - if constexpr (HAS_BETA) { | 331 | + yTmpAddr, idx); |
| 320 | - Add(xRegFp32, xRegFp32, betaRegFp32, maskReg); | ||
| 321 | - } | ||
| 322 | - if constexpr (HAS_SMOOTH_SCALE) { | ||
| 323 | - Mul(yRegFp32, xRegFp32, smoothScaleRegFp32, maskReg); | ||
| 324 | - DataCopy<float>(yTmpAddr + idx * V_LENGTH, yRegFp32, maskReg); | ||
| 325 | - Abs(yRegFp32, yRegFp32, maskReg); // VF abs is zeroing mode | ||
| 326 | - Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 327 | - } else { | ||
| 328 | - DataCopy<float>(yTmpAddr + idx * V_LENGTH, xRegFp32, maskReg); | ||
| 329 | - Abs(yRegFp32, xRegFp32, maskReg); // VF abs is zeroing mode | ||
| 330 | - Max(scaleReg, scaleReg, yRegFp32, maskRegFull); // Using full mask | ||
| 331 | - } | ||
| 332 | } | 332 | } |
| 333 | ReduceMax(scaleReg, scaleReg, maskRegFull); | 333 | ReduceMax(scaleReg, scaleReg, maskRegFull); |
| 334 | Max(scaleReg, scaleReg, scaleLastReg, maskRegOne); | 334 | Max(scaleReg, scaleReg, scaleLastReg, maskRegOne); |
| @@ -156,25 +156,8 @@ private: | |||
| 156 | 156 | ||
| 157 | __aicore__ inline void Compute(int64_t curTileBLen) | 157 | __aicore__ inline void Compute(int64_t curTileBLen) |
| 158 | { | 158 | { |
| 159 | - LocalTensor<T> x = xQueue_.DeQue<T>(); | 159 | + InferComputeImpl<BatchNormV3InferLastChannelContinuousA, T>( |
| 160 | - LocalTensor<T> y = yQueue_.AllocTensor<T>(); | 160 | + *this, xQueue_, yQueue_, betaFp32Buf_, gammaFp32Buf_, meanFp32Buf_, rstdFp32Buf_, curTileBLen); |
| 161 | - LocalTensor<float> betaFp32 = betaFp32Buf_.Get<float>(); | ||
| 162 | - LocalTensor<float> gammaFp32 = gammaFp32Buf_.Get<float>(); | ||
| 163 | - LocalTensor<float> meanFp32 = meanFp32Buf_.Get<float>(); | ||
| 164 | - LocalTensor<float> rstdFp32 = rstdFp32Buf_.Get<float>(); | ||
| 165 | - | ||
| 166 | - __local_mem__ T* xLocal = (__local_mem__ T*)x.GetPhyAddr(); | ||
| 167 | - __local_mem__ T* yLocal = (__local_mem__ T*)y.GetPhyAddr(); | ||
| 168 | - __local_mem__ float* betaFp32Local = (__local_mem__ float*)betaFp32.GetPhyAddr(); | ||
| 169 | - __local_mem__ float* gammaFp32Local = (__local_mem__ float*)gammaFp32.GetPhyAddr(); | ||
| 170 | - __local_mem__ float* meanFp32Local = (__local_mem__ float*)meanFp32.GetPhyAddr(); | ||
| 171 | - __local_mem__ float* rstdFp32Local = (__local_mem__ float*)rstdFp32.GetPhyAddr(); | ||
| 172 | - | ||
| 173 | - VFNormalize(xLocal, gammaFp32Local, betaFp32Local, meanFp32Local, rstdFp32Local, yLocal, curTileBLen); | ||
| 174 | - | ||
| 175 | - yQueue_.EnQue(y); | ||
| 176 | - | ||
| 177 | - xQueue_.FreeTensor<T>(x); | ||
| 178 | } | 161 | } |
| 179 | 162 | ||
| 180 | __aicore__ inline void VFPrepareParamCache(__local_mem__ T_GAMMA* gammaLocal, __local_mem__ T_GAMMA* betaLocal, | 163 | __aicore__ inline void VFPrepareParamCache(__local_mem__ T_GAMMA* gammaLocal, __local_mem__ T_GAMMA* betaLocal, |
| @@ -212,6 +195,7 @@ private: | |||
| 212 | } | 195 | } |
| 213 | } | 196 | } |
| 214 | 197 | ||
| 198 | +public: | ||
| 215 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, | 199 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, |
| 216 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, | 200 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, |
| 217 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, int64_t curTileBLen) | 201 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, int64_t curTileBLen) |
| @@ -249,6 +233,7 @@ private: | |||
| 249 | } | 233 | } |
| 250 | } | 234 | } |
| 251 | 235 | ||
| 236 | +private: | ||
| 252 | template <typename T_SRC> | 237 | template <typename T_SRC> |
| 253 | __aicore__ inline void LoadParamForDtypeT(__local_mem__ T_SRC* src, RegTensor<float>& dst, MaskReg& preg, | 238 | __aicore__ inline void LoadParamForDtypeT(__local_mem__ T_SRC* src, RegTensor<float>& dst, MaskReg& preg, |
| 254 | uint32_t offset) | 239 | uint32_t offset) |
| @@ -144,28 +144,11 @@ private: | |||
| 144 | 144 | ||
| 145 | __aicore__ inline void Compute(int64_t curTileBLen, int64_t curTileALen) | 145 | __aicore__ inline void Compute(int64_t curTileBLen, int64_t curTileALen) |
| 146 | { | 146 | { |
| 147 | - LocalTensor<T> x = xQueue_.DeQue<T>(); | 147 | + InferComputeImpl<BatchNormV3InferLastChannelSmallA, T>(*this, xQueue_, yQueue_, betaFp32Buf_, gammaFp32Buf_, |
| 148 | - LocalTensor<T> y = yQueue_.AllocTensor<T>(); | 148 | + meanFp32Buf_, rstdFp32Buf_, curTileBLen * curTileALen); |
| 149 | - LocalTensor<float> betaFp32 = betaFp32Buf_.Get<float>(); | ||
| 150 | - LocalTensor<float> gammaFp32 = gammaFp32Buf_.Get<float>(); | ||
| 151 | - LocalTensor<float> meanFp32 = meanFp32Buf_.Get<float>(); | ||
| 152 | - LocalTensor<float> rstdFp32 = rstdFp32Buf_.Get<float>(); | ||
| 153 | - | ||
| 154 | - __local_mem__ T* xLocal = (__local_mem__ T*)x.GetPhyAddr(); | ||
| 155 | - __local_mem__ T* yLocal = (__local_mem__ T*)y.GetPhyAddr(); | ||
| 156 | - __local_mem__ float* betaFp32Local = (__local_mem__ float*)betaFp32.GetPhyAddr(); | ||
| 157 | - __local_mem__ float* gammaFp32Local = (__local_mem__ float*)gammaFp32.GetPhyAddr(); | ||
| 158 | - __local_mem__ float* meanFp32Local = (__local_mem__ float*)meanFp32.GetPhyAddr(); | ||
| 159 | - __local_mem__ float* rstdFp32Local = (__local_mem__ float*)rstdFp32.GetPhyAddr(); | ||
| 160 | - | ||
| 161 | - VFNormalize(xLocal, gammaFp32Local, betaFp32Local, meanFp32Local, rstdFp32Local, yLocal, | ||
| 162 | - curTileBLen * curTileALen); | ||
| 163 | - | ||
| 164 | - yQueue_.EnQue(y); | ||
| 165 | - | ||
| 166 | - xQueue_.FreeTensor<T>(x); | ||
| 167 | } | 149 | } |
| 168 | 150 | ||
| 151 | +public: | ||
| 169 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, | 152 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, |
| 170 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, | 153 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, |
| 171 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, uint32_t curElemLen) | 154 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, uint32_t curElemLen) |
| @@ -205,6 +188,7 @@ private: | |||
| 205 | } | 188 | } |
| 206 | } | 189 | } |
| 207 | 190 | ||
| 191 | +private: | ||
| 208 | __aicore__ inline void CopyOutY(int64_t yGmOffset, int64_t curTileBLen, int64_t curTileALen) | 192 | __aicore__ inline void CopyOutY(int64_t yGmOffset, int64_t curTileBLen, int64_t curTileALen) |
| 209 | { | 193 | { |
| 210 | LocalTensor<T> y = yQueue_.DeQue<T>(); | 194 | LocalTensor<T> y = yQueue_.DeQue<T>(); |
| @@ -134,27 +134,11 @@ private: | |||
| 134 | 134 | ||
| 135 | __aicore__ inline void Compute(int64_t curTileB0Len) | 135 | __aicore__ inline void Compute(int64_t curTileB0Len) |
| 136 | { | 136 | { |
| 137 | - LocalTensor<T> x = xQueue_.DeQue<T>(); | 137 | + InferComputeImpl<BatchNormV3InferSmallAB1, T>(*this, xQueue_, yQueue_, betaFp32Buf_, gammaFp32Buf_, |
| 138 | - LocalTensor<T> y = yQueue_.AllocTensor<T>(); | 138 | + meanFp32Buf_, rstdFp32Buf_, curTileB0Len); |
| 139 | - LocalTensor<float> betaFp32 = betaFp32Buf_.Get<float>(); | ||
| 140 | - LocalTensor<float> gammaFp32 = gammaFp32Buf_.Get<float>(); | ||
| 141 | - LocalTensor<float> meanFp32 = meanFp32Buf_.Get<float>(); | ||
| 142 | - LocalTensor<float> rstdFp32 = rstdFp32Buf_.Get<float>(); | ||
| 143 | - | ||
| 144 | - __local_mem__ T* xLocal = (__local_mem__ T*)x.GetPhyAddr(); | ||
| 145 | - __local_mem__ T* yLocal = (__local_mem__ T*)y.GetPhyAddr(); | ||
| 146 | - __local_mem__ float* betaFp32Local = (__local_mem__ float*)betaFp32.GetPhyAddr(); | ||
| 147 | - __local_mem__ float* gammaFp32Local = (__local_mem__ float*)gammaFp32.GetPhyAddr(); | ||
| 148 | - __local_mem__ float* meanFp32Local = (__local_mem__ float*)meanFp32.GetPhyAddr(); | ||
| 149 | - __local_mem__ float* rstdFp32Local = (__local_mem__ float*)rstdFp32.GetPhyAddr(); | ||
| 150 | - | ||
| 151 | - VFNormalize(xLocal, gammaFp32Local, betaFp32Local, meanFp32Local, rstdFp32Local, yLocal, curTileB0Len); | ||
| 152 | - | ||
| 153 | - yQueue_.EnQue(y); | ||
| 154 | - | ||
| 155 | - xQueue_.FreeTensor<T>(x); | ||
| 156 | } | 139 | } |
| 157 | 140 | ||
| 141 | +public: | ||
| 158 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, | 142 | __aicore__ inline void VFNormalize(__local_mem__ T* xLocal, __local_mem__ float* gammaFp32Local, |
| 159 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, | 143 | __local_mem__ float* betaFp32Local, __local_mem__ float* meanFp32Local, |
| 160 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, | 144 | __local_mem__ float* rstdFp32Local, __local_mem__ T* yLocal, |
| @@ -196,6 +180,7 @@ private: | |||
| 196 | } | 180 | } |
| 197 | } | 181 | } |
| 198 | 182 | ||
| 183 | +private: | ||
| 199 | __aicore__ inline uint32_t GetSmallAB1ParamCacheElemLen() const | 184 | __aicore__ inline uint32_t GetSmallAB1ParamCacheElemLen() const |
| 200 | { | 185 | { |
| 201 | uint32_t abLen = static_cast<uint32_t>(tilingData_->totalALen * tilingData_->totalB1Len); | 186 | uint32_t abLen = static_cast<uint32_t>(tilingData_->totalALen * tilingData_->totalB1Len); |
| @@ -59,6 +59,32 @@ struct RLessThanParams { | |||
| 59 | uint32_t remainderTailOffset3; | 59 | uint32_t remainderTailOffset3; |
| 60 | }; | 60 | }; |
| 61 | 61 | ||
| 62 | +template <typename Self, typename T> | ||
| 63 | +__aicore__ inline void InferComputeImpl(Self& self, TQue<QuePosition::VECIN, 1>& xQueue, | ||
| 64 | + TQue<QuePosition::VECOUT, 1>& yQueue, TBuf<TPosition::VECCALC>& betaBuf, | ||
| 65 | + TBuf<TPosition::VECCALC>& gammaBuf, TBuf<TPosition::VECCALC>& meanBuf, | ||
| 66 | + TBuf<TPosition::VECCALC>& rstdBuf, int64_t vfLen) | ||
| 67 | +{ | ||
| 68 | + LocalTensor<T> x = xQueue.template DeQue<T>(); | ||
| 69 | + LocalTensor<T> y = yQueue.template AllocTensor<T>(); | ||
| 70 | + LocalTensor<float> betaFp32 = betaBuf.template Get<float>(); | ||
| 71 | + LocalTensor<float> gammaFp32 = gammaBuf.template Get<float>(); | ||
| 72 | + LocalTensor<float> meanFp32 = meanBuf.template Get<float>(); | ||
| 73 | + LocalTensor<float> rstdFp32 = rstdBuf.template Get<float>(); | ||
| 74 | + | ||
| 75 | + __local_mem__ T* xLocal = (__local_mem__ T*)x.GetPhyAddr(); | ||
| 76 | + __local_mem__ T* yLocal = (__local_mem__ T*)y.GetPhyAddr(); | ||
| 77 | + __local_mem__ float* betaFp32Local = (__local_mem__ float*)betaFp32.GetPhyAddr(); | ||
| 78 | + __local_mem__ float* gammaFp32Local = (__local_mem__ float*)gammaFp32.GetPhyAddr(); | ||
| 79 | + __local_mem__ float* meanFp32Local = (__local_mem__ float*)meanFp32.GetPhyAddr(); | ||
| 80 | + __local_mem__ float* rstdFp32Local = (__local_mem__ float*)rstdFp32.GetPhyAddr(); | ||
| 81 | + | ||
| 82 | + self.VFNormalize(xLocal, gammaFp32Local, betaFp32Local, meanFp32Local, rstdFp32Local, yLocal, vfLen); | ||
| 83 | + | ||
| 84 | + yQueue.EnQue(y); | ||
| 85 | + xQueue.template FreeTensor<T>(x); | ||
| 86 | +} | ||
| 87 | + | ||
| 62 | __aicore__ inline RLessThanParams GetRLessThanParams(uint32_t scaleCoef, uint32_t currentANumAlign, uint32_t r1) | 88 | __aicore__ inline RLessThanParams GetRLessThanParams(uint32_t scaleCoef, uint32_t currentANumAlign, uint32_t r1) |
| 63 | { | 89 | { |
| 64 | RLessThanParams params; | 90 | RLessThanParams params; |
| @@ -726,36 +752,6 @@ __aicore__ inline void CalculateRLessThanVF(__local_mem__ float* xInUb, __local_ | |||
| 726 | } | 752 | } |
| 727 | } | 753 | } |
| 728 | 754 | ||
| 729 | -__aicore__ inline void TwoRowAddPartialMeanWithTail(RegTensor<float>& dst, __local_mem__ float* input, | ||
| 730 | - __local_mem__ float* tCount, MaskReg& preg, uint32_t offset1, | ||
| 731 | - uint32_t offset2, uint32_t offset3, uint32_t offset4, | ||
| 732 | - uint32_t offset5, uint32_t offset6, uint32_t offset7, | ||
| 733 | - uint32_t offset8, RegTensor<float>& rem, RegTensor<float>& nextRow, | ||
| 734 | - RegTensor<float>& remNextRow, RegTensor<float>& dstCount, | ||
| 735 | - RegTensor<float>& remCount, RegTensor<float>& nextRowCount, | ||
| 736 | - RegTensor<float>& remNextRowCount, float n) | ||
| 737 | -{ | ||
| 738 | - DataCopy(dst, ((__local_mem__ float*)(input) + (offset1))); | ||
| 739 | - DataCopy(rem, ((__local_mem__ float*)(input) + (offset2))); | ||
| 740 | - DataCopy<float, LoadDist::DIST_BRC_B32>(dstCount, ((__local_mem__ float*)(tCount) + (offset5))); | ||
| 741 | - DataCopy<float, LoadDist::DIST_BRC_B32>(remCount, ((__local_mem__ float*)(tCount) + (offset6))); | ||
| 742 | - Mul(dst, dst, dstCount, preg); | ||
| 743 | - Mul(rem, rem, remCount, preg); | ||
| 744 | - Muls(dst, dst, n, preg); | ||
| 745 | - Muls(rem, rem, n, preg); | ||
| 746 | - Add(dst, dst, rem, preg); | ||
| 747 | - DataCopy(nextRow, ((__local_mem__ float*)(input) + (offset3))); | ||
| 748 | - DataCopy(remNextRow, ((__local_mem__ float*)(input) + (offset4))); | ||
| 749 | - DataCopy<float, LoadDist::DIST_BRC_B32>(nextRowCount, ((__local_mem__ float*)(tCount) + (offset7))); | ||
| 750 | - DataCopy<float, LoadDist::DIST_BRC_B32>(remNextRowCount, ((__local_mem__ float*)(tCount) + (offset8))); | ||
| 751 | - Mul(nextRow, nextRow, nextRowCount, preg); | ||
| 752 | - Mul(remNextRow, remNextRow, remNextRowCount, preg); | ||
| 753 | - Muls(nextRow, nextRow, n, preg); | ||
| 754 | - Muls(remNextRow, remNextRow, n, preg); | ||
| 755 | - Add(nextRow, nextRow, remNextRow, preg); | ||
| 756 | - Add(dst, dst, nextRow, preg); | ||
| 757 | -} | ||
| 758 | - | ||
| 759 | __aicore__ inline void TwoRowAddPartialMean(RegTensor<float>& dst, __local_mem__ float* input, | 755 | __aicore__ inline void TwoRowAddPartialMean(RegTensor<float>& dst, __local_mem__ float* input, |
| 760 | __local_mem__ float* tCount, MaskReg& preg, uint32_t offset1, | 756 | __local_mem__ float* tCount, MaskReg& preg, uint32_t offset1, |
| 761 | uint32_t offset2, uint32_t offset5, uint32_t offset6, RegTensor<float>& rem, | 757 | uint32_t offset2, uint32_t offset5, uint32_t offset6, RegTensor<float>& rem, |
| @@ -772,6 +768,28 @@ __aicore__ inline void TwoRowAddPartialMean(RegTensor<float>& dst, __local_mem__ | |||
| 772 | Add(dst, dst, rem, preg); | 768 | Add(dst, dst, rem, preg); |
| 773 | } | 769 | } |
| 774 | 770 | ||
| 771 | +__aicore__ inline void TwoRowAddPartialMeanWithTail(RegTensor<float>& dst, __local_mem__ float* input, | ||
| 772 | + __local_mem__ float* tCount, MaskReg& preg, uint32_t offset1, | ||
| 773 | + uint32_t offset2, uint32_t offset3, uint32_t offset4, | ||
| 774 | + uint32_t offset5, uint32_t offset6, uint32_t offset7, | ||
| 775 | + uint32_t offset8, RegTensor<float>& rem, RegTensor<float>& nextRow, | ||
| 776 | + RegTensor<float>& remNextRow, RegTensor<float>& dstCount, | ||
| 777 | + RegTensor<float>& remCount, RegTensor<float>& nextRowCount, | ||
| 778 | + RegTensor<float>& remNextRowCount, float n) | ||
| 779 | +{ | ||
| 780 | + TwoRowAddPartialMean(dst, input, tCount, preg, offset1, offset2, offset5, offset6, rem, dstCount, remCount, n); | ||
| 781 | + DataCopy(nextRow, ((__local_mem__ float*)(input) + (offset3))); | ||
| 782 | + DataCopy(remNextRow, ((__local_mem__ float*)(input) + (offset4))); | ||
| 783 | + DataCopy<float, LoadDist::DIST_BRC_B32>(nextRowCount, ((__local_mem__ float*)(tCount) + (offset7))); | ||
| 784 | + DataCopy<float, LoadDist::DIST_BRC_B32>(remNextRowCount, ((__local_mem__ float*)(tCount) + (offset8))); | ||
| 785 | + Mul(nextRow, nextRow, nextRowCount, preg); | ||
| 786 | + Mul(remNextRow, remNextRow, remNextRowCount, preg); | ||
| 787 | + Muls(nextRow, nextRow, n, preg); | ||
| 788 | + Muls(remNextRow, remNextRow, n, preg); | ||
| 789 | + Add(nextRow, nextRow, remNextRow, preg); | ||
| 790 | + Add(dst, dst, nextRow, preg); | ||
| 791 | +} | ||
| 792 | + | ||
| 775 | __aicore__ inline void TwoRowAddPartialVar(RegTensor<float>& dst, __local_mem__ float* tmpMean, | 793 | __aicore__ inline void TwoRowAddPartialVar(RegTensor<float>& dst, __local_mem__ float* tmpMean, |
| 776 | __local_mem__ float* tmpM2, __local_mem__ float* tCount, MaskReg& preg, | 794 | __local_mem__ float* tmpM2, __local_mem__ float* tCount, MaskReg& preg, |
| 777 | uint32_t offset1, uint32_t offset2, uint32_t offset5, uint32_t offset6, | 795 | uint32_t offset1, uint32_t offset2, uint32_t offset5, uint32_t offset6, |
| @@ -226,14 +226,10 @@ public: | |||
| 226 | } | 226 | } |
| 227 | } | 227 | } |
| 228 | 228 | ||
| 229 | - __aicore__ inline void CalcDgamma(uint32_t inputOffset, uint32_t currentCols, bool isWithPad) | 229 | + __aicore__ inline void ComputePreDgamma(LocalTensor<DY_TYPE>& dyLocal, LocalTensor<X_TYPE>& xLocal, |
| 230 | + LocalTensor<RSTD_TYPE>& rstdLocal, LocalTensor<float>& dgammaOutLocal, | ||
| 231 | + uint32_t currentCols, __local_mem__ float*& dgammaOutAddr) | ||
| 230 | { | 232 | { |
| 231 | - LocalTensor<RSTD_TYPE> rstdLocal = rstdQueue_.template AllocTensor<RSTD_TYPE>(); | ||
| 232 | - LocalTensor<DY_TYPE> dyLocal = dyQueue_.template AllocTensor<DY_TYPE>(); | ||
| 233 | - LocalTensor<X_TYPE> xLocal = xQueue_.template AllocTensor<X_TYPE>(); | ||
| 234 | - LocalTensor<float> dgammaOutLocal = dgammaQueue_.template AllocTensor<float>(); | ||
| 235 | - | ||
| 236 | - CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, 0); | ||
| 237 | xQueue_.EnQue(xLocal); | 233 | xQueue_.EnQue(xLocal); |
| 238 | rstdQueue_.EnQue(rstdLocal); | 234 | rstdQueue_.EnQue(rstdLocal); |
| 239 | dyQueue_.EnQue(dyLocal); | 235 | dyQueue_.EnQue(dyLocal); |
| @@ -245,12 +241,24 @@ public: | |||
| 245 | __local_mem__ DY_TYPE* dyAddr = (__local_mem__ DY_TYPE*)dyLocal[0].GetPhyAddr(); | 241 | __local_mem__ DY_TYPE* dyAddr = (__local_mem__ DY_TYPE*)dyLocal[0].GetPhyAddr(); |
| 246 | __local_mem__ X_TYPE* xAddr = (__local_mem__ X_TYPE*)xLocal[0].GetPhyAddr(); | 242 | __local_mem__ X_TYPE* xAddr = (__local_mem__ X_TYPE*)xLocal[0].GetPhyAddr(); |
| 247 | __local_mem__ RSTD_TYPE* rstdAddr = (__local_mem__ RSTD_TYPE*)rstdLocal[0].GetPhyAddr(); | 243 | __local_mem__ RSTD_TYPE* rstdAddr = (__local_mem__ RSTD_TYPE*)rstdLocal[0].GetPhyAddr(); |
| 248 | - __local_mem__ float* dgammaOutAddr = (__local_mem__ float*)dgammaOutLocal[0].GetPhyAddr(); | 244 | + dgammaOutAddr = (__local_mem__ float*)dgammaOutLocal[0].GetPhyAddr(); |
| 249 | 245 | ||
| 250 | VFCalcPreDgamma(dyAddr, xAddr, rstdAddr, dgammaOutAddr, currentCols, rowsPerUB_); | 246 | VFCalcPreDgamma(dyAddr, xAddr, rstdAddr, dgammaOutAddr, currentCols, rowsPerUB_); |
| 251 | dyQueue_.FreeTensor(dyLocal); | 247 | dyQueue_.FreeTensor(dyLocal); |
| 252 | xQueue_.FreeTensor(xLocal); | 248 | xQueue_.FreeTensor(xLocal); |
| 253 | rstdQueue_.FreeTensor(rstdLocal); | 249 | rstdQueue_.FreeTensor(rstdLocal); |
| 250 | + } | ||
| 251 | + | ||
| 252 | + __aicore__ inline void CalcDgamma(uint32_t inputOffset, uint32_t currentCols, bool isWithPad) | ||
| 253 | + { | ||
| 254 | + LocalTensor<RSTD_TYPE> rstdLocal = rstdQueue_.template AllocTensor<RSTD_TYPE>(); | ||
| 255 | + LocalTensor<DY_TYPE> dyLocal = dyQueue_.template AllocTensor<DY_TYPE>(); | ||
| 256 | + LocalTensor<X_TYPE> xLocal = xQueue_.template AllocTensor<X_TYPE>(); | ||
| 257 | + LocalTensor<float> dgammaOutLocal = dgammaQueue_.template AllocTensor<float>(); | ||
| 258 | + | ||
| 259 | + CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, 0); | ||
| 260 | + __local_mem__ float* dgammaOutAddr; | ||
| 261 | + ComputePreDgamma(dyLocal, xLocal, rstdLocal, dgammaOutLocal, currentCols, dgammaOutAddr); | ||
| 254 | if (isWithPad) { | 262 | if (isWithPad) { |
| 255 | uint32_t mainRows = rows_ - tiling_->rowsTailDG; | 263 | uint32_t mainRows = rows_ - tiling_->rowsTailDG; |
| 256 | VFDuplicateRows(dgammaOutAddr, vlFp32_, rows_ * vlFp32_); | 264 | VFDuplicateRows(dgammaOutAddr, vlFp32_, rows_ * vlFp32_); |
| @@ -279,23 +287,8 @@ public: | |||
| 279 | uint32_t rstdOffset = i * rowsPerUB_; | 287 | uint32_t rstdOffset = i * rowsPerUB_; |
| 280 | 288 | ||
| 281 | CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, rstdOffset); | 289 | CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, rstdOffset); |
| 282 | - xQueue_.EnQue(xLocal); | 290 | + __local_mem__ float* dgammaOutAddr; |
| 283 | - rstdQueue_.EnQue(rstdLocal); | 291 | + ComputePreDgamma(dyLocal, xLocal, rstdLocal, dgammaOutLocal, currentCols, dgammaOutAddr); |
| 284 | - dyQueue_.EnQue(dyLocal); | ||
| 285 | - | ||
| 286 | - dyLocal = dyQueue_.template DeQue<DY_TYPE>(); | ||
| 287 | - xLocal = xQueue_.template DeQue<X_TYPE>(); | ||
| 288 | - rstdLocal = rstdQueue_.template DeQue<RSTD_TYPE>(); | ||
| 289 | - | ||
| 290 | - __local_mem__ DY_TYPE* dyAddr = (__local_mem__ DY_TYPE*)dyLocal[0].GetPhyAddr(); | ||
| 291 | - __local_mem__ X_TYPE* xAddr = (__local_mem__ X_TYPE*)xLocal[0].GetPhyAddr(); | ||
| 292 | - __local_mem__ RSTD_TYPE* rstdAddr = (__local_mem__ RSTD_TYPE*)rstdLocal[0].GetPhyAddr(); | ||
| 293 | - __local_mem__ float* dgammaOutAddr = (__local_mem__ float*)dgammaOutLocal[0].GetPhyAddr(); | ||
| 294 | - | ||
| 295 | - VFCalcPreDgamma(dyAddr, xAddr, rstdAddr, dgammaOutAddr, currentCols, rowsPerUB_); | ||
| 296 | - dyQueue_.FreeTensor(dyLocal); | ||
| 297 | - xQueue_.FreeTensor(xLocal); | ||
| 298 | - rstdQueue_.FreeTensor(rstdLocal); | ||
| 299 | 292 | ||
| 300 | VFBinaryReduceSumWithoutTail(dgammaOutAddr, currentCols, rowsPerUB_); | 293 | VFBinaryReduceSumWithoutTail(dgammaOutAddr, currentCols, rowsPerUB_); |
| 301 | UpdateCache(binaryAddCacheLocal, dgammaOutAddr, cacheID, vlFp32_); | 294 | UpdateCache(binaryAddCacheLocal, dgammaOutAddr, cacheID, vlFp32_); |
| @@ -314,23 +307,8 @@ public: | |||
| 314 | int64_t cacheID = GetCacheID(i); | 307 | int64_t cacheID = GetCacheID(i); |
| 315 | 308 | ||
| 316 | CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, rstdOffset); | 309 | CopyInputsToUB(dyLocal, xLocal, rstdLocal, inputOffset, currentCols, rowsPerUB_, rstdOffset); |
| 317 | - xQueue_.EnQue(xLocal); | 310 | + __local_mem__ float* dgammaOutAddr; |
| 318 | - rstdQueue_.EnQue(rstdLocal); | 311 | + ComputePreDgamma(dyLocal, xLocal, rstdLocal, dgammaOutLocal, currentCols, dgammaOutAddr); |
| 319 | - dyQueue_.EnQue(dyLocal); | ||
| 320 | - | ||
| 321 | - dyLocal = dyQueue_.template DeQue<DY_TYPE>(); | ||
| 322 | - xLocal = xQueue_.template DeQue<X_TYPE>(); | ||
| 323 | - rstdLocal = rstdQueue_.template DeQue<RSTD_TYPE>(); | ||
| 324 | - | ||
| 325 | - __local_mem__ DY_TYPE* dyAddr = (__local_mem__ DY_TYPE*)dyLocal[0].GetPhyAddr(); | ||
| 326 | - __local_mem__ X_TYPE* xAddr = (__local_mem__ X_TYPE*)xLocal[0].GetPhyAddr(); | ||
| 327 | - __local_mem__ RSTD_TYPE* rstdAddr = (__local_mem__ RSTD_TYPE*)rstdLocal[0].GetPhyAddr(); | ||
| 328 | - __local_mem__ float* dgammaOutAddr = (__local_mem__ float*)dgammaOutLocal[0].GetPhyAddr(); | ||
| 329 | - | ||
| 330 | - VFCalcPreDgamma(dyAddr, xAddr, rstdAddr, dgammaOutAddr, currentCols, rowsPerUB_); | ||
| 331 | - dyQueue_.FreeTensor(dyLocal); | ||
| 332 | - xQueue_.FreeTensor(xLocal); | ||
| 333 | - rstdQueue_.FreeTensor(rstdLocal); | ||
| 334 | 312 | ||
| 335 | // 处理累加的尾块 | 313 | // 处理累加的尾块 |
| 336 | LocalTensor<RSTD_TYPE> rstdLocal1 = rstdQueue_.template AllocTensor<RSTD_TYPE>(); | 314 | LocalTensor<RSTD_TYPE> rstdLocal1 = rstdQueue_.template AllocTensor<RSTD_TYPE>(); |
| @@ -512,4 +490,4 @@ private: | |||
| 512 | const RmsNormGradRegbaseTilingData* tiling_; | 490 | const RmsNormGradRegbaseTilingData* tiling_; |
| 513 | }; | 491 | }; |
| 514 | } // namespace RmsNormGrad | 492 | } // namespace RmsNormGrad |
| 515 | -#endif // RMS_NORM_GRAD_REGBASE_DGAMMA_H | 493 | +#endif // RMS_NORM_GRAD_REGBASE_DGAMMA_H |