514
csrc/attention/common/op_kernel/CopyInL1.h
Normal file
514
csrc/attention/common/op_kernel/CopyInL1.h
Normal file
@@ -0,0 +1,514 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file CopyInL1.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef COPYINL1_H
|
||||
#define COPYINL1_H
|
||||
|
||||
enum class KVLAYOUT
|
||||
{
|
||||
BNBD, // [blockNums, headNum, blockSize, headDim]
|
||||
BBH, // [blockNums, blockSize, headNum * headDim]
|
||||
NZ // [blockNums, headNum, d1, blockSize, d0], d1 = headDim / d0, d0 = 32 (block byte) / sizeof(KV_T)
|
||||
};
|
||||
|
||||
struct CopyParam{
|
||||
uint32_t width;
|
||||
uint32_t height;
|
||||
uint32_t orgWidth;
|
||||
};
|
||||
|
||||
struct PAShape{
|
||||
uint32_t blockNum;
|
||||
uint32_t blockSize;
|
||||
uint32_t headNum; // 一般为kv的head num
|
||||
uint32_t headDim; // mla下rope为64, 非rope为512
|
||||
uint32_t maxblockNumPerBatch; // block table 每一行的最大个数
|
||||
uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d
|
||||
uint32_t copyRowNum;
|
||||
uint32_t copyRowNumAlign;
|
||||
uint32_t pageStride;
|
||||
};
|
||||
|
||||
struct Position{
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx;
|
||||
uint32_t s2Offset;
|
||||
uint32_t dIdx; // N轴被切,对应D轴被切
|
||||
};
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmCopyInToL1(LocalTensor<L1Type>& L1Tensor, GlobalTensor<L1Type>& GmTensor, const CopyParam& mmCopyParam)
|
||||
{
|
||||
Nd2NzParams Gm2L1Nd2NzParams;
|
||||
Gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
|
||||
Gm2L1Nd2NzParams.nValue = mmCopyParam.height; // 单个ND矩阵的实际行数,单位为元素个数
|
||||
Gm2L1Nd2NzParams.dValue = mmCopyParam.width; // 单个ND矩阵的实际列数(vD),单位为元素个数
|
||||
Gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
|
||||
Gm2L1Nd2NzParams.srcDValue = mmCopyParam.orgWidth; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
|
||||
Gm2L1Nd2NzParams.dstNzC0Stride = (Gm2L1Nd2NzParams.nValue + 15) >> 4 << 4; // 转换为NZ矩阵后,相邻Block起始地址之间的偏移, 单位为Block个数
|
||||
Gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
|
||||
Gm2L1Nd2NzParams.dstNzMatrixStride = 0; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
|
||||
DataCopy(L1Tensor, GmTensor, Gm2L1Nd2NzParams);
|
||||
}
|
||||
|
||||
// 场景:key、value GM to L1
|
||||
// GM按ND格式存储
|
||||
// L1按NZ格式存储
|
||||
// GM的行、列、列的stride(D or ND)BNSD 和 BSH的区别
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void DataCopyGmNDToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
uint32_t rowAct,
|
||||
uint32_t rowAlign,
|
||||
uint32_t col, // D
|
||||
uint32_t colStride) // D or N*D
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = rowAct; // 行数
|
||||
|
||||
nd2nzPara.dValue = col;
|
||||
nd2nzPara.srcDValue = colStride;
|
||||
nd2nzPara.dstNzC0Stride = rowAlign;
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
DataCopy(l1Tensor, gmTensor, nd2nzPara);
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void DataCopyGmScaleNDToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
uint32_t rowAct,
|
||||
uint32_t rowAlign,
|
||||
uint32_t col, // D
|
||||
uint32_t colStride) // D or N*D
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = rowAct;
|
||||
|
||||
nd2nzPara.dValue = col;
|
||||
nd2nzPara.srcDValue = colStride;
|
||||
nd2nzPara.dstNzC0Stride = rowAlign;
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = nd2nzPara.nValue;
|
||||
|
||||
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
|
||||
GlobalTensor<bfloat16_t> gmTensorCast;
|
||||
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
|
||||
DataCopy(l1TensorCast, gmTensorCast, nd2nzPara);
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void DataCopyGmScaleDNToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
uint32_t rowAct,
|
||||
uint32_t rowAlign,
|
||||
uint32_t col,
|
||||
uint32_t colStride)
|
||||
{
|
||||
Dn2NzParams dn2nzPara;
|
||||
dn2nzPara.dnNum = 1;
|
||||
dn2nzPara.nValue = col / 2;
|
||||
dn2nzPara.dValue = rowAct;
|
||||
dn2nzPara.srcDValue = colStride / 2;
|
||||
dn2nzPara.dstNzC0Stride = dn2nzPara.nValue;
|
||||
dn2nzPara.dstNzNStride = 1;
|
||||
dn2nzPara.srcDnMatrixStride = 0;
|
||||
dn2nzPara.dstNzMatrixStride = dn2nzPara.nValue;
|
||||
|
||||
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
|
||||
GlobalTensor<bfloat16_t> gmTensorCast;
|
||||
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
|
||||
DataCopy(l1TensorCast, gmTensorCast, dn2nzPara);
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void DataCopyGmNZToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
uint32_t rowAct,
|
||||
uint32_t dstRowStride,
|
||||
uint32_t srcRowStride,
|
||||
uint32_t col)
|
||||
{
|
||||
uint32_t blockElementCnt = 32U / sizeof(L1Type);
|
||||
if constexpr (IsSameType<L1Type, int4b_t>::value) {
|
||||
blockElementCnt = 64U;
|
||||
}
|
||||
DataCopyParams intriParams;
|
||||
intriParams.blockCount = col / blockElementCnt;
|
||||
intriParams.blockLen = rowAct;
|
||||
intriParams.dstStride = dstRowStride;
|
||||
intriParams.srcStride = srcRowStride;
|
||||
DataCopy(l1Tensor, gmTensor, intriParams);
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmCopyInToL1HasRopePANoContinue(LocalTensor<L1Type>& nopeTensor, LocalTensor<L1Type>& ropeTensor,
|
||||
GlobalTensor<L1Type>& nopeGmTensor, GlobalTensor<L1Type>& ropeGmTensor,
|
||||
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
|
||||
const PAShape &shape,
|
||||
const PAShape &ropeShape,
|
||||
const Position &startPos)
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
|
||||
uint32_t curS2Idx = startPos.s2Offset;
|
||||
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
|
||||
// ropeshape的M方向与nopeshape保持一样, 此处只判断nopeshape的
|
||||
while(copyFinishRowCnt < shape.copyRowNum){
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
|
||||
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
|
||||
//计算可以拷贝行数
|
||||
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
|
||||
if (shape.pageStride > 0) {
|
||||
offset = idInBlockTable * shape.pageStride;
|
||||
}
|
||||
uint64_t keyRopeOffset = idInBlockTable * ropeShape.blockSize * ropeShape.headNum * ropeShape.headDim;
|
||||
if (ropeShape.pageStride > 0) {
|
||||
keyRopeOffset = idInBlockTable * ropeShape.pageStride;
|
||||
}
|
||||
|
||||
if (kvLayout == KVLAYOUT::NZ) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.blockSize * ropeShape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * ropeShape.blockSize;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
|
||||
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
|
||||
|
||||
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
|
||||
DataCopyGmNZToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, (ropeShape.copyRowNumAlign - copyRowCnt), (ropeShape.blockSize - copyRowCnt), ropeShape.actHeadDim);
|
||||
} else {
|
||||
uint64_t dStride = shape.headDim;
|
||||
uint64_t dRopeStride = ropeShape.headDim;
|
||||
if (kvLayout == KVLAYOUT::BBH) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim) + remainRowCnt * ropeShape.headDim * ropeShape.headNum;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
dRopeStride = ropeShape.headDim * ropeShape.headNum;
|
||||
} else{
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim * ropeShape.blockSize) + remainRowCnt * ropeShape.headDim;
|
||||
}
|
||||
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
uint32_t srcDValue = dStride;
|
||||
uint32_t dRopeValue = ropeShape.actHeadDim;
|
||||
uint32_t srcRopeDValue = dRopeStride;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
|
||||
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
|
||||
|
||||
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
|
||||
DataCopyGmNDToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dRopeValue, srcRopeDValue);
|
||||
}
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmCopyInToL1HasRopePA(LocalTensor<L1Type>& nopeTensor, LocalTensor<L1Type>& ropeTensor,
|
||||
GlobalTensor<L1Type>& nopeGmTensor, GlobalTensor<L1Type>& ropeGmTensor,
|
||||
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
|
||||
const PAShape &shape,
|
||||
const PAShape &ropeShape,
|
||||
const Position &startPos)
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
|
||||
uint32_t curS2Idx = startPos.s2Offset;
|
||||
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
|
||||
// ropeshape的M方向与nopeshape保持一样, 此处只判断nopeshape的
|
||||
while(copyFinishRowCnt < shape.copyRowNum){
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
|
||||
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
|
||||
//计算可以拷贝行数
|
||||
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
|
||||
uint64_t keyRopeOffset = idInBlockTable * ropeShape.blockSize * ropeShape.headNum * ropeShape.headDim;
|
||||
if (kvLayout == KVLAYOUT::NZ) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.blockSize * ropeShape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * ropeShape.blockSize;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
|
||||
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
|
||||
|
||||
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
|
||||
DataCopyGmNZToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, (ropeShape.copyRowNumAlign - copyRowCnt), (ropeShape.blockSize - copyRowCnt), ropeShape.actHeadDim);
|
||||
} else {
|
||||
uint64_t dStride = shape.headDim;
|
||||
uint64_t dRopeStride = ropeShape.headDim;
|
||||
if (kvLayout == KVLAYOUT::BBH) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim) + remainRowCnt * ropeShape.headDim * ropeShape.headNum;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
dRopeStride = ropeShape.headDim * ropeShape.headNum;
|
||||
} else{
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
|
||||
keyRopeOffset += static_cast<uint64_t>(startPos.n2Idx * ropeShape.headDim * ropeShape.blockSize) + remainRowCnt * ropeShape.headDim;
|
||||
}
|
||||
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
uint32_t srcDValue = dStride;
|
||||
uint32_t dRopeValue = ropeShape.actHeadDim;
|
||||
uint32_t srcRopeDValue = dRopeStride;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = nopeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = nopeGmTensor[offset];
|
||||
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
|
||||
|
||||
LocalTensor<L1Type> tmpRopeDstTensor = ropeTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpRopeSrcTensor = ropeGmTensor[keyRopeOffset];
|
||||
DataCopyGmNDToL1(tmpRopeDstTensor, tmpRopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dRopeValue, srcRopeDValue);
|
||||
}
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmCopyInToL1PA(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
|
||||
const PAShape &shape, const Position &startPos)
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch; // 块表的基偏移量
|
||||
uint32_t curS2Idx = startPos.s2Offset;
|
||||
uint32_t blockElementCnt = 32U / sizeof(L1Type); // 每个块的元素数量
|
||||
while(copyFinishRowCnt < shape.copyRowNum){
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
|
||||
uint64_t remainRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上获取的编号
|
||||
//计算可以拷贝行数
|
||||
uint32_t copyRowCnt = shape.blockSize - remainRowCnt; // 一次只能处理一个Block
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum){
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
|
||||
if (kvLayout == KVLAYOUT::NZ) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
|
||||
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
|
||||
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
|
||||
} else {
|
||||
uint64_t dStride = shape.headDim;
|
||||
if (kvLayout == KVLAYOUT::BBH) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
} else {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
|
||||
}
|
||||
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
uint32_t srcDValue = dStride;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
|
||||
DataCopyGmNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
|
||||
}
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmScaleCopyInToL1PAForND(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
|
||||
const PAShape &shape, const Position &startPos)
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
|
||||
uint32_t curS2Idx = startPos.s2Offset;
|
||||
constexpr uint32_t blockElementCnt = 32U / sizeof(L1Type);
|
||||
while(copyFinishRowCnt < shape.copyRowNum) {
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize;
|
||||
uint64_t remainRowCnt = curS2Idx % shape.blockSize;
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
|
||||
uint32_t copyRowCnt = shape.blockSize - remainRowCnt;
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt;
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim;
|
||||
if (kvLayout == KVLAYOUT::NZ) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
|
||||
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
|
||||
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
|
||||
} else {
|
||||
uint64_t dStride = shape.headDim;
|
||||
if (kvLayout == KVLAYOUT::BBH) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
} else {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
|
||||
}
|
||||
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
uint32_t srcDValue = dStride;
|
||||
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset * 2];
|
||||
DataCopyGmScaleNDToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, copyRowCnt, dValue, srcDValue);
|
||||
}
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template<typename L1Type>
|
||||
__aicore__ inline void GmScaleCopyInToL1PAForDN(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
GlobalTensor<int32_t>& blockTableGm, KVLAYOUT kvLayout,
|
||||
const PAShape &shape, const Position &startPos)
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
|
||||
uint32_t curS2Idx = startPos.s2Offset;
|
||||
constexpr uint32_t blockElementCnt = 32U / sizeof(L1Type);
|
||||
while(copyFinishRowCnt < shape.copyRowNum) {
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize;
|
||||
uint64_t remainRowCnt = curS2Idx % shape.blockSize;
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
|
||||
uint32_t copyRowCnt = shape.blockSize - remainRowCnt;
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt;
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim;
|
||||
if (kvLayout == KVLAYOUT::NZ) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.blockSize * shape.headDim) + remainRowCnt * blockElementCnt + startPos.dIdx * shape.blockSize;
|
||||
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
|
||||
DataCopyGmNZToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, (shape.copyRowNumAlign - copyRowCnt), (shape.blockSize - copyRowCnt), shape.actHeadDim);
|
||||
} else {
|
||||
uint64_t dStride = shape.headDim;
|
||||
if (kvLayout == KVLAYOUT::BBH) {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim) + remainRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
} else {
|
||||
offset += static_cast<uint64_t>(startPos.n2Idx * shape.headDim * shape.blockSize) + remainRowCnt * shape.headDim + startPos.dIdx;
|
||||
}
|
||||
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
uint32_t srcDValue = dStride;
|
||||
LocalTensor<L1Type> tmpNopeDstTensor = l1Tensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<L1Type> tmpNopeSrcTensor = gmTensor[offset];
|
||||
|
||||
DataCopyGmScaleDNToL1(tmpNopeDstTensor, tmpNopeSrcTensor, copyRowCnt, copyRowCnt, dValue, srcDValue);
|
||||
}
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template<typename INPUT_T>
|
||||
__aicore__ inline void CopyToL1Nd2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
|
||||
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
|
||||
{
|
||||
Nd2NzParams gm2L1Nd2NzParams;
|
||||
gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
|
||||
gm2L1Nd2NzParams.nValue = nValue; // 单个ND矩阵的实际行数,单位为元素个数
|
||||
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
|
||||
gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
|
||||
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__) || (__NPU_ARCH__ == 5102)
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value || IsSameType<INPUT_T, int8_t>::value) {
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 31) >> 5 << 5;
|
||||
} else {
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 15) >> 4 << 4;
|
||||
}
|
||||
#else
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (nValue + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移, 单位为Block个数
|
||||
#endif
|
||||
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
|
||||
gm2L1Nd2NzParams.dstNzMatrixStride = 0; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
|
||||
DataCopy(l1Tensor, gmTensor, gm2L1Nd2NzParams);
|
||||
}
|
||||
|
||||
template<typename INPUT_T>
|
||||
__aicore__ inline void CopyScaleToL1Nd2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
|
||||
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
|
||||
{
|
||||
Nd2NzParams gm2L1Nd2NzParams;
|
||||
gm2L1Nd2NzParams.ndNum = 1; // ND矩阵的个数
|
||||
gm2L1Nd2NzParams.nValue = nValue / 2; // 单个ND矩阵的实际行数,单位为元素个数
|
||||
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
|
||||
gm2L1Nd2NzParams.srcNdMatrixStride = 0; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = nValue / 2; // NZ矩阵相邻Block起始地址之间的偏移, 单位为Block个数
|
||||
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
|
||||
gm2L1Nd2NzParams.dstNzMatrixStride = gm2L1Nd2NzParams.nValue; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
|
||||
|
||||
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
|
||||
GlobalTensor<bfloat16_t> gmTensorCast;
|
||||
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
|
||||
DataCopy(l1TensorCast, gmTensorCast, gm2L1Nd2NzParams);
|
||||
}
|
||||
|
||||
template<typename INPUT_T>
|
||||
__aicore__ inline void CopyScaleToL1Dn2Nz(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
|
||||
uint32_t nValue, uint32_t dValue, uint32_t srcDValue)
|
||||
{
|
||||
Dn2NzParams gm2L1Dn2NzParams;
|
||||
gm2L1Dn2NzParams.dnNum = 1; // ND矩阵的个数
|
||||
gm2L1Dn2NzParams.nValue = nValue / 2; // 单个DN矩阵的实际列数,单位为元素个数
|
||||
gm2L1Dn2NzParams.dValue = dValue; // 单个DN矩阵的实际行数,单位为元素个数
|
||||
gm2L1Dn2NzParams.srcDnMatrixStride = 0; // 相邻Dn矩阵起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Dn2NzParams.srcDValue = srcDValue / 2; // 同一个Dn矩阵中相邻行起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Dn2NzParams.dstNzC0Stride = nValue / 2;
|
||||
gm2L1Dn2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
|
||||
gm2L1Dn2NzParams.dstNzMatrixStride = gm2L1Dn2NzParams.nValue; // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
|
||||
|
||||
LocalTensor<bfloat16_t> l1TensorCast = l1Tensor.template ReinterpretCast<bfloat16_t>();
|
||||
GlobalTensor<bfloat16_t> gmTensorCast;
|
||||
gmTensorCast.SetGlobalBuffer(((__gm__ bfloat16_t*)(gmTensor.GetPhyAddr())));
|
||||
DataCopy(l1TensorCast, gmTensorCast, gm2L1Dn2NzParams);
|
||||
}
|
||||
|
||||
template<typename INPUT_T>
|
||||
__aicore__ inline void CopyToL1Nd2NzGS1Merge(const LocalTensor<INPUT_T> &l1Tensor, const GlobalTensor<INPUT_T> &gmTensor,
|
||||
uint32_t ndNum, uint32_t nValue, uint32_t dValue, uint32_t srcNdMatrixStride, uint32_t srcDValue, uint32_t dstNzC0Stride) // BSNGD 合轴拷贝
|
||||
{
|
||||
Nd2NzParams gm2L1Nd2NzParams;
|
||||
gm2L1Nd2NzParams.ndNum = ndNum; // ND矩阵的个数
|
||||
gm2L1Nd2NzParams.nValue = nValue; // 单个ND矩阵的实际行数,单位为元素个数
|
||||
gm2L1Nd2NzParams.dValue = dValue; // 单个ND矩阵的实际列数,单位为元素个数
|
||||
gm2L1Nd2NzParams.srcNdMatrixStride = srcNdMatrixStride; // 相邻ND矩阵起始地址之间的偏移, 单位为元素个数
|
||||
gm2L1Nd2NzParams.srcDValue = srcDValue; // 同一个ND矩阵中相邻行起始地址之间的偏移, 单位为元素个数
|
||||
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__) || (__NPU_ARCH__ == 5102)
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value || IsSameType<INPUT_T, int8_t>::value) {
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 31) >> 5 << 5; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,32对齐
|
||||
} else {
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,16对齐
|
||||
}
|
||||
#else
|
||||
gm2L1Nd2NzParams.dstNzC0Stride = (dstNzC0Stride + 15) >> 4 << 4; // NZ矩阵相邻Block起始地址之间的偏移,单位为Block个数,16对齐
|
||||
#endif
|
||||
gm2L1Nd2NzParams.dstNzNStride = 1; // 转换为NZ矩阵后,ND之间相邻两行在NZ矩阵中起始地址之间的偏移, 单位为Block个数
|
||||
gm2L1Nd2NzParams.dstNzMatrixStride = nValue * 32 / sizeof(INPUT_T); // 两个NZ矩阵,起始地址之间的偏移, 单位为元素数量
|
||||
DataCopy(l1Tensor, gmTensor, gm2L1Nd2NzParams);
|
||||
}
|
||||
#endif
|
||||
56
csrc/attention/common/op_kernel/FixpipeOut.h
Normal file
56
csrc/attention/common/op_kernel/FixpipeOut.h
Normal file
@@ -0,0 +1,56 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file FixpipeOut.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef FIXPIPEOUT_H
|
||||
#define FIXPIPEOUT_H
|
||||
|
||||
constexpr FixpipeConfig PFA_CFG_ROW_MAJOR_UB = {CO2Layout::ROW_MAJOR, true}; // ROW_MAJOR: 使能NZ2ND,输出数据格式为ND格式; true: 用于用户指定目的地址的位置是否是UB
|
||||
constexpr FixpipeConfig PFA_CFG_ROW_MAJOR_GM = {CO2Layout::ROW_MAJOR, false}; // ROW_MAJOR: 使能NZ2ND,输出数据格式为ND格式; true: 用于用户指定目的地址的位置是否是UB
|
||||
constexpr FixpipeConfig FA_CFG_NZ_UB = {CO2Layout::NZ, true}; // 不使能NZ2ND,输出数据格式为NZ格式; true: 用于用户指定目的地址的位置是否是UB
|
||||
|
||||
struct fixpipeOutParams {
|
||||
uint32_t fixpOutMSize;
|
||||
uint32_t fixpOutNSize;
|
||||
};
|
||||
|
||||
template<typename mmOutputType, typename computeType, typename l0cType>
|
||||
__aicore__ inline void FixpipeMmCopyOutToUB(LocalTensor<mmOutputType>& mmResUb, LocalTensor<l0cType>& L0CTensor, const fixpipeOutParams& fixpOutParam)
|
||||
{
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> L0C2UbFixpParams; // L0C->UB
|
||||
L0C2UbFixpParams.nSize = (fixpOutParam.fixpOutNSize + 7) >> 3 << 3; // L0C上的bmm1结果矩阵N方向的size大小;同mmadParams.n;8个元素(32B)对齐
|
||||
L0C2UbFixpParams.mSize = (fixpOutParam.fixpOutMSize + 1) >> 1 << 1; // 有效数据不足16行,只需输出部分行即可;L0C上的bmm1结果矩阵M方向的size大小必须是偶数
|
||||
L0C2UbFixpParams.srcStride = ((L0C2UbFixpParams.mSize + 15) >> 4) << 4; // L0C上matmul结果相邻连续数据片断间隔(前面一个数据块的头与后面数据块的头的间隔),单位为16 *sizeof(T) //源NZ矩阵中相邻Z排布的起始地址偏移
|
||||
L0C2UbFixpParams.dstStride = (L0C2UbFixpParams.nSize + 15) >> 4 << 4; // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到,ND方案(S1 * S2)时脏数据用mask剔除
|
||||
L0C2UbFixpParams.dualDstCtl = 1; // 双目标模式,按M维度拆分, M / 2 * N写入每个UB,M必须为2的倍数
|
||||
L0C2UbFixpParams.params.ndNum = 1;
|
||||
L0C2UbFixpParams.params.srcNdStride = 0;
|
||||
L0C2UbFixpParams.params.dstNdStride = 0;
|
||||
Fixpipe<mmOutputType, computeType, PFA_CFG_ROW_MAJOR_UB>(mmResUb, L0CTensor, L0C2UbFixpParams); // 将matmul结果从L0C搬运到UB
|
||||
}
|
||||
|
||||
template<typename mmOutputType, typename computeType, typename l0cType>
|
||||
__aicore__ inline void FixpipeMmCopyOutToGm(GlobalTensor<mmOutputType>& mmResGm,LocalTensor<l0cType>& L0CTensor, const fixpipeOutParams& fixpOutParam)
|
||||
{
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> L0C2GmFixpParams; // L0C->Gm
|
||||
L0C2GmFixpParams.nSize = (fixpOutParam.fixpOutNSize + 7) >> 3 << 3; // L0C上的bmm1结果矩阵N方向的size大小;同mmadParams.n;8个元素(32B)对齐;分档计算且vector1中通过mask筛选出实际有效值
|
||||
L0C2GmFixpParams.mSize = (fixpOutParam.fixpOutMSize + 1) >> 1 << 1; // 有效数据不足16行,只需输出部分行即可;L0C上的bmm1结果矩阵M方向的size大小;同mmadParams.m
|
||||
L0C2GmFixpParams.srcStride = ((L0C2GmFixpParams.mSize + 15) >> 4) << 4; // L0C上bmm1结果相邻连续数据片断间隔(前面一个数据块的头与后面数据块的头的间隔)
|
||||
L0C2GmFixpParams.dstStride = (L0C2GmFixpParams.nSize + 15) >> 4 << 4; // mmResGm上两行之间的间隔
|
||||
L0C2GmFixpParams.dualDstCtl = 1;
|
||||
L0C2GmFixpParams.params.ndNum = 1;
|
||||
L0C2GmFixpParams.params.srcNdStride = 0;
|
||||
L0C2GmFixpParams.params.dstNdStride = 0;
|
||||
Fixpipe<mmOutputType, computeType, PFA_CFG_ROW_MAJOR_GM>(mmResGm, L0CTensor, L0C2GmFixpParams); // 将matmul结果从L0C搬运到Gm
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,139 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned128_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// no update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<T, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
|
||||
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
|
||||
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
// maxUb is [S1, 1], BRC_B32 is reading one fp32 element and broadcast it to all 64 vreg element
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
ProcessVec1NoUpdateImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,143 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned128_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_max_tmp, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_max_tmp, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// update, originN == 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateImpl128VF <T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED128_UPDATE_SFA_H
|
||||
@@ -0,0 +1,149 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned128_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateGeneralImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN, uint32_t pltTailN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b8 = AscendC::MicroAPI::CreateMask<T2, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<float>(pltOriTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_reduce_n =
|
||||
AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::VL8>();
|
||||
|
||||
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
|
||||
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_max_tmp, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_all_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_all_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, 64 < originN <= 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateGeneralImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (65*16*2/32),单位block, low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
const uint32_t oriTailN = originN - floatRepSize;
|
||||
const uint32_t tailN = s2BaseSize - floatRepSize;
|
||||
uint32_t pltOriTailN = oriTailN;
|
||||
uint32_t pltTailN = tailN;
|
||||
|
||||
ProcessVec1NoUpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriTailN, pltTailN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,159 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned128_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateGeneralImpl128VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriTailN,
|
||||
uint32_t pltTailN, uint32_t pltN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x_unroll_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_even;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_odd;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_odd_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_odd_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 = AscendC::MicroAPI::CreateMask<uint16_t,
|
||||
AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_n_b16 = AscendC::MicroAPI::UpdateMask<uint16_t>(pltN);
|
||||
AscendC::MicroAPI::MaskReg preg_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltTailN);
|
||||
AscendC::MicroAPI::MaskReg preg_ori_tail_n = AscendC::MicroAPI::UpdateMask<T>(pltOriTailN);
|
||||
|
||||
AscendC::MicroAPI::Duplicate(vreg_min, minValue);
|
||||
// x_max = max(src, axis=-1, keepdims=True); x_max = Max(x_max, inMax)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x_unroll, srcUb + floatRepSize + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_all); // Muls(scale)
|
||||
AscendC::MicroAPI::Muls(vreg_input_x_unroll, vreg_input_x_unroll, scale, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::Select(vreg_input_x_unroll_new, vreg_input_x_unroll, vreg_min, preg_ori_tail_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + floatRepSize + i * s2BaseSize, vreg_input_x_unroll_new, preg_tail_n);
|
||||
AscendC::MicroAPI::Max(vreg_max_tmp, vreg_input_x, vreg_input_x_unroll_new, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_cur_max, vreg_max_tmp, preg_all);
|
||||
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2); // 获取新的max[s1, 1]
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧max的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_DINTLV_B32>(
|
||||
vreg_input_x, vreg_input_x_unroll, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_even, vreg_input_x, vreg_max_brc, preg_all);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp_odd, vreg_input_x_unroll, vreg_max_brc, preg_all);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Add(vreg_exp_sum, vreg_exp_even, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_bf16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_bf16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_bf16, (RegTensor<uint16_t>&)vreg_exp_even_bf16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_bf16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_bf16, blockStride, repeatStride, preg_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_even_fp16, vreg_exp_even, preg_all);
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitOne>(vreg_exp_odd_fp16, vreg_exp_odd, preg_all);
|
||||
AscendC::MicroAPI::Or((RegTensor<uint16_t>&)vreg_exp_fp16, (RegTensor<uint16_t>&)vreg_exp_even_fp16,
|
||||
(RegTensor<uint16_t>&)vreg_exp_odd_fp16, preg_all_b16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_exp_fp16, blockStride, repeatStride, preg_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
|
||||
// update, 64 < originN <= 128
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateGeneralImpl128(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
const uint32_t oriTailN = originN - floatRepSize;
|
||||
const uint32_t tailN = s2BaseSize - floatRepSize;
|
||||
uint32_t pltOriTailN = oriTailN;
|
||||
uint32_t pltTailN = tailN;
|
||||
uint32_t pltN = s2BaseSize;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateGeneralImpl128VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride,
|
||||
m, scale, minValue, pltOriTailN, pltTailN, pltN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SFA_H
|
||||
@@ -0,0 +1,129 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_unaligned64_no_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1NoUpdateImpl64VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * expSumUb, __ubuf__ T * maxUb, __ubuf__ T * maxUbStart,
|
||||
__ubuf__ T * srcUb, const uint32_t blockStride, const uint32_t repeatStride,
|
||||
const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN, uint32_t pltSrcN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_min;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<float>(pltSrcN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
|
||||
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n); // Muls(scale)
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_input_max, vreg_input_x, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), vreg_input_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)maxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, maxUbStart + i);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
|
||||
vreg_exp_bf16, vreg_exp_bf16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
|
||||
vreg_exp_fp16, vreg_exp_fp16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)expSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
// no update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1NoUpdateImpl64(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUbStart = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
uint32_t pltOriginalN = originN;
|
||||
uint32_t pltSrcN = s2BaseSize;
|
||||
|
||||
ProcessVec1NoUpdateImpl64VF<T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, expSumUb, maxUb, maxUbStart, srcUb, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriginalN, pltSrcN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SFA_H
|
||||
@@ -0,0 +1,141 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_aligned64_update_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
// update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 128, uint32_t s2BaseSize = 128>
|
||||
__simd_vf__ void ProcessVec1UpdateImpl64VF(
|
||||
__ubuf__ T2 * expUb, __ubuf__ T * srcUb, __ubuf__ T * inMaxUb,
|
||||
__ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, __ubuf__ T * tmpMaxUb2, const uint32_t blockStride,
|
||||
const uint32_t repeatStride, const uint16_t m, const T scale, const T minValue, uint32_t pltOriginalN,
|
||||
uint32_t pltSrcN)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_x;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_tmp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_in_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_new;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_max_brc;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_cur_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
|
||||
// bfloat16_t
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_exp_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_even_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_dst_odd_bf16;
|
||||
// half
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_exp_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_even_fp16;
|
||||
AscendC::MicroAPI::RegTensor<half> vreg_dst_odd_fp16;
|
||||
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_max;
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_all_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg preg_ori_src_n = AscendC::MicroAPI::UpdateMask<T>(pltOriginalN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n = AscendC::MicroAPI::UpdateMask<T>(pltSrcN);
|
||||
AscendC::MicroAPI::MaskReg preg_src_n_b16 =
|
||||
AscendC::MicroAPI::CreateMask<uint16_t, AscendC::MicroAPI::MaskPattern::H>();
|
||||
|
||||
// x_max = max(src, axis=-1, keepdims=True)
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::Muls(vreg_input_x, vreg_input_x, scale, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)srcUb + i * s2BaseSize, vreg_input_x, preg_src_n);
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::MAX, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_cur_max, vreg_input_x, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), vreg_cur_max, ureg_max, 1);
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpMaxUb), ureg_max, 0);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_in_max, inMaxUb);
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
AscendC::MicroAPI::LoadAlign(vreg_cur_max, tmpMaxUb2);
|
||||
AscendC::MicroAPI::Max(vreg_max_new, vreg_cur_max, vreg_in_max, preg_all); // 计算新、旧的最大值
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)tmpMaxUb2, vreg_max_new, preg_all);
|
||||
|
||||
AscendC::MicroAPI::LocalMemBar<MemType::VEC_STORE, MemType::VEC_LOAD>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
vreg_max_brc, tmpMaxUb2 + i);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_x, srcUb + i * s2BaseSize);
|
||||
AscendC::MicroAPI::ExpSub(vreg_exp, vreg_input_x, vreg_max_brc, preg_ori_src_n);
|
||||
|
||||
// x_sum = sum(x_exp, axis=-1, keepdims=True)
|
||||
AscendC::MicroAPI::Reduce<MicroAPI::ReduceType::SUM, float, float, MicroAPI::MaskMergeMode::ZEROING>(
|
||||
vreg_exp_sum, vreg_exp, preg_ori_src_n);
|
||||
AscendC::MicroAPI::StoreUnAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), vreg_exp_sum, ureg_exp_sum, 1);
|
||||
|
||||
if constexpr (IsSameType<T2, bfloat16_t>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_bf16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_bf16, vreg_dst_odd_bf16,
|
||||
vreg_exp_bf16, vreg_exp_bf16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_bf16, blockStride, repeatStride, preg_src_n_b16);
|
||||
} else if constexpr (IsSameType<T2, half>::value) {
|
||||
AscendC::MicroAPI::Cast<T2, T, castTraitZero>(vreg_exp_fp16, vreg_exp, preg_all_b16);
|
||||
AscendC::MicroAPI::DeInterleave(vreg_dst_even_fp16, vreg_dst_odd_fp16,
|
||||
vreg_exp_fp16, vreg_exp_fp16);
|
||||
AscendC::MicroAPI::StoreAlign<T2, MicroAPI::DataCopyMode::DATA_BLOCK_COPY,
|
||||
MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T2 *&)expUb), vreg_dst_even_fp16, blockStride, repeatStride, preg_src_n_b16);
|
||||
}
|
||||
}
|
||||
AscendC::MicroAPI::StoreUnAlignPost<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
((__ubuf__ T *&)tmpExpSumUb), ureg_exp_sum, 0);
|
||||
}
|
||||
|
||||
|
||||
// update, originN <= 64
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128>
|
||||
__aicore__ inline void ProcessVec1UpdateImpl64(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
// 写的时候固定用65或者33的stride去写,因为正向目前使能settail之后mm2的s1方向必须算满128或者64行
|
||||
// stride, high 16bits: blockStride (m*16*2/32), low 16bits: repeatStride (1)
|
||||
const uint32_t blockStride = s1BaseSize >> 1 | 0x1;
|
||||
const uint32_t repeatStride = 1;
|
||||
uint32_t pltOriginalN = originN;
|
||||
uint32_t pltSrcN = s2BaseSize;
|
||||
|
||||
__ubuf__ T2 * expUb = (__ubuf__ T2*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
__ubuf__ T * tmpMaxUb2 = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
ProcessVec1UpdateImpl64VF <T, T2, s1BaseSize, s2BaseSize>(
|
||||
expUb, srcUb, inMaxUb, tmpExpSumUb, tmpMaxUb, tmpMaxUb2, blockStride, repeatStride, m, scale, minValue,
|
||||
pltOriginalN, pltSrcN);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_ALIGNED64_UPDATE_SFA_H
|
||||
112
csrc/attention/common/op_kernel/arch35/vf/vf_basic_block_utils.h
Normal file
112
csrc/attention/common/op_kernel/arch35/vf/vf_basic_block_utils.h
Normal file
@@ -0,0 +1,112 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_basic_block_utils.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UTILS_H
|
||||
#define VF_BASIC_BLOCK_UTILS_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_basic_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
namespace FaVectorApi {
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
constexpr uint32_t halfRepSize = 128;
|
||||
constexpr uint32_t blockBytesU8 = 32;
|
||||
constexpr float fp8e4m3MaxValue = 448.0f;
|
||||
constexpr float int8MaxValue = 127.0f;
|
||||
constexpr float hifp8MaxValue = 32768.0f;
|
||||
constexpr float floatEps = 2.220446049250313e-16;
|
||||
/* **************************************************************************************************
|
||||
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
|
||||
* ************************************************************************************************* */
|
||||
using namespace MicroAPI;
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitZero = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitOne = {
|
||||
AscendC::MicroAPI::RegLayout::ONE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitTwo = {
|
||||
AscendC::MicroAPI::RegLayout::TWO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitThree = {
|
||||
AscendC::MicroAPI::RegLayout::THREE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_ROUND,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintZero = {
|
||||
AscendC::MicroAPI::RegLayout::ZERO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintOne = {
|
||||
AscendC::MicroAPI::RegLayout::ONE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintTwo = {
|
||||
AscendC::MicroAPI::RegLayout::TWO,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
constexpr static AscendC::MicroAPI::CastTrait castTraitRintThree = {
|
||||
AscendC::MicroAPI::RegLayout::THREE,
|
||||
AscendC::MicroAPI::SatMode::SAT,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING,
|
||||
AscendC::RoundMode::CAST_RINT,
|
||||
};
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, fp8e4m3MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P_INT8(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, int8MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
|
||||
#define USE_MLA_FULLQUANT_V1_P_HIFP8(vreg_exp, vreg_rowmax_p, MaskReg) \
|
||||
do { \
|
||||
Muls(vreg_exp, vreg_exp, hifp8MaxValue, MaskReg); \
|
||||
Div(vreg_exp, vreg_exp, vreg_rowmax_p, MaskReg); \
|
||||
} while (0)
|
||||
} // namespace
|
||||
|
||||
#endif // VF_BASIC_BLOCK_UTILS_H
|
||||
727
csrc/attention/common/op_kernel/arch35/vf/vf_flashupdate_new.h
Normal file
727
csrc/attention/common/op_kernel/arch35/vf/vf_flashupdate_new.h
Normal file
@@ -0,0 +1,727 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_flashupdate_new.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MY_FLASH_UPDATE_NEW_INTERFACE_H
|
||||
#define MY_FLASH_UPDATE_NEW_INTERFACE_H
|
||||
|
||||
#include "kernel_tensor.h"
|
||||
|
||||
namespace FaVectorApi {
|
||||
// bf16->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraitFp16_32_update = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
constexpr uint16_t REDUCE_SIZE = 1;
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void FlashUpdateBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_row_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
|
||||
// dstTensor = preTensor * expMaxTensor + curTensor
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
|
||||
if constexpr (isMlaFullQuant) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
|
||||
}
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
|
||||
}
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
/* **************************************************************************************************
|
||||
* FlashUpdate, fp32
|
||||
* ************************************************************************************************* */
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateBasic(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
|
||||
|
||||
FlashUpdateBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
|
||||
dstUb, curUb, preUb, expMaxUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__simd_vf__ inline void FlashUpdateGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
// dstTensor = preTensor * expMaxTensor + curTensor
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_add, preg_all);
|
||||
}
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_add, preg_tail_d);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__aicore__ inline void FlashUpdateGeneral(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = static_cast<uint32_t>(tailD);
|
||||
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
FlashUpdateGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
|
||||
dstUb, curUb, preUb, expMaxUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
|
||||
}
|
||||
|
||||
/*
|
||||
* @ingroup FlashUpdate
|
||||
* @brief compute, dstTensor = preTensor * expMaxTensor + curTensor
|
||||
* @param [out] dstTensor, output LocalTensor
|
||||
* @param [in] curTensor, input LocalTensor
|
||||
* @param [in] preTensor, input LocalTensor
|
||||
* @param [in] expMaxTensor, input LocalTensor
|
||||
* @param [in] m, input rows
|
||||
* @param [in] d, input columns, should be 32 bytes aligned
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdate, T must be float");
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
if constexpr(srcD % floatRepSize == 0) {
|
||||
FlashUpdateBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor,
|
||||
m, d, deScaleV, deScaleVPre);
|
||||
} else {
|
||||
|
||||
FlashUpdateGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(dstTensor, curTensor, preTensor, expMaxTensor, m, d,
|
||||
deScaleV, deScaleVPre);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void FlashUpdateLastBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, __ubuf__ float * rowMaxUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_row_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<half> vreg_cast;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
|
||||
constexpr float int8MaxValueRec = 1 / 127.0f;
|
||||
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_row_max, rowMaxUb + i * reduceSize);
|
||||
}
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
Mul(vreg_input_cur, vreg_input_cur, vreg_row_max, preg_all);
|
||||
}
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
|
||||
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
|
||||
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
|
||||
} else {
|
||||
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
|
||||
}
|
||||
}
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateLastBasic(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ float * rowMaxUb = (__ubuf__ T*)rowMaxTensor.GetPhyAddr();
|
||||
|
||||
FlashUpdateLastBasicVF<T, INPUT_T, OUTPUT_T, srcD, reduceSize, isUpdatePre, isMlaFullQuant>(
|
||||
dstUb, curUb, preUb, expMaxUb, expSumUb, rowMaxUb, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__simd_vf__ inline void FlashUpdateLastGeneralVF(__ubuf__ float * dstUb, __ubuf__ float * curUb,
|
||||
__ubuf__ float * preUb, __ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, const uint16_t m, const uint16_t d,
|
||||
const float deScaleV, const float deScaleVPre, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_input_pre;
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_mul;
|
||||
RegTensor<float> vreg_add;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<half> vreg_cast;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MaskReg preg_tail_d = UpdateMask<float>(tmpTailD);
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + j * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_all);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_all);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_all);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_all);
|
||||
}
|
||||
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
LoadAlign(vreg_input_pre, preUb + i * d + dLoops * floatRepSize);
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + dLoops * floatRepSize);
|
||||
Mul(vreg_mul, vreg_exp_max, vreg_input_pre, preg_tail_d);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
if constexpr (isUpdatePre) {
|
||||
Muls(vreg_mul, vreg_mul, deScaleVPre, preg_all);
|
||||
}
|
||||
}
|
||||
Add(vreg_add, vreg_mul, vreg_input_cur, preg_tail_d);
|
||||
Div(vreg_div, vreg_add, vreg_exp_sum, preg_tail_d);
|
||||
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + dLoops * floatRepSize, vreg_div, preg_tail_d);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t reduceSize, bool isUpdatePre>
|
||||
__aicore__ inline void FlashUpdateLastGeneral(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * preUb = (__ubuf__ T*)preTensor.GetPhyAddr();
|
||||
__ubuf__ float * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = tailD;
|
||||
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
FlashUpdateLastGeneralVF<T, INPUT_T, OUTPUT_T, reduceSize, isUpdatePre>(
|
||||
dstUb, curUb, preUb, expMaxUb, expSumUb, m, d, deScaleV, deScaleVPre, pltTailD, hasTail);
|
||||
}
|
||||
|
||||
/*
|
||||
* @ingroup FlashUpdateLast
|
||||
* @brief compute, dstTensor = (preTensor * expMaxTensor + curTensor) / expSumTensor
|
||||
* @param [out] dstTensor, output LocalTensor
|
||||
* @param [in] curTensor, input LocalTensor
|
||||
* @param [in] preTensor, input LocalTensor
|
||||
* @param [in] expMaxTensor, input LocalTensor
|
||||
* @param [in] expSumTensor, input LocalTensor
|
||||
* @param [in] m, input rows
|
||||
* @param [in] d, input columns, 32 bytes align
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, bool isUpdatePre, bool isMlaFullQuant>
|
||||
__aicore__ inline void FlashUpdateLastNew(const LocalTensor<T>& dstTensor,
|
||||
const LocalTensor<T>& curTensor, const LocalTensor<T>& preTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& rowMaxTensor, const LocalTensor<T>& expSumTensor,
|
||||
uint16_t m, uint16_t d, const float deScaleV, const float deScaleVPre)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdateLast, T must be float");
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
if constexpr(srcD % floatRepSize == 0) {
|
||||
FlashUpdateLastBasic<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE, isUpdatePre, isMlaFullQuant>(
|
||||
dstTensor, curTensor, preTensor, expMaxTensor, rowMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
|
||||
} else {
|
||||
FlashUpdateLastGeneral<T, INPUT_T, OUTPUT_T, REDUCE_SIZE, isUpdatePre>(
|
||||
dstTensor, curTensor, preTensor, expMaxTensor, expSumTensor, m, d, deScaleV, deScaleVPre);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
|
||||
__simd_vf__ inline void LastDivNewVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * expSumUb,
|
||||
const uint16_t m, const uint16_t d, const float deScaleV)
|
||||
{
|
||||
RegTensor<float> vreg_input_cur;
|
||||
RegTensor<float> vreg_div;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t dLoops = d >> 6;
|
||||
constexpr float fp8e4m3MaxValueRec = 1 / 448.0f;
|
||||
constexpr float int8MaxValueRec = 1 / 127.0f;
|
||||
constexpr float hifp8MaxValueRec = 1 / 32768.0f;
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
uint32_t sreg_init = d;
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * REDUCE_SIZE);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
MaskReg preg_update = UpdateMask<float>(sreg_init);
|
||||
|
||||
LoadAlign(vreg_input_cur, curUb + i * d + j * floatRepSize);
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value ||
|
||||
IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_t>::value ||
|
||||
IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_input_cur, vreg_input_cur, deScaleV, preg_all);
|
||||
}
|
||||
Div(vreg_div, vreg_input_cur, vreg_exp_sum, preg_update);
|
||||
if constexpr (isMlaFullQuant) {
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e4m3fn_t>::value) {
|
||||
Muls(vreg_div, vreg_div, fp8e4m3MaxValueRec, preg_all);
|
||||
} else if constexpr (IsSameType<INPUT_T, int8_t>::value) {
|
||||
Muls(vreg_div, vreg_div, int8MaxValueRec, preg_all);
|
||||
} else {
|
||||
Muls(vreg_div, vreg_div, hifp8MaxValueRec, preg_all);
|
||||
}
|
||||
}
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_div, preg_update);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// dstTensor = curTensor / expSumTensor, curTensor: [64,128], expSumTensor: [64,8]
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD, bool isMlaFullQuant>
|
||||
__aicore__ inline void LastDivNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& expSumTensor, const uint16_t m, const uint16_t d, const float deScaleV)
|
||||
{
|
||||
__ubuf__ float * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ float * curUb = (__ubuf__ T*)curTensor.GetPhyAddr();
|
||||
__ubuf__ float * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
|
||||
LastDivNewVF<T, INPUT_T, OUTPUT_T, srcD, isMlaFullQuant>(dstUb, curUb, expSumUb, m, d, deScaleV);
|
||||
}
|
||||
|
||||
template <typename T, uint32_t srcD>
|
||||
__simd_vf__ inline void InvalidLineUpdateVF(__ubuf__ T * dstUb, __ubuf__ T * srcUb, __ubuf__ T * maxUb,
|
||||
const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
|
||||
{
|
||||
RegTensor<float> vreg_invalid_value;
|
||||
RegTensor<float> vreg_max;
|
||||
RegTensor<float> vreg_input;
|
||||
RegTensor<float> vreg_input_brc;
|
||||
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
MaskReg preg_compare;
|
||||
const uint16_t dLoops = d >> 6;
|
||||
|
||||
Duplicate(vreg_invalid_value, invalidValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_max, maxUb + i);
|
||||
Compares<T, CMPMODE::EQ>(preg_compare, vreg_max, minValue, preg_all);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
LoadAlign(vreg_input, srcUb + i * d + j * floatRepSize);
|
||||
Select(vreg_input_brc, vreg_invalid_value, vreg_input, preg_compare);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * d + j * floatRepSize, vreg_input_brc, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, uint32_t srcD>
|
||||
__aicore__ inline void InvalidLineUpdate(const LocalTensor<T>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& maxTensor, const uint16_t m, const uint16_t d, const T minValue, const T invalidValue)
|
||||
{
|
||||
__ubuf__ T * dstUb = (__ubuf__ T*)dstTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcUb = (__ubuf__ T*)srcTensor.GetPhyAddr();
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
uint16_t dLoops = d >> 6;
|
||||
|
||||
InvalidLineUpdateVF<T, srcD>(dstUb, srcUb, maxUb, m, d, minValue, invalidValue);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void ComputeLseOutputVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ T *dstUb, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<T> vregResFinal;
|
||||
MicroAPI::RegTensor<float> vregMinValue;
|
||||
MicroAPI::RegTensor<float> vregInfValue;
|
||||
MicroAPI::MaskReg pregCompare;
|
||||
constexpr uint32_t dealRows = 8;
|
||||
constexpr uint32_t floatRepSize = 64; // 64: 一个寄存器存64个float
|
||||
constexpr float infValue = 3e+99; // 3e+99 for float inf
|
||||
constexpr uint32_t tmpMin = 0xFF7FFFFF;
|
||||
float minValue = *((float*)&tmpMin);
|
||||
uint16_t updateLoops = dealCount / dealRows;
|
||||
uint16_t tailLSize = dealCount % dealRows * 8;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailLSize);
|
||||
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
|
||||
MicroAPI::Duplicate<float, float>(vregInfValue, infValue);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + (i * dealRows));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + (i * dealRows));
|
||||
|
||||
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregAll);
|
||||
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
|
||||
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + (i * floatRepSize), vregResFinal, pregAll);
|
||||
}
|
||||
|
||||
if (tailLSize != 0) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregSum, srcSumUb + dealRows * updateLoops);
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_E2B_B32>(vregMax, srcMaxUb + dealRows * updateLoops);
|
||||
|
||||
MicroAPI::Log<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSum, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, vregMax, pregTail);
|
||||
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregTail);
|
||||
MicroAPI::Select<T>(vregResFinal, vregInfValue, vregRes, pregCompare);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(dstUb + floatRepSize * updateLoops, vregResFinal, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void ComputeLseOutputVF(const LocalTensor<T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
|
||||
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * dstUb = (__ubuf__ T *)dstTensor.GetPhyAddr();
|
||||
|
||||
ComputeLseOutputVF<T>(srcSumUb, srcMaxUb, dstUb, dealCount);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void SinkSubExpAddVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, const T sinkValue, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<T> vregSink;
|
||||
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
|
||||
uint16_t updateLoops = dealCount / floatRepSize;
|
||||
uint16_t tailSize = dealCount % floatRepSize;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailSize);
|
||||
|
||||
//mask
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
|
||||
Duplicate(vregSink, sinkValue);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregAll);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
|
||||
}
|
||||
|
||||
for (uint16_t i = 0; i < tailSize; i = i + tailSize) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSink, vregMax, pregTail);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void SinkSubExpAddVF(const LocalTensor<T>& softmaxSumTensor, const LocalTensor<T>& softmaxMaxTensor,
|
||||
const T sinkValue, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
|
||||
SinkSubExpAddVF<T>(srcSumUb, srcMaxUb, sinkValue, dealCount);
|
||||
}
|
||||
|
||||
template <typename T, typename SINK_T>
|
||||
__simd_vf__ inline void SinkSubExpAddGSFusedVF(__ubuf__ T *srcSumUb, __ubuf__ T *srcMaxUb, __ubuf__ uint16_t *sinkUb, const uint32_t dealCount)
|
||||
{
|
||||
MicroAPI::RegTensor<T> vregSum;
|
||||
MicroAPI::RegTensor<T> vregMax;
|
||||
MicroAPI::RegTensor<T> vregRes;
|
||||
MicroAPI::RegTensor<SINK_T> vregSink;
|
||||
MicroAPI::RegTensor<T> vregSinkCast;
|
||||
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
|
||||
uint16_t updateLoops = dealCount / floatRepSize;
|
||||
uint16_t tailSize = dealCount % floatRepSize;
|
||||
uint32_t pltTail = static_cast<uint32_t>(tailSize);
|
||||
|
||||
//mask
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregTail = MicroAPI::UpdateMask<T>(pltTail);
|
||||
MicroAPI::MaskReg pregSinkAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_UNPACK_B16>((MicroAPI::RegTensor<uint16_t>&)vregSink, sinkUb);
|
||||
MicroAPI::Cast<T, SINK_T, castTraitFp16_32_update>(vregSinkCast, vregSink, pregSinkAll);
|
||||
|
||||
for (uint16_t i = 0; i < updateLoops; ++i) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (i * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (i * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregAll);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregAll);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregAll);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (i * floatRepSize), vregSum, pregAll);
|
||||
}
|
||||
|
||||
if (tailSize != 0) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregSum, srcSumUb + (updateLoops * floatRepSize));
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregMax, srcMaxUb + (updateLoops * floatRepSize));
|
||||
|
||||
MicroAPI::Sub<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregSinkCast, vregMax, pregTail);
|
||||
MicroAPI::Exp<T, MicroAPI::MaskMergeMode::ZEROING>(vregRes, vregRes, pregTail);
|
||||
MicroAPI::Add<T, MicroAPI::MaskMergeMode::ZEROING>(vregSum, vregSum, vregRes, pregTail);
|
||||
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(srcSumUb + (updateLoops * floatRepSize), vregSum, pregTail);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename SINK_T>
|
||||
__aicore__ inline void SinkSubExpAddGSFusedVF(const LocalTensor<SINK_T>& dstTensor, const LocalTensor<T>& softmaxSumTensor,
|
||||
const LocalTensor<T>& softmaxMaxTensor, uint32_t dealCount)
|
||||
{
|
||||
__ubuf__ T * srcSumUb = (__ubuf__ T *)softmaxSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * srcMaxUb = (__ubuf__ T *)softmaxMaxTensor.GetPhyAddr();
|
||||
__ubuf__ uint16_t * dstUb = (__ubuf__ uint16_t *)dstTensor.GetPhyAddr();
|
||||
|
||||
SinkSubExpAddGSFusedVF<T, SINK_T>(srcSumUb, srcMaxUb, dstUb, dealCount);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void RowInvalidUpdateVF(__ubuf__ T *finalUb, __ubuf__ float *maxUb, const uint16_t m,
|
||||
const uint16_t d, int64_t dSize, const uint32_t pltTailD, const uint16_t hasTail)
|
||||
{
|
||||
constexpr uint16_t floatRepSize = 64; // 64: 一个寄存器可以存储64个float类型数据
|
||||
const uint16_t dLoops = d / floatRepSize;
|
||||
|
||||
|
||||
constexpr uint32_t tmpZero = 0x00000000; // zero value of fp16 and fp32
|
||||
const T zeroValue = *((T*)&tmpZero);
|
||||
constexpr uint32_t tmpMin = 0xFF7FFFFF; // min value of float
|
||||
const float minValue = *((float*)&tmpMin);
|
||||
MicroAPI::RegTensor<float> vregMinValue;
|
||||
MicroAPI::RegTensor<T> vregZeroValue;
|
||||
MicroAPI::RegTensor<float> vregMax;
|
||||
MicroAPI::RegTensor<T> vregFinal;
|
||||
MicroAPI::RegTensor<T> vregFinalNew;
|
||||
|
||||
MicroAPI::MaskReg pregAll = MicroAPI::CreateMask<T, MicroAPI::MaskPattern::ALL>();
|
||||
uint32_t tmpTailD = pltTailD;
|
||||
MicroAPI::MaskReg pregTailD = MicroAPI::UpdateMask<T>(tmpTailD);
|
||||
MicroAPI::MaskReg pregCompare;
|
||||
|
||||
MicroAPI::Duplicate<float, float>(vregMinValue, minValue);
|
||||
MicroAPI::Duplicate<T, T>(vregZeroValue, zeroValue);
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_BRC_B32>(vregMax, maxUb + i);
|
||||
MicroAPI::Compare<float, CMPMODE::EQ>(pregCompare, vregMax, vregMinValue, pregAll);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + j * floatRepSize);
|
||||
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + j * floatRepSize,
|
||||
vregFinalNew, pregAll);
|
||||
}
|
||||
for (uint16_t t = 0; t < hasTail; ++t) {
|
||||
MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_NORM>(vregFinal, finalUb + i * dSize + dLoops * floatRepSize);
|
||||
MicroAPI::Select<T>(vregFinalNew, vregZeroValue, vregFinal, pregCompare);
|
||||
MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(finalUb + i * dSize + dLoops * floatRepSize,
|
||||
vregFinalNew, pregTailD);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void RowInvalidUpdateVF(const LocalTensor<T>& finalTensor, const LocalTensor<float>& maxTensor,
|
||||
const uint16_t m, const uint16_t d, int64_t dSize)
|
||||
{
|
||||
__ubuf__ T * finalUb = (__ubuf__ T*)finalTensor.GetPhyAddr();
|
||||
__ubuf__ float * maxUb = (__ubuf__ float*)maxTensor.GetPhyAddr();
|
||||
|
||||
constexpr uint16_t floatRepSize = 64;
|
||||
const uint16_t tailD = d % floatRepSize;
|
||||
uint32_t pltTailD = static_cast<uint32_t>(tailD);
|
||||
uint16_t hasTail = 0;
|
||||
if (tailD > 0) {
|
||||
hasTail = 1;
|
||||
}
|
||||
|
||||
RowInvalidUpdateVF<T>(finalUb, maxUb, m, d, dSize, pltTailD, hasTail);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // MY_FLASH_UPDATE_INTERFACE_H
|
||||
@@ -0,0 +1,164 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vf_mul_sel_softmaxflashv2_cast_nz_sfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
#define MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
|
||||
#include "vf_basic_block_aligned128_no_update_sfa.h"
|
||||
#include "vf_basic_block_aligned128_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned64_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned64_no_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned128_no_update_sfa.h"
|
||||
#include "vf_basic_block_unaligned128_update_sfa.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace FaVectorApi {
|
||||
/* **************************************************************************************************
|
||||
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
|
||||
* ************************************************************************************************* */
|
||||
using AscendC::LocalTensor;
|
||||
|
||||
enum class OriginNRange {
|
||||
EQ_128_SFA = 0, // originN == 128, better performance than GT_64_AND_LTE_128 (s2BaseSize=128)
|
||||
GT_0_AND_LTE_64_SFA, // 0 < originN <= 64 (s2BaseSize <= 64 or tail s2)
|
||||
GT_64_AND_LTE_128_SFA, // 64 < originN <= 128, support for non-alignment (s2BaseSize=128)
|
||||
N_INVALID_SFA
|
||||
};
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1NoUpdate(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
|
||||
ProcessVec1NoUpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
|
||||
ProcessVec1NoUpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
|
||||
ProcessVec1NoUpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1Update(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
if constexpr (oriNRange == OriginNRange::EQ_128_SFA) {
|
||||
ProcessVec1UpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SFA) {
|
||||
ProcessVec1UpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SFA) {
|
||||
ProcessVec1UpdateGeneralImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename T2, bool isUpdate = false, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128,
|
||||
OriginNRange oriNRange = OriginNRange::EQ_128_SFA>
|
||||
__aicore__ inline void ProcessVec1Vf(
|
||||
const LocalTensor<T2>& dstTensor, const LocalTensor<T>& srcTensor,
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor, const LocalTensor<T>& inMaxTensor,
|
||||
const LocalTensor<T>& sharedTmpBuffer, const uint16_t m, const uint32_t originN, const T scale, const T minValue)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF mul_sel_softmaxflashv2_cast_nz, T must be float");
|
||||
static_assert((IsSameType<T2, half>::value || IsSameType<T2, bfloat16_t>::value),
|
||||
"VF mul_sel_softmaxflashv2_cast_nz, T2 must be half or bfloat16");
|
||||
|
||||
if constexpr (!isUpdate) {
|
||||
ProcessVec1NoUpdate<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else {
|
||||
ProcessVec1Update<T, T2, s1BaseSize, s2BaseSize, oriNRange>(
|
||||
dstTensor, srcTensor, expSumTensor, maxTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void UpdateExpSumAndExpMaxVF(__ubuf__ T * maxUb, __ubuf__ T * inMaxUb, __ubuf__ T * expMaxUb,
|
||||
__ubuf__ T * expSumUb, __ubuf__ T * inExpSumUb, __ubuf__ T * tmpExpSumUb, __ubuf__ T * tmpMaxUb, const uint32_t m)
|
||||
{
|
||||
RegTensor<float> vreg_input_x;
|
||||
RegTensor<float> vreg_input_x_unroll;
|
||||
RegTensor<float> vreg_max;
|
||||
RegTensor<float> vreg_in_max;
|
||||
RegTensor<float> vreg_exp_sum;
|
||||
RegTensor<float> vreg_in_exp_sum;
|
||||
RegTensor<float> vreg_exp_max;
|
||||
RegTensor<float> vreg_exp_sum_brc;
|
||||
RegTensor<float> vreg_exp_sum_update;
|
||||
MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
// 注意:当m大于64的时候需要开启循环
|
||||
LoadAlign(vreg_max, tmpMaxUb);
|
||||
LoadAlign(vreg_in_max, inMaxUb);
|
||||
FusedExpSub(vreg_exp_max, vreg_in_max, vreg_max, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)expMaxUb, vreg_exp_max, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)maxUb, vreg_max, preg_all);
|
||||
LoadAlign(vreg_in_exp_sum, inExpSumUb);
|
||||
|
||||
// x_sum = exp_max * insum + x_sum
|
||||
LoadAlign(vreg_exp_sum_brc, tmpExpSumUb);
|
||||
Mul(vreg_exp_sum_update, vreg_exp_max, vreg_in_exp_sum, preg_all);
|
||||
Add(vreg_exp_sum_update, vreg_exp_sum_update, vreg_exp_sum_brc, preg_all);
|
||||
StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)expSumUb, vreg_exp_sum_update, preg_all);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void SFAUpdateExpSumAndExpMax(
|
||||
const LocalTensor<T>& expSumTensor, const LocalTensor<T>& maxTensor,
|
||||
const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& inExpSumTensor,
|
||||
const LocalTensor<T>& inMaxTensor, const LocalTensor<T>& sharedTmpBuffer, const uint32_t m)
|
||||
{
|
||||
__ubuf__ T * maxUb = (__ubuf__ T*)maxTensor.GetPhyAddr();
|
||||
__ubuf__ T * inMaxUb = (__ubuf__ T*)inMaxTensor.GetPhyAddr();
|
||||
|
||||
__ubuf__ T * expMaxUb = (__ubuf__ T*)expMaxTensor.GetPhyAddr();
|
||||
__ubuf__ T * expSumUb = (__ubuf__ T*)expSumTensor.GetPhyAddr();
|
||||
__ubuf__ T * inExpSumUb = (__ubuf__ T*)inExpSumTensor.GetPhyAddr();
|
||||
|
||||
__ubuf__ T * tmpExpSumUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr();
|
||||
__ubuf__ T * tmpMaxUb = (__ubuf__ T*)sharedTmpBuffer.GetPhyAddr() + 64;
|
||||
|
||||
UpdateExpSumAndExpMaxVF<T>(maxUb, inMaxUb, expMaxUb, expSumUb, inExpSumUb, tmpExpSumUb, tmpMaxUb, m);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__simd_vf__ inline void DuplicateSumWithR0VF(__ubuf__ T * sumUb, const T R0, uint32_t m) {
|
||||
AscendC::MicroAPI::RegTensor<T> vreg_sum;
|
||||
AscendC::MicroAPI::MaskReg preg_m = AscendC::MicroAPI::UpdateMask<T>(m);
|
||||
AscendC::MicroAPI::UnalignRegForStore ureg;
|
||||
AscendC::MicroAPI::Duplicate<T, MicroAPI::MaskMergeMode::ZEROING, T>(vreg_sum, R0, preg_m);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(sumUb, vreg_sum, preg_m);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void DuplicateSumWithR0(const LocalTensor<T>& sumTensor, const T R0, uint32_t m)
|
||||
{
|
||||
__ubuf__ T * sumUb = (__ubuf__ T*)sumTensor.GetPhyAddr();
|
||||
DuplicateSumWithR0VF<T>(sumUb, R0, m);
|
||||
}
|
||||
} // namespace
|
||||
#endif // MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SFA_INTERFACE_H
|
||||
292
csrc/attention/common/op_kernel/buffer.h
Normal file
292
csrc/attention/common/op_kernel/buffer.h
Normal file
@@ -0,0 +1,292 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file buffer.h
|
||||
* \brief同步管理
|
||||
*/
|
||||
#ifndef BUFFER_H
|
||||
#define BUFFER_H
|
||||
#include<type_traits>
|
||||
#include"lib/matmul_intf.h"
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_basic_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
using namespace AscendC;
|
||||
namespace fa_base_matmul {
|
||||
__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum;
|
||||
#define MAKE_ID ((++idCounterNum) % 11)
|
||||
|
||||
// 核间同步中,AIC(flagId 0-10)对应AIV0(flagId 0-10),对应AIV1(flagId 16-26)
|
||||
#define AIV0_AIV1_OFFSET 16
|
||||
|
||||
enum class BufferType {
|
||||
L1 = 0,
|
||||
L0A = 1,
|
||||
L0B = 2,
|
||||
L0C = 3,
|
||||
UB = 4,
|
||||
GM = 5,
|
||||
C2 = 6,
|
||||
};
|
||||
|
||||
enum class SyncType {
|
||||
NO_SYNC,
|
||||
INNER_CORE_SYNC,
|
||||
CROSS_CORE_SYNC_FORWARD,
|
||||
CROSS_CORE_SYNC_BOTH,
|
||||
CROSS_CORE_SYNC_BACKWARD,
|
||||
};
|
||||
|
||||
constexpr uint32_t INVALID_CROSS_CORE_EVENT_ID = 16;
|
||||
static constexpr uint64_t CROSS_CORE_SYNC_MODE = 4;
|
||||
|
||||
template<BufferType Type>
|
||||
struct BufferInfo{
|
||||
// Cons 消费者,Prod 生产者
|
||||
__aicore__ const static constexpr HardEvent ConsWaitProdStatus() {
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return HardEvent::MTE2_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return HardEvent::MTE1_M;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return HardEvent::MTE1_M;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return HardEvent::M_FIX;
|
||||
} else if constexpr (Type == BufferType::C2) {
|
||||
return HardEvent::MTE1_M;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ const static constexpr HardEvent ProdWaitConsStatus() {
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return HardEvent::MTE1_MTE2;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return HardEvent::M_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return HardEvent::M_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return HardEvent::FIX_M;
|
||||
} else if constexpr (Type == BufferType::C2) {
|
||||
return HardEvent::M_MTE1;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ const static constexpr TPosition GetTPosition() {
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return TPosition::A1;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return TPosition::A2;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return TPosition::B2;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return TPosition::CO1;
|
||||
} else if constexpr (Type == BufferType::UB) {
|
||||
return TPosition::VECIN;
|
||||
} else if constexpr (Type == BufferType::GM) {
|
||||
return TPosition::GM;
|
||||
} else if constexpr (Type == BufferType::C2) {
|
||||
return TPosition::C2;
|
||||
}
|
||||
}
|
||||
|
||||
static constexpr HardEvent EventP2C = ConsWaitProdStatus(); // 生产者到消费者方向的HardEvent:消费者等生产者提供/生产者通知消费者已生成
|
||||
static constexpr HardEvent EventC2P = ProdWaitConsStatus(); // 消费者到生产者方向的HardEvent:生产者等消费者消耗/消费者通知生产者已消耗’
|
||||
static constexpr TPosition Position = GetTPosition();
|
||||
};
|
||||
|
||||
// buffer绑定生产者、消费者关系
|
||||
// L1 buffer的生产者为MTE2或者MTE3,消费者为MTE1
|
||||
// L0A buffer的生产者为MTE1,消费者为M
|
||||
// L0B buffer的生产者为MTE1,消费者为M
|
||||
// L0C buffer的生产者为M,消费者为FIX
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class Buffer {
|
||||
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
|
||||
|
||||
template <typename T>
|
||||
using TargetTensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<T>, LocalTensor<T>>;
|
||||
public:
|
||||
__aicore__ inline Buffer() {}
|
||||
__aicore__ inline Buffer(TensorType tensor, uint32_t size) {
|
||||
tensor_ = tensor;
|
||||
size_ = size;
|
||||
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_FORWARD) {
|
||||
id0_ = MAKE_ID;
|
||||
id1_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
|
||||
id0_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
id1_ = MAKE_ID;
|
||||
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
|
||||
id0_ = MAKE_ID;
|
||||
id1_ = MAKE_ID;
|
||||
} else {
|
||||
id0_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
id1_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Init() {
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
|
||||
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
|
||||
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void UnInit() {
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 确保只能被调用一次
|
||||
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<HardEvent EventType>
|
||||
__aicore__ inline void Wait() {
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
WaitFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 消费者等待生产者完成生产
|
||||
} else {
|
||||
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 生产者等待消费者完成消费
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<HardEvent EventType>
|
||||
__aicore__ inline void Set() {
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
SetFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 生产者通知消费者已完成生产
|
||||
} else {
|
||||
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 消费者通知生产者已完成消费
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SetEventID() {
|
||||
if ASCEND_IS_AIC {
|
||||
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
|
||||
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
|
||||
}
|
||||
}
|
||||
|
||||
template<HardEvent EventType>
|
||||
__aicore__ inline TEventID GetEventID() {
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
return p2cEventId_; // 生产者通知消费者已完成生产
|
||||
} else {
|
||||
return c2pEventId_; // 消费者通知生产者已完成消费
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<bool isReuse = false>
|
||||
__aicore__ inline void WaitCrossCore() {
|
||||
if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
|
||||
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id1_);
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id1_ + AIV0_AIV1_OFFSET);
|
||||
} else {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE2>(id0_);
|
||||
}
|
||||
} else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) {
|
||||
// AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id1_);
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id1_ + AIV0_AIV1_OFFSET);
|
||||
} else {
|
||||
if constexpr (isReuse) {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id0_);
|
||||
} else {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_V>(id0_);
|
||||
}
|
||||
}
|
||||
} else if constexpr (bufferType == BufferType::L1) {
|
||||
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id0_);
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id0_ + AIV0_AIV1_OFFSET);
|
||||
} else {
|
||||
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
|
||||
CrossCoreWaitFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<bool isReuse = false>
|
||||
__aicore__ inline void SetCrossCore() {
|
||||
if constexpr (bufferType == BufferType::GM && syncType == SyncType::CROSS_CORE_SYNC_BACKWARD) {
|
||||
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_);
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_ + AIV0_AIV1_OFFSET);
|
||||
} else {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
|
||||
}
|
||||
} else if constexpr (bufferType == BufferType::UB || bufferType == BufferType::GM) {
|
||||
// AIC属于生产者,AIV属于消费者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_);
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_FIX>(id0_ + AIV0_AIV1_OFFSET);
|
||||
} else {
|
||||
if constexpr (isReuse) {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id1_);
|
||||
} else {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_V>(id1_);
|
||||
}
|
||||
}
|
||||
} else if constexpr (bufferType == BufferType::L1) {
|
||||
// AIC属于消费者,AIV属于生产者,且一个AIC对应两个AIV
|
||||
if ASCEND_IS_AIC {
|
||||
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id1_);
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE1>(id1_ + AIV0_AIV1_OFFSET);
|
||||
}
|
||||
} else {
|
||||
CrossCoreSetFlag<CROSS_CORE_SYNC_MODE, PIPE_MTE3>(id0_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__aicore__ inline TargetTensorType<T> GetTensor() {
|
||||
return tensor_.template ReinterpretCast<T>();
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__aicore__ inline TargetTensorType<T> GetTensor(uint64_t startindex) {
|
||||
TargetTensorType<T> tmpTensor = tensor_.template ReinterpretCast<T>();
|
||||
return tmpTensor[startindex];
|
||||
}
|
||||
|
||||
private:
|
||||
TensorType tensor_;
|
||||
uint32_t size_;
|
||||
TEventID p2cEventId_;
|
||||
TEventID c2pEventId_;
|
||||
uint32_t id0_; // 用作正向同步:生产者通知消费者,或者消费者等待生产者;
|
||||
uint32_t id1_; // 用作反向同步:消费者通知生产者,或者生产者等待消费者;
|
||||
};
|
||||
}
|
||||
#endif
|
||||
61
csrc/attention/common/op_kernel/buffer_manager.h
Normal file
61
csrc/attention/common/op_kernel/buffer_manager.h
Normal file
@@ -0,0 +1,61 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file buffer_manager.h
|
||||
* \brief buffer内存管理
|
||||
*/
|
||||
#ifndef BUFFER_MANAGER_H
|
||||
#define BUFFER_MANAGER_H
|
||||
|
||||
#if (__NPU_ARCH__ == 5102)
|
||||
#include "buffer_mix_core.h"
|
||||
#else
|
||||
#include "buffer.h"
|
||||
#endif
|
||||
|
||||
// L1 TPosition::A1
|
||||
// L0A TPosition::A2
|
||||
// L0B TPosition::B2
|
||||
// L0C TPosition::CO1
|
||||
// UB TPosition::VECIN
|
||||
namespace fa_base_matmul {
|
||||
template<BufferType bufferType>
|
||||
class BufferManager {
|
||||
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
|
||||
public:
|
||||
__aicore__ inline void Init(TPipe *pipe, uint32_t size) {
|
||||
static_assert(bufferType != BufferType::GM, "GM should use workspace.");
|
||||
TBuf<BufferInfo<bufferType>::Position> tbuf;
|
||||
pipe->InitBuffer(tbuf, size);
|
||||
mem_ = tbuf.template Get<uint8_t>();
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(__gm__ uint8_t* workspace) {
|
||||
static_assert(bufferType == BufferType::GM, "BufferType should be GM.");
|
||||
mem_.SetGlobalBuffer((__gm__ uint8_t*)workspace);
|
||||
}
|
||||
|
||||
template<SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
__aicore__ inline Buffer<bufferType, syncType> AllocBuffer(uint32_t size) {
|
||||
TensorType temp = mem_[offset_];
|
||||
offset_ += size;
|
||||
return Buffer<bufferType, syncType>(temp, size);
|
||||
}
|
||||
|
||||
template<SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
__aicore__ inline void FreeBuffer(Buffer<bufferType, syncType> &buffer){
|
||||
}
|
||||
private:
|
||||
uint32_t offset_ = 0;
|
||||
TensorType mem_;
|
||||
};
|
||||
}
|
||||
#endif
|
||||
216
csrc/attention/common/op_kernel/buffer_mix_core.h
Normal file
216
csrc/attention/common/op_kernel/buffer_mix_core.h
Normal file
@@ -0,0 +1,216 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file buffer_mix_core.h
|
||||
* \brief同步管理
|
||||
*/
|
||||
#ifndef BUFFER_MIX_CORE_H
|
||||
#define BUFFER_MIX_CORE_H
|
||||
#include <type_traits>
|
||||
#include "lib/matmul_intf.h"
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_basic_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
using namespace AscendC;
|
||||
namespace fa_base_matmul {
|
||||
__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum;
|
||||
#define MAKE_ID ((++idCounterNum) % 11)
|
||||
|
||||
// 核间同步中,AIC(flagId 0-10)对应AIV0(flagId 0-10),对应AIV1(flagId 16-26)
|
||||
#define AIV0_AIV1_OFFSET 16
|
||||
|
||||
enum class BufferType {
|
||||
L1 = 0,
|
||||
L0A = 1,
|
||||
L0B = 2,
|
||||
L0C = 3,
|
||||
UB = 4,
|
||||
GM = 5,
|
||||
};
|
||||
|
||||
enum class SyncType {
|
||||
NO_SYNC,
|
||||
INNER_CORE_SYNC,
|
||||
CROSS_CORE_SYNC_FORWARD,
|
||||
CROSS_CORE_SYNC_BOTH,
|
||||
};
|
||||
|
||||
constexpr uint32_t INVALID_CROSS_CORE_EVENT_ID = 16;
|
||||
static constexpr uint64_t CROSS_CORE_SYNC_MODE = 4;
|
||||
|
||||
template <BufferType Type>
|
||||
struct BufferInfo {
|
||||
// Cons 消费者,Prod 生产者
|
||||
__aicore__ const static constexpr HardEvent ConsWaitProdStatus()
|
||||
{
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return HardEvent::MTE2_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return HardEvent::MTE1_M;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return HardEvent::MTE1_M;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return HardEvent::M_FIX;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ const static constexpr HardEvent ProdWaitConsStatus()
|
||||
{
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return HardEvent::MTE1_MTE2;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return HardEvent::M_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return HardEvent::M_MTE1;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return HardEvent::FIX_M;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ const static constexpr TPosition GetTPosition()
|
||||
{
|
||||
if constexpr (Type == BufferType::L1) {
|
||||
return TPosition::A1;
|
||||
} else if constexpr (Type == BufferType::L0A) {
|
||||
return TPosition::A2;
|
||||
} else if constexpr (Type == BufferType::L0B) {
|
||||
return TPosition::B2;
|
||||
} else if constexpr (Type == BufferType::L0C) {
|
||||
return TPosition::CO1;
|
||||
} else if constexpr (Type == BufferType::UB) {
|
||||
return TPosition::VECIN;
|
||||
} else if constexpr (Type == BufferType::GM) {
|
||||
return TPosition::GM;
|
||||
}
|
||||
}
|
||||
|
||||
static constexpr HardEvent EventP2C =
|
||||
ConsWaitProdStatus(); // 生产者到消费者方向的HardEvent:消费者等生产者提供/生产者通知消费者已生成
|
||||
static constexpr HardEvent EventC2P =
|
||||
ProdWaitConsStatus(); // 消费者到生产者方向的HardEvent:生产者等消费者消耗/消费者通知生产者已消耗’
|
||||
static constexpr TPosition Position = GetTPosition();
|
||||
};
|
||||
|
||||
// buffer绑定生产者、消费者关系
|
||||
// L1 buffer的生产者为MTE2或者MTE3,消费者为MTE1
|
||||
// L0A buffer的生产者为MTE1,消费者为M
|
||||
// L0B buffer的生产者为MTE1,消费者为M
|
||||
// L0C buffer的生产者为M,消费者为FIX
|
||||
template <BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class Buffer {
|
||||
using TensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<uint8_t>, LocalTensor<uint8_t>>;
|
||||
|
||||
template <typename T>
|
||||
using TargetTensorType = std::conditional_t<bufferType == BufferType::GM, GlobalTensor<T>, LocalTensor<T>>;
|
||||
|
||||
public:
|
||||
__aicore__ inline Buffer()
|
||||
{
|
||||
}
|
||||
__aicore__ inline Buffer(TensorType tensor, uint32_t size)
|
||||
{
|
||||
tensor_ = tensor;
|
||||
size_ = size;
|
||||
if constexpr (syncType == SyncType::CROSS_CORE_SYNC_FORWARD) {
|
||||
id0_ = MAKE_ID;
|
||||
id1_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
} else if constexpr (syncType == SyncType::CROSS_CORE_SYNC_BOTH) {
|
||||
id0_ = MAKE_ID;
|
||||
id1_ = MAKE_ID;
|
||||
} else {
|
||||
id0_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
id1_ = INVALID_CROSS_CORE_EVENT_ID;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Init()
|
||||
{
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
|
||||
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
|
||||
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void UnInit()
|
||||
{
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 确保只能被调用一次
|
||||
GetTPipePtr()->ReleaseEventID<BufferInfo<bufferType>::EventC2P>(c2pEventId_);
|
||||
}
|
||||
}
|
||||
|
||||
template <HardEvent EventType>
|
||||
__aicore__ inline void Wait()
|
||||
{
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
WaitFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 消费者等待生产者完成生产
|
||||
} else {
|
||||
WaitFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 生产者等待消费者完成消费
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <HardEvent EventType>
|
||||
__aicore__ inline void Set()
|
||||
{
|
||||
if constexpr (syncType == SyncType::INNER_CORE_SYNC) {
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
SetFlag<BufferInfo<bufferType>::EventP2C>(p2cEventId_); // 生产者通知消费者已完成生产
|
||||
} else {
|
||||
SetFlag<BufferInfo<bufferType>::EventC2P>(c2pEventId_); // 消费者通知生产者已完成消费
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void SetEventID()
|
||||
{
|
||||
p2cEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventP2C>(); // 确保只能被调用一次
|
||||
c2pEventId_ = GetTPipePtr()->AllocEventID<BufferInfo<bufferType>::EventC2P>();
|
||||
}
|
||||
|
||||
template <HardEvent EventType>
|
||||
__aicore__ inline TEventID GetEventID()
|
||||
{
|
||||
if constexpr (EventType == BufferInfo<bufferType>::EventP2C) {
|
||||
return p2cEventId_; // 生产者通知消费者已完成生产
|
||||
} else {
|
||||
return c2pEventId_; // 消费者通知生产者已完成消费
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline TargetTensorType<T> GetTensor()
|
||||
{
|
||||
return tensor_.template ReinterpretCast<T>();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline TargetTensorType<T> GetTensor(uint64_t startindex)
|
||||
{
|
||||
TargetTensorType<T> tmpTensor = tensor_.template ReinterpretCast<T>();
|
||||
return tmpTensor[startindex];
|
||||
}
|
||||
|
||||
private:
|
||||
TensorType tensor_;
|
||||
uint32_t size_;
|
||||
TEventID p2cEventId_;
|
||||
TEventID c2pEventId_;
|
||||
uint32_t id0_; // 用作正向同步:生产者通知消费者,或者消费者等待生产者;
|
||||
uint32_t id1_; // 用作反向同步:消费者通知生产者,或者生产者等待消费者;
|
||||
};
|
||||
} // namespace fa_base_matmul
|
||||
#endif
|
||||
407
csrc/attention/common/op_kernel/buffers_policy.h
Normal file
407
csrc/attention/common/op_kernel/buffers_policy.h
Normal file
@@ -0,0 +1,407 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file buffers_policy.h
|
||||
* \brief 综合管理buffer的内存和同步
|
||||
*/
|
||||
#ifndef BUFFERS_POLICY_H
|
||||
#define BUFFERS_POLICY_H
|
||||
|
||||
#include "buffer_manager.h"
|
||||
#define NUM_2 2
|
||||
#define NUM_3 3
|
||||
#define NUM_4 4
|
||||
// Q复用 KV复用
|
||||
// 申请单块buffer
|
||||
namespace fa_base_matmul {
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class BuffersPolicySingleBuffer {
|
||||
public:
|
||||
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size){
|
||||
buffer_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
buffer_.Init();
|
||||
}
|
||||
|
||||
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager){
|
||||
buffer_.UnInit();
|
||||
bufferManager.FreeBuffer(buffer_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &Get(){
|
||||
return buffer_;
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetPre(){
|
||||
return Get();
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetReused(){
|
||||
return Get();
|
||||
}
|
||||
private:
|
||||
Buffer<bufferType, syncType> buffer_;
|
||||
};
|
||||
|
||||
// 申请2个buffer,乒乓轮转
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class BuffersPolicyDB {
|
||||
public:
|
||||
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size){
|
||||
ping_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
pong_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
|
||||
ping_.Init();
|
||||
pong_.Init();
|
||||
}
|
||||
|
||||
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager){
|
||||
ping_.UnInit();
|
||||
pong_.UnInit();
|
||||
|
||||
bufferManager.FreeBuffer(ping_);
|
||||
bufferManager.FreeBuffer(pong_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &Get() {
|
||||
if (flag1_) { // 1
|
||||
flag1_ = 0;
|
||||
return ping_;
|
||||
} else { // 0
|
||||
flag1_ = 1;
|
||||
return pong_;
|
||||
}
|
||||
}
|
||||
|
||||
// 需要与Get联用, 首次调用Get,第二次调用GetPre(Q复用)
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetPre() {
|
||||
if (flag1_) { // 0->1
|
||||
return pong_;
|
||||
} else { // 1->0
|
||||
return ping_;
|
||||
}
|
||||
}
|
||||
|
||||
// 需要与Get,GetPre联用, 首次调用Get,第二次调用GetPre,第三次复用时GetReused(KV复用)
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
|
||||
if (flag2_ == 0) {
|
||||
flag2_ = 1;
|
||||
return pong_;
|
||||
} else {
|
||||
flag2_ = 0;
|
||||
return ping_;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetReused(bool isNextS2IdxNoChange) {
|
||||
if (isNextS2IdxNoChange) {
|
||||
if (flag2_ == 0) {
|
||||
return pong_;
|
||||
} else {
|
||||
return ping_;
|
||||
}
|
||||
} else {
|
||||
return GetReused();
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Buffer<bufferType, syncType> ping_;
|
||||
Buffer<bufferType, syncType> pong_;
|
||||
uint32_t flag1_ = 0;
|
||||
uint32_t flag2_ = 0;
|
||||
};
|
||||
|
||||
// 申请3个buffer, 轮转
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class BuffersPolicy3buff {
|
||||
public:
|
||||
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
|
||||
a_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
b_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
c_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
|
||||
a_.Init();
|
||||
b_.Init();
|
||||
c_.Init();
|
||||
}
|
||||
|
||||
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
|
||||
a_.UnInit();
|
||||
b_.UnInit();
|
||||
c_.UnInit();
|
||||
|
||||
bufferManager.FreeBuffer(a_);
|
||||
bufferManager.FreeBuffer(b_);
|
||||
bufferManager.FreeBuffer(c_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &Get() {
|
||||
if (flag1_ == 0) {
|
||||
flag1_ = 1;
|
||||
return a_;
|
||||
} else if (flag1_ == 1) {
|
||||
flag1_ = NUM_2;
|
||||
return b_;
|
||||
} else {
|
||||
flag1_ = 0;
|
||||
return c_;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetVec() { // mixcore architecture
|
||||
if (flag1_vec1_ == 0) {
|
||||
flag1_vec1_ = 1;
|
||||
return a_;
|
||||
} else if (flag1_vec1_ == 1) {
|
||||
flag1_vec1_ = NUM_2;
|
||||
return b_;
|
||||
} else {
|
||||
flag1_vec1_ = 0;
|
||||
return c_;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetCube() { // mixcore architecture
|
||||
if (flag1_bmm2_ == 0) {
|
||||
flag1_bmm2_ = 1;
|
||||
return a_;
|
||||
} else if (flag1_bmm2_ == 1) {
|
||||
flag1_bmm2_ = NUM_2;
|
||||
return b_;
|
||||
} else {
|
||||
flag1_bmm2_ = 0;
|
||||
return c_;
|
||||
}
|
||||
}
|
||||
|
||||
// Q复用
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetPre() {
|
||||
if (flag1_ == 0) {
|
||||
return c_;
|
||||
} else if (flag1_ == 1) {
|
||||
return a_;
|
||||
} else {
|
||||
return b_;
|
||||
}
|
||||
}
|
||||
|
||||
// KV复用
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
|
||||
if (flag2_ == 0) {
|
||||
flag2_ = 1;
|
||||
return a_;
|
||||
} else if (flag2_ == 1){
|
||||
flag2_ = NUM_2;
|
||||
return b_;
|
||||
} else {
|
||||
flag2_ = 0;
|
||||
return c_;
|
||||
}
|
||||
}
|
||||
private:
|
||||
Buffer<bufferType, syncType> a_;
|
||||
Buffer<bufferType, syncType> b_;
|
||||
Buffer<bufferType, syncType> c_;
|
||||
uint32_t flag1_ = 0;
|
||||
uint32_t flag1_vec1_ = 0;
|
||||
uint32_t flag1_bmm2_ = 0;
|
||||
uint32_t flag2_ = 0;
|
||||
};
|
||||
|
||||
// 申请4个buffer + kv复用
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class BuffersPolicy4buff {
|
||||
public:
|
||||
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
|
||||
a_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
b_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
c_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
d_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
|
||||
a_.Init();
|
||||
b_.Init();
|
||||
c_.Init();
|
||||
d_.Init();
|
||||
}
|
||||
|
||||
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
|
||||
a_.UnInit();
|
||||
b_.UnInit();
|
||||
c_.UnInit();
|
||||
d_.UnInit();
|
||||
|
||||
bufferManager.FreeBuffer(a_);
|
||||
bufferManager.FreeBuffer(b_);
|
||||
bufferManager.FreeBuffer(c_);
|
||||
bufferManager.FreeBuffer(d_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &Get(uint32_t id) {
|
||||
uint32_t flag = id % 4;
|
||||
if (flag == 0) {
|
||||
return a_;
|
||||
} else if (flag == 1) {
|
||||
return b_;
|
||||
} else if (flag == 2) { // 2:c_
|
||||
return c_;
|
||||
} else {
|
||||
return d_;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &Get() {
|
||||
auto& buffer = Get(head_);
|
||||
head_++;
|
||||
return buffer;
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetReused() {
|
||||
auto& buffer = Get(used_);
|
||||
used_ = (used_ - tail_ + 1) % (head_ - tail_) + tail_;
|
||||
return buffer;
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetFree() {
|
||||
if (tail_ == used_) {
|
||||
used_++;
|
||||
}
|
||||
auto& buffer = Get(tail_);
|
||||
tail_++;
|
||||
return buffer;
|
||||
}
|
||||
private:
|
||||
Buffer<bufferType, syncType> a_;
|
||||
Buffer<bufferType, syncType> b_;
|
||||
Buffer<bufferType, syncType> c_;
|
||||
Buffer<bufferType, syncType> d_;
|
||||
uint32_t tail_ = 0; // 表示当前正在使用的buffer队列队尾
|
||||
uint32_t head_ = 0; // 表示当前正在使用的buffer队列队首+1
|
||||
uint32_t used_ = 0; // 表示当前正在使用的buffer,于首尾间,左闭右开
|
||||
};
|
||||
|
||||
template<BufferType bufferType, SyncType syncType = SyncType::INNER_CORE_SYNC>
|
||||
class Matrix2x2BufferPolicy { // 4buffer
|
||||
// 二维buffer管理,地址行优先,使用列优先
|
||||
// MracBuffer:memory address with row first, alloc/use/free with column first
|
||||
public:
|
||||
__aicore__ inline void Init(BufferManager<bufferType> &bufferManager, uint32_t size) {
|
||||
bufferM0k0_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
bufferM0k1_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
bufferM1k0_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
bufferM1k1_ = bufferManager.template AllocBuffer<syncType>(size);
|
||||
|
||||
bufferM0k0_.Init();
|
||||
bufferM0k1_.Init();
|
||||
bufferM1k0_.Init();
|
||||
bufferM1k1_.Init();
|
||||
}
|
||||
|
||||
__aicore__ inline void Uninit(BufferManager<bufferType> &bufferManager) {
|
||||
bufferM0k0_.UnInit();
|
||||
bufferM0k1_.UnInit();
|
||||
bufferM1k0_.UnInit();
|
||||
bufferM1k1_.UnInit();
|
||||
|
||||
bufferManager.FreeBuffer(bufferM0k0_);
|
||||
bufferManager.FreeBuffer(bufferM0k1_);
|
||||
bufferManager.FreeBuffer(bufferM1k0_);
|
||||
bufferManager.FreeBuffer(bufferM1k1_);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetMExtent(int32_t mExtent) {
|
||||
aIdx_ = -1;
|
||||
amIdx_ = (amIdx_ + mSize_ - 1) % mSize_; // 翻转 0->1, 1->0
|
||||
akIdx_ = 0;
|
||||
|
||||
uIdx_ = -1;
|
||||
umIdx_ = (umIdx_ + mSize_ - 1) % mSize_;
|
||||
ukIdx_ = 0;
|
||||
|
||||
fIdx_ = -1;
|
||||
fmIdx_ = (fmIdx_ + mSize_ - 1) % mSize_;
|
||||
fkIdx_ = 0;
|
||||
|
||||
mExtent_ = mExtent;
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &AllocNext() {
|
||||
aIdx_++;
|
||||
return GetBuffer(aIdx_, amIdx_, akIdx_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &ReuseNext() {
|
||||
uIdx_++;
|
||||
return GetBuffer(uIdx_, umIdx_, ukIdx_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &FreeNext() {
|
||||
fIdx_++;
|
||||
return GetBuffer(fIdx_, fmIdx_, fkIdx_);
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &PeekNextK() { // 在Alloc阶段使用,k方向取下一个
|
||||
return PeekBuffer(amIdx_, (1 - akIdx_)); // k翻转
|
||||
}
|
||||
private:
|
||||
__aicore__ inline Buffer<bufferType, syncType> &GetBuffer(int32_t xIdx, int32_t &mIdx, int32_t &kIdx) {
|
||||
// xIdx为入参,表示当前alloc/use/free的idx,mIdx和kIdx为下标出参,移动到下一个buffer并获取
|
||||
mIdx = (mIdx + mExtent_ - 1) % mExtent_;
|
||||
kIdx = (xIdx / mExtent_) % kSize_;
|
||||
if (mIdx == 0 && kIdx == 0) {
|
||||
return bufferM0k0_;
|
||||
} else if (mIdx == 0 && kIdx == 1) {
|
||||
return bufferM0k1_;
|
||||
} else if (mIdx == 1 && kIdx == 0) {
|
||||
return bufferM1k0_;
|
||||
} else { // 该分支条件为:mIdx == 1 && kIdx == 1
|
||||
return bufferM1k1_;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline Buffer<bufferType, syncType> &PeekBuffer(int32_t mIdx, int32_t kIdx) {
|
||||
// 只访问buffer,不进行下标移动
|
||||
if (mIdx == 0 && kIdx == 0) {
|
||||
return bufferM0k0_;
|
||||
} else if (mIdx == 0 && kIdx == 1) {
|
||||
return bufferM0k1_;
|
||||
} else if ((mIdx == 1) && (kIdx == 0)) {
|
||||
return bufferM1k0_;
|
||||
} else { // mIdx == 1 && kIdx == 1
|
||||
return bufferM1k1_;
|
||||
}
|
||||
}
|
||||
|
||||
Buffer<bufferType, syncType> bufferM0k0_;
|
||||
Buffer<bufferType, syncType> bufferM0k1_;
|
||||
Buffer<bufferType, syncType> bufferM1k0_;
|
||||
Buffer<bufferType, syncType> bufferM1k1_;
|
||||
int32_t mSize_ = 2; // m的总buffer数
|
||||
int32_t kSize_ = 2; // k的总buffer数
|
||||
|
||||
// Alloc
|
||||
int32_t aIdx_ = -1; // 当前第几次Alloc Buffer
|
||||
int32_t amIdx_ = 0; // 当前Alloc Buffer的m下标
|
||||
int32_t akIdx_ = 0; // 当前Alloc Buffer的k下标
|
||||
|
||||
// Reuse
|
||||
int32_t uIdx_ = -1;
|
||||
int32_t umIdx_ = 0;
|
||||
int32_t ukIdx_ = 0;
|
||||
|
||||
// Free
|
||||
int32_t fIdx_ = -1;
|
||||
int32_t fmIdx_ = 0;
|
||||
int32_t fkIdx_ = 0;
|
||||
|
||||
int32_t mExtent_ = 0; // m实际使用的大小,可以为1或者2
|
||||
};
|
||||
}
|
||||
#endif
|
||||
1158
csrc/attention/common/op_kernel/matmul.h
Normal file
1158
csrc/attention/common/op_kernel/matmul.h
Normal file
File diff suppressed because it is too large
Load Diff
35
csrc/attention/common/op_kernel/memcopy/fa_gm_tensor.h
Normal file
35
csrc/attention/common/op_kernel/memcopy/fa_gm_tensor.h
Normal file
@@ -0,0 +1,35 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file fa_gm_tensor.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef FA_GM_TENSOR_H
|
||||
#define FA_GM_TENSOR_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_vec_intf.h"
|
||||
#include "kernel_cube_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
#include "gm_layout.h"
|
||||
#include "offset_calculator_v2.h"
|
||||
|
||||
using AscendC::GlobalTensor;
|
||||
|
||||
template <typename Q_T, GmFormat FORMAT, typename ACTLEN_T = uint64_t>
|
||||
struct FaGmTensor {
|
||||
GlobalTensor<Q_T> gmTensor;
|
||||
OffsetCalculator<FORMAT, ACTLEN_T> offsetCalculator;
|
||||
};
|
||||
|
||||
#endif
|
||||
43
csrc/attention/common/op_kernel/memcopy/fa_l1_tensor.h
Normal file
43
csrc/attention/common/op_kernel/memcopy/fa_l1_tensor.h
Normal file
@@ -0,0 +1,43 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file fa_l1_tensor.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef FA_L1_TENSOR_H
|
||||
#define FA_L1_TENSOR_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_vec_intf.h"
|
||||
#include "kernel_cube_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
using AscendC::LocalTensor;
|
||||
|
||||
enum class L1Format {
|
||||
NZ = 0
|
||||
};
|
||||
|
||||
enum class ScaleTrans {
|
||||
NO_TRANS = 0,
|
||||
ND2NZ = 1,
|
||||
DN2NZ = 2
|
||||
};
|
||||
|
||||
template <typename Q_T, L1Format FORMAT>
|
||||
struct FaL1Tensor {
|
||||
LocalTensor<Q_T> tensor;
|
||||
uint32_t rowCount;
|
||||
};
|
||||
|
||||
#endif
|
||||
26
csrc/attention/common/op_kernel/memcopy/gm_coord.h
Normal file
26
csrc/attention/common/op_kernel/memcopy/gm_coord.h
Normal file
@@ -0,0 +1,26 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file gm_coord.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef GM_COORD_H
|
||||
#define GM_COORD_H
|
||||
|
||||
struct GmCoord {
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx;
|
||||
uint32_t gS1Idx;
|
||||
uint32_t dIdx;
|
||||
uint32_t gS1DealSize;
|
||||
uint32_t dDealSize;
|
||||
};
|
||||
#endif
|
||||
427
csrc/attention/common/op_kernel/memcopy/gm_layout.h
Normal file
427
csrc/attention/common/op_kernel/memcopy/gm_layout.h
Normal file
@@ -0,0 +1,427 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file gm_layout.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef GM_LAYOUT_H
|
||||
#define GM_LAYOUT_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_vec_intf.h"
|
||||
#include "kernel_cube_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
// ----------------------------------------------GmLayout--------------------------------
|
||||
enum class GmFormat {
|
||||
BSNGD = 0,
|
||||
BNGSD = 1,
|
||||
NGBSD = 2,
|
||||
TNGD = 3,
|
||||
NGTD = 4,
|
||||
BSND = 5,
|
||||
BNSD = 6,
|
||||
TND = 7,
|
||||
NTD = 8,
|
||||
PA_BnBsND = 9,
|
||||
PA_BnNBsD = 10,
|
||||
PA_NZ = 11,
|
||||
NGD = 12, // post_quant
|
||||
ND = 13, //antiquant no PA
|
||||
BS2 = 14,
|
||||
BNS2 = 15,
|
||||
PA_BnBs = 16, //antiquant PA
|
||||
PA_BnNBs = 17,
|
||||
BN2GS1S2 = 18, //PSE_GmFormat
|
||||
SBNGD = 19,
|
||||
SBND = 20,
|
||||
NTGD = 21,
|
||||
PA_NZ_K_SCALE = 22,
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct GmLayout {
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BSNGD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, g, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t gStride = dStride * d;
|
||||
uint64_t nStride = gStride * g;
|
||||
uint64_t sStride = nStride * n;
|
||||
uint64_t bStride = sStride * s;
|
||||
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BNGSD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, g, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t sStride = dStride * d;
|
||||
uint64_t gStride = sStride * s;
|
||||
uint64_t nStride = gStride * g;
|
||||
uint64_t bStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::NGBSD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, g, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t sStride = dStride * d;
|
||||
uint64_t bStride = sStride * s;
|
||||
uint64_t gStride = bStride * b;
|
||||
uint64_t nStride = gStride * g;
|
||||
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::TNGD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
|
||||
shape = AscendC::MakeShape(t, n, g, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t gStride = dStride * d;
|
||||
uint64_t nStride = gStride * g;
|
||||
uint64_t tStride = nStride * n;
|
||||
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::NGTD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
|
||||
shape = AscendC::MakeShape(t, n, g, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t tStride = dStride * d;
|
||||
uint64_t gStride = tStride * t;
|
||||
uint64_t nStride = gStride * g;
|
||||
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::NTGD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t g, uint32_t d) {
|
||||
shape = AscendC::MakeShape(t, n, g, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t gStride = dStride * d;
|
||||
uint64_t tStride = gStride * g;
|
||||
uint64_t nStride = tStride * t;
|
||||
stride = AscendC::MakeStride(tStride, nStride, gStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BSND> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t nStride = dStride * d;
|
||||
uint64_t sStride = nStride * n;
|
||||
uint64_t bStride = sStride * s;
|
||||
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BNSD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t sStride = dStride * d;
|
||||
uint64_t nStride = sStride * s;
|
||||
uint64_t bStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::TND> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t d) {
|
||||
shape = AscendC::MakeShape(t, n, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t nStride = dStride * d;
|
||||
uint64_t tStride = nStride * n;
|
||||
stride = AscendC::MakeStride(tStride, nStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::NTD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t t, uint32_t n, uint32_t d) {
|
||||
shape = AscendC::MakeShape(t, n, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t tStride = dStride * d;
|
||||
uint64_t nStride = tStride * t;
|
||||
stride = AscendC::MakeStride(tStride, nStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_BnBsND> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d) {
|
||||
shape = AscendC::MakeShape(n, blockSize, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t nStride = dStride * d;
|
||||
uint64_t bsStride = nStride * n;
|
||||
uint64_t bnStride = bsStride * blockSize;
|
||||
stride = AscendC::MakeStride(bnStride, nStride, bsStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_BnNBsD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d) {
|
||||
shape = AscendC::MakeShape(n, blockSize, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t bsStride = dStride * d;
|
||||
uint64_t nStride = bsStride * blockSize;
|
||||
uint64_t bnStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bnStride, nStride, bsStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_NZ> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize, uint32_t d1, uint32_t d0) {
|
||||
shape = AscendC::MakeShape(n, d1, blockSize, d0);
|
||||
uint64_t d0Stride = 1;
|
||||
uint64_t bsStride = d0Stride * d0;
|
||||
uint64_t d1Stride = bsStride * blockSize;
|
||||
uint64_t nStride = d1Stride * d1;
|
||||
uint64_t bnStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bnStride, nStride, d1Stride, bsStride, d0Stride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_NZ_K_SCALE> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize1, uint32_t d, uint32_t blockSize0) {
|
||||
shape = AscendC::MakeShape(n, blockSize1, d, blockSize0);
|
||||
uint64_t bs0Stride = 1;
|
||||
uint64_t dStride = bs0Stride * blockSize0;
|
||||
uint64_t bs1Stride = dStride * d;
|
||||
uint64_t nStride = bs1Stride * blockSize1;
|
||||
uint64_t bnStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bnStride, nStride, bs1Stride, dStride, bs0Stride);
|
||||
}
|
||||
};
|
||||
|
||||
// post_quant
|
||||
template <>
|
||||
struct GmLayout<GmFormat::NGD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t g, uint32_t d) {
|
||||
shape = AscendC::MakeShape(n, g, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t gStride = dStride * d;
|
||||
uint64_t nStride = gStride * g;
|
||||
stride = AscendC::MakeStride(nStride, gStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
//antiquant
|
||||
template <>
|
||||
struct GmLayout<GmFormat::ND> {
|
||||
AscendC::Shape<uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t d) {
|
||||
shape = AscendC::MakeShape(n, d);
|
||||
|
||||
uint64_t dStride = 1;
|
||||
uint64_t nStride = dStride * d; //headDim
|
||||
stride = AscendC::MakeStride(nStride, dStride);
|
||||
}
|
||||
};
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BS2> {
|
||||
AscendC::Shape<uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t s) {
|
||||
shape = AscendC::MakeShape(b, s);
|
||||
|
||||
uint64_t sStride = 1;
|
||||
uint64_t bStride = sStride * s;
|
||||
|
||||
stride = AscendC::MakeStride(bStride, sStride);
|
||||
}
|
||||
};
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BNS2> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s) {
|
||||
shape = AscendC::MakeShape(b, n, s);
|
||||
|
||||
uint64_t sStride = 1;
|
||||
uint64_t nStride = sStride * s;
|
||||
uint64_t bStride = nStride * n;
|
||||
|
||||
stride = AscendC::MakeStride(bStride, nStride, sStride);
|
||||
}
|
||||
};
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_BnBs> {
|
||||
AscendC::Shape<uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t blockSize) {
|
||||
shape = AscendC::MakeShape(blockSize);
|
||||
|
||||
uint64_t bsStride = 1;
|
||||
uint64_t bnStride = bsStride * blockSize;
|
||||
stride = AscendC::MakeStride(bnStride, bsStride);
|
||||
}
|
||||
};
|
||||
template <>
|
||||
struct GmLayout<GmFormat::PA_BnNBs> {
|
||||
AscendC::Shape<uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t n, uint32_t blockSize) {
|
||||
shape = AscendC::MakeShape(n, blockSize);
|
||||
|
||||
uint64_t bsStride = 1;
|
||||
uint64_t nStride = bsStride * blockSize;
|
||||
uint64_t bnStride = nStride * n; //blockSize * kvHeadNum
|
||||
stride = AscendC::MakeStride(bnStride, nStride, bsStride);
|
||||
}
|
||||
};
|
||||
|
||||
//PSE_GmLayout
|
||||
template <>
|
||||
struct GmLayout<GmFormat::BN2GS1S2> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s1, uint32_t s2)
|
||||
{
|
||||
shape = AscendC::MakeShape(b, n, g, s1, s2);
|
||||
uint64_t s2Stride = 1;
|
||||
uint64_t s1Stride = s2Stride * s2;
|
||||
uint64_t gStride = s1Stride * s1;
|
||||
uint64_t nStride = gStride * g;
|
||||
uint64_t bStride = nStride * n;
|
||||
stride = AscendC::MakeStride(bStride, nStride, gStride, s1Stride, s2Stride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::SBNGD> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t g, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, g, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t gStride = dStride * d;
|
||||
uint64_t nStride = gStride * g;
|
||||
uint64_t bStride = nStride * n;
|
||||
uint64_t sStride = bStride * b;
|
||||
stride = AscendC::MakeStride(bStride, nStride, gStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayout<GmFormat::SBND> {
|
||||
AscendC::Shape<uint32_t, uint32_t, uint32_t, uint32_t> shape;
|
||||
AscendC::Stride<uint64_t, uint64_t, uint64_t, uint64_t> stride;
|
||||
|
||||
__aicore__ inline GmLayout() = default;
|
||||
__aicore__ inline void MakeLayout(uint32_t b, uint32_t n, uint32_t s, uint32_t d) {
|
||||
shape = AscendC::MakeShape(b, n, s, d);
|
||||
uint64_t dStride = 1;
|
||||
uint64_t nStride = dStride * d;
|
||||
uint64_t bStride = nStride * n;
|
||||
uint64_t sStride = bStride * b;
|
||||
stride = AscendC::MakeStride(bStride, nStride, sStride, dStride);
|
||||
}
|
||||
};
|
||||
|
||||
#endif
|
||||
1104
csrc/attention/common/op_kernel/memcopy/offset_calculator_v2.h
Normal file
1104
csrc/attention/common/op_kernel/memcopy/offset_calculator_v2.h
Normal file
File diff suppressed because it is too large
Load Diff
140
csrc/attention/common/op_kernel/memcopy/parser.h
Normal file
140
csrc/attention/common/op_kernel/memcopy/parser.h
Normal file
@@ -0,0 +1,140 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file parser.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef PARSER_H
|
||||
#define PARSER_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_vec_intf.h"
|
||||
#include "kernel_cube_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
using AscendC::GlobalTensor;
|
||||
|
||||
// ----------------------------------------------ActualSeqLensParser--------------------------------
|
||||
enum class ActualSeqLensMode
|
||||
{
|
||||
BY_BATCH = 0,
|
||||
ACCUM = 1,
|
||||
};
|
||||
|
||||
template <ActualSeqLensMode MODE, typename ACTLEN_T = uint64_t>
|
||||
class ActualSeqLensParser {
|
||||
};
|
||||
|
||||
template <typename ACTLEN_T>
|
||||
class ActualSeqLensParser<ActualSeqLensMode::ACCUM, ACTLEN_T> {
|
||||
public:
|
||||
__aicore__ inline ActualSeqLensParser() = default;
|
||||
|
||||
__aicore__ inline void Init(GlobalTensor<ACTLEN_T> actualSeqLengthsGm, uint32_t actualLenDims,
|
||||
uint64_t defaultVal = 0)
|
||||
{
|
||||
this->actualSeqLengthsGm = actualSeqLengthsGm;
|
||||
this->actualLenDims = actualLenDims;
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetTBase(uint32_t bIdx) const
|
||||
{
|
||||
if (bIdx == 0) {
|
||||
return 0;
|
||||
}
|
||||
return actualSeqLengthsGm.GetValue(bIdx - 1);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetMxVscaleTBase(uint32_t bIdx) const
|
||||
{
|
||||
if (bIdx == 0) {
|
||||
return 0;
|
||||
}
|
||||
uint64_t vScaleTBaseOffset = 0;
|
||||
for (uint32_t idx = 0; idx < bIdx; idx++) {
|
||||
vScaleTBaseOffset += ((GetActualSeqLength(idx) + 63) >> 6);
|
||||
}
|
||||
return vScaleTBaseOffset;
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetActualSeqLength(uint32_t bIdx) const
|
||||
{
|
||||
if (bIdx == 0) {
|
||||
return actualSeqLengthsGm.GetValue(0);
|
||||
}
|
||||
return (actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1));
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetTSize() const
|
||||
{
|
||||
return actualSeqLengthsGm.GetValue(actualLenDims - 1);
|
||||
}
|
||||
private:
|
||||
GlobalTensor<ACTLEN_T> actualSeqLengthsGm;
|
||||
uint32_t actualLenDims;
|
||||
};
|
||||
|
||||
template <typename ACTLEN_T>
|
||||
class ActualSeqLensParser<ActualSeqLensMode::BY_BATCH, ACTLEN_T> {
|
||||
public:
|
||||
__aicore__ inline ActualSeqLensParser() = default;
|
||||
|
||||
__aicore__ inline void Init(GlobalTensor<ACTLEN_T> actualSeqLengthsGm, uint32_t actualLenDims, uint64_t defaultVal)
|
||||
{
|
||||
this->actualSeqLengthsGm = actualSeqLengthsGm;
|
||||
this->actualLenDims = actualLenDims;
|
||||
this->defaultVal = defaultVal;
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetActualSeqLength(uint32_t bIdx) const
|
||||
{
|
||||
if (actualLenDims == 0) {
|
||||
return defaultVal;
|
||||
}
|
||||
if (actualLenDims == 1) {
|
||||
return actualSeqLengthsGm.GetValue(0);
|
||||
}
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t GetActualLenDims() const
|
||||
{
|
||||
return actualLenDims;
|
||||
}
|
||||
private:
|
||||
GlobalTensor<ACTLEN_T> actualSeqLengthsGm;
|
||||
uint32_t actualLenDims = 0;
|
||||
uint64_t defaultVal = 0;
|
||||
};
|
||||
|
||||
// ----------------------------------------------BlockTableParser--------------------------------
|
||||
class BlockTableParser {
|
||||
public:
|
||||
__aicore__ inline BlockTableParser() = default;
|
||||
|
||||
__aicore__ inline void Init(GlobalTensor<int32_t> blockTableGm, uint32_t maxblockNumPerBatch)
|
||||
{
|
||||
this->blockTableGm = blockTableGm;
|
||||
this->maxblockNumPerBatch = maxblockNumPerBatch;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t GetBlockIdx(uint32_t bIdx, uint32_t blockIdxInBatch) const
|
||||
{
|
||||
return blockTableGm.GetValue(bIdx * maxblockNumPerBatch + blockIdxInBatch);
|
||||
}
|
||||
private:
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
uint32_t maxblockNumPerBatch;
|
||||
};
|
||||
|
||||
#endif
|
||||
31
csrc/attention/common/op_kernel/offset_calculator.h
Normal file
31
csrc/attention/common/op_kernel/offset_calculator.h
Normal file
@@ -0,0 +1,31 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file offset_calculator.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OFFSET_CALCULATOR_H
|
||||
#define OFFSET_CALCULATOR_H
|
||||
|
||||
#if ASC_DEVKIT_MAJOR >= 9
|
||||
#include "kernel_basic_intf.h"
|
||||
#else
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
#include "memcopy/gm_layout.h"
|
||||
#include "memcopy/parser.h"
|
||||
#include "memcopy/offset_calculator_v2.h"
|
||||
#include "memcopy/fa_gm_tensor.h"
|
||||
#include "memcopy/fa_l1_tensor.h"
|
||||
#include "memcopy/gm_coord.h"
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user