已合并
部分算子vf整改 #8979
部分算子vf整改 #8979
已合并
xieshengwei1024创建于 18 天前
14 个文件变更+432-449
@@ -208,7 +208,7 @@ __simd_vf__ inline void VFProcessGroupIndexSmallVf(__ubuf__ T* yLocalAddr, __ubu
208 Adds(x, x, static_cast<T>(0), pregLoop);208 Adds(x, x, static_cast<T>(0), pregLoop);
209 Add(sum, sum, x, pregMain);209 Add(sum, sum, x, pregMain);
210 }210 }
211- ReduceSum(sum, sum, pregMain);211+ Reduce<ReduceType::SUM>(sum, sum, pregMain);
212 if (withUbReduce) {212 if (withUbReduce) {
213 RegTensor<T> origin;213 RegTensor<T> origin;
214 LoadAlign(origin, yLocalAddr);214 LoadAlign(origin, yLocalAddr);
@@ -256,7 +256,7 @@ __simd_vf__ inline void VFProcessGroupIndexLargeVf(__ubuf__ T* yLocalAddr, __ubu
256 Add(sum0, sum0, sum1, pregMain);256 Add(sum0, sum0, sum1, pregMain);
257 Add(sum2, sum2, sum3, pregMain);257 Add(sum2, sum2, sum3, pregMain);
258 Add(sum0, sum0, sum2, pregMain);258 Add(sum0, sum0, sum2, pregMain);
259- ReduceSum(sum0, sum0, pregMain);259+ Reduce<ReduceType::SUM>(sum0, sum0, pregMain);
260 if (withUbReduce) {260 if (withUbReduce) {
261 RegTensor<T> origin;261 RegTensor<T> origin;
262 LoadAlign(origin, yLocalAddr);262 LoadAlign(origin, yLocalAddr);
@@ -210,9 +210,9 @@ __simd_callee__ inline void FP16Convert(AscendC::Reg::RegTensor<half>& output, A
210 AscendC::Reg::Duplicate(specialValueTensor, specialValue);210 AscendC::Reg::Duplicate(specialValueTensor, specialValue);
211 AscendC::Reg::Duplicate(newMantissa, NEW_MANTISSA);211 AscendC::Reg::Duplicate(newMantissa, NEW_MANTISSA);
212 AscendC::Reg::And(andResult, (AscendC::Reg::RegTensor<uint16_t>&)input, specialValueTensor, mask);212 AscendC::Reg::And(andResult, (AscendC::Reg::RegTensor<uint16_t>&)input, specialValueTensor, mask);
213- AscendC::Reg::CompareScalar<uint16_t, CMPMODE::GT>(nonzeroMask, andResult, 0, mask);213+ AscendC::Reg::Compares<uint16_t, CMPMODE::GT>(nonzeroMask, andResult, 0, mask);
214- AscendC::Reg::CompareScalar<uint16_t, CMPMODE::LT>(specialMask, andResult, NEW_MANTISSA, mask);214+ AscendC::Reg::Compares<uint16_t, CMPMODE::LT>(specialMask, andResult, NEW_MANTISSA, mask);
215- AscendC::Reg::MaskAnd(specialMask, specialMask, nonzeroMask, mask);215+ AscendC::Reg::And(specialMask, specialMask, nonzeroMask, mask);
216 AscendC::Reg::Or(newValue, (AscendC::Reg::RegTensor<uint16_t>&)input, newMantissa, mask);216 AscendC::Reg::Or(newValue, (AscendC::Reg::RegTensor<uint16_t>&)input, newMantissa, mask);
217 AscendC::Reg::Select<uint16_t>((AscendC::Reg::RegTensor<uint16_t>&)output, newValue,217 AscendC::Reg::Select<uint16_t>((AscendC::Reg::RegTensor<uint16_t>&)output, newValue,
218 (AscendC::Reg::RegTensor<uint16_t>&)input, specialMask);218 (AscendC::Reg::RegTensor<uint16_t>&)input, specialMask);
@@ -318,7 +318,7 @@ __simd_vf__ inline void VFComputeMaxExpMXFP4Vf(__ubuf__ T* srcAddr, __ubuf__ uin
318 }318 }
319 319 
320 AscendC::Reg::Max(vdMaxExp, vdExpExtract0, vdExpExtract1, scaleMask1);320 AscendC::Reg::Max(vdMaxExp, vdExpExtract0, vdExpExtract1, scaleMask1);
321- AscendC::Reg::ReduceMaxWithDataBlock(vdMaxExp, vdMaxExp, scaleMask1);321+ AscendC::Reg::ReduceDataBlock<ReduceType::MAX>(vdMaxExp, vdMaxExp, scaleMask1);
322 322 
323 AscendC::Reg::StoreUnAlign<uint16_t, AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(maxExpAddr, vdMaxExp, u1,323 AscendC::Reg::StoreUnAlign<uint16_t, AscendC::Reg::PostLiteral::POST_MODE_UPDATE>(maxExpAddr, vdMaxExp, u1,
324 scaleNum);324 scaleNum);
@@ -581,15 +581,15 @@ __simd_vf__ inline void VFProcessSwigluGroupQuantVf(__ubuf__ T0* yLocalAddr, __u
581 }581 }
582 Muls(xAbsLeft, xLeft, 0.0f, pregMain);582 Muls(xAbsLeft, xLeft, 0.0f, pregMain);
583 Compare<float, CMPMODE::NE>(compareLeft, xAbsLeft, xAbsLeft, pregMain);583 Compare<float, CMPMODE::NE>(compareLeft, xAbsLeft, xAbsLeft, pregMain);
584- MaskNot(compareLeft, compareLeft, pregMain);584+ Not(compareLeft, compareLeft, pregMain);
585 Abs(xAbsLeft, xLeft, compareLeft);585 Abs(xAbsLeft, xLeft, compareLeft);
586- ReduceMax(scale0, xAbsLeft, pregMain);586+ Reduce<ReduceType::MAX>(scale0, xAbsLeft, pregMain);
587 587 
588 Muls(xAbsRight, xRight, 0.0f, pregMain);588 Muls(xAbsRight, xRight, 0.0f, pregMain);
589 Compare<float, CMPMODE::NE>(compareRight, xAbsRight, xAbsRight, pregMain);589 Compare<float, CMPMODE::NE>(compareRight, xAbsRight, xAbsRight, pregMain);
590- MaskNot(compareRight, compareRight, pregMain);590+ Not(compareRight, compareRight, pregMain);
591 Abs(xAbsRight, xRight, compareRight);591 Abs(xAbsRight, xRight, compareRight);
592- ReduceMax(scale1, xAbsRight, pregMain);592+ Reduce<ReduceType::MAX>(scale1, xAbsRight, pregMain);
593 Max(scale, scale0, scale1, pregMerge);593 Max(scale, scale0, scale1, pregMerge);
594 Maxs(clampScale, scale, 0.0001f, pregMerge); // amax594 Maxs(clampScale, scale, 0.0001f, pregMerge); // amax
595 595 
@@ -753,7 +753,7 @@ __simd_vf__ inline void VFProcessGroupIndexSmallVf(__ubuf__ T* yLocalAddr, __ubu
753 Adds(x, x, static_cast<T>(0), pregLoop);753 Adds(x, x, static_cast<T>(0), pregLoop);
754 Add(sum, sum, x, pregMain);754 Add(sum, sum, x, pregMain);
755 }755 }
756- ReduceSum(sum, sum, pregMain);756+ Reduce<ReduceType::SUM>(sum, sum, pregMain);
757 if (withUbReduce) {757 if (withUbReduce) {
758 RegTensor<T> origin;758 RegTensor<T> origin;
759 LoadAlign(origin, yLocalAddr);759 LoadAlign(origin, yLocalAddr);
@@ -801,7 +801,7 @@ __simd_vf__ inline void VFProcessGroupIndexLargeVf(__ubuf__ T* yLocalAddr, __ubu
801 Add(sum0, sum0, sum1, pregMain);801 Add(sum0, sum0, sum1, pregMain);
802 Add(sum2, sum2, sum3, pregMain);802 Add(sum2, sum2, sum3, pregMain);
803 Add(sum0, sum0, sum2, pregMain);803 Add(sum0, sum0, sum2, pregMain);
804- ReduceSum(sum0, sum0, pregMain);804+ Reduce<ReduceType::SUM>(sum0, sum0, pregMain);
805 if (withUbReduce) {805 if (withUbReduce) {
806 RegTensor<T> origin;806 RegTensor<T> origin;
807 LoadAlign(origin, yLocalAddr);807 LoadAlign(origin, yLocalAddr);
@@ -942,7 +942,7 @@ __simd_vf__ inline void VFProcessSwigluMxFp8InvScaleVf(__ubuf__ T0* yOriginLocal
942 // fp32场景,32个数对应4个Block;先做一次Max,接着ReduceMaxWithBlock942 // fp32场景,32个数对应4个Block;先做一次Max,接着ReduceMaxWithBlock
943 // 然后DeInterLeave并且Max,得到每4个Block的最大值943 // 然后DeInterLeave并且Max,得到每4个Block的最大值
944 Max(yMax0, yLayout0, yLayout1, pregMain0); // 4 --> 2944 Max(yMax0, yLayout0, yLayout1, pregMain0); // 4 --> 2
945- ReduceMaxWithDataBlock(yMax1, yMax0, pregMain0);945+ ReduceDataBlock<ReduceType::MAX>(yMax1, yMax0, pregMain0);
946 DeInterleave(yMax1Layout0, yMax1Layout1, yMax1, yMax1);946 DeInterleave(yMax1Layout0, yMax1Layout1, yMax1, yMax1);
947 947 
948 Max(scale, yMax1Layout0, yMax1Layout1, pregMerge);948 Max(scale, yMax1Layout0, yMax1Layout1, pregMerge);
@@ -1,12 +1,11 @@
1/**1/**
2- * This program is free software, you can redistribute it and/or modify.2+ * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd.
3- * Copyright (c) 2025 Huawei Technologies Co., Ltd.3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * This file is a part of the CANN Open Software.4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Licensed under CANN Open Software License Agreement Version 2.0 (the "License").
6 * Please refer to the License for details. You may not use this file except in compliance with the License.5 * Please refer to the License for details. You may not use this file except in compliance with the License.
7- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
8- * BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. See LICENSE in the root of7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
9- * the software repository for the full text of the License.8+ * See LICENSE in the root of the software repository for the full text of the License.
10 */9 */
11 10 
12/*!11/*!
@@ -52,10 +51,10 @@ public:
52 LocalTensor<T> outLocal, int64_t dataCount)51 LocalTensor<T> outLocal, int64_t dataCount)
53 {52 {
54 // input + scalar * tensor1 * tensor253 // input + scalar * tensor1 * tensor2
55- __local_mem__ T* inUbAddr = (__ubuf__ T*)inLocal.GetPhyAddr();54+ __ubuf__ T* inUbAddr = (__ubuf__ T*)inLocal.GetPhyAddr();
56- __local_mem__ T* tensorOneUbAddr = (__ubuf__ T*)tensorOneLocal.GetPhyAddr();55+ __ubuf__ T* tensorOneUbAddr = (__ubuf__ T*)tensorOneLocal.GetPhyAddr();
57- __local_mem__ T* tensorTwoUbAddr = (__ubuf__ T*)tensorTwoLocal.GetPhyAddr();56+ __ubuf__ T* tensorTwoUbAddr = (__ubuf__ T*)tensorTwoLocal.GetPhyAddr();
58- __local_mem__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();57+ __ubuf__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();
59 58 
60 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);59 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);
61 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);60 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);
@@ -88,13 +87,13 @@ public:
88 RegTensor<T> tensorTwoReg;87 RegTensor<T> tensorTwoReg;
89 for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) {88 for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) {
90 maskReg = UpdateMask<float>(sreg);89 maskReg = UpdateMask<float>(sreg);
91- DataCopy(inReg, inUbAddr + i * dataCountPerLoop);90+ LoadAlign(inReg, inUbAddr + i * dataCountPerLoop);
92- DataCopy(tensorOneReg, tensorOneUbAddr + i * dataCountPerLoop);91+ LoadAlign(tensorOneReg, tensorOneUbAddr + i * dataCountPerLoop);
93- DataCopy(tensorTwoReg, tensorTwoUbAddr + i * dataCountPerLoop);92+ LoadAlign(tensorTwoReg, tensorTwoUbAddr + i * dataCountPerLoop);
94 Mul(tensorOneReg, tensorOneReg, tensorTwoReg, maskReg);93 Mul(tensorOneReg, tensorOneReg, tensorTwoReg, maskReg);
95 Muls(tensorOneReg, tensorOneReg, scalarVal, maskReg);94 Muls(tensorOneReg, tensorOneReg, scalarVal, maskReg);
96 Add(inReg, inReg, tensorOneReg, maskReg);95 Add(inReg, inReg, tensorOneReg, maskReg);
97- DataCopy(outUbAddr + i * dataCountPerLoop, inReg, maskReg);96+ StoreAlign(outUbAddr + i * dataCountPerLoop, inReg, maskReg);
98 }97 }
99 }98 }
100 }99 }
@@ -106,4 +105,4 @@ private:
106};105};
107} // namespace ForeachAddcmulScalar106} // namespace ForeachAddcmulScalar
108 107 
109-#endif // FOREACH_ADDCMUL_SCALAR_REGBASE_H108+#endif // FOREACH_ADDCMUL_SCALAR_REGBASE_H
@@ -43,8 +43,8 @@ public:
43 __aicore__ inline void Compute(LocalTensor<T> tensorLocal, LocalTensor<T> outLocal, int64_t tensorIndex,43 __aicore__ inline void Compute(LocalTensor<T> tensorLocal, LocalTensor<T> outLocal, int64_t tensorIndex,
44 int64_t dataCount)44 int64_t dataCount)
45 {45 {
46- __local_mem__ T* inUbAddr = (__ubuf__ T*)tensorLocal.GetPhyAddr();46+ __ubuf__ T* inUbAddr = (__ubuf__ T*)tensorLocal.GetPhyAddr();
47- __local_mem__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();47+ __ubuf__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();
48 48 
49 float scaleVal = float(inScalarGM_.GetValue(tensorIndex));49 float scaleVal = float(inScalarGM_.GetValue(tensorIndex));
50 50 
@@ -73,4 +73,4 @@ private:
73};73};
74} // namespace ForeachDivScalarList74} // namespace ForeachDivScalarList
75 75 
76-#endif // FOREACH_DIV_SCALAR_LIST_REGBASE_H76+#endif // FOREACH_DIV_SCALAR_LIST_REGBASE_H
@@ -49,9 +49,9 @@ public:
49 {49 {
50 // tensor1 + weight * (tensor2 - tensor1)50 // tensor1 + weight * (tensor2 - tensor1)
51 // tensor2 + (tensor2 - tensor1) * (weight - 1)51 // tensor2 + (tensor2 - tensor1) * (weight - 1)
52- __local_mem__ T* tensorOneUbAddr = (__ubuf__ T*)tensorOneLocal.GetPhyAddr();52+ __ubuf__ T* tensorOneUbAddr = (__ubuf__ T*)tensorOneLocal.GetPhyAddr();
53- __local_mem__ T* tensorTwoUbAddr = (__ubuf__ T*)tensorTwoLocal.GetPhyAddr();53+ __ubuf__ T* tensorTwoUbAddr = (__ubuf__ T*)tensorTwoLocal.GetPhyAddr();
54- __local_mem__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();54+ __ubuf__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();
55 55 
56 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);56 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);
57 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);57 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);
@@ -43,8 +43,8 @@ public:
43 __aicore__ inline void Compute(LocalTensor<T> tensorLocal, LocalTensor<T> outLocal, int64_t tensorIndex,43 __aicore__ inline void Compute(LocalTensor<T> tensorLocal, LocalTensor<T> outLocal, int64_t tensorIndex,
44 int64_t dataCount)44 int64_t dataCount)
45 {45 {
46- __local_mem__ T* inUbAddr = (__ubuf__ T*)tensorLocal.GetPhyAddr();46+ __ubuf__ T* inUbAddr = (__ubuf__ T*)tensorLocal.GetPhyAddr();
47- __local_mem__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();47+ __ubuf__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();
48 48 
49 using scalarCalcType = typename Conditional<AscendC::IsSameType<ScalarT, int32_t>::value, int32_t, float>::type;49 using scalarCalcType = typename Conditional<AscendC::IsSameType<ScalarT, int32_t>::value, int32_t, float>::type;
50 scalarCalcType scaleVal = scalarCalcType(inScalarGM_.GetValue(tensorIndex));50 scalarCalcType scaleVal = scalarCalcType(inScalarGM_.GetValue(tensorIndex));
@@ -69,9 +69,9 @@ public:
69 RegTensor<T> inReg;69 RegTensor<T> inReg;
70 for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) {70 for (uint16_t i = 0; i < (uint16_t)repeatTimes; i++) {
71 maskReg = UpdateMask<float>(sreg);71 maskReg = UpdateMask<float>(sreg);
72- DataCopy(inReg, inUbAddr + i * dataCountPerLoop);72+ LoadAlign(inReg, inUbAddr + i * dataCountPerLoop);
73 Muls(inReg, inReg, scaleVal, maskReg);73 Muls(inReg, inReg, scaleVal, maskReg);
74- DataCopy(outUbAddr + i * dataCountPerLoop, inReg, maskReg);74+ StoreAlign(outUbAddr + i * dataCountPerLoop, inReg, maskReg);
75 }75 }
76 }76 }
77 }77 }
@@ -82,4 +82,4 @@ private:
82};82};
83} // namespace ForeachMulScalarList83} // namespace ForeachMulScalarList
84 84 
85-#endif // FOREACH_MUL_SCALAR_LIST_REGBASE_H85+#endif // FOREACH_MUL_SCALAR_LIST_REGBASE_H
@@ -189,10 +189,10 @@ __aicore__ inline void ForeachNonFiniteCheckAndUnscaleNDRegbase<T>::Compute(uint
189 LocalTensor<T> computeInLT = copyInQueue_.DeQue<T>();189 LocalTensor<T> computeInLT = copyInQueue_.DeQue<T>();
190 LocalTensor<T> computeOutLT = copyOutQueue_.AllocTensor<T>();190 LocalTensor<T> computeOutLT = copyOutQueue_.AllocTensor<T>();
191 191 
192- __local_mem__ T* inUbAddr = (__ubuf__ T*)computeInLT.GetPhyAddr();192+ __ubuf__ T* inUbAddr = (__ubuf__ T*)computeInLT.GetPhyAddr();
193- __local_mem__ T* outUbAddr = (__ubuf__ T*)computeOutLT.GetPhyAddr();193+ __ubuf__ T* outUbAddr = (__ubuf__ T*)computeOutLT.GetPhyAddr();
194- __local_mem__ float* minUbAddr = (__ubuf__ float*)minLocal_.GetPhyAddr();194+ __ubuf__ float* minUbAddr = (__ubuf__ float*)minLocal_.GetPhyAddr();
195- __local_mem__ float* maxUbAddr = (__ubuf__ float*)maxLocal_.GetPhyAddr();195+ __ubuf__ float* maxUbAddr = (__ubuf__ float*)maxLocal_.GetPhyAddr();
196 196 
197 uint32_t dataCountPerLoop = platform::GetVRegSize() / sizeof(float);197 uint32_t dataCountPerLoop = platform::GetVRegSize() / sizeof(float);
198 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);198 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);
@@ -229,14 +229,14 @@ __aicore__ inline void ForeachNonFiniteCheckAndUnscaleNDRegbase<T>::Compute(uint
229 Muls(y, x, invScaleVal, maskReg);229 Muls(y, x, invScaleVal, maskReg);
230 ops::StoreOneTensorForDtypeT<T>(outUbAddr, y, maskReg, (repeatTimes - 1) * dataCountPerLoop);230 ops::StoreOneTensorForDtypeT<T>(outUbAddr, y, maskReg, (repeatTimes - 1) * dataCountPerLoop);
231 }231 }
232- DataCopy(lastMin, (__local_mem__ float*)(minUbAddr));232+ LoadAlign(lastMin, (__ubuf__ float*)(minUbAddr));
233- DataCopy(lastMax, (__local_mem__ float*)(maxUbAddr));233+ LoadAlign(lastMax, (__ubuf__ float*)(maxUbAddr));
234- ReduceMax(max, max, pregMain);234+ Reduce<ReduceType::MAX>(max, max, pregMain);
235- ReduceMin(min, min, pregMain);235+ Reduce<ReduceType::MIN>(min, min, pregMain);
236 Max(lastMax, lastMax, max, pregMain);236 Max(lastMax, lastMax, max, pregMain);
237 Min(lastMin, lastMin, min, pregMain);237 Min(lastMin, lastMin, min, pregMain);
238- DataCopy<float, StoreDist::DIST_FIRST_ELEMENT_B32>(((__local_mem__ float*)minUbAddr), lastMin, pregMerge);238+ StoreAlign<float, StoreDist::DIST_FIRST_ELEMENT_B32>(((__ubuf__ float*)minUbAddr), lastMin, pregMerge);
239- DataCopy<float, StoreDist::DIST_FIRST_ELEMENT_B32>(((__local_mem__ float*)maxUbAddr), lastMax, pregMerge);239+ StoreAlign<float, StoreDist::DIST_FIRST_ELEMENT_B32>(((__ubuf__ float*)maxUbAddr), lastMax, pregMerge);
240 }240 }
241 copyOutQueue_.EnQue(computeOutLT);241 copyOutQueue_.EnQue(computeOutLT);
242 copyInQueue_.FreeTensor(computeInLT);242 copyInQueue_.FreeTensor(computeInLT);
@@ -280,4 +280,4 @@ __aicore__ inline __gm__ T* ForeachNonFiniteCheckAndUnscaleNDRegbase<T>::GetTens
280 280 
281} // namespace ForeachNonFiniteCheckAndUnscaleRegbase281} // namespace ForeachNonFiniteCheckAndUnscaleRegbase
282 282 
283-#endif // FOREACH_NON_FINITE_CHECK_AND_UNSCALE_REGBASE_H283+#endif // FOREACH_NON_FINITE_CHECK_AND_UNSCALE_REGBASE_H
@@ -137,8 +137,8 @@ public:
137 template <bool isMulitNum, uint8_t modelOrd>137 template <bool isMulitNum, uint8_t modelOrd>
138 __aicore__ inline void ReduceComputeImpl(LocalTensor<float> inLocal, LocalTensor<T> outLocal, uint16_t dataCount)138 __aicore__ inline void ReduceComputeImpl(LocalTensor<float> inLocal, LocalTensor<T> outLocal, uint16_t dataCount)
139 {139 {
140- __local_mem__ float* inUbAddr = (__ubuf__ float*)inLocal.GetPhyAddr();140+ __ubuf__ float* inUbAddr = (__ubuf__ float*)inLocal.GetPhyAddr();
141- __local_mem__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();141+ __ubuf__ T* outUbAddr = (__ubuf__ T*)outLocal.GetPhyAddr();
142 142 
143 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);143 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);
144 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);144 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);
@@ -168,9 +168,9 @@ public:
168 }168 }
169 if constexpr (isMulitNum) {169 if constexpr (isMulitNum) {
170 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {170 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {
171- ReduceMax<float>(inRegToFloat, reduceFloat, pregMain);171+ Reduce<ReduceType::MAX>(inRegToFloat, reduceFloat, pregMain);
172 } else {172 } else {
173- ReduceSum<float>(inRegToFloat, reduceFloat, pregMain);173+ Reduce<ReduceType::SUM>(inRegToFloat, reduceFloat, pregMain);
174 }174 }
175 }175 }
176 if constexpr (modelOrd == TWO_SCALAR_NORM_MODEL_CODE) {176 if constexpr (modelOrd == TWO_SCALAR_NORM_MODEL_CODE) {
@@ -196,8 +196,8 @@ public:
196 __aicore__ inline void DoCompute(LocalTensor<T> inLocal, LocalTensor<float> outLocal, int64_t dataCount,196 __aicore__ inline void DoCompute(LocalTensor<T> inLocal, LocalTensor<float> outLocal, int64_t dataCount,
197 uint16_t outOffset)197 uint16_t outOffset)
198 {198 {
199- __local_mem__ T* inUbAddr = (__ubuf__ T*)inLocal.GetPhyAddr();199+ __ubuf__ T* inUbAddr = (__ubuf__ T*)inLocal.GetPhyAddr();
200- __local_mem__ float* outUbAddr = (__ubuf__ float*)outLocal.GetPhyAddr();200+ __ubuf__ float* outUbAddr = (__ubuf__ float*)outLocal.GetPhyAddr();
201 201 
202 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);202 uint32_t dataCountPerLoop = VL_SIZE / sizeof(float);
203 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);203 uint16_t repeatTimes = CeilDivision(dataCount, dataCountPerLoop);
@@ -226,24 +226,24 @@ public:
226 }226 }
227 }227 }
228 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {228 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {
229- ReduceMax(reduceFloat, reduceFloat, pregMain);229+ Reduce<ReduceType::MAX>(reduceFloat, reduceFloat, pregMain);
230 } else {230 } else {
231- ReduceSum(reduceFloat, reduceFloat, pregMain);231+ Reduce<ReduceType::SUM>(reduceFloat, reduceFloat, pregMain);
232 }232 }
233 MicroAPI::MaskReg pregMarge = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::VL1>();233 MicroAPI::MaskReg pregMarge = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::VL1>();
234 if constexpr (isFirst) {234 if constexpr (isFirst) {
235- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(235+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
236- (__local_mem__ float*)(outUbAddr + outOffset), reduceFloat, pregMarge);236+ (__ubuf__ float*)(outUbAddr + outOffset), reduceFloat, pregMarge);
237 } else {237 } else {
238- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(238+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(outFloat,
239- outFloat, (__local_mem__ float*)(outUbAddr + outOffset));239+ (__ubuf__ float*)(outUbAddr + outOffset));
240 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {240 if constexpr (modelOrd == POSITIVE_INF_SCALAR_NORM_MODEL_CODE) {
241 Max(outFloat, outFloat, reduceFloat, pregMarge);241 Max(outFloat, outFloat, reduceFloat, pregMarge);
242 } else {242 } else {
243 Add(outFloat, outFloat, reduceFloat, pregMarge);243 Add(outFloat, outFloat, reduceFloat, pregMarge);
244 }244 }
245- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(245+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
246- (__local_mem__ float*)(outUbAddr + outOffset), outFloat, pregMarge);246+ (__ubuf__ float*)(outUbAddr + outOffset), outFloat, pregMarge);
247 }247 }
248 }248 }
249 }249 }
@@ -358,4 +358,4 @@ protected:
358 358 
359} // namespace ForeachNorm359} // namespace ForeachNorm
360 360 
361-#endif // FOREACH_NORM_N_D_H361+#endif // FOREACH_NORM_N_D_H
@@ -78,15 +78,15 @@ struct CalcBCEWithLogitsV2 : public Vec::ElemwiseQuaternaryOP<T, T, T, T, T> {
78 for (uint16_t loop = 0; loop < (uint16_t)repeatTimes; loop++) {78 for (uint16_t loop = 0; loop < (uint16_t)repeatTimes; loop++) {
79 pregUp = MicroAPI::UpdateMask<T>(totalLen);79 pregUp = MicroAPI::UpdateMask<T>(totalLen);
80 AscendC::MicroAPI::Duplicate(regOne, (T)1.0f, pregUp);80 AscendC::MicroAPI::Duplicate(regOne, (T)1.0f, pregUp);
81- MicroAPI::DataCopy<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regX, xAddr, (int32_t)oneRepeat);81+ MicroAPI::LoadAlign<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regX, xAddr, (int32_t)oneRepeat);
82- MicroAPI::DataCopy<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regY, yAddr, (int32_t)oneRepeat);82+ MicroAPI::LoadAlign<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regY, yAddr, (int32_t)oneRepeat);
83 if constexpr (HAS_WEIGHT) {83 if constexpr (HAS_WEIGHT) {
84- MicroAPI::DataCopy<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regWeight, weightAddr,84+ MicroAPI::LoadAlign<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regWeight, weightAddr,
85- (int32_t)oneRepeat);85+ (int32_t)oneRepeat);
86 }86 }
87 if constexpr (HAS_POS_WEIGHT) {87 if constexpr (HAS_POS_WEIGHT) {
88- MicroAPI::DataCopy<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regPosWeight, posWeightAddr,88+ MicroAPI::LoadAlign<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(regPosWeight, posWeightAddr,
89- (int32_t)oneRepeat);89+ (int32_t)oneRepeat);
90 }90 }
91 91 
92 MicroAPI::Mins(regMinVal, regX, (T)0.0f, pregUp);92 MicroAPI::Mins(regMinVal, regX, (T)0.0f, pregUp);
@@ -113,8 +113,8 @@ struct CalcBCEWithLogitsV2 : public Vec::ElemwiseQuaternaryOP<T, T, T, T, T> {
113 MicroAPI::Mul(regLoss, regLoss, regWeight, pregUp);113 MicroAPI::Mul(regLoss, regLoss, regWeight, pregUp);
114 }114 }
115 115 
116- MicroAPI::DataCopy<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(lossAddr, regLoss, (int32_t)oneRepeat,116+ MicroAPI::StoreAlign<T, MicroAPI::PostLiteral::POST_MODE_UPDATE>(lossAddr, regLoss, (int32_t)oneRepeat,
117- pregUp);117+ pregUp);
118 }118 }
119 }119 }
120#endif120#endif
@@ -219,4 +219,4 @@ struct SigmoidCEWithLogitsV2 {
219 using OpDag = DAGSch<Outputs, void, MemCfg>;219 using OpDag = DAGSch<Outputs, void, MemCfg>;
220};220};
221} // namespace SigmoidCrossEntropyWithLogitsV2221} // namespace SigmoidCrossEntropyWithLogitsV2
222-#endif // ASCENDC_SIGMOID_CROSS_ENTROPY_WITH_LOGITS_V2_DAG_H_222+#endif // ASCENDC_SIGMOID_CROSS_ENTROPY_WITH_LOGITS_V2_DAG_H_
@@ -20,7 +20,8 @@ using namespace AscendC;
20using namespace AscendC::MicroAPI;20using namespace AscendC::MicroAPI;
21using AscendC::MicroAPI::MaskReg;21using AscendC::MicroAPI::MaskReg;
22using AscendC::MicroAPI::RegTensor;22using AscendC::MicroAPI::RegTensor;
23-using AscendC::MicroAPI::UnalignReg;23+using AscendC::MicroAPI::UnalignRegForLoad;
24+using AscendC::MicroAPI::UnalignRegForStore;
24static constexpr int32_t BLOCK_SIZE = 32;25static constexpr int32_t BLOCK_SIZE = 32;
25static constexpr int32_t FP32_ONE_REPEAT = 64;26static constexpr int32_t FP32_ONE_REPEAT = 64;
26static constexpr int32_t FLOAT_BYTE_SIZE = 4;27static constexpr int32_t FLOAT_BYTE_SIZE = 4;
@@ -97,57 +98,55 @@ __aicore__ inline uint32_t RoundDown(uint32_t x)
97}98}
98 99 
99template <typename T>100template <typename T>
100-__aicore__ inline void LoadInputData(RegTensor<float>& dst, __local_mem__ T* src, MaskReg pregLoop, uint32_t srcOffset)101+__aicore__ inline void LoadInputData(RegTensor<float>& dst, __ubuf__ T* src, MaskReg pregLoop, uint32_t srcOffset)
101{102{
102 if constexpr (IsSameType<T, float>::value) {103 if constexpr (IsSameType<T, float>::value) {
103- DataCopy(dst, src + srcOffset);104+ LoadAlign(dst, src + srcOffset);
104 } else {105 } else {
105 RegTensor<T> tmp;106 RegTensor<T> tmp;
106- DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(tmp, src + srcOffset);107+ LoadAlign<T, AscendC::MicroAPI::LoadDist::DIST_UNPACK_B16>(tmp, src + srcOffset);
107 Cast<float, T, castTraitB162B32Even>(dst, tmp, pregLoop);108 Cast<float, T, castTraitB162B32Even>(dst, tmp, pregLoop);
108 }109 }
109}110}
110 111 
111template <typename T, bool hasGamma, bool hasBeta>112template <typename T, bool hasGamma, bool hasBeta>
112-__aicore__ inline void LoadGammaAndBetaData(RegTensor<float>& gamma, RegTensor<float>& beta,113+__aicore__ inline void LoadGammaAndBetaData(RegTensor<float>& gamma, RegTensor<float>& beta, __ubuf__ T* gammaLocal,
113- __local_mem__ T* gammaLocal, __local_mem__ T* betaLocal, MaskReg pregLoop,114+ __ubuf__ T* betaLocal, MaskReg pregLoop, uint32_t srcOffset)
114- uint32_t srcOffset)
115{115{
116 if constexpr (IsSameType<T, float>::value) {116 if constexpr (IsSameType<T, float>::value) {
117 if constexpr (hasGamma) {117 if constexpr (hasGamma) {
118- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(gamma, gammaLocal + srcOffset);118+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(gamma, gammaLocal + srcOffset);
119 }119 }
120 if constexpr (hasBeta) {120 if constexpr (hasBeta) {
121- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(beta, betaLocal + srcOffset);121+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(beta, betaLocal + srcOffset);
122 }122 }
123 } else {123 } else {
124 if constexpr (hasGamma) {124 if constexpr (hasGamma) {
125 RegTensor<T> gammaB16;125 RegTensor<T> gammaB16;
126- DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(gammaB16, gammaLocal + srcOffset);126+ LoadAlign<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(gammaB16, gammaLocal + srcOffset);
127 Cast<float, T, castTraitB162B32Even>(gamma, gammaB16, pregLoop);127 Cast<float, T, castTraitB162B32Even>(gamma, gammaB16, pregLoop);
128 }128 }
129 if constexpr (hasBeta) {129 if constexpr (hasBeta) {
130 RegTensor<T> betaB16;130 RegTensor<T> betaB16;
131- DataCopy<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(betaB16, betaLocal + srcOffset);131+ LoadAlign<T, AscendC::MicroAPI::LoadDist::DIST_BRC_B16>(betaB16, betaLocal + srcOffset);
132 Cast<float, T, castTraitB162B32Even>(beta, betaB16, pregLoop);132 Cast<float, T, castTraitB162B32Even>(beta, betaB16, pregLoop);
133 }133 }
134 }134 }
135}135}
136 136 
137template <typename T>137template <typename T>
138-__aicore__ inline void StoreOutputData(__local_mem__ T* dst, RegTensor<float>& src, MaskReg pregLoop,138+__aicore__ inline void StoreOutputData(__ubuf__ T* dst, RegTensor<float>& src, MaskReg pregLoop, uint32_t dstOffset)
139- uint32_t dstOffset)
140{139{
141 if constexpr (IsSameType<T, float>::value) {140 if constexpr (IsSameType<T, float>::value) {
142- DataCopy(dst + dstOffset, src, pregLoop);141+ StoreAlign(dst + dstOffset, src, pregLoop);
143 } else {142 } else {
144 RegTensor<T> tmpB16;143 RegTensor<T> tmpB16;
145 Cast<T, float, castTraitB322B16Even>(tmpB16, src, pregLoop);144 Cast<T, float, castTraitB322B16Even>(tmpB16, src, pregLoop);
146- DataCopy<T, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(dst + dstOffset, tmpB16, pregLoop);145+ StoreAlign<T, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(dst + dstOffset, tmpB16, pregLoop);
147 }146 }
148}147}
149 148 
150-__aicore__ inline void DichotomyAdd(RegTensor<float>& dstReg, __local_mem__ float* src, uint16_t outerLoop,149+__aicore__ inline void DichotomyAdd(RegTensor<float>& dstReg, __ubuf__ float* src, uint16_t outerLoop,
151 uint16_t innerLoop, uint32_t lastNum)150 uint16_t innerLoop, uint32_t lastNum)
152{151{
153 RegTensor<float> tmpReg1;152 RegTensor<float> tmpReg1;
@@ -158,17 +157,17 @@ __aicore__ inline void DichotomyAdd(RegTensor<float>& dstReg, __local_mem__ floa
158 for (uint16_t k = 0; k < outerLoop; k++) {157 for (uint16_t k = 0; k < outerLoop; k++) {
159 innerLoop = innerLoop / DICHOTOMY_ADD_COEFF;158 innerLoop = innerLoop / DICHOTOMY_ADD_COEFF;
160 for (uint16_t i = 0; i < innerLoop; i++) {159 for (uint16_t i = 0; i < innerLoop; i++) {
161- DataCopy(tmpReg1, src + i * VL_FP32);160+ LoadAlign(tmpReg1, src + i * VL_FP32);
162- DataCopy(tmpReg2, src + (i + innerLoop) * VL_FP32);161+ LoadAlign(tmpReg2, src + (i + innerLoop) * VL_FP32);
163 Add(tmpReg3, tmpReg1, tmpReg2, pregMain);162 Add(tmpReg3, tmpReg1, tmpReg2, pregMain);
164- DataCopy(src + i * VL_FP32, tmpReg3, pregMain);163+ StoreAlign(src + i * VL_FP32, tmpReg3, pregMain);
165 }164 }
166 LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();165 LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
167 }166 }
168 uint32_t sreg0 = lastNum;167 uint32_t sreg0 = lastNum;
169 MaskReg pregLoop = UpdateMask<float>(sreg0);168 MaskReg pregLoop = UpdateMask<float>(sreg0);
170- DataCopy(tmpReg3, src);169+ LoadAlign(tmpReg3, src);
171- ReduceSum(dstReg, tmpReg3, pregLoop);170+ Reduce<ReduceType::SUM>(dstReg, tmpReg3, pregLoop);
172}171}
173 172 
174__aicore__ inline void CalRstdByHighPrecision(RegTensor<float>& var, RegTensor<float>& rstd, float epsilon)173__aicore__ inline void CalRstdByHighPrecision(RegTensor<float>& var, RegTensor<float>& rstd, float epsilon)
@@ -211,16 +210,15 @@ __aicore__ inline void CalRstdByHighPrecision(RegTensor<float>& var, RegTensor<f
211 Mula(s, var, r, pregMerge); // s + x * t210 Mula(s, var, r, pregMerge); // s + x * t
212 Mul(s, s, rstd, pregMerge); // e * y211 Mul(s, s, rstd, pregMerge); // e * y
213 Mula(rstd, s, scalar1, pregMerge); // y + y * e * 0.5212 Mula(rstd, s, scalar1, pregMerge); // y + y * e * 0.5
214- CompareScalar(cmpReg1, var, POS_INF, pregMerge);213+ Compares(cmpReg1, var, POS_INF, pregMerge);
215 Select(rstd, scalar3, rstd, cmpReg1);214 Select(rstd, scalar3, rstd, cmpReg1);
216- CompareScalar(cmpReg2, var, ZERO, pregMerge);215+ Compares(cmpReg2, var, ZERO, pregMerge);
217 Select(rstd, scalar2, rstd, cmpReg2);216 Select(rstd, scalar2, rstd, cmpReg2);
218}217}
219 218 
220template <typename T>219template <typename T>
221-__aicore__ inline void VFInnerWelfordParallelUpdateWithInit(__local_mem__ T* x1Local, __local_mem__ float* tmpMeanLocal,220+__aicore__ inline void VFInnerWelfordParallelUpdateWithInit(__ubuf__ T* x1Local, __ubuf__ float* tmpMeanLocal,
222- __local_mem__ float* tmpVarLocal, uint64_t calLen,221+ __ubuf__ float* tmpVarLocal, uint64_t calLen, float scale)
223- float scale)
224{222{
225 uint16_t loopCount = CeilDiv(calLen, VL_FP32);223 uint16_t loopCount = CeilDiv(calLen, VL_FP32);
226 __VEC_SCOPE__224 __VEC_SCOPE__
@@ -241,13 +239,13 @@ __aicore__ inline void VFInnerWelfordParallelUpdateWithInit(__local_mem__ T* x1L
241 Sub(delta1, x1, tmpMean, pregLoop);239 Sub(delta1, x1, tmpMean, pregLoop);
242 Muls(delta2, delta1, scale, pregLoop);240 Muls(delta2, delta1, scale, pregLoop);
243 Add(tmpMean, tmpMean, delta2, pregLoop);241 Add(tmpMean, tmpMean, delta2, pregLoop);
244- DataCopy(tmpMeanLocal + i * VL_FP32, tmpMean, pregLoop);242+ StoreAlign(tmpMeanLocal + i * VL_FP32, tmpMean, pregLoop);
245 243 
246 Duplicate(tmpVar, 0.0, pregLoop);244 Duplicate(tmpVar, 0.0, pregLoop);
247 Sub(delta3, x1, tmpMean, pregLoop);245 Sub(delta3, x1, tmpMean, pregLoop);
248 Mul(delat4, delta1, delta3, pregLoop);246 Mul(delat4, delta1, delta3, pregLoop);
249 Add(tmpVar, tmpVar, delat4, pregLoop);247 Add(tmpVar, tmpVar, delat4, pregLoop);
250- DataCopy(tmpVarLocal + i * VL_FP32, tmpVar, pregLoop);248+ StoreAlign(tmpVarLocal + i * VL_FP32, tmpVar, pregLoop);
251 }249 }
252 }250 }
253}251}
@@ -262,8 +260,8 @@ __aicore__ inline void VFInnerWelfordParallelUpdateWithInit(__local_mem__ T* x1L
262 return count, mean, var260 return count, mean, var
263*/261*/
264template <typename T>262template <typename T>
265-__aicore__ inline void VFInnerWelfordParallelUpdate(__local_mem__ T* x1Local, __local_mem__ float* tmpMeanLocal,263+__aicore__ inline void VFInnerWelfordParallelUpdate(__ubuf__ T* x1Local, __ubuf__ float* tmpMeanLocal,
266- __local_mem__ float* tmpVarLocal, uint64_t calLen, float scale)264+ __ubuf__ float* tmpVarLocal, uint64_t calLen, float scale)
267{265{
268 uint16_t loopCount = CeilDiv(calLen, VL_FP32);266 uint16_t loopCount = CeilDiv(calLen, VL_FP32);
269 __VEC_SCOPE__267 __VEC_SCOPE__
@@ -280,24 +278,24 @@ __aicore__ inline void VFInnerWelfordParallelUpdate(__local_mem__ T* x1Local, __
280 for (uint16_t i = 0; i < loopCount; i++) {278 for (uint16_t i = 0; i < loopCount; i++) {
281 pregLoop = UpdateMask<float>(sreg0);279 pregLoop = UpdateMask<float>(sreg0);
282 LoadInputData<T>(x1, x1Local, pregLoop, i * VL_FP32);280 LoadInputData<T>(x1, x1Local, pregLoop, i * VL_FP32);
283- DataCopy(tmpMean, tmpMeanLocal + i * VL_FP32);281+ LoadAlign(tmpMean, tmpMeanLocal + i * VL_FP32);
284 Sub(delta1, x1, tmpMean, pregLoop);282 Sub(delta1, x1, tmpMean, pregLoop);
285 Muls(delta2, delta1, scale, pregLoop);283 Muls(delta2, delta1, scale, pregLoop);
286 Add(tmpMean, tmpMean, delta2, pregLoop);284 Add(tmpMean, tmpMean, delta2, pregLoop);
287- DataCopy(tmpMeanLocal + i * VL_FP32, tmpMean, pregLoop);285+ StoreAlign(tmpMeanLocal + i * VL_FP32, tmpMean, pregLoop);
288 286 
289- DataCopy(tmpVar, tmpVarLocal + i * VL_FP32);287+ LoadAlign(tmpVar, tmpVarLocal + i * VL_FP32);
290 Sub(delta3, x1, tmpMean, pregLoop);288 Sub(delta3, x1, tmpMean, pregLoop);
291 Mul(delat4, delta1, delta3, pregLoop);289 Mul(delat4, delta1, delta3, pregLoop);
292 Add(tmpVar, tmpVar, delat4, pregLoop);290 Add(tmpVar, tmpVar, delat4, pregLoop);
293- DataCopy(tmpVarLocal + i * VL_FP32, tmpVar, pregLoop);291+ StoreAlign(tmpVarLocal + i * VL_FP32, tmpVar, pregLoop);
294 }292 }
295 }293 }
296}294}
297 295 
298template <typename T>296template <typename T>
299-__aicore__ inline void VFWelfordParallelUpdate(__local_mem__ T* x1Local, __local_mem__ float* tmpMeanLocal,297+__aicore__ inline void VFWelfordParallelUpdate(__ubuf__ T* x1Local, __ubuf__ float* tmpMeanLocal,
300- __local_mem__ float* tmpVarLocal, uint64_t curLoop, uint64_t calLen,298+ __ubuf__ float* tmpVarLocal, uint64_t curLoop, uint64_t calLen,
301 float scale)299 float scale)
302{300{
303 if (curLoop == 0) {301 if (curLoop == 0) {
@@ -318,10 +316,9 @@ __aicore__ inline void VFWelfordParallelUpdate(__local_mem__ T* x1Local, __local
318 welford采用二分累加计算mean和variance, 基本逻辑为:316 welford采用二分累加计算mean和variance, 基本逻辑为:
319 先将尾块折叠到整块上,整尾块vadd之后,做一次vcadd回刷到UB上,剩余整块直接vcadd回刷到UB上,最后对UB上的结果做完全二分对折317 先将尾块折叠到整块上,整尾块vadd之后,做一次vcadd回刷到UB上,剩余整块直接vcadd回刷到UB上,最后对UB上的结果做完全二分对折
320*/318*/
321-__aicore__ inline void VFWelfordParallelFinalizeAlign(__local_mem__ float* meanLocal, __local_mem__ float* rstdLocal,319+__aicore__ inline void VFWelfordParallelFinalizeAlign(__ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
322- __local_mem__ float* tmpMeanLocal,320+ __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
323- __local_mem__ float* tmpVarLocal,321+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount,
324- __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,
325 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,322 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
326 uint32_t dichotomyAddLastNum, uint32_t offset, float reduceScale,323 uint32_t dichotomyAddLastNum, uint32_t offset, float reduceScale,
327 float scale, float cnt, float eps)324 float scale, float cnt, float eps)
@@ -354,27 +351,27 @@ __aicore__ inline void VFWelfordParallelFinalizeAlign(__local_mem__ float* meanL
354 // PART1: 整尾块合并351 // PART1: 整尾块合并
355 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {352 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
356 pregLoop = UpdateMask<float>(sreg0);353 pregLoop = UpdateMask<float>(sreg0);
357- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);354+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
358- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);355+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
359 Muls(dichotomyAddMeanL, dichotomyAddMeanL, scale, pregMain);356 Muls(dichotomyAddMeanL, dichotomyAddMeanL, scale, pregMain);
360 Muls(dichotomyAddMeanR, dichotomyAddMeanR, scale, pregLoop);357 Muls(dichotomyAddMeanR, dichotomyAddMeanR, scale, pregLoop);
361 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);358 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
362- ReduceSum(mean, sumMean, pregMain);359+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
363- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,360+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,
364- pregMerge);361+ pregMerge);
365 }362 }
366 363 
367 // PART2: 整块剩余部分vcadd回刷UB364 // PART2: 整块剩余部分vcadd回刷UB
368 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {365 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
369- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);366+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);
370 Muls(dichotomyAddMeanL, dichotomyAddMeanL, scale, pregMain);367 Muls(dichotomyAddMeanL, dichotomyAddMeanL, scale, pregMain);
371- ReduceSum(mean, dichotomyAddMeanL, pregMain);368+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
372- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(369+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
373 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, mean, pregMerge);370 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, mean, pregMerge);
374 }371 }
375 372 
376 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);373 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
377- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);374+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);
378 375 
379 Duplicate(one, float(1.0), pregMain);376 Duplicate(one, float(1.0), pregMain);
380 Duplicate(mean, mean, pregMain);377 Duplicate(mean, mean, pregMain);
@@ -382,45 +379,45 @@ __aicore__ inline void VFWelfordParallelFinalizeAlign(__local_mem__ float* meanL
382 // PART1: 整尾块合并379 // PART1: 整尾块合并
383 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {380 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
384 pregLoop = UpdateMask<float>(sreg0);381 pregLoop = UpdateMask<float>(sreg0);
385- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);382+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
386 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);383 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
387 Mul(deltaL, deltaL, deltaL, pregMain);384 Mul(deltaL, deltaL, deltaL, pregMain);
388 Muls(deltaL, deltaL, cnt, pregMain);385 Muls(deltaL, deltaL, cnt, pregMain);
389- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);386+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);
390 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);387 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
391 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);388 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
392 389 
393- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);390+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
394 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);391 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
395 Mul(deltaR, deltaR, deltaR, pregLoop);392 Mul(deltaR, deltaR, deltaR, pregLoop);
396 Muls(deltaR, deltaR, cnt, pregLoop);393 Muls(deltaR, deltaR, cnt, pregLoop);
397- DataCopy(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);394+ LoadAlign(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);
398 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);395 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
399 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);396 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
400 397 
401 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);398 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
402- ReduceSum(var, sumVar, pregMain);399+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
403- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,400+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,
404- pregMerge);401+ pregMerge);
405 }402 }
406 403 
407 // PART2: 整块剩余部分vcadd回刷UB404 // PART2: 整块剩余部分vcadd回刷UB
408 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {405 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
409- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);406+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);
410 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);407 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
411 Mul(deltaL, deltaL, deltaL, pregMain);408 Mul(deltaL, deltaL, deltaL, pregMain);
412 Muls(deltaL, deltaL, cnt, pregMain);409 Muls(deltaL, deltaL, cnt, pregMain);
413- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);410+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderLoopCount) * VL_FP32);
414 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);411 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
415 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);412 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
416- ReduceSum(var, dichotomyAddVarL, pregMain);413+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
417- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(414+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
418 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, var, pregMerge);415 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, var, pregMerge);
419 }416 }
420 417 
421 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);418 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
422 CalRstdByHighPrecision(var, rstd, eps);419 CalRstdByHighPrecision(var, rstd, eps);
423- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);420+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);
424 }421 }
425}422}
426 423 
@@ -448,10 +445,9 @@ __aicore__ inline void VFWelfordParallelFinalizeAlign(__local_mem__ float* meanL
448 445 
449// welford整块大于等于二分累加整块446// welford整块大于等于二分累加整块
450__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation1(447__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation1(
451- __local_mem__ float* meanLocal, __local_mem__ float* rstdLocal, __local_mem__ float* tmpMeanLocal,448+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal, __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
452- __local_mem__ float* tmpVarLocal, __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,449+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount, uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
453- uint32_t dichotomyAddPower, uint32_t dichotomyAddK, uint32_t dichotomyAddLastNum, uint32_t offset,450+ uint32_t dichotomyAddLastNum, uint32_t offset, uint32_t tailSize, float reduceScale, float cnt, float eps)
454- uint32_t tailSize, float reduceScale, float cnt, float eps)
455{451{
456 float tailCnt = cnt + float(1.0);452 float tailCnt = cnt + float(1.0);
457 float coeff = tailCnt / cnt;453 float coeff = tailCnt / cnt;
@@ -498,14 +494,14 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation1(
498 494 
499 // 整块使用tailCountScale,尾块使用tailCountScale495 // 整块使用tailCountScale,尾块使用tailCountScale
500 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {496 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
501- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);497+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
502- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);498+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
503 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);499 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
504 Muls(dichotomyAddMeanR, dichotomyAddMeanR, tailCountScale, pregMain);500 Muls(dichotomyAddMeanR, dichotomyAddMeanR, tailCountScale, pregMain);
505 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);501 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
506- ReduceSum(mean, sumMean, pregMain);502+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
507- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,503+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,
508- pregMerge);504+ pregMerge);
509 }505 }
510 506 
511 // 处理welford第一次非对齐点, 整块使用tailCountScale,尾块部分使用tailCountScale, 部分使用countScale507 // 处理welford第一次非对齐点, 整块使用tailCountScale,尾块部分使用tailCountScale, 部分使用countScale
@@ -514,144 +510,145 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation1(
514 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {510 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
515 pregLoop = UpdateMask<float>(sreg0);511 pregLoop = UpdateMask<float>(sreg0);
516 pregLoop1 = UpdateMask<float>(sreg1);512 pregLoop1 = UpdateMask<float>(sreg1);
517- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);513+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);
518- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);514+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
519 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);515 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
520 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);516 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);
521 Muls(tmp, dichotomyAddMeanR, coeff, pregLoop1);517 Muls(tmp, dichotomyAddMeanR, coeff, pregLoop1);
522- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanR, tmp, pregLoop1);518+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanR, tmp, pregLoop1);
523 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);519 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
524- ReduceSum(mean, sumMean, pregMain);520+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
525- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(521+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
526 dichotomyAddLocal + i + welfordDiffLoopCount, mean, pregMerge);522 dichotomyAddLocal + i + welfordDiffLoopCount, mean, pregMerge);
527 }523 }
528 524 
529 // 整块使用tailCountScale,尾块使用countScale525 // 整块使用tailCountScale,尾块使用countScale
530 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {526 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
531 pregLoop = UpdateMask<float>(sreg0);527 pregLoop = UpdateMask<float>(sreg0);
532- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);528+ LoadAlign(dichotomyAddMeanL,
533- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +529+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
534- dichotomyAddPower);530+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 +
531+ welfordDiffReminderAlign + dichotomyAddPower);
535 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);532 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
536 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);533 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);
537 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);534 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
538- ReduceSum(mean, sumMean, pregMain);535+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
539- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(536+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
540 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, mean, pregMerge);537 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, mean, pregMerge);
541 }538 }
542 // PART2: 整块剩余部分vcadd回刷UB,使用tailCountScale539 // PART2: 整块剩余部分vcadd回刷UB,使用tailCountScale
543 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {540 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
544- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);541+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
545 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);542 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
546- ReduceSum(mean, dichotomyAddMeanL, pregMain);543+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
547- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(544+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
548 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, mean, pregMerge);545 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, mean, pregMerge);
549 }546 }
550 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);547 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
551- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);548+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);
552 549 
553 // 计算rstd550 // 计算rstd
554 Duplicate(one, float(1.0), pregMain);551 Duplicate(one, float(1.0), pregMain);
555 Duplicate(mean, mean, pregMain);552 Duplicate(mean, mean, pregMain);
556 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {553 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
557- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);554+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
558 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);555 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
559 Mul(deltaL, deltaL, deltaL, pregMain);556 Mul(deltaL, deltaL, deltaL, pregMain);
560 Muls(deltaL, deltaL, tailCnt, pregMain);557 Muls(deltaL, deltaL, tailCnt, pregMain);
561- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);558+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
562 Sub(deltaR, dichotomyAddMeanR, mean, pregMain);559 Sub(deltaR, dichotomyAddMeanR, mean, pregMain);
563 Mul(deltaR, deltaR, deltaR, pregMain);560 Mul(deltaR, deltaR, deltaR, pregMain);
564 Muls(deltaR, deltaR, tailCnt, pregMain);561 Muls(deltaR, deltaR, tailCnt, pregMain);
565 562 
566- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);563+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);
567 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);564 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
568 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);565 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
569- DataCopy(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);566+ LoadAlign(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);
570 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregMain);567 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregMain);
571 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregMain);568 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregMain);
572 569 
573 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);570 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
574- ReduceSum(var, sumVar, pregMain);571+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
575- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,572+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,
576- pregMerge);573+ pregMerge);
577 }574 }
578 sreg0 = dichotomyAddReminder - welfordDiffLoopCount * VL_FP32;575 sreg0 = dichotomyAddReminder - welfordDiffLoopCount * VL_FP32;
579 sreg1 = welfordDiffReminder;576 sreg1 = welfordDiffReminder;
580 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {577 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
581 pregLoop = UpdateMask<float>(sreg0);578 pregLoop = UpdateMask<float>(sreg0);
582 pregLoop1 = UpdateMask<float>(sreg1);579 pregLoop1 = UpdateMask<float>(sreg1);
583- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);580+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);
584 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);581 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
585 Mul(deltaL, deltaL, deltaL, pregMain);582 Mul(deltaL, deltaL, deltaL, pregMain);
586 Muls(deltaL, deltaL, tailCnt, pregMain);583 Muls(deltaL, deltaL, tailCnt, pregMain);
587- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);584+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
588 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);585 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
589 Mul(deltaR, deltaR, deltaR, pregLoop);586 Mul(deltaR, deltaR, deltaR, pregLoop);
590 Muls(deltaR, deltaR, cnt, pregLoop);587 Muls(deltaR, deltaR, cnt, pregLoop);
591 Muls(tmp, deltaR, coeff, pregLoop1);588 Muls(tmp, deltaR, coeff, pregLoop1);
592- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaR, tmp, pregLoop1);589+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaR, tmp, pregLoop1);
593 590 
594- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32);591+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32);
595 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);592 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
596 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);593 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
597- DataCopy(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);594+ LoadAlign(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
598 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);595 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
599 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);596 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
600 597 
601 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);598 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
602- ReduceSum(var, sumVar, pregMain);599+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
603- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(600+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
604 dichotomyAddLocal + i + welfordDiffLoopCount, var, pregMerge);601 dichotomyAddLocal + i + welfordDiffLoopCount, var, pregMerge);
605 }602 }
606 603 
607 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {604 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
608 pregLoop = UpdateMask<float>(sreg0);605 pregLoop = UpdateMask<float>(sreg0);
609- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);606+ LoadAlign(dichotomyAddMeanL,
607+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
610 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);608 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
611 Mul(deltaL, deltaL, deltaL, pregMain);609 Mul(deltaL, deltaL, deltaL, pregMain);
612 Muls(deltaL, deltaL, tailCnt, pregMain);610 Muls(deltaL, deltaL, tailCnt, pregMain);
613- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +611+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 +
614- dichotomyAddPower);612+ welfordDiffReminderAlign + dichotomyAddPower);
615 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);613 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
616 Mul(deltaR, deltaR, deltaR, pregLoop);614 Mul(deltaR, deltaR, deltaR, pregLoop);
617 Muls(deltaR, deltaR, cnt, pregLoop);615 Muls(deltaR, deltaR, cnt, pregLoop);
618 616 
619- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);617+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
620 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);618 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
621 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);619 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
622- DataCopy(dichotomyAddVarR,620+ LoadAlign(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +
623- tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign + dichotomyAddPower);621+ dichotomyAddPower);
624 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);622 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
625 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);623 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
626 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);624 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
627- ReduceSum(var, sumVar, pregMain);625+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
628- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(626+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
629 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, var, pregMerge);627 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, var, pregMerge);
630 }628 }
631 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {629 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
632- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);630+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
633 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);631 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
634 Mul(deltaL, deltaL, deltaL, pregMain);632 Mul(deltaL, deltaL, deltaL, pregMain);
635 Muls(deltaL, deltaL, tailCnt, pregMain);633 Muls(deltaL, deltaL, tailCnt, pregMain);
636- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);634+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
637 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);635 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
638 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);636 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
639- ReduceSum(var, dichotomyAddVarL, pregMain);637+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
640- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(638+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
641 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, var, pregMerge);639 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, var, pregMerge);
642 }640 }
643 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);641 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
644 CalRstdByHighPrecision(var, rstd, eps);642 CalRstdByHighPrecision(var, rstd, eps);
645- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);643+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);
646 }644 }
647}645}
648 646 
649// welford整块小于二分累加整块,并且小于等于二分累加尾块向上对齐647// welford整块小于二分累加整块,并且小于等于二分累加尾块向上对齐
650__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation2(648__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation2(
651- __local_mem__ float* meanLocal, __local_mem__ float* rstdLocal, __local_mem__ float* tmpMeanLocal,649+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal, __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
652- __local_mem__ float* tmpVarLocal, __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,650+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount, uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
653- uint32_t dichotomyAddPower, uint32_t dichotomyAddK, uint32_t dichotomyAddLastNum, uint32_t offset,651+ uint32_t dichotomyAddLastNum, uint32_t offset, uint32_t tailSize, float reduceScale, float cnt, float eps)
654- uint32_t tailSize, float reduceScale, float cnt, float eps)
655{652{
656 float tailCnt = cnt + float(1.0);653 float tailCnt = cnt + float(1.0);
657 float coeff = tailCnt / cnt;654 float coeff = tailCnt / cnt;
@@ -698,14 +695,14 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation2(
698 695 
699 // 整块使用tailCountScale,尾块使用countScale696 // 整块使用tailCountScale,尾块使用countScale
700 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {697 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
701- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);698+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
702- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);699+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
703 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);700 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
704 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregMain);701 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregMain);
705 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);702 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
706- ReduceSum(mean, sumMean, pregMain);703+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
707- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,704+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,
708- pregMerge);705+ pregMerge);
709 }706 }
710 707 
711 // 处理welford第一次非对齐点, 尾块使用countScale,整块部分使用tailCountScale, 部分使用countScale708 // 处理welford第一次非对齐点, 尾块使用countScale,整块部分使用tailCountScale, 部分使用countScale
@@ -714,145 +711,146 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation2(
714 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {711 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
715 pregLoop = UpdateMask<float>(sreg0);712 pregLoop = UpdateMask<float>(sreg0);
716 pregLoop1 = UpdateMask<float>(sreg1);713 pregLoop1 = UpdateMask<float>(sreg1);
717- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);714+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);
718- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);715+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
719 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);716 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);
720 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);717 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);
721 Muls(tmp, dichotomyAddMeanL, coeff, pregLoop1);718 Muls(tmp, dichotomyAddMeanL, coeff, pregLoop1);
722- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanL, tmp, pregLoop1);719+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanL, tmp, pregLoop1);
723 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);720 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
724- ReduceSum(mean, sumMean, pregMain);721+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
725- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(722+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
726 dichotomyAddLocal + i + welfordDiffLoopCount, mean, pregMerge);723 dichotomyAddLocal + i + welfordDiffLoopCount, mean, pregMerge);
727 }724 }
728 725 
729 // 整块使用countScale,尾块使用countScale726 // 整块使用countScale,尾块使用countScale
730 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {727 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
731 pregLoop = UpdateMask<float>(sreg0);728 pregLoop = UpdateMask<float>(sreg0);
732- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);729+ LoadAlign(dichotomyAddMeanL,
733- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +730+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
734- dichotomyAddPower);731+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 +
732+ welfordDiffReminderAlign + dichotomyAddPower);
735 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);733 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);
736 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);734 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);
737 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);735 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
738- ReduceSum(mean, sumMean, pregMain);736+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
739- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(737+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
740 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, mean, pregMerge);738 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, mean, pregMerge);
741 }739 }
742 // PART2: 整块剩余部分vcadd回刷UB,使用countScale740 // PART2: 整块剩余部分vcadd回刷UB,使用countScale
743 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {741 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
744- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);742+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
745 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);743 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);
746- ReduceSum(mean, dichotomyAddMeanL, pregMain);744+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
747- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(745+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
748 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, mean, pregMerge);746 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, mean, pregMerge);
749 }747 }
750 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);748 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
751- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);749+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);
752 750 
753 // 计算rstd751 // 计算rstd
754 Duplicate(one, float(1.0), pregMain);752 Duplicate(one, float(1.0), pregMain);
755 Duplicate(mean, mean, pregMain);753 Duplicate(mean, mean, pregMain);
756 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {754 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
757- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);755+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
758 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);756 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
759 Mul(deltaL, deltaL, deltaL, pregMain);757 Mul(deltaL, deltaL, deltaL, pregMain);
760 Muls(deltaL, deltaL, tailCnt, pregMain);758 Muls(deltaL, deltaL, tailCnt, pregMain);
761- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);759+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
762 Sub(deltaR, dichotomyAddMeanR, mean, pregMain);760 Sub(deltaR, dichotomyAddMeanR, mean, pregMain);
763 Mul(deltaR, deltaR, deltaR, pregMain);761 Mul(deltaR, deltaR, deltaR, pregMain);
764 Muls(deltaR, deltaR, cnt, pregMain);762 Muls(deltaR, deltaR, cnt, pregMain);
765 763 
766- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);764+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);
767 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);765 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
768 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);766 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
769- DataCopy(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);767+ LoadAlign(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);
770 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregMain);768 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregMain);
771 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregMain);769 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregMain);
772 770 
773 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);771 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
774- ReduceSum(var, sumVar, pregMain);772+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
775- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,773+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,
776- pregMerge);774+ pregMerge);
777 }775 }
778 sreg0 = dichotomyAddReminder - welfordDiffLoopCount * VL_FP32;776 sreg0 = dichotomyAddReminder - welfordDiffLoopCount * VL_FP32;
779 sreg1 = welfordDiffReminder;777 sreg1 = welfordDiffReminder;
780 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {778 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
781 pregLoop = UpdateMask<float>(sreg0);779 pregLoop = UpdateMask<float>(sreg0);
782 pregLoop1 = UpdateMask<float>(sreg1);780 pregLoop1 = UpdateMask<float>(sreg1);
783- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);781+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32);
784 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);782 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
785 Mul(deltaL, deltaL, deltaL, pregMain);783 Mul(deltaL, deltaL, deltaL, pregMain);
786 Muls(deltaL, deltaL, cnt, pregMain);784 Muls(deltaL, deltaL, cnt, pregMain);
787- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);785+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
788 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);786 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
789 Mul(deltaR, deltaR, deltaR, pregLoop);787 Mul(deltaR, deltaR, deltaR, pregLoop);
790 Muls(deltaR, deltaR, cnt, pregLoop);788 Muls(deltaR, deltaR, cnt, pregLoop);
791 Muls(tmp, deltaL, coeff, pregLoop1);789 Muls(tmp, deltaL, coeff, pregLoop1);
792- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaL, tmp, pregLoop1);790+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaL, tmp, pregLoop1);
793 791 
794- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32);792+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32);
795 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);793 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
796 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);794 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
797- DataCopy(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);795+ LoadAlign(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddPower);
798 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);796 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
799 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);797 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
800 798 
801 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);799 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
802- ReduceSum(var, sumVar, pregMain);800+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
803- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(801+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
804 dichotomyAddLocal + i + welfordDiffLoopCount, var, pregMerge);802 dichotomyAddLocal + i + welfordDiffLoopCount, var, pregMerge);
805 }803 }
806 804 
807 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {805 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
808 pregLoop = UpdateMask<float>(sreg0);806 pregLoop = UpdateMask<float>(sreg0);
809- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);807+ LoadAlign(dichotomyAddMeanL,
808+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
810 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);809 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
811 Mul(deltaL, deltaL, deltaL, pregMain);810 Mul(deltaL, deltaL, deltaL, pregMain);
812 Muls(deltaL, deltaL, cnt, pregMain);811 Muls(deltaL, deltaL, cnt, pregMain);
813- DataCopy(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +812+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 +
814- dichotomyAddPower);813+ welfordDiffReminderAlign + dichotomyAddPower);
815 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);814 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
816 Mul(deltaR, deltaR, deltaR, pregLoop);815 Mul(deltaR, deltaR, deltaR, pregLoop);
817 Muls(deltaR, deltaR, cnt, pregLoop);816 Muls(deltaR, deltaR, cnt, pregLoop);
818 817 
819- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);818+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign);
820 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);819 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
821 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);820 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
822- DataCopy(dichotomyAddVarR,821+ LoadAlign(dichotomyAddVarR, tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign +
823- tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + welfordDiffReminderAlign + dichotomyAddPower);822+ dichotomyAddPower);
824 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);823 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
825 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);824 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
826 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);825 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
827- ReduceSum(var, sumVar, pregMain);826+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
828- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(827+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
829 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, var, pregMerge);828 dichotomyAddLocal + i + welfordDiffLoopCount + welfordReminderLoopCount, var, pregMerge);
830 }829 }
831 830 
832 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {831 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
833- DataCopy(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);832+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
834 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);833 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
835 Mul(deltaL, deltaL, deltaL, pregMain);834 Mul(deltaL, deltaL, deltaL, pregMain);
836 Muls(deltaL, deltaL, cnt, pregMain);835 Muls(deltaL, deltaL, cnt, pregMain);
837- DataCopy(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);836+ LoadAlign(dichotomyAddVarL, tmpVarLocal + (i + dichotomyAddReminderRealLoopCount) * VL_FP32);
838 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);837 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
839 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);838 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
840- ReduceSum(var, dichotomyAddVarL, pregMain);839+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
841- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(840+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
842 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, var, pregMerge);841 dichotomyAddLocal + dichotomyAddReminderRealLoopCount + i, var, pregMerge);
843 }842 }
844 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);843 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
845 CalRstdByHighPrecision(var, rstd, eps);844 CalRstdByHighPrecision(var, rstd, eps);
846- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);845+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);
847 }846 }
848}847}
849 848 
850// 场景3:welford整块小于二分累加整块,并且大于二分累加尾块向上对齐849// 场景3:welford整块小于二分累加整块,并且大于二分累加尾块向上对齐
851__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation3(850__aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation3(
852- __local_mem__ float* meanLocal, __local_mem__ float* rstdLocal, __local_mem__ float* tmpMeanLocal,851+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal, __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
853- __local_mem__ float* tmpVarLocal, __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,852+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount, uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
854- uint32_t dichotomyAddPower, uint32_t dichotomyAddK, uint32_t dichotomyAddLastNum, uint32_t offset,853+ uint32_t dichotomyAddLastNum, uint32_t offset, uint32_t tailSize, float reduceScale, float cnt, float eps)
855- uint32_t tailSize, float reduceScale, float cnt, float eps)
856{854{
857 float tailCnt = cnt + float(1.0);855 float tailCnt = cnt + float(1.0);
858 float coeff = tailCnt / cnt;856 float coeff = tailCnt / cnt;
@@ -900,50 +898,50 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation3(
900 // 整块使用tailCountScale, 尾块使用CountScale898 // 整块使用tailCountScale, 尾块使用CountScale
901 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {899 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
902 pregLoop = UpdateMask<float>(sreg0);900 pregLoop = UpdateMask<float>(sreg0);
903- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);901+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
904- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);902+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
905 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);903 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
906 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);904 Muls(dichotomyAddMeanR, dichotomyAddMeanR, countScale, pregLoop);
907 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);905 Add(sumMean, dichotomyAddMeanL, dichotomyAddMeanR, pregMain);
908- ReduceSum(mean, sumMean, pregMain);906+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
909- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,907+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, mean,
910- pregMerge);908+ pregMerge);
911 }909 }
912 910 
913 // 剩余整块需要拆分成多部分911 // 剩余整块需要拆分成多部分
914 // 整块剩余部分回刷UB,整块使用tailCountScale912 // 整块剩余部分回刷UB,整块使用tailCountScale
915 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {913 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
916- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);914+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);
917 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);915 Muls(dichotomyAddMeanL, dichotomyAddMeanL, tailCountScale, pregMain);
918- ReduceSum(mean, dichotomyAddMeanL, pregMain);916+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
919- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(917+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
920 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, mean, pregMerge);918 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, mean, pregMerge);
921 }919 }
922 920 
923 sreg0 = welfordDiffReminder;921 sreg0 = welfordDiffReminder;
924 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {922 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
925 pregLoop = UpdateMask<float>(sreg0);923 pregLoop = UpdateMask<float>(sreg0);
926- DataCopy(dichotomyAddMeanL,924+ LoadAlign(dichotomyAddMeanL,
927- tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);925+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);
928 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);926 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);
929 Muls(tmp, dichotomyAddMeanL, coeff, pregLoop);927 Muls(tmp, dichotomyAddMeanL, coeff, pregLoop);
930- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanL, tmp, pregLoop);928+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(dichotomyAddMeanL, tmp, pregLoop);
931- ReduceSum(mean, dichotomyAddMeanL, pregMain);929+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
932- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(930+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
933 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + i, mean, pregMerge);931 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + i, mean, pregMerge);
934 }932 }
935 933 
936 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {934 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
937- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddPowerOffset);935+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddPowerOffset);
938 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);936 Muls(dichotomyAddMeanL, dichotomyAddMeanL, countScale, pregMain);
939- ReduceSum(mean, dichotomyAddMeanL, pregMain);937+ Reduce<ReduceType::SUM>(mean, dichotomyAddMeanL, pregMain);
940- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(938+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
941 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + welfordReminderLoopCount + i,939 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + welfordReminderLoopCount + i,
942 mean, pregMerge);940 mean, pregMerge);
943 }941 }
944 942 
945 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);943 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
946- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);944+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + offset, mean, pregMerge);
947 945 
948 // 计算rstd946 // 计算rstd
949 Duplicate(one, float(1.0), pregMain);947 Duplicate(one, float(1.0), pregMain);
@@ -952,85 +950,84 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlignSituation3(
952 sreg0 = dichotomyAddReminder;950 sreg0 = dichotomyAddReminder;
953 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {951 for (uint16_t i = 0; i < dichotomyAddReminderLoopCount; i++) {
954 pregLoop = UpdateMask<float>(sreg0);952 pregLoop = UpdateMask<float>(sreg0);
955- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);953+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32);
956 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);954 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
957 Mul(deltaL, deltaL, deltaL, pregMain);955 Mul(deltaL, deltaL, deltaL, pregMain);
958 Muls(deltaL, deltaL, tailCnt, pregMain);956 Muls(deltaL, deltaL, tailCnt, pregMain);
959- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);957+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32);
960 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);958 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
961 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);959 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
962 960 
963- DataCopy(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);961+ LoadAlign(dichotomyAddMeanR, tmpMeanLocal + i * VL_FP32 + dichotomyAddPower);
964 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);962 Sub(deltaR, dichotomyAddMeanR, mean, pregLoop);
965 Mul(deltaR, deltaR, deltaR, pregLoop);963 Mul(deltaR, deltaR, deltaR, pregLoop);
966 Muls(deltaR, deltaR, cnt, pregLoop);964 Muls(deltaR, deltaR, cnt, pregLoop);
967- DataCopy(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);965+ LoadAlign(dichotomyAddVarR, tmpVarLocal + i * VL_FP32 + dichotomyAddPower);
968 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);966 Add(dichotomyAddVarR, dichotomyAddVarR, deltaR, pregLoop);
969 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);967 Muls(dichotomyAddVarR, dichotomyAddVarR, reduceScale, pregLoop);
970 968 
971 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);969 Add(sumVar, dichotomyAddVarL, dichotomyAddVarR, pregMain);
972- ReduceSum(var, sumVar, pregMain);970+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
973- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,971+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + i, var,
974- pregMerge);972+ pregMerge);
975 }973 }
976 974 
977 // 整块剩余部分回刷UB,整块使用tailCountScale975 // 整块剩余部分回刷UB,整块使用tailCountScale
978 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {976 for (uint16_t i = 0; i < welfordDiffLoopCount; i++) {
979- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);977+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);
980 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);978 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
981 Mul(deltaL, deltaL, deltaL, pregMain);979 Mul(deltaL, deltaL, deltaL, pregMain);
982 Muls(deltaL, deltaL, tailCnt, pregMain);980 Muls(deltaL, deltaL, tailCnt, pregMain);
983- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);981+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32 + dichotomyAddReminderRoundUp);
984 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);982 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
985 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);983 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
986- ReduceSum(var, dichotomyAddVarL, pregMain);984+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
987- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(985+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
988 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, var, pregMerge);986 dichotomyAddLocal + dichotomyAddReminderLoopCount + i, var, pregMerge);
989 }987 }
990 988 
991 sreg0 = welfordDiffReminder;989 sreg0 = welfordDiffReminder;
992 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {990 for (uint16_t i = 0; i < welfordReminderLoopCount; i++) {
993 pregLoop = UpdateMask<float>(sreg0);991 pregLoop = UpdateMask<float>(sreg0);
994- DataCopy(dichotomyAddMeanL,992+ LoadAlign(dichotomyAddMeanL,
995- tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);993+ tmpMeanLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);
996 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);994 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
997 Mul(deltaL, deltaL, deltaL, pregMain);995 Mul(deltaL, deltaL, deltaL, pregMain);
998 Muls(deltaL, deltaL, cnt, pregMain);996 Muls(deltaL, deltaL, cnt, pregMain);
999 Muls(tmp, deltaL, coeff, pregLoop);997 Muls(tmp, deltaL, coeff, pregLoop);
1000- Copy<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaL, tmp, pregLoop);998+ Move<float, AscendC::MicroAPI::MaskMergeMode::MERGING>(deltaL, tmp, pregLoop);
1001- DataCopy(dichotomyAddVarL,999+ LoadAlign(dichotomyAddVarL,
1002- tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);1000+ tmpVarLocal + (i + welfordDiffLoopCount) * VL_FP32 + dichotomyAddReminderRoundUp);
1003 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);1001 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
1004 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);1002 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
1005- ReduceSum(var, dichotomyAddVarL, pregMain);1003+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
1006- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(1004+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
1007 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + i, var, pregMerge);1005 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + i, var, pregMerge);
1008 }1006 }
1009 1007 
1010 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {1008 for (uint16_t i = 0; i < dichotomyAddPowerRemainLoopCount; i++) {
1011- DataCopy(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddPowerOffset);1009+ LoadAlign(dichotomyAddMeanL, tmpMeanLocal + i * VL_FP32 + dichotomyAddPowerOffset);
1012 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);1010 Sub(deltaL, dichotomyAddMeanL, mean, pregMain);
1013 Mul(deltaL, deltaL, deltaL, pregMain);1011 Mul(deltaL, deltaL, deltaL, pregMain);
1014 Muls(deltaL, deltaL, cnt, pregMain);1012 Muls(deltaL, deltaL, cnt, pregMain);
1015- DataCopy(dichotomyAddVarL, tmpVarLocal + i * VL_FP32 + dichotomyAddPowerOffset);1013+ LoadAlign(dichotomyAddVarL, tmpVarLocal + i * VL_FP32 + dichotomyAddPowerOffset);
1016 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);1014 Add(dichotomyAddVarL, dichotomyAddVarL, deltaL, pregMain);
1017 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);1015 Muls(dichotomyAddVarL, dichotomyAddVarL, reduceScale, pregMain);
1018- ReduceSum(var, dichotomyAddVarL, pregMain);1016+ Reduce<ReduceType::SUM>(var, dichotomyAddVarL, pregMain);
1019- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(1017+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
1020 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + welfordReminderLoopCount + i,1018 dichotomyAddLocal + dichotomyAddReminderLoopCount + welfordDiffLoopCount + welfordReminderLoopCount + i,
1021 var, pregMerge);1019 var, pregMerge);
1022 }1020 }
1023 1021 
1024 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);1022 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
1025 CalRstdByHighPrecision(var, rstd, eps);1023 CalRstdByHighPrecision(var, rstd, eps);
1026- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);1024+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + offset, rstd, pregMerge);
1027 }1025 }
1028}1026}
1029 1027 
1030-__aicore__ inline void VFWelfordParallelFinalizeNonAlign(__local_mem__ float* meanLocal, __local_mem__ float* rstdLocal,1028+__aicore__ inline void VFWelfordParallelFinalizeNonAlign(__ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1031- __local_mem__ float* tmpMeanLocal,1029+ __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
1032- __local_mem__ float* tmpVarLocal,1030+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount,
1033- __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,
1034 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,1031 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
1035 uint32_t dichotomyAddLastNum, uint32_t offset,1032 uint32_t dichotomyAddLastNum, uint32_t offset,
1036 uint32_t tailSize, float reduceScale, float cnt, float eps)1033 uint32_t tailSize, float reduceScale, float cnt, float eps)
@@ -1055,9 +1052,9 @@ __aicore__ inline void VFWelfordParallelFinalizeNonAlign(__local_mem__ float* me
1055 offset, tailSize, reduceScale, cnt, eps);1052 offset, tailSize, reduceScale, cnt, eps);
1056}1053}
1057 1054 
1058-__aicore__ inline void VFWelfordParallelFinalize(__local_mem__ float* meanLocal, __local_mem__ float* rstdLocal,1055+__aicore__ inline void VFWelfordParallelFinalize(__ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1059- __local_mem__ float* tmpMeanLocal, __local_mem__ float* tmpVarLocal,1056+ __ubuf__ float* tmpMeanLocal, __ubuf__ float* tmpVarLocal,
1060- __local_mem__ float* dichotomyAddLocal, uint32_t reduceCount,1057+ __ubuf__ float* dichotomyAddLocal, uint32_t reduceCount,
1061 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,1058 uint32_t dichotomyAddPower, uint32_t dichotomyAddK,
1062 uint32_t dichotomyAddLastNum, uint32_t offset, uint32_t tailSize,1059 uint32_t dichotomyAddLastNum, uint32_t offset, uint32_t tailSize,
1063 float reduceScale, float scale, float cnt, float eps,1060 float reduceScale, float scale, float cnt, float eps,
@@ -1076,12 +1073,11 @@ __aicore__ inline void VFWelfordParallelFinalize(__local_mem__ float* meanLocal,
1076}1073}
1077 1074 
1078template <typename T>1075template <typename T>
1079-__aicore__ inline void CalMeanAndRstdByDichotomyAdd(__local_mem__ T* xLocal, __local_mem__ float* meanLocal,1076+__aicore__ inline void CalMeanAndRstdByDichotomyAdd(__ubuf__ T* xLocal, __ubuf__ float* meanLocal,
1080- __local_mem__ float* rstdLocal,1077+ __ubuf__ float* rstdLocal, __ubuf__ float* dichotomyAddLocal,
1081- __local_mem__ float* dichotomyAddLocal, uint16_t numPerCoreProcess,1078+ uint16_t numPerCoreProcess, uint32_t dichotomyAddPower,
1082- uint32_t dichotomyAddPower, uint32_t dichotomyAddK,1079+ uint32_t dichotomyAddK, uint32_t dichotomyAddLastNum,
1083- uint32_t dichotomyAddLastNum, uint64_t powerOfTwoForReduce,1080+ uint64_t powerOfTwoForReduce, uint64_t reduceCount, float eps)
1084- uint64_t reduceCount, float eps)
1085{1081{
1086 uint32_t dichotomyAddReminder = reduceCount - dichotomyAddPower;1082 uint32_t dichotomyAddReminder = reduceCount - dichotomyAddPower;
1087 uint16_t dichotomyAddReminderLoopCount = CeilDiv(dichotomyAddReminder, VL_FP32);1083 uint16_t dichotomyAddReminderLoopCount = CeilDiv(dichotomyAddReminder, VL_FP32);
@@ -1116,9 +1112,9 @@ __aicore__ inline void CalMeanAndRstdByDichotomyAdd(__local_mem__ T* xLocal, __l
1116 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);1112 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);
1117 Muls(dichotomyAddR, dichotomyAddR, n, pregLoop);1113 Muls(dichotomyAddR, dichotomyAddR, n, pregLoop);
1118 Add(sumMean, dichotomyAddL, dichotomyAddR, pregMain);1114 Add(sumMean, dichotomyAddL, dichotomyAddR, pregMain);
1119- ReduceSum(mean, sumMean, pregMain);1115+ Reduce<ReduceType::SUM>(mean, sumMean, pregMain);
1120- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + j, mean,1116+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + j, mean,
1121- pregMerge);1117+ pregMerge);
1122 }1118 }
1123 1119 
1124 // 整块剩余部分vcadd回刷UB1120 // 整块剩余部分vcadd回刷UB
@@ -1126,14 +1122,14 @@ __aicore__ inline void CalMeanAndRstdByDichotomyAdd(__local_mem__ T* xLocal, __l
1126 LoadInputData<T>(dichotomyAddL, xLocal, pregMain,1122 LoadInputData<T>(dichotomyAddL, xLocal, pregMain,
1127 i * elemNumAlign + (j + dichotomyAddReminderLoopCount) * VL_FP32);1123 i * elemNumAlign + (j + dichotomyAddReminderLoopCount) * VL_FP32);
1128 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);1124 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);
1129- ReduceSum(mean, dichotomyAddL, pregMain);1125+ Reduce<ReduceType::SUM>(mean, dichotomyAddL, pregMain);
1130- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(1126+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
1131 dichotomyAddLocal + dichotomyAddReminderLoopCount + j, mean, pregMerge);1127 dichotomyAddLocal + dichotomyAddReminderLoopCount + j, mean, pregMerge);
1132 }1128 }
1133 1129 
1134 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);1130 DichotomyAdd(mean, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
1135 Muls(mean, mean, nCorrectionFactor, pregMerge);1131 Muls(mean, mean, nCorrectionFactor, pregMerge);
1136- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + i, mean, pregMerge);1132+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + i, mean, pregMerge);
1137 // 计算rstd1133 // 计算rstd
1138 Duplicate(one, float(1.0), pregMain);1134 Duplicate(one, float(1.0), pregMain);
1139 Duplicate(mean, mean, pregMain);1135 Duplicate(mean, mean, pregMain);
@@ -1149,9 +1145,9 @@ __aicore__ inline void CalMeanAndRstdByDichotomyAdd(__local_mem__ T* xLocal, __l
1149 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);1145 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);
1150 Muls(dichotomyAddR, dichotomyAddR, n, pregLoop);1146 Muls(dichotomyAddR, dichotomyAddR, n, pregLoop);
1151 Add(sumVar, dichotomyAddL, dichotomyAddR, pregMain);1147 Add(sumVar, dichotomyAddL, dichotomyAddR, pregMain);
1152- ReduceSum(var, sumVar, pregMain);1148+ Reduce<ReduceType::SUM>(var, sumVar, pregMain);
1153- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + j, var,1149+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(dichotomyAddLocal + j, var,
1154- pregMerge);1150+ pregMerge);
1155 }1151 }
1156 1152 
1157 // 整块剩余部分vcadd回刷UB1153 // 整块剩余部分vcadd回刷UB
@@ -1161,23 +1157,23 @@ __aicore__ inline void CalMeanAndRstdByDichotomyAdd(__local_mem__ T* xLocal, __l
1161 Sub(dichotomyAddL, dichotomyAddL, mean, pregMain);1157 Sub(dichotomyAddL, dichotomyAddL, mean, pregMain);
1162 Mul(dichotomyAddL, dichotomyAddL, dichotomyAddL, pregMain);1158 Mul(dichotomyAddL, dichotomyAddL, dichotomyAddL, pregMain);
1163 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);1159 Muls(dichotomyAddL, dichotomyAddL, n, pregMain);
1164- ReduceSum(var, dichotomyAddL, pregMain);1160+ Reduce<ReduceType::SUM>(var, dichotomyAddL, pregMain);
1165- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(1161+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(
1166 dichotomyAddLocal + dichotomyAddReminderLoopCount + j, var, pregMerge);1162 dichotomyAddLocal + dichotomyAddReminderLoopCount + j, var, pregMerge);
1167 }1163 }
1168 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);1164 DichotomyAdd(var, dichotomyAddLocal, dichotomyAddK, innerLoopCountOrigin, dichotomyAddLastNum);
1169 Muls(var, var, nCorrectionFactor, pregMerge);1165 Muls(var, var, nCorrectionFactor, pregMerge);
1170 CalRstdByHighPrecision(var, rstd, eps);1166 CalRstdByHighPrecision(var, rstd, eps);
1171- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + i, rstd, pregMerge);1167+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + i, rstd, pregMerge);
1172 }1168 }
1173 }1169 }
1174}1170}
1175 1171 
1176// R轴小于641172// R轴小于64
1177template <typename T>1173template <typename T>
1178-__aicore__ inline void CalMeanAndRstdSpecial(__local_mem__ T* xLocal, __local_mem__ float* meanLocal,1174+__aicore__ inline void CalMeanAndRstdSpecial(__ubuf__ T* xLocal, __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1179- __local_mem__ float* rstdLocal, uint16_t numPerCoreProcess,1175+ uint16_t numPerCoreProcess, uint64_t powerOfTwoForReduce,
1180- uint64_t powerOfTwoForReduce, uint64_t reduceCount, float eps)1176+ uint64_t reduceCount, float eps)
1181{1177{
1182 uint32_t elemNumAlign = RoundUp<T>(reduceCount);1178 uint32_t elemNumAlign = RoundUp<T>(reduceCount);
1183 float n = static_cast<float>(1) / static_cast<float>(powerOfTwoForReduce);1179 float n = static_cast<float>(1) / static_cast<float>(powerOfTwoForReduce);
@@ -1199,28 +1195,27 @@ __aicore__ inline void CalMeanAndRstdSpecial(__local_mem__ T* xLocal, __local_me
1199 pregLoop = UpdateMask<float>(sreg0);1195 pregLoop = UpdateMask<float>(sreg0);
1200 LoadInputData<T>(x, xLocal, pregLoop, i * elemNumAlign);1196 LoadInputData<T>(x, xLocal, pregLoop, i * elemNumAlign);
1201 Muls(xScale, x, n, pregLoop);1197 Muls(xScale, x, n, pregLoop);
1202- ReduceSum(mean, xScale, pregLoop);1198+ Reduce<ReduceType::SUM>(mean, xScale, pregLoop);
1203 Muls(mean, mean, nCorrectionFactor, pregMerge);1199 Muls(mean, mean, nCorrectionFactor, pregMerge);
1204- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + i, mean, pregMerge);1200+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(meanLocal + i, mean, pregMerge);
1205 1201 
1206 Duplicate(mean, mean, pregMain);1202 Duplicate(mean, mean, pregMain);
1207 Sub(x, x, mean, pregLoop);1203 Sub(x, x, mean, pregLoop);
1208 Mul(x, x, x, pregLoop);1204 Mul(x, x, x, pregLoop);
1209 Muls(xScale, x, n, pregLoop);1205 Muls(xScale, x, n, pregLoop);
1210- ReduceSum(var, xScale, pregLoop);1206+ Reduce<ReduceType::SUM>(var, xScale, pregLoop);
1211 Muls(var, var, nCorrectionFactor, pregMerge);1207 Muls(var, var, nCorrectionFactor, pregMerge);
1212 CalRstdByHighPrecision(var, rstd, eps);1208 CalRstdByHighPrecision(var, rstd, eps);
1213- DataCopy<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + i, rstd, pregMerge);1209+ StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_FIRST_ELEMENT_B32>(rstdLocal + i, rstd, pregMerge);
1214 }1210 }
1215 }1211 }
1216}1212}
1217 1213 
1218template <typename T>1214template <typename T>
1219-__aicore__ inline void CalMeanAndRstd(__local_mem__ T* xLocal, __local_mem__ float* meanLocal,1215+__aicore__ inline void CalMeanAndRstd(__ubuf__ T* xLocal, __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1220- __local_mem__ float* rstdLocal, __local_mem__ float* dichotomyAddLocal,1216+ __ubuf__ float* dichotomyAddLocal, uint16_t numPerCoreProcess,
1221- uint16_t numPerCoreProcess, uint32_t dichotomyAddPower, uint32_t dichotomyAddK,1217+ uint32_t dichotomyAddPower, uint32_t dichotomyAddK, uint32_t dichotomyAddLastNum,
1222- uint32_t dichotomyAddLastNum, uint64_t powerOfTwoForReduce, uint64_t reduceCount,1218+ uint64_t powerOfTwoForReduce, uint64_t reduceCount, float eps)
1223- float eps)
1224{1219{
1225 if (dichotomyAddPower >= VL_FP32) {1220 if (dichotomyAddPower >= VL_FP32) {
1226 CalMeanAndRstdByDichotomyAdd(xLocal, meanLocal, rstdLocal, dichotomyAddLocal, numPerCoreProcess,1221 CalMeanAndRstdByDichotomyAdd(xLocal, meanLocal, rstdLocal, dichotomyAddLocal, numPerCoreProcess,
@@ -1267,9 +1262,9 @@ __aicore__ inline void VFInnerNormalize(RegTensor<float>& x, RegTensor<float>& m
1267}1262}
1268 1263 
1269template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>1264template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>
1270-__aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1265+__aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal,
1271- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1266+ __ubuf__ T2* betaLocal, __ubuf__ float* meanLocal,
1272- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1267+ __ubuf__ float* rstdLocal, __ubuf__ T1* yLocal,
1273 uint16_t rowsCount, int32_t reduceCount)1268 uint16_t rowsCount, int32_t reduceCount)
1274{1269{
1275 uint16_t VL = GetVLSize<T1>();1270 uint16_t VL = GetVLSize<T1>();
@@ -1291,11 +1286,11 @@ __aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal,
1291 MaskReg pregLoop;1286 MaskReg pregLoop;
1292 MaskReg pregMain = CreateMask<T1, AscendC::MicroAPI::MaskPattern::ALL>();1287 MaskReg pregMain = CreateMask<T1, AscendC::MicroAPI::MaskPattern::ALL>();
1293 1288 
1294- UnalignReg uSrc;1289+ UnalignRegForLoad uSrc;
1295- UnalignReg uDst;1290+ UnalignRegForStore uDst;
1296- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal);1291+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal);
1297- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal);1292+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal);
1298- DataCopyUnAlignPre<T1>(uSrc, xLocal);1293+ LoadUnAlignPre<T1>(uSrc, xLocal);
1299 for (uint16_t i = 0; i < rowsCount; i++) {1294 for (uint16_t i = 0; i < rowsCount; i++) {
1300 LoadGammaAndBetaData<T2, hasGamma, hasBeta>(gamma, beta, gammaLocal, betaLocal, pregMain, i);1295 LoadGammaAndBetaData<T2, hasGamma, hasBeta>(gamma, beta, gammaLocal, betaLocal, pregMain, i);
1301 if constexpr (IsSameType<T1, half>::value || IsSameType<T1, bfloat16_t>::value) {1296 if constexpr (IsSameType<T1, half>::value || IsSameType<T1, bfloat16_t>::value) {
@@ -1304,7 +1299,7 @@ __aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal,
1304 RegTensor<T1> yOddTmp;1299 RegTensor<T1> yOddTmp;
1305 RegTensor<T1> yTmp;1300 RegTensor<T1> yTmp;
1306 for (uint16_t j = 0; j < loopCount; j++) {1301 for (uint16_t j = 0; j < loopCount; j++) {
1307- DataCopyUnAlign(xTmp, uSrc, xLocal, VL);1302+ LoadUnAlign(xTmp, uSrc, xLocal, VL);
1308 Cast<float, T1, castTraitB162B32Even>(xEven, xTmp, pregMain);1303 Cast<float, T1, castTraitB162B32Even>(xEven, xTmp, pregMain);
1309 Cast<float, T1, castTraitB162B32Odd>(xOdd, xTmp, pregMain);1304 Cast<float, T1, castTraitB162B32Odd>(xOdd, xTmp, pregMain);
1310 if constexpr (activateSilu) {1305 if constexpr (activateSilu) {
@@ -1318,12 +1313,12 @@ __aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal,
1318 Cast<T1, float, castTraitB322B16Odd>(yOddTmp, yOdd, pregMain);1313 Cast<T1, float, castTraitB322B16Odd>(yOddTmp, yOdd, pregMain);
1319 Or((RegTensor<int16_t>&)yTmp, (RegTensor<int16_t>&)yEvenTmp, (RegTensor<int16_t>&)yOddTmp,1314 Or((RegTensor<int16_t>&)yTmp, (RegTensor<int16_t>&)yEvenTmp, (RegTensor<int16_t>&)yOddTmp,
1320 pregMain);1315 pregMain);
1321- DataCopyUnAlign(yLocal, yTmp, uDst, VL);1316+ StoreUnAlign(yLocal, yTmp, uDst, VL);
1322 }1317 }
1323 uint32_t sreg0 = tailNum;1318 uint32_t sreg0 = tailNum;
1324 for (uint16_t k = 0; k < tailLoop; k++) {1319 for (uint16_t k = 0; k < tailLoop; k++) {
1325 pregLoop = UpdateMask<half>(sreg0);1320 pregLoop = UpdateMask<half>(sreg0);
1326- DataCopyUnAlign(xTmp, uSrc, xLocal, tailNum);1321+ LoadUnAlign(xTmp, uSrc, xLocal, tailNum);
1327 Cast<float, T1, castTraitB162B32Even>(xEven, xTmp, pregLoop);1322 Cast<float, T1, castTraitB162B32Even>(xEven, xTmp, pregLoop);
1328 Cast<float, T1, castTraitB162B32Odd>(xOdd, xTmp, pregLoop);1323 Cast<float, T1, castTraitB162B32Odd>(xOdd, xTmp, pregLoop);
1329 if constexpr (activateSilu) {1324 if constexpr (activateSilu) {
@@ -1337,41 +1332,41 @@ __aicore__ inline void VFInnerNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal,
1337 Cast<T1, float, castTraitB322B16Odd>(yOddTmp, yOdd, pregLoop);1332 Cast<T1, float, castTraitB322B16Odd>(yOddTmp, yOdd, pregLoop);
1338 Or((RegTensor<int16_t>&)yTmp, (RegTensor<int16_t>&)yEvenTmp, (RegTensor<int16_t>&)yOddTmp,1333 Or((RegTensor<int16_t>&)yTmp, (RegTensor<int16_t>&)yEvenTmp, (RegTensor<int16_t>&)yOddTmp,
1339 pregLoop);1334 pregLoop);
1340- DataCopyUnAlign(yLocal, yTmp, uDst, tailNum);1335+ StoreUnAlign(yLocal, yTmp, uDst, tailNum);
1341 }1336 }
1342- DataCopyUnAlignPost(yLocal, uDst, 0);1337+ StoreUnAlignPost(yLocal, uDst, 0);
1343 } else {1338 } else {
1344 for (uint16_t j = 0; j < loopCount; j++) {1339 for (uint16_t j = 0; j < loopCount; j++) {
1345- DataCopyUnAlign(x, uSrc, xLocal, VL_FP32);1340+ LoadUnAlign(x, uSrc, xLocal, VL_FP32);
1346 if constexpr (activateSilu) {1341 if constexpr (activateSilu) {
1347 VFInnerNormalizeAndSwish<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregMain);1342 VFInnerNormalizeAndSwish<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregMain);
1348 } else {1343 } else {
1349 VFInnerNormalize<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregMain);1344 VFInnerNormalize<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregMain);
1350 }1345 }
1351- DataCopyUnAlign(yLocal, y, uDst, VL_FP32);1346+ StoreUnAlign(yLocal, y, uDst, VL_FP32);
1352 }1347 }
1353 uint32_t sreg0 = tailNum;1348 uint32_t sreg0 = tailNum;
1354 for (uint16_t k = 0; k < tailLoop; k++) {1349 for (uint16_t k = 0; k < tailLoop; k++) {
1355 pregLoop = UpdateMask<float>(sreg0);1350 pregLoop = UpdateMask<float>(sreg0);
1356- DataCopyUnAlign(x, uSrc, xLocal, tailNum);1351+ LoadUnAlign(x, uSrc, xLocal, tailNum);
1357 if constexpr (activateSilu) {1352 if constexpr (activateSilu) {
1358 VFInnerNormalizeAndSwish<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregLoop);1353 VFInnerNormalizeAndSwish<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregLoop);
1359 } else {1354 } else {
1360 VFInnerNormalize<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregLoop);1355 VFInnerNormalize<hasGamma, hasBeta>(x, mean, rstd, gamma, beta, y, pregLoop);
1361 }1356 }
1362- DataCopyUnAlign(yLocal, y, uDst, tailNum);1357+ StoreUnAlign(yLocal, y, uDst, tailNum);
1363 }1358 }
1364- DataCopyUnAlignPost(yLocal, uDst, 0);1359+ StoreUnAlignPost(yLocal, uDst, 0);
1365 }1360 }
1366 }1361 }
1367 }1362 }
1368}1363}
1369 1364 
1370template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>1365template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>
1371-__aicore__ inline void VFInnerNormalizeAndSwishAlign(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1366+__aicore__ inline void VFInnerNormalizeAndSwishAlign(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal,
1372- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1367+ __ubuf__ T2* betaLocal, __ubuf__ float* meanLocal,
1373- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1368+ __ubuf__ float* rstdLocal, __ubuf__ T1* yLocal, uint16_t rowsCount,
1374- uint16_t rowsCount, int32_t reduceCount)1369+ int32_t reduceCount)
1375{1370{
1376 uint16_t loopCount = CeilDiv(reduceCount, VL_FP32);1371 uint16_t loopCount = CeilDiv(reduceCount, VL_FP32);
1377 uint32_t reduceCountAlign = RoundUp<T1>(reduceCount);1372 uint32_t reduceCountAlign = RoundUp<T1>(reduceCount);
@@ -1385,8 +1380,8 @@ __aicore__ inline void VFInnerNormalizeAndSwishAlign(__local_mem__ T1* xLocal, _
1385 RegTensor<float> y;1380 RegTensor<float> y;
1386 MaskReg pregLoop;1381 MaskReg pregLoop;
1387 MaskReg pregMain = CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();1382 MaskReg pregMain = CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
1388- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal);1383+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal);
1389- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal);1384+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal);
1390 for (uint16_t i = 0; i < rowsCount; i++) {1385 for (uint16_t i = 0; i < rowsCount; i++) {
1391 uint32_t sreg0 = reduceCount;1386 uint32_t sreg0 = reduceCount;
1392 LoadGammaAndBetaData<T2, hasGamma, hasBeta>(gamma, beta, gammaLocal, betaLocal, pregMain, i);1387 LoadGammaAndBetaData<T2, hasGamma, hasBeta>(gamma, beta, gammaLocal, betaLocal, pregMain, i);
@@ -1405,10 +1400,10 @@ __aicore__ inline void VFInnerNormalizeAndSwishAlign(__local_mem__ T1* xLocal, _
1405}1400}
1406 1401 
1407template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>1402template <typename T1, typename T2, bool activateSilu, bool hasGamma, bool hasBeta>
1408-__aicore__ inline void VFInnerNormalizeAndSwishFold(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1403+__aicore__ inline void VFInnerNormalizeAndSwishFold(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal,
1409- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1404+ __ubuf__ T2* betaLocal, __ubuf__ float* meanLocal,
1410- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1405+ __ubuf__ float* rstdLocal, __ubuf__ T1* yLocal, uint16_t groupNums,
1411- uint16_t groupNums, uint16_t rowsCount, int32_t reduceCount)1406+ uint16_t rowsCount, int32_t reduceCount)
1412{1407{
1413 uint16_t loopCount = CeilDiv(reduceCount, VL_FP32);1408 uint16_t loopCount = CeilDiv(reduceCount, VL_FP32);
1414 uint32_t reduceCountAlign = RoundUp<T1>(reduceCount);1409 uint32_t reduceCountAlign = RoundUp<T1>(reduceCount);
@@ -1423,8 +1418,8 @@ __aicore__ inline void VFInnerNormalizeAndSwishFold(__local_mem__ T1* xLocal, __
1423 MaskReg pregLoop;1418 MaskReg pregLoop;
1424 for (uint16_t i = 0; i < groupNums; i++) {1419 for (uint16_t i = 0; i < groupNums; i++) {
1425 for (uint16_t j = 0; j < rowsCount; j++) {1420 for (uint16_t j = 0; j < rowsCount; j++) {
1426- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal + i * rowsCount + j);1421+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(rstd, rstdLocal + i * rowsCount + j);
1427- DataCopy<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal + i * rowsCount + j);1422+ LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(mean, meanLocal + i * rowsCount + j);
1428 uint32_t sreg0 = reduceCount;1423 uint32_t sreg0 = reduceCount;
1429 for (uint16_t k = 0; k < loopCount; k++) {1424 for (uint16_t k = 0; k < loopCount; k++) {
1430 pregLoop = UpdateMask<float>(sreg0);1425 pregLoop = UpdateMask<float>(sreg0);
@@ -1450,11 +1445,10 @@ __aicore__ inline void VFInnerNormalizeAndSwishFold(__local_mem__ T1* xLocal, __
1450}1445}
1451 1446 
1452template <typename T1, typename T2>1447template <typename T1, typename T2>
1453-__aicore__ inline void VFNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1448+__aicore__ inline void VFNormalizeAndSwishUnAlign(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal, __ubuf__ T2* betaLocal,
1454- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1449+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1455- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1450+ __ubuf__ T1* yLocal, uint16_t rowsCount, int32_t reduceCount,
1456- uint16_t rowsCount, int32_t reduceCount, bool activateSilu,1451+ bool activateSilu, bool hasGamma, bool hasBeta)
1457- bool hasGamma, bool hasBeta)
1458{1452{
1459 if (activateSilu) {1453 if (activateSilu) {
1460 if (hasGamma && hasBeta) {1454 if (hasGamma && hasBeta) {
@@ -1488,11 +1482,10 @@ __aicore__ inline void VFNormalizeAndSwishUnAlign(__local_mem__ T1* xLocal, __lo
1488}1482}
1489 1483 
1490template <typename T1, typename T2>1484template <typename T1, typename T2>
1491-__aicore__ inline void VFNormalizeAndSwishAlign(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1485+__aicore__ inline void VFNormalizeAndSwishAlign(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal, __ubuf__ T2* betaLocal,
1492- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1486+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1493- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1487+ __ubuf__ T1* yLocal, uint16_t rowsCount, int32_t reduceCount,
1494- uint16_t rowsCount, int32_t reduceCount, bool activateSilu,1488+ bool activateSilu, bool hasGamma, bool hasBeta)
1495- bool hasGamma, bool hasBeta)
1496{1489{
1497 if (activateSilu) {1490 if (activateSilu) {
1498 if (hasGamma && hasBeta) {1491 if (hasGamma && hasBeta) {
@@ -1526,11 +1519,10 @@ __aicore__ inline void VFNormalizeAndSwishAlign(__local_mem__ T1* xLocal, __loca
1526}1519}
1527 1520 
1528template <typename T1, typename T2>1521template <typename T1, typename T2>
1529-__aicore__ inline void VFNormalizeAndSwishFold(__local_mem__ T1* xLocal, __local_mem__ T2* gammaLocal,1522+__aicore__ inline void VFNormalizeAndSwishFold(__ubuf__ T1* xLocal, __ubuf__ T2* gammaLocal, __ubuf__ T2* betaLocal,
1530- __local_mem__ T2* betaLocal, __local_mem__ float* meanLocal,1523+ __ubuf__ float* meanLocal, __ubuf__ float* rstdLocal,
1531- __local_mem__ float* rstdLocal, __local_mem__ T1* yLocal,1524+ __ubuf__ T1* yLocal, uint16_t groupNums, uint16_t rowsCount,
1532- uint16_t groupNums, uint16_t rowsCount, int32_t reduceCount,1525+ int32_t reduceCount, bool activateSilu, bool hasGamma, bool hasBeta)
1533- bool activateSilu, bool hasGamma, bool hasBeta)
1534{1526{
1535 if (activateSilu) {1527 if (activateSilu) {
1536 if (hasGamma && hasBeta) {1528 if (hasGamma && hasBeta) {
@@ -1596,7 +1588,7 @@ __aicore__ inline void CopyGammaAndBeta2UBByNDDMA(const GlobalTensor<T>& gammaGm
1596 const uint16_t numGroups, const uint32_t shapeD, const uint16_t hwNum,1588 const uint16_t numGroups, const uint32_t shapeD, const uint16_t hwNum,
1597 const uint32_t eleNumAlign, bool hasGamma = true, bool hasBeta = true)1589 const uint32_t eleNumAlign, bool hasGamma = true, bool hasBeta = true)
1598{1590{
1599- MultiCopyLoopInfo<GAMMA_BETA_UB_DIM> loopInfo;1591+ NdDmaLoopInfo<GAMMA_BETA_UB_DIM> loopInfo;
1600 loopInfo.loopSize[INDEX_0] = numGroups;1592 loopInfo.loopSize[INDEX_0] = numGroups;
1601 loopInfo.loopSrcStride[INDEX_0] = shapeD;1593 loopInfo.loopSrcStride[INDEX_0] = shapeD;
1602 loopInfo.loopDstStride[INDEX_0] = eleNumAlign;1594 loopInfo.loopDstStride[INDEX_0] = eleNumAlign;
@@ -1610,8 +1602,8 @@ __aicore__ inline void CopyGammaAndBeta2UBByNDDMA(const GlobalTensor<T>& gammaGm
1610 loopInfo.loopDstStride[INDEX_2] = 1;1602 loopInfo.loopDstStride[INDEX_2] = 1;
1611 1603 
1612 T constValue = 0;1604 T constValue = 0;
1613- static constexpr MultiCopyConfig config = {false};1605+ static constexpr NdDmaConfig config = {false};
1614- MultiCopyParams<T, GAMMA_BETA_UB_DIM> paramsMain = {loopInfo, constValue};1606+ NdDmaParams<T, GAMMA_BETA_UB_DIM> paramsMain = {loopInfo, constValue};
1615 1607 
1616 if (hasGamma) {1608 if (hasGamma) {
1617 DataCopy<T, GAMMA_BETA_UB_DIM, config>(gammaTensor, gammaGm, paramsMain);1609 DataCopy<T, GAMMA_BETA_UB_DIM, config>(gammaTensor, gammaGm, paramsMain);
@@ -1674,10 +1666,10 @@ __aicore__ inline void ProcessMeanAndRstd(LocalTensor<float>& meanTensor, LocalT
1674 if constexpr (IsSameType<T1, float>::value) {1666 if constexpr (IsSameType<T1, float>::value) {
1675 CopyMeanAndRstd2Gm<float>(meanGm[gmOffset], rstdGm[gmOffset], meanTensor, rstdTensor, 1, curNumPerCore);1667 CopyMeanAndRstd2Gm<float>(meanGm[gmOffset], rstdGm[gmOffset], meanTensor, rstdTensor, 1, curNumPerCore);
1676 } else {1668 } else {
1677- __local_mem__ T1* meanOutLocal = (__local_mem__ T1*)meanOutTensor.GetPhyAddr();1669+ __ubuf__ T1* meanOutLocal = (__ubuf__ T1*)meanOutTensor.GetPhyAddr();
1678- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor.GetPhyAddr();1670+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor.GetPhyAddr();
1679- __local_mem__ T1* rstdOutLocal = (__local_mem__ T1*)rstdOutTensor.GetPhyAddr();1671+ __ubuf__ T1* rstdOutLocal = (__ubuf__ T1*)rstdOutTensor.GetPhyAddr();
1680- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor.GetPhyAddr();1672+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor.GetPhyAddr();
1681 uint16_t loopCount = CeilDiv(curNumPerCore, VL_FP32);1673 uint16_t loopCount = CeilDiv(curNumPerCore, VL_FP32);
1682 __VEC_SCOPE__1674 __VEC_SCOPE__
1683 {1675 {
@@ -1689,14 +1681,14 @@ __aicore__ inline void ProcessMeanAndRstd(LocalTensor<float>& meanTensor, LocalT
1689 RegTensor<T1> rstdOut;1681 RegTensor<T1> rstdOut;
1690 for (uint16_t i = 0; i < loopCount; i++) {1682 for (uint16_t i = 0; i < loopCount; i++) {
1691 pregLoop = UpdateMask<float>(sreg0);1683 pregLoop = UpdateMask<float>(sreg0);
1692- DataCopy(mean, meanLocal + i * VL_FP32);1684+ LoadAlign(mean, meanLocal + i * VL_FP32);
1693- DataCopy(rstd, rstdLocal + i * VL_FP32);1685+ LoadAlign(rstd, rstdLocal + i * VL_FP32);
1694 Cast<T1, float, castTraitB322B16Even>(meanOut, mean, pregLoop);1686 Cast<T1, float, castTraitB322B16Even>(meanOut, mean, pregLoop);
1695 Cast<T1, float, castTraitB322B16Even>(rstdOut, rstd, pregLoop);1687 Cast<T1, float, castTraitB322B16Even>(rstdOut, rstd, pregLoop);
1696- DataCopy<T1, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(meanOutLocal + i * VL_FP32, meanOut,1688+ StoreAlign<T1, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(meanOutLocal + i * VL_FP32, meanOut,
1697- pregLoop);1689+ pregLoop);
1698- DataCopy<T1, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(rstdOutLocal + i * VL_FP32, rstdOut,1690+ StoreAlign<T1, AscendC::MicroAPI::StoreDist::DIST_PACK_B32>(rstdOutLocal + i * VL_FP32, rstdOut,
1699- pregLoop);1691+ pregLoop);
1700 }1692 }
1701 }1693 }
1702 event_t eventIdVToMte3 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));1694 event_t eventIdVToMte3 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::V_MTE3));
@@ -97,7 +97,7 @@ private:
97 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());97 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());
98 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());98 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
99 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());99 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
100- __local_mem__ float* dichotomyAddLocal = (__local_mem__ float*)dichotomyAddTensor.GetPhyAddr();100+ __ubuf__ float* dichotomyAddLocal = (__ubuf__ float*)dichotomyAddTensor.GetPhyAddr();
101 for (int64_t i = 0; i < numPerCoreExtent; i++) {101 for (int64_t i = 0; i < numPerCoreExtent; i++) {
102 if (i == numPerCoreExtent - 1) {102 if (i == numPerCoreExtent - 1) {
103 numPerCoreProcess = numPerCoreTail;103 numPerCoreProcess = numPerCoreTail;
@@ -111,9 +111,9 @@ private:
111 CopyX2UB<T1>(xGm[xGmOffset], xTensor[xUbOffset], numPerCoreProcess, elemNum);111 CopyX2UB<T1>(xGm[xGmOffset], xTensor[xUbOffset], numPerCoreProcess, elemNum);
112 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);112 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
113 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);113 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
114- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset].GetPhyAddr();114+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset].GetPhyAddr();
115- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[onceNumPerCore * i].GetPhyAddr();115+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[onceNumPerCore * i].GetPhyAddr();
116- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[onceNumPerCore * i].GetPhyAddr();116+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[onceNumPerCore * i].GetPhyAddr();
117 if (i > 1) {117 if (i > 1) {
118 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);118 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);
119 }119 }
@@ -152,20 +152,17 @@ private:
152 __aicore__ inline void NormalizeAndSwishCommon(uint32_t xUbOffset, uint32_t numPerCoreoffset,152 __aicore__ inline void NormalizeAndSwishCommon(uint32_t xUbOffset, uint32_t numPerCoreoffset,
153 int64_t numPerCoreProcess, uint32_t numPerCoreLoop)153 int64_t numPerCoreProcess, uint32_t numPerCoreLoop)
154 {154 {
155- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset].GetPhyAddr();155+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset].GetPhyAddr();
156- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[xUbOffset].GetPhyAddr();156+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[xUbOffset].GetPhyAddr();
157 for (int64_t i = 0; i < numPerCoreProcess; i++) {157 for (int64_t i = 0; i < numPerCoreProcess; i++) {
158 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset + i) % numGroups) * shapeD;158 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset + i) % numGroups) * shapeD;
159 uint64_t betaOffset = gammaOffset;159 uint64_t betaOffset = gammaOffset;
160- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();160+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();
161- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();161+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();
162- __local_mem__ T2* gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor[gammaOffset].GetPhyAddr() :162+ __ubuf__ T2* gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor[gammaOffset].GetPhyAddr() : nullptr;
163- nullptr;163+ __ubuf__ T2* betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;
164- __local_mem__ T2* betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;164+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[numPerCoreLoop * onceNumPerCore + i].GetPhyAddr();
165- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[numPerCoreLoop * onceNumPerCore + i]165+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore + i].GetPhyAddr();
166- .GetPhyAddr();
167- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore + i]
168- .GetPhyAddr();
169 VFNormalizeAndSwishUnAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, shapeD,166 VFNormalizeAndSwishUnAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, shapeD,
170 hwNum, activateSilu, hasGamma, hasBeta);167 hwNum, activateSilu, hasGamma, hasBeta);
171 }168 }
@@ -188,12 +185,12 @@ private:
188 uint64_t curNumPerCoreTailDownAlign = (curNumPerCoreTail / numGroups) * numGroups;185 uint64_t curNumPerCoreTailDownAlign = (curNumPerCoreTail / numGroups) * numGroups;
189 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset) % numGroups) * elemNumAlign;186 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset) % numGroups) * elemNumAlign;
190 uint64_t betaOffset = gammaOffset;187 uint64_t betaOffset = gammaOffset;
191- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset].GetPhyAddr();188+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset].GetPhyAddr();
192- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[xUbOffset].GetPhyAddr();189+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[xUbOffset].GetPhyAddr();
193- __local_mem__ T2* gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor[gammaOffset].GetPhyAddr() : nullptr;190+ __ubuf__ T2* gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor[gammaOffset].GetPhyAddr() : nullptr;
194- __local_mem__ T2* betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;191+ __ubuf__ T2* betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;
195- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[numPerCoreLoop * onceNumPerCore].GetPhyAddr();192+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[numPerCoreLoop * onceNumPerCore].GetPhyAddr();
196- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore].GetPhyAddr();193+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore].GetPhyAddr();
197 // case1194 // case1
198 if (curNumPerCoreTailDownAlign < curNumPerCoreHeadUpAlign) {195 if (curNumPerCoreTailDownAlign < curNumPerCoreHeadUpAlign) {
199 VFNormalizeAndSwishFold(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, 1,196 VFNormalizeAndSwishFold(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, 1,
@@ -214,8 +211,8 @@ private:
214 yOutLocal = yOutLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead) * elemNumAlign;211 yOutLocal = yOutLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead) * elemNumAlign;
215 meanLocal = meanLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead);212 meanLocal = meanLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead);
216 rstdLocal = rstdLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead);213 rstdLocal = rstdLocal + (curNumPerCoreHeadUpAlign - curNumPerCoreHead);
217- gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor.GetPhyAddr() : nullptr;214+ gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor.GetPhyAddr() : nullptr;
218- betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor.GetPhyAddr() : nullptr;215+ betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor.GetPhyAddr() : nullptr;
219 if (groupsCount > 0) {216 if (groupsCount > 0) {
220 VFNormalizeAndSwishFold(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, groupsCount,217 VFNormalizeAndSwishFold(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, groupsCount,
221 numGroups, shapeD * hwNum, activateSilu, hasGamma, hasBeta);218 numGroups, shapeD * hwNum, activateSilu, hasGamma, hasBeta);
@@ -96,7 +96,7 @@ private:
96 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());96 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());
97 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());97 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
98 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());98 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
99- __local_mem__ float* dichotomyAddLocal = (__local_mem__ float*)dichotomyAddTensor.GetPhyAddr();99+ __ubuf__ float* dichotomyAddLocal = (__ubuf__ float*)dichotomyAddTensor.GetPhyAddr();
100 for (int64_t i = 0; i < numPerCoreExtent; i++) {100 for (int64_t i = 0; i < numPerCoreExtent; i++) {
101 if (i == numPerCoreExtent - 1) {101 if (i == numPerCoreExtent - 1) {
102 numPerCoreProcess = numPerCoreTail;102 numPerCoreProcess = numPerCoreTail;
@@ -110,9 +110,9 @@ private:
110 CopyX2UB<T1>(xGm[xGmOffset], xTensor[xUbOffset], numPerCoreProcess, elemNum);110 CopyX2UB<T1>(xGm[xGmOffset], xTensor[xUbOffset], numPerCoreProcess, elemNum);
111 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);111 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
112 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);112 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
113- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset].GetPhyAddr();113+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset].GetPhyAddr();
114- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[onceNumPerCore * i].GetPhyAddr();114+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[onceNumPerCore * i].GetPhyAddr();
115- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[onceNumPerCore * i].GetPhyAddr();115+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[onceNumPerCore * i].GetPhyAddr();
116 if (i > 1) {116 if (i > 1) {
117 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);117 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);
118 }118 }
@@ -155,14 +155,12 @@ private:
155 for (int64_t i = 0; i < numPerCoreProcess; i++) {155 for (int64_t i = 0; i < numPerCoreProcess; i++) {
156 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset + i) % numGroups) * shapeD;156 uint64_t gammaOffset = ((blockIdx * tiling->numPerCore + numPerCoreoffset + i) % numGroups) * shapeD;
157 uint64_t betaOffset = gammaOffset;157 uint64_t betaOffset = gammaOffset;
158- __local_mem__ T1* xLocal = (__local_mem__ T1*)xTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();158+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xTensor[xUbOffset + i * elemNumAlign].GetPhyAddr();
159- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[outputUbOffset + i * elemNumAlign].GetPhyAddr();159+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[outputUbOffset + i * elemNumAlign].GetPhyAddr();
160- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[numPerCoreLoop * onceNumPerCore + i]160+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[numPerCoreLoop * onceNumPerCore + i].GetPhyAddr();
161- .GetPhyAddr();161+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore + i].GetPhyAddr();
162- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[numPerCoreLoop * onceNumPerCore + i]162+ __ubuf__ T2* gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor.GetPhyAddr() : nullptr;
163- .GetPhyAddr();163+ __ubuf__ T2* betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor.GetPhyAddr() : nullptr;
164- __local_mem__ T2* gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor.GetPhyAddr() : nullptr;
165- __local_mem__ T2* betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor.GetPhyAddr() : nullptr;
166 if (i > 0) {164 if (i > 0) {
167 WaitFlag<HardEvent::V_MTE2>(eventIDVToMte2);165 WaitFlag<HardEvent::V_MTE2>(eventIDVToMte2);
168 }166 }
@@ -98,11 +98,11 @@ private:
98 98 
99 __aicore__ inline void CalMeanAndRstdByWelford(uint64_t curNumPerCore, uint64_t curInnerNumPerCore)99 __aicore__ inline void CalMeanAndRstdByWelford(uint64_t curNumPerCore, uint64_t curInnerNumPerCore)
100 {100 {
101- __local_mem__ float* tmpMeanLocal = (__local_mem__ float*)tMeanTensor.GetPhyAddr();101+ __ubuf__ float* tmpMeanLocal = (__ubuf__ float*)tMeanTensor.GetPhyAddr();
102- __local_mem__ float* tmpVarLocal = (__local_mem__ float*)tVarTensor.GetPhyAddr();102+ __ubuf__ float* tmpVarLocal = (__ubuf__ float*)tVarTensor.GetPhyAddr();
103- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor.GetPhyAddr();103+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor.GetPhyAddr();
104- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor.GetPhyAddr();104+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor.GetPhyAddr();
105- __local_mem__ float* dichotomyAddLocal = (__local_mem__ float*)dichotomyAddTensor.GetPhyAddr();105+ __ubuf__ float* dichotomyAddLocal = (__ubuf__ float*)dichotomyAddTensor.GetPhyAddr();
106 uint64_t xGmOffset = blockIdx * tiling->numPerCore * elemNum;106 uint64_t xGmOffset = blockIdx * tiling->numPerCore * elemNum;
107 uint32_t welfordLen = parallelN;107 uint32_t welfordLen = parallelN;
108 count = 0;108 count = 0;
@@ -123,7 +123,7 @@ private:
123 welfordLen);123 welfordLen);
124 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);124 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
125 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);125 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
126- __local_mem__ T1* x1Local = (__local_mem__ T1*)xPhase1Tensor[xPhase1Offset].GetPhyAddr();126+ __ubuf__ T1* x1Local = (__ubuf__ T1*)xPhase1Tensor[xPhase1Offset].GetPhyAddr();
127 count = count + 1;127 count = count + 1;
128 float scale = static_cast<float>(1.0) / static_cast<float>(count);128 float scale = static_cast<float>(1.0) / static_cast<float>(count);
129 VFWelfordParallelUpdate<T1>(x1Local, tmpMeanLocal, tmpVarLocal, i, welfordLen, scale);129 VFWelfordParallelUpdate<T1>(x1Local, tmpMeanLocal, tmpVarLocal, i, welfordLen, scale);
@@ -208,13 +208,12 @@ private:
208 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);208 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
209 uint64_t gammaOffset = gammaBaseOffset + i * (processSize / hwNumAlign);209 uint64_t gammaOffset = gammaBaseOffset + i * (processSize / hwNumAlign);
210 uint64_t betaOffset = gammaOffset;210 uint64_t betaOffset = gammaOffset;
211- __local_mem__ T1* xLocal = (__local_mem__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();211+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();
212- __local_mem__ T2* gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor[gammaOffset].GetPhyAddr() :212+ __ubuf__ T2* gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor[gammaOffset].GetPhyAddr() : nullptr;
213- nullptr;213+ __ubuf__ T2* betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;
214- __local_mem__ T2* betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;214+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();
215- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();215+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();
216- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();216+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[inputUbOffset].GetPhyAddr();
217- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[inputUbOffset].GetPhyAddr();
218 if (i > 1) {217 if (i > 1) {
219 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);218 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);
220 }219 }
@@ -278,14 +277,12 @@ private:
278 CopyX2UB(xGm[inputOffset], xPhase2Tensor[inputUbOffset], 1, copyLen);277 CopyX2UB(xGm[inputOffset], xPhase2Tensor[inputUbOffset], 1, copyLen);
279 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);278 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
280 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);279 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
281- __local_mem__ T1* xLocal = (__local_mem__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();280+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();
282- __local_mem__ T2* gammaLocal = hasGamma ? (__local_mem__ T2*)gammaTensor[gammaOffset].GetPhyAddr() :281+ __ubuf__ T2* gammaLocal = hasGamma ? (__ubuf__ T2*)gammaTensor[gammaOffset].GetPhyAddr() : nullptr;
283- nullptr;282+ __ubuf__ T2* betaLocal = hasBeta ? (__ubuf__ T2*)betaTensor[betaOffset].GetPhyAddr() : nullptr;
284- __local_mem__ T2* betaLocal = hasBeta ? (__local_mem__ T2*)betaTensor[betaOffset].GetPhyAddr() :283+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();
285- nullptr;284+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();
286- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();285+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[inputUbOffset].GetPhyAddr();
287- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();
288- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[inputUbOffset].GetPhyAddr();
289 if (extent > 1) {286 if (extent > 1) {
290 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);287 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);
291 }288 }
@@ -98,11 +98,11 @@ private:
98 98 
99 __aicore__ inline void CalMeanAndRstdByWelford(uint64_t curNumPerCore, uint64_t curInnerNumPerCore)99 __aicore__ inline void CalMeanAndRstdByWelford(uint64_t curNumPerCore, uint64_t curInnerNumPerCore)
100 {100 {
101- __local_mem__ float* tmpMeanLocal = (__local_mem__ float*)tMeanTensor.GetPhyAddr();101+ __ubuf__ float* tmpMeanLocal = (__ubuf__ float*)tMeanTensor.GetPhyAddr();
102- __local_mem__ float* tmpVarLocal = (__local_mem__ float*)tVarTensor.GetPhyAddr();102+ __ubuf__ float* tmpVarLocal = (__ubuf__ float*)tVarTensor.GetPhyAddr();
103- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor.GetPhyAddr();103+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor.GetPhyAddr();
104- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor.GetPhyAddr();104+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor.GetPhyAddr();
105- __local_mem__ float* dichotomyAddLocal = (__local_mem__ float*)dichotomyAddTensor.GetPhyAddr();105+ __ubuf__ float* dichotomyAddLocal = (__ubuf__ float*)dichotomyAddTensor.GetPhyAddr();
106 uint64_t xGmOffset = blockIdx * tiling->numPerCore * elemNum;106 uint64_t xGmOffset = blockIdx * tiling->numPerCore * elemNum;
107 uint32_t welfordLen = parallelN;107 uint32_t welfordLen = parallelN;
108 count = 0;108 count = 0;
@@ -123,7 +123,7 @@ private:
123 welfordLen);123 welfordLen);
124 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);124 SetFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
125 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);125 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
126- __local_mem__ T1* x1Local = (__local_mem__ T1*)xPhase1Tensor[xPhase1Offset].GetPhyAddr();126+ __ubuf__ T1* x1Local = (__ubuf__ T1*)xPhase1Tensor[xPhase1Offset].GetPhyAddr();
127 count = count + 1;127 count = count + 1;
128 float scale = static_cast<float>(1.0) / static_cast<float>(count);128 float scale = static_cast<float>(1.0) / static_cast<float>(count);
129 VFWelfordParallelUpdate<T1>(x1Local, tmpMeanLocal, tmpVarLocal, i, welfordLen, scale);129 VFWelfordParallelUpdate<T1>(x1Local, tmpMeanLocal, tmpVarLocal, i, welfordLen, scale);
@@ -175,10 +175,10 @@ private:
175 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());175 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
176 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());176 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
177 177 
178- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();178+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();
179- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();179+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();
180- __local_mem__ T2* gammaLocal = (__local_mem__ T2*)gammaTensor.GetPhyAddr();180+ __ubuf__ T2* gammaLocal = (__ubuf__ T2*)gammaTensor.GetPhyAddr();
181- __local_mem__ T2* betaLocal = (__local_mem__ T2*)betaTensor.GetPhyAddr();181+ __ubuf__ T2* betaLocal = (__ubuf__ T2*)betaTensor.GetPhyAddr();
182 for (int64_t i = 0; i < loopNum; i++) {182 for (int64_t i = 0; i < loopNum; i++) {
183 uint64_t inputGmOffset = inputBaseOffset + hwNum * rowsCount * i + elemNum * curNumPerCore;183 uint64_t inputGmOffset = inputBaseOffset + hwNum * rowsCount * i + elemNum * curNumPerCore;
184 bool isPing = (i % BUFFER_NUM) == 0;184 bool isPing = (i % BUFFER_NUM) == 0;
@@ -203,8 +203,8 @@ private:
203 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);203 WaitFlag<HardEvent::MTE2_V>(isPing ? eventIDMte2ToVPing : eventIDMte2ToVPong);
204 SetFlag<HardEvent::MTE2_V>(eventIDMte2ToV);204 SetFlag<HardEvent::MTE2_V>(eventIDMte2ToV);
205 WaitFlag<HardEvent::MTE2_V>(eventIDMte2ToV);205 WaitFlag<HardEvent::MTE2_V>(eventIDMte2ToV);
206- __local_mem__ T1* xLocal = (__local_mem__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();206+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();
207- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[inputUbOffset].GetPhyAddr();207+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[inputUbOffset].GetPhyAddr();
208 VFNormalizeAndSwishAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, rowsCount,208 VFNormalizeAndSwishAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, rowsCount,
209 reduceCount, activateSilu, hasGamma, hasBeta);209 reduceCount, activateSilu, hasGamma, hasBeta);
210 SetFlag<HardEvent::V_MTE3>(isPing ? eventIDVToMte3Ping : eventIDVToMte3Pong);210 SetFlag<HardEvent::V_MTE3>(isPing ? eventIDVToMte3Ping : eventIDVToMte3Pong);
@@ -254,10 +254,10 @@ private:
254 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());254 auto eventIDVToMte2Pong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE2>());
255 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());255 auto eventIDMte3ToVPing = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
256 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());256 auto eventIDMte3ToVPong = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>());
257- __local_mem__ float* meanLocal = (__local_mem__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();257+ __ubuf__ float* meanLocal = (__ubuf__ float*)meanTensor[curInnerNumPerCore].GetPhyAddr();
258- __local_mem__ float* rstdLocal = (__local_mem__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();258+ __ubuf__ float* rstdLocal = (__ubuf__ float*)rstdTensor[curInnerNumPerCore].GetPhyAddr();
259- __local_mem__ T2* gammaLocal = (__local_mem__ T2*)gammaTensor.GetPhyAddr();259+ __ubuf__ T2* gammaLocal = (__ubuf__ T2*)gammaTensor.GetPhyAddr();
260- __local_mem__ T2* betaLocal = (__local_mem__ T2*)betaTensor.GetPhyAddr();260+ __ubuf__ T2* betaLocal = (__ubuf__ T2*)betaTensor.GetPhyAddr();
261 for (int64_t i = 0; i < loopNum; i++) { // for D261 for (int64_t i = 0; i < loopNum; i++) { // for D
262 int64_t copyLen = totalSize;262 int64_t copyLen = totalSize;
263 uint64_t gammaOffset = gammaBaseOffset + i;263 uint64_t gammaOffset = gammaBaseOffset + i;
@@ -286,8 +286,8 @@ private:
286 if (extent > 1) {286 if (extent > 1) {
287 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);287 WaitFlag<HardEvent::MTE3_V>(isPing ? eventIDMte3ToVPing : eventIDMte3ToVPong);
288 }288 }
289- __local_mem__ T1* xLocal = (__local_mem__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();289+ __ubuf__ T1* xLocal = (__ubuf__ T1*)xPhase2Tensor[inputUbOffset].GetPhyAddr();
290- __local_mem__ T1* yOutLocal = (__local_mem__ T1*)yTensor[inputUbOffset].GetPhyAddr();290+ __ubuf__ T1* yOutLocal = (__ubuf__ T1*)yTensor[inputUbOffset].GetPhyAddr();
291 int32_t reduceCount = copyLen;291 int32_t reduceCount = copyLen;
292 VFNormalizeAndSwishAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, 1,292 VFNormalizeAndSwishAlign<T1, T2>(xLocal, gammaLocal, betaLocal, meanLocal, rstdLocal, yOutLocal, 1,
293 copyLen, activateSilu, hasGamma, hasBeta);293 copyLen, activateSilu, hasGamma, hasBeta);