已合并
refactor(modulate): A5(arch35) Reg 加载存储接口升级为 LoadAlign/StoreAlign #9990
refactor(modulate): A5(arch35) Reg 加载存储接口升级为 LoadAlign/StoreAlign #9990
已合并
Nice try创建于 9月7日
共 1 个文件变更+20-20
@@ -207,10 +207,10 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeScaleShif
207 const LocalTensor<T>& yLocal, const LocalTensor<T>& xLocal, const LocalTensor<T>& scaleLocal,207 const LocalTensor<T>& yLocal, const LocalTensor<T>& xLocal, const LocalTensor<T>& scaleLocal,
208 const LocalTensor<T>& shiftLocal, const uint64_t& rows, const uint64_t& calCount)208 const LocalTensor<T>& shiftLocal, const uint64_t& rows, const uint64_t& calCount)
209{209{
210- __local_mem__ T* xAddr = (__local_mem__ T*)xLocal.GetPhyAddr();210+ __ubuf__ T* xAddr = (__ubuf__ T*)xLocal.GetPhyAddr();
211- __local_mem__ T* scaleAddr = (__local_mem__ T*)scaleLocal.GetPhyAddr();211+ __ubuf__ T* scaleAddr = (__ubuf__ T*)scaleLocal.GetPhyAddr();
212- __local_mem__ T* shiftAddr = (__local_mem__ T*)shiftLocal.GetPhyAddr();212+ __ubuf__ T* shiftAddr = (__ubuf__ T*)shiftLocal.GetPhyAddr();
213- __local_mem__ T* yAddr = (__local_mem__ T*)yLocal.GetPhyAddr();213+ __ubuf__ T* yAddr = (__ubuf__ T*)yLocal.GetPhyAddr();
214 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;214 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;
215 uint32_t dtypeSize = sizeof(CAST_T);215 uint32_t dtypeSize = sizeof(CAST_T);
216 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;216 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;
@@ -238,15 +238,15 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeScaleShif
238 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);238 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);
239 }239 }
240 } else {240 } else {
241- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(scaleReg, scaleAddr + i * VL);241+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(scaleReg, scaleAddr + i * VL);
242- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(shiftReg, shiftAddr + i * VL);242+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(shiftReg, shiftAddr + i * VL);
243 Reg::Adds(scaleReg, scaleReg, 1.0f, preg);243 Reg::Adds(scaleReg, scaleReg, 1.0f, preg);
244 for (uint16_t row = 0; row < rowLen; row++) {244 for (uint16_t row = 0; row < rowLen; row++) {
245 uint64_t rowOffset = row * calCountAlign;245 uint64_t rowOffset = row * calCountAlign;
246- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);246+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);
247 Reg::Mul(xReg, xReg, scaleReg, preg);247 Reg::Mul(xReg, xReg, scaleReg, preg);
248 Reg::Add(yReg, xReg, shiftReg, preg);248 Reg::Add(yReg, xReg, shiftReg, preg);
249- Reg::DataCopy<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);249+ Reg::StoreAlign<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);
250 }250 }
251 }251 }
252 }252 }
@@ -268,9 +268,9 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeScale(con
268 const uint64_t& rows,268 const uint64_t& rows,
269 const uint64_t& calCount)269 const uint64_t& calCount)
270{270{
271- __local_mem__ T* xAddr = (__local_mem__ T*)xLocal.GetPhyAddr();271+ __ubuf__ T* xAddr = (__ubuf__ T*)xLocal.GetPhyAddr();
272- __local_mem__ T* scaleAddr = (__local_mem__ T*)scaleLocal.GetPhyAddr();272+ __ubuf__ T* scaleAddr = (__ubuf__ T*)scaleLocal.GetPhyAddr();
273- __local_mem__ T* yAddr = (__local_mem__ T*)yLocal.GetPhyAddr();273+ __ubuf__ T* yAddr = (__ubuf__ T*)yLocal.GetPhyAddr();
274 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;274 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;
275 uint32_t dtypeSize = sizeof(CAST_T);275 uint32_t dtypeSize = sizeof(CAST_T);
276 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;276 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;
@@ -296,13 +296,13 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeScale(con
296 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);296 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);
297 }297 }
298 } else {298 } else {
299- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(scaleReg, scaleAddr + i * VL);299+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(scaleReg, scaleAddr + i * VL);
300 Reg::Adds(scaleReg, scaleReg, 1.0f, preg);300 Reg::Adds(scaleReg, scaleReg, 1.0f, preg);
301 for (uint16_t row = 0; row < static_cast<uint16_t>(rows); row++) {301 for (uint16_t row = 0; row < static_cast<uint16_t>(rows); row++) {
302 uint64_t rowOffset = row * calCountAlign;302 uint64_t rowOffset = row * calCountAlign;
303- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);303+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);
304 Reg::Mul(yReg, xReg, scaleReg, preg);304 Reg::Mul(yReg, xReg, scaleReg, preg);
305- Reg::DataCopy<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);305+ Reg::StoreAlign<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);
306 }306 }
307 }307 }
308 }308 }
@@ -324,9 +324,9 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeShift(con
324 const uint64_t& rows,324 const uint64_t& rows,
325 const uint64_t& calCount)325 const uint64_t& calCount)
326{326{
327- __local_mem__ T* xAddr = (__local_mem__ T*)xLocal.GetPhyAddr();327+ __ubuf__ T* xAddr = (__ubuf__ T*)xLocal.GetPhyAddr();
328- __local_mem__ T* shiftAddr = (__local_mem__ T*)shiftLocal.GetPhyAddr();328+ __ubuf__ T* shiftAddr = (__ubuf__ T*)shiftLocal.GetPhyAddr();
329- __local_mem__ T* yAddr = (__local_mem__ T*)yLocal.GetPhyAddr();329+ __ubuf__ T* yAddr = (__ubuf__ T*)yLocal.GetPhyAddr();
330 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;330 using CAST_T = std::conditional_t<std::is_same_v<T, bfloat16_t>, float, T>;
331 uint32_t dtypeSize = sizeof(CAST_T);331 uint32_t dtypeSize = sizeof(CAST_T);
332 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;332 uint16_t VL = AscendC::VECTOR_REG_WIDTH / dtypeSize;
@@ -351,12 +351,12 @@ __aicore__ inline void ModulateBaseKernel<T, isScale, isShift>::ComputeShift(con
351 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);351 ops::StoreOneTensorForDtypeT<T>(yAddr, yReg, preg, rowOffset + i * VL);
352 }352 }
353 } else {353 } else {
354- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(shiftReg, shiftAddr + i * VL);354+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(shiftReg, shiftAddr + i * VL);
355 for (uint16_t row = 0; row < static_cast<uint16_t>(rows); row++) {355 for (uint16_t row = 0; row < static_cast<uint16_t>(rows); row++) {
356 uint64_t rowOffset = row * calCountAlign;356 uint64_t rowOffset = row * calCountAlign;
357- Reg::DataCopy<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);357+ Reg::LoadAlign<T, Reg::LoadDist::DIST_NORM>(xReg, xAddr + rowOffset + i * VL);
358 Reg::Add(yReg, xReg, shiftReg, preg);358 Reg::Add(yReg, xReg, shiftReg, preg);
359- Reg::DataCopy<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);359+ Reg::StoreAlign<T, Reg::StoreDist::DIST_NORM>(yAddr + rowOffset + i * VL, yReg, preg);
360 }360 }
361 }361 }
362 }362 }