已合并
cleancode: norm 类算子代码去重 #7977
cleancode: norm 类算子代码去重 #7977
已合并
rk创建于 7月27日
共 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 cast60 // 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 cast81 // 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 cast97 // 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 x2107 // 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 
157template <typename T>128template <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+ 
169template <typename T_X, typename T_GAMMA, typename T_SMOOTH_SCALE = float, bool HAS_SMOOTH_SCALE = true,193template <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() const184 __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 RmsNormGrad492} // namespace RmsNormGrad
515-#endif // RMS_NORM_GRAD_REGBASE_DGAMMA_H493+#endif // RMS_NORM_GRAD_REGBASE_DGAMMA_H