@@ -0,0 +1,266 @@
|
||||
/**
|
||||
* 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 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;
|
||||
};
|
||||
|
||||
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 = BaseApi::Align16Func(gm2L1Nd2NzParams.nValue); // 转换为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 DataCopyGmNZToL1(LocalTensor<L1Type>& l1Tensor, GlobalTensor<L1Type>& gmTensor,
|
||||
uint32_t rowAct, // 实际需要拷贝的行数
|
||||
uint32_t dstRowStride,
|
||||
uint32_t srcRowStride,
|
||||
uint32_t col) // D
|
||||
{
|
||||
// 4bit场景下,blockElementCnt * 2
|
||||
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 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 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__)
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_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 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__)
|
||||
if constexpr (IsSameType<INPUT_T, fp8_e5m2_t>::value || IsSameType<INPUT_T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<INPUT_T, hifloat8_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
|
||||
@@ -0,0 +1,55 @@
|
||||
/**
|
||||
* 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 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
|
||||
|
||||
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,259 @@
|
||||
/**
|
||||
* 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.h
|
||||
* \brief同步管理
|
||||
*/
|
||||
#ifndef BUFFER_H
|
||||
#define BUFFER_H
|
||||
#include<type_traits>
|
||||
#include"lib/matmul_intf.h"
|
||||
#include"kernel_event.h"
|
||||
#include"kernel_common.h"
|
||||
#include"kernel_tpipe.h"
|
||||
using namespace AscendC;
|
||||
namespace fa_base_matmul {
|
||||
__BLOCK_LOCAL__ __inline__ uint32_t idCounterNum;
|
||||
#define MAKE_ID ((++idCounterNum) % 16)
|
||||
|
||||
// 核间同步中,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,
|
||||
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::GM) {
|
||||
return HardEvent::MTE2_S;
|
||||
}
|
||||
}
|
||||
|
||||
__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::GM) {
|
||||
return HardEvent::S_MTE2;
|
||||
}
|
||||
}
|
||||
|
||||
__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_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 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 {
|
||||
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_);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__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 {
|
||||
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
|
||||
@@ -0,0 +1,57 @@
|
||||
/**
|
||||
* 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_manager.h
|
||||
* \brief buffer内存管理
|
||||
*/
|
||||
#ifndef BUFFER_MANAGER_H
|
||||
#define BUFFER_MANAGER_H
|
||||
|
||||
#include "buffer.h"
|
||||
|
||||
// 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
|
||||
@@ -0,0 +1,409 @@
|
||||
/**
|
||||
* 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 buffers_policy.h
|
||||
* \brief 综合管理buffer的内存和同步
|
||||
*/
|
||||
#ifndef BUFFERS_POLICY_H
|
||||
#define BUFFERS_POLICY_H
|
||||
|
||||
#include "buffer_manager.h"
|
||||
#include "buffer.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
|
||||
@@ -0,0 +1,637 @@
|
||||
/**
|
||||
* 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 matmul.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MATMUL_H
|
||||
#define MATMUL_H
|
||||
#include "buffers_policy.h"
|
||||
using namespace AscendC;
|
||||
namespace fa_base_matmul {
|
||||
|
||||
constexpr uint32_t UNITFLAG_DISABLE = 0;
|
||||
constexpr uint32_t UNITFLAG_ENABLE = 2;
|
||||
constexpr uint32_t UNITFLAG_EN_OUTER_LAST = 3;
|
||||
static constexpr uint32_t FP16_ONE_FRACTAL_ELEMENT = 16; // 一个分形512B,16*16个fp16
|
||||
static constexpr uint32_t INT4_ONE_FRACTAL_ELEMENT = 64; // 一个分形512B,16*64个fp16
|
||||
static constexpr uint32_t ONE_FRACTAL_H_ELEMENT = 16; // 一个分形512B,height方向为16个element
|
||||
static constexpr uint32_t ONE_FRACTAL_W_BYTE = 32; // 一个分形512B,weight方向为32B
|
||||
static constexpr uint32_t LOAD3D_L1W_SIZE = 16;
|
||||
static constexpr uint32_t MMAD_MN_SIZE_10 = 10;
|
||||
static constexpr uint8_t LOAD3D_STRIDE_W = 1;
|
||||
static constexpr uint8_t LOAD3D_STRIDE_H = 1;
|
||||
static constexpr uint8_t LOAD3D_FILTER_W = 1;
|
||||
static constexpr uint8_t LOAD3D_FILTER_H = 1;
|
||||
static constexpr uint8_t LOAD3D_DILA_FILTER_W = 1;
|
||||
static constexpr uint8_t LOAD3D_DILA_FILTER_H = 1;
|
||||
static constexpr uint32_t K_STEP_ALIGN_BASE = 2;
|
||||
static constexpr uint32_t M_STEP_ALIGN_BASE = 2;
|
||||
|
||||
struct MMParam {
|
||||
uint32_t singleM;
|
||||
uint32_t singleN;
|
||||
uint32_t singleK;
|
||||
bool isLeftTranspose;
|
||||
bool isRightTranspose;
|
||||
bool cmatrixInitVal = true;
|
||||
bool isOutKFisrt = true; // 默认值为true, true:在L1切K轴的场景中,表示首轮K
|
||||
uint32_t unitFlag = 0; // 0:disable: 不配置unitFlag
|
||||
// 2:enable: 行为在切K接口中(MatmulK),会将mmadParams.unitFlag设置为 0b10
|
||||
// 3:enable: 行为在切K接口中(MatmulK),在k的最后一轮循环,会将mmadParams.unitFlag设置为 0b11
|
||||
// 外部使用时,在外层k循环的最后一轮将该参数配置为3
|
||||
uint32_t realM = 0; // bmm2以s1realsize为M轴,不赋值时不影响现有代码逻辑
|
||||
};
|
||||
|
||||
enum class ABLayout {
|
||||
MK = 0,
|
||||
KM = 1,
|
||||
KN = 2,
|
||||
NK = 3,
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T AlignUp(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
#if ((__CCE_AICORE__ == 310) || (defined __DAV_310R6__) || (__NPU_ARCH__ == 5102))
|
||||
template <typename T>
|
||||
__aicore__ inline uint32_t GetBlockNum(uint32_t size) {
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
return ((size + 7) >> 3 << 3) >> 3;
|
||||
} else if constexpr ((IsSameType<T, fp8_e5m2_t>::value ||
|
||||
IsSameType<T, fp8_e4m3fn_t>::value ||
|
||||
IsSameType<T, hifloat8_t>::value ||
|
||||
IsSameType<T, int8_t>::value)) {
|
||||
return ((size + 31) >> 5 << 5) >> 5;
|
||||
} else {
|
||||
return ((size + 15) >> 4 << 4) >> 4;
|
||||
}
|
||||
}
|
||||
// L1->L0A + 切k/切M/全载
|
||||
template <typename T>
|
||||
__aicore__ inline void LoadDataToL0A(LocalTensor<T>& aL0Tensor, const LocalTensor<T>& aL1Tensor,
|
||||
const MMParam& mmParam, uint64_t L1Aoffset, uint32_t kSplitSize, uint32_t mSplitSize)
|
||||
{
|
||||
LoadData2DParamsV2 loadData2DParamsA; // 基础API LoadData的参数结构体
|
||||
loadData2DParamsA.mStartPosition = 0; // 以M*K矩阵为例,源矩阵M轴方向的起始位置,单位为16 element
|
||||
loadData2DParamsA.kStartPosition = 0; // 以M*K矩阵为例,源矩阵K轴方向的起始位置,单位为32B
|
||||
loadData2DParamsA.ifTranspose = mmParam.isLeftTranspose; // 是否启用转置功能,对每个分型矩阵进行转置
|
||||
if (loadData2DParamsA.ifTranspose) {
|
||||
loadData2DParamsA.mStep = ((kSplitSize + 15) >> 4 << 4) / 16; // 以M*K矩阵为例,源矩阵M轴方向搬运长度(S1向上对齐分形(512B),16*16个f16->向上对齐16),单位为16 element,取值范围:mStep属于[0,255]
|
||||
if constexpr (IsSameType<T, fp8_e5m2_t>::value || IsSameType<T, fp8_e4m3fn_t>::value || IsSameType<T, hifloat8_t>::value) {
|
||||
loadData2DParamsA.mStep = (loadData2DParamsA.mStep + 1) >> 1 << 1;
|
||||
}
|
||||
loadData2DParamsA.kStep = GetBlockNum<T>(mSplitSize); // 以M*K矩阵为例,源矩阵K轴方向搬运长度(qkD个f16),单位为32B,取值范围:nStep属于[0,255]
|
||||
} else {
|
||||
loadData2DParamsA.mStep = ((mSplitSize + 15) >> 4 << 4) / 16; // 以M*K矩阵为例,源矩阵M轴方向搬运长度(S1向上对齐分形(512B),16*16个f16->向上对齐16),单位为16 element,取值范围:mStep属于[0,255]
|
||||
loadData2DParamsA.kStep = GetBlockNum<T>(kSplitSize); // 以M*K矩阵为例,源矩阵K轴方向搬运长度(qkD个f16),单位为32B,取值范围:nStep属于[0,255]
|
||||
}
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
if (loadData2DParamsA.ifTranspose) {
|
||||
loadData2DParamsA.kStep = CeilAlign(loadData2DParamsA.kStep, K_STEP_ALIGN_BASE);
|
||||
}
|
||||
}
|
||||
if constexpr (IsSameType<T, fp8_e5m2_t>::value || IsSameType<T, fp8_e4m3fn_t>::value || IsSameType<T, hifloat8_t>::value) {
|
||||
// 配合ub->L1使用 256 * 32 / 256
|
||||
// 64搬运对齐
|
||||
loadData2DParamsA.srcStride = loadData2DParamsA.ifTranspose ? ((kSplitSize + 63) >> 6 << 6) / 16 : ((mSplitSize + 31) >> 5 << 5) / 16; // 以M*K矩阵为例,源矩阵K方向前一个分形起始地址与后一个分形起始地址的间隔,单位:512B
|
||||
} else {
|
||||
loadData2DParamsA.srcStride = loadData2DParamsA.ifTranspose ? ((mmParam.singleK + 15) >> 4 << 4) / 16 : loadData2DParamsA.mStep;
|
||||
}
|
||||
if (mmParam.realM != 0) {
|
||||
loadData2DParamsA.mStep = ((mmParam.realM + 15) >> 4 << 4) / 16;
|
||||
}
|
||||
loadData2DParamsA.dstStride = loadData2DParamsA.ifTranspose ? (mSplitSize + 15) / 16 : loadData2DParamsA.mStep;
|
||||
LoadData(aL0Tensor, aL1Tensor[L1Aoffset], loadData2DParamsA);
|
||||
}
|
||||
|
||||
// L1->L0B + 切k/切M/全载
|
||||
template <typename T>
|
||||
__aicore__ inline void LoadDataToL0B(LocalTensor<T>& bL0Tensor, const LocalTensor<T>& bL1Tensor,
|
||||
const MMParam& mmParam, uint64_t L1Boffset, uint32_t kSplitSize, uint32_t nSplitSize, int nLoops = 1)
|
||||
{
|
||||
LoadData2DParamsV2 loadData2DParamsB; // 基础API LoadData的参数结构体
|
||||
loadData2DParamsB.mStartPosition = 0; // 以M*K矩阵为例,源矩阵M轴方向的起始位置,单位为16 element
|
||||
loadData2DParamsB.kStartPosition = 0; // 以M*K矩阵为例,源矩阵K轴方向的起始位置,单位为32B
|
||||
loadData2DParamsB.ifTranspose = !mmParam.isRightTranspose; // 是否启用转置功能,对每个分型矩阵进行转置
|
||||
if (loadData2DParamsB.ifTranspose) {
|
||||
loadData2DParamsB.mStep = ((kSplitSize + 15) >> 4 << 4) / 16; // 以M*K矩阵为例,源矩阵M轴方向搬运长度(S1向上对齐分形(512B),16*16个f16->向上对齐16),单位为16 element,取值范围:mStep属于[0,255]
|
||||
loadData2DParamsB.kStep = GetBlockNum<T>(nSplitSize); // 以M*K矩阵为例,源矩阵K轴方向搬运长度(qkD个f16),单位为32B,取值范围:nStep属于[0,255]
|
||||
} else {
|
||||
loadData2DParamsB.mStep = ((nSplitSize + 15) >> 4 << 4) / 16; // 以M*K矩阵为例,源矩阵M轴方向搬运长度(S1向上对齐分形(512B),16*16个f16->向上对齐16),单位为16 element,取值范围:mStep属于[0,255]
|
||||
loadData2DParamsB.kStep = GetBlockNum<T>(kSplitSize); // 以M*K矩阵为例,源矩阵K轴方向搬运长度(qkD个f16),单位为32B,取值范围:nStep属于[0,255]
|
||||
}
|
||||
if constexpr (IsSameType<T, float>::value) {
|
||||
if (loadData2DParamsB.ifTranspose) {
|
||||
loadData2DParamsB.kStep = CeilAlign(loadData2DParamsB.kStep, K_STEP_ALIGN_BASE);
|
||||
}
|
||||
}
|
||||
if constexpr (IsSameType<T, fp8_e5m2_t>::value || IsSameType<T, fp8_e4m3fn_t>::value || IsSameType<T, hifloat8_t>::value) {
|
||||
if (loadData2DParamsB.ifTranspose) {
|
||||
loadData2DParamsB.srcStride = ((kSplitSize + 31) >> 5 << 5) / 16;
|
||||
} else {
|
||||
loadData2DParamsB.srcStride = ((nSplitSize + 31) >> 5 << 5) / 16;
|
||||
}
|
||||
} else {
|
||||
loadData2DParamsB.srcStride = loadData2DParamsB.ifTranspose ? (((mmParam.singleK + 15) >> 4 << 4) / 16) : (((mmParam.singleN + 15 ) >> 4 << 4) / 16); // 以M*K矩阵为例,源矩阵K方向前一个分形起始地址与后一个分形起始地址的间隔,单位:512B
|
||||
}
|
||||
loadData2DParamsB.dstStride = loadData2DParamsB.ifTranspose ? (nSplitSize + 15) / 16 : loadData2DParamsB.mStep; // 以M*K矩阵为例,目标矩阵K方向前一个分形起始地址与后一个分形起始地址的间隔,单位:512B
|
||||
if constexpr (IsSameType<T, fp8_e5m2_t>::value || IsSameType<T, fp8_e4m3fn_t>::value || IsSameType<T, hifloat8_t>::value) {
|
||||
if (loadData2DParamsB.ifTranspose) {
|
||||
uint32_t l0bLoop = (loadData2DParamsB.mStep + 1) >> 1;
|
||||
loadData2DParamsB.mStep = M_STEP_ALIGN_BASE;
|
||||
uint64_t dstOffset = 0;
|
||||
uint64_t dstAddrStride = (nSplitSize + 15) / 16 * 16 * 32;
|
||||
uint16_t oriMStep = loadData2DParamsB.mStartPosition;
|
||||
for (uint32_t idx = 0; idx < l0bLoop; ++idx) {
|
||||
loadData2DParamsB.mStartPosition = oriMStep + M_STEP_ALIGN_BASE * idx;
|
||||
LoadData(bL0Tensor[dstOffset], bL1Tensor[L1Boffset], loadData2DParamsB);
|
||||
dstOffset += dstAddrStride;
|
||||
}
|
||||
} else {
|
||||
LoadData(bL0Tensor, bL1Tensor[L1Boffset], loadData2DParamsB);
|
||||
}
|
||||
} else {
|
||||
LoadData(bL0Tensor, bL1Tensor[L1Boffset], loadData2DParamsB);
|
||||
}
|
||||
}
|
||||
#else
|
||||
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
|
||||
template <typename T, ABLayout AL>
|
||||
__aicore__ inline void LoadDataToL0A(LocalTensor<T>& aL0Tensor, const LocalTensor<T>& aL1Tensor,
|
||||
const MMParam& mmParam, uint64_t L1Aoffset, uint32_t kSplitSize,
|
||||
uint32_t mSplitSize)
|
||||
{
|
||||
if constexpr (AL == ABLayout::MK) {
|
||||
LoadData3DParamsV2<T> loadData3DParams;
|
||||
loadData3DParams.l1H = mSplitSize / LOAD3D_L1W_SIZE; // 源操作数height
|
||||
loadData3DParams.l1W = LOAD3D_L1W_SIZE; // 源操作数weight
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParams.mExtension = mSplitSize; // 在目的操作数height维度的传输长度
|
||||
loadData3DParams.kExtension = kSplitSize; // 在目的操作数width维度的传输长度
|
||||
loadData3DParams.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParams.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParams.strideW = 1; // 卷积核在源操作数width维度滑动的步长
|
||||
loadData3DParams.strideH = 1; // 卷积核在源操作数height维度滑动的步长
|
||||
loadData3DParams.filterW = 1; // 卷积核width
|
||||
loadData3DParams.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParams.filterH = 1; // 卷积核height
|
||||
loadData3DParams.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParams.dilationFilterW = 1; // 卷积核width膨胀系数
|
||||
loadData3DParams.dilationFilterH = 1; // 卷积核height膨胀系数
|
||||
loadData3DParams.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置
|
||||
loadData3DParams.fMatrixCtrl = 0;
|
||||
loadData3DParams.channelSize = kSplitSize; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
LoadData<T, LOAD3DV2_CONFIG>(aL0Tensor, aL1Tensor[L1Aoffset], loadData3DParams);
|
||||
} else if constexpr (AL == ABLayout::KM) {
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0; // 分型矩阵ID,表明搬运起始位置为源操作数中第0个分型
|
||||
loadData2DParams.repeatTimes = (kSplitSize / ONE_FRACTAL_H_ELEMENT) * (mmParam.singleM /
|
||||
(ONE_FRACTAL_W_BYTE / sizeof(T))); // 迭代次数,每个迭代可以处理512B数据
|
||||
loadData2DParams.srcStride = 1; // 相邻迭代间,源操作数前一个分型和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.dstGap = 0; // 相邻迭代间,目的操作数前一个分型的结束地址和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.ifTranspose = true;
|
||||
LoadData(aL0Tensor, aL1Tensor[L1Aoffset], loadData2DParams);
|
||||
}
|
||||
}
|
||||
|
||||
// L1→L0B + 切K/切N/全载
|
||||
template <typename T, ABLayout BL>
|
||||
__aicore__ inline void LoadDataToL0B(LocalTensor<T>& bL0Tensor, const LocalTensor<T>& bL1Tensor,
|
||||
const MMParam& mmParam, uint64_t L1Boffset, uint32_t kSplitSize,
|
||||
uint32_t nSplitSize)
|
||||
{
|
||||
if constexpr (BL == ABLayout::KN) {
|
||||
LoadData3DParamsV2<T> loadData3DParams;
|
||||
loadData3DParams.l1H = kSplitSize / LOAD3D_L1W_SIZE; // 源操作数height
|
||||
loadData3DParams.l1W = LOAD3D_L1W_SIZE; // 源操作数weight=16,目的height=l1H*L1W
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParams.mExtension = kSplitSize; // 在目的操作数height维度的传输长度
|
||||
loadData3DParams.kExtension = nSplitSize; // 在目的操作数width维度的传输长度
|
||||
loadData3DParams.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParams.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParams.strideW = LOAD3D_STRIDE_W;
|
||||
loadData3DParams.strideH = LOAD3D_STRIDE_H;
|
||||
loadData3DParams.filterW = LOAD3D_FILTER_W;
|
||||
loadData3DParams.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParams.filterH = LOAD3D_FILTER_H;
|
||||
loadData3DParams.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParams.dilationFilterW = LOAD3D_DILA_FILTER_W; // 卷积核width膨胀系数
|
||||
loadData3DParams.dilationFilterH = LOAD3D_DILA_FILTER_H; // 卷积核height膨胀系数
|
||||
loadData3DParams.enTranspose = 1; // 是否启用转置功能
|
||||
loadData3DParams.fMatrixCtrl = 0; // 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1
|
||||
loadData3DParams.channelSize = nSplitSize; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
LoadData<T, LOAD3DV2_CONFIG>(bL0Tensor, bL1Tensor[L1Boffset], loadData3DParams);
|
||||
} else if constexpr (BL == ABLayout::NK) {
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0;
|
||||
loadData2DParams.repeatTimes = (nSplitSize + (ONE_FRACTAL_H_ELEMENT - 1)) / ONE_FRACTAL_H_ELEMENT *
|
||||
(kSplitSize / (ONE_FRACTAL_W_BYTE / sizeof(T))); // 迭代次数,每个迭代可以处理512B数据
|
||||
loadData2DParams.srcStride = 1;
|
||||
loadData2DParams.dstGap = 0;
|
||||
loadData2DParams.ifTranspose = false;
|
||||
LoadData(bL0Tensor, bL1Tensor[L1Boffset], loadData2DParams);
|
||||
}
|
||||
}
|
||||
|
||||
#endif
|
||||
// 全载
|
||||
// 外部L1切入K时,需要传入cmatrixInitVal的标记
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL, typename L0AType, typename L0BType>
|
||||
__aicore__ inline void MatmulFull(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
L0AType &aL0BuffsDb,
|
||||
L0BType &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
struct MMParam ¶m)
|
||||
{
|
||||
Buffer<BufferType::L0A> l0aBuffer = aL0BuffsDb.Get();
|
||||
l0aBuffer.Wait<HardEvent::M_MTE1>();
|
||||
LocalTensor<A> L0ATensor = l0aBuffer.GetTensor<A>();
|
||||
LoadDataToL0A(L0ATensor, aL1Tensor, param, 0, param.singleK, param.singleM);
|
||||
l0aBuffer.Set<HardEvent::MTE1_M>();
|
||||
|
||||
Buffer<BufferType::L0B> l0bBuffer = bL0BuffsDb.Get();
|
||||
l0bBuffer.Wait<HardEvent::M_MTE1>();
|
||||
LocalTensor<B> L0BTensor = l0bBuffer.GetTensor<B>();
|
||||
LoadDataToL0B(L0BTensor, bL1Tensor, param, 0, param.singleK, param.singleN);
|
||||
l0bBuffer.Set<HardEvent::MTE1_M>();
|
||||
|
||||
l0aBuffer.Wait<HardEvent::MTE1_M>();
|
||||
l0bBuffer.Wait<HardEvent::MTE1_M>();
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = param.singleM;
|
||||
if (param.realM != 0) {
|
||||
mmadParams.m = param.realM;
|
||||
}
|
||||
mmadParams.n = param.singleN;
|
||||
mmadParams.k = param.singleK;
|
||||
mmadParams.cmatrixInitVal = param.isOutKFisrt;
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.unitFlag = param.unitFlag;
|
||||
if (mmadParams.m == 1) {
|
||||
mmadParams.m = 16;
|
||||
}
|
||||
|
||||
Mmad(cL0Tensor, L0ATensor, L0BTensor, mmadParams);
|
||||
|
||||
l0aBuffer.Set<HardEvent::M_MTE1>();
|
||||
l0bBuffer.Set<HardEvent::M_MTE1>();
|
||||
}
|
||||
|
||||
// 切K
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL, typename L0AType, typename L0BType>
|
||||
__aicore__ inline void MatmulK(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
L0AType &aL0BuffsDb,
|
||||
L0BType &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
const MMParam ¶m)
|
||||
{
|
||||
uint32_t kLoops = (param.singleK + baseK - 1) / baseK; // 尾块处理
|
||||
uint32_t tailSize = param.singleK % baseK;
|
||||
uint32_t tailK = tailSize ? tailSize : baseK;
|
||||
uint64_t L1Aoffset = param.isLeftTranspose ? baseK << 4 : ((param.singleM + 15) >> 4 << 4) * baseK; // 要对齐
|
||||
uint64_t L1Boffset = param.isRightTranspose ? ((param.singleN + 15) >> 4 << 4) * baseK : baseK << 4;
|
||||
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__)
|
||||
if constexpr (IsSameType<A, fp8_e5m2_t>::value || IsSameType<A, fp8_e4m3fn_t>::value || IsSameType<A, hifloat8_t>::value) {
|
||||
L1Aoffset = ((param.singleM + 31) >> 5 << 5) * baseK;
|
||||
L1Boffset = ((param.singleN + 31) >> 5 << 5) * baseK;
|
||||
}
|
||||
if constexpr (IsSameType<A, float32_t>::value) {
|
||||
L1Aoffset = param.isLeftTranspose ? baseK << 3 : ((param.singleM + 15) >> 4 << 4) * baseK;
|
||||
L1Boffset = param.isRightTranspose ? ((param.singleN + 15) >> 4 << 4) * baseK : baseK << 3;
|
||||
}
|
||||
#endif
|
||||
|
||||
for (uint32_t k = 0; k < kLoops; k++) {
|
||||
uint32_t tileK = (k == (kLoops - 1)) ? tailK : baseK;
|
||||
Buffer<BufferType::L0A> l0aBuffer = aL0BuffsDb.Get();
|
||||
l0aBuffer.Wait<HardEvent::M_MTE1>(); // mte1等Matmul:上一轮matmul完成后才能搬运新数据到L0A
|
||||
LocalTensor<A> L0ATensor = l0aBuffer.GetTensor<A>();
|
||||
LoadDataToL0A(L0ATensor, aL1Tensor, param, k * L1Aoffset, tileK, param.singleM);
|
||||
l0aBuffer.Set<HardEvent::MTE1_M>(); // mte1搬运完后,通知可以开始matmul
|
||||
|
||||
Buffer<BufferType::L0B> l0bBuffer = bL0BuffsDb.Get();
|
||||
l0bBuffer.Wait<HardEvent::M_MTE1>(); // mte1等Matmul:上一轮matmul完成后才能搬运新数据到L0B
|
||||
LocalTensor<B> L0BTensor = l0bBuffer.GetTensor<B>();
|
||||
uint64_t loopNum = param.isRightTranspose ? 1 : kLoops;
|
||||
LoadDataToL0B(L0BTensor, bL1Tensor, param, k * L1Boffset, tileK, param.singleN, loopNum);
|
||||
l0bBuffer.Set<HardEvent::MTE1_M>(); // mte1搬运完后,通知可以开始matmul
|
||||
|
||||
l0aBuffer.Wait<HardEvent::MTE1_M>(); // matmul等mte1:L0A数据搬运完成后才能开始matmul
|
||||
l0bBuffer.Wait<HardEvent::MTE1_M>(); // matmul等mte1:L0B数据搬运完成后才能开始matmul
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = param.singleM;
|
||||
if (param.realM != 0) {
|
||||
mmadParams.m = param.realM;
|
||||
}
|
||||
mmadParams.n = param.singleN;
|
||||
mmadParams.k = tileK;
|
||||
if (mmadParams.m == 1) { // m等于1或默认开GEMV模式,文档上没有写怎么关闭GEMV,所以规避当做矩阵运算
|
||||
mmadParams.m = 16;
|
||||
}
|
||||
mmadParams.cmatrixInitVal = param.isOutKFisrt && (k == 0);
|
||||
mmadParams.cmatrixSource = false;
|
||||
if (param.unitFlag != 0) {
|
||||
mmadParams.unitFlag = (param.unitFlag == UNITFLAG_EN_OUTER_LAST) && (k == kLoops - 1) ?
|
||||
UNITFLAG_EN_OUTER_LAST : UNITFLAG_ENABLE;
|
||||
}
|
||||
Mmad(cL0Tensor, L0ATensor, L0BTensor, mmadParams);
|
||||
|
||||
l0aBuffer.Set<HardEvent::M_MTE1>(); // matmul完成后,通知mte1可以开始搬运新数据到L0A
|
||||
l0bBuffer.Set<HardEvent::M_MTE1>(); // matmul完成后,通知mte1可以开始搬运新数据到L0B
|
||||
}
|
||||
}
|
||||
|
||||
// 切N
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL, typename L0AType, typename L0BType>
|
||||
__aicore__ inline void MatmulN(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
L0AType &aL0BuffsDb,
|
||||
L0BType &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
struct MMParam ¶m)
|
||||
{
|
||||
uint32_t nLoops = (param.singleN + baseN - 1) / baseN; // 尾块处理
|
||||
uint32_t tailSize = param.singleN % baseN;
|
||||
uint32_t tailN = tailSize ? tailSize : baseN;
|
||||
uint64_t L1Boffset = param.isRightTranspose ? (baseN << 4) : ((param.singleK + 15) >> 4 << 4) * baseN;
|
||||
#if (__CCE_AICORE__ == 310) || (defined __DAV_310R6__)
|
||||
if constexpr (IsSameType<A, fp8_e5m2_t>::value || IsSameType<A, fp8_e4m3fn_t>::value || IsSameType<A, hifloat8_t>::value) {
|
||||
L1Boffset = ((param.singleK + 31) >> 5 << 5) * baseN;
|
||||
}
|
||||
#endif
|
||||
uint64_t L0Coffset = ((param.singleM + 15) >> 4 << 4) * baseN;
|
||||
if (param.realM != 0) {
|
||||
L0Coffset = ((param.realM + 15) >> 4 << 4) * baseN;
|
||||
}
|
||||
|
||||
Buffer<BufferType::L0A> l0aBuffer = aL0BuffsDb.Get();
|
||||
l0aBuffer.Wait<HardEvent::M_MTE1>(); // mte1等Matmul:上一轮matmul完成后才能搬运新数据到L0A
|
||||
LocalTensor<A> L0ATensor = l0aBuffer.GetTensor<A>();
|
||||
LoadDataToL0A(L0ATensor, aL1Tensor, param, 0, param.singleK, param.singleM);
|
||||
l0aBuffer.Set<HardEvent::MTE1_M>(); // mte1搬运完后,通知可以matmul
|
||||
l0aBuffer.Wait<HardEvent::MTE1_M>(); // matmul等mte1:L0A数据搬运完成后才能开始matmul
|
||||
|
||||
for (uint32_t n = 0; n < nLoops; n++) {
|
||||
uint32_t tileN = (n == (nLoops - 1)) ? tailN : baseN;
|
||||
|
||||
Buffer<BufferType::L0B> l0bBuffer = bL0BuffsDb.Get();
|
||||
l0bBuffer.Wait<HardEvent::M_MTE1>(); // mte1等Matmul:上一轮matmul完成后才能搬运新数据到L0B
|
||||
LocalTensor<B> L0BTensor = l0bBuffer.GetTensor<B>();
|
||||
uint64_t loopNum = param.isRightTranspose ? nLoops : 1;
|
||||
LoadDataToL0B(L0BTensor, bL1Tensor, param, n * L1Boffset, param.singleK, tileN, loopNum);
|
||||
l0bBuffer.Set<HardEvent::MTE1_M>(); // mte1搬运完后,通知可以开始matmul
|
||||
|
||||
l0bBuffer.Wait<HardEvent::MTE1_M>(); // matmul等mte1:L0B数据搬运完成后才能开始matmul
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = param.singleM;
|
||||
if (param.realM != 0) {
|
||||
mmadParams.m = param.realM;
|
||||
}
|
||||
mmadParams.n = tileN;
|
||||
mmadParams.k = param.singleK;
|
||||
if (mmadParams.m == 1) {
|
||||
mmadParams.m = 16;
|
||||
}
|
||||
mmadParams.cmatrixInitVal = param.isOutKFisrt;
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.unitFlag = param.unitFlag;
|
||||
Mmad(cL0Tensor[n * L0Coffset], L0ATensor, L0BTensor, mmadParams);
|
||||
|
||||
l0bBuffer.Set<HardEvent::M_MTE1>(); // matmul完成后,通知mte1可以开始搬运新数据到L0B
|
||||
}
|
||||
l0aBuffer.Set<HardEvent::M_MTE1>(); // matmul完成后,通知mte1可以开始搬运新数据到L0A
|
||||
}
|
||||
|
||||
// 切M
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL>
|
||||
__aicore__ inline void MatmulKM(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
BuffersPolicyDB<BufferType::L0A> &aL0BuffsDb,
|
||||
BuffersPolicyDB<BufferType::L0B> &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
struct MMParam ¶m)
|
||||
{
|
||||
uint32_t mLoops = (param.singleM + baseM - 1) / baseN; // 尾块处理
|
||||
uint32_t kLoops = (param.singleK + baseK - 1) / baseK; // 尾块处理
|
||||
uint32_t mSplitSize = (mLoops == 1) ? param.singleM : baseM;
|
||||
uint32_t mplitTailSize = (param.singleM % baseM) ? (param.singleM % baseM) : mSplitSize;
|
||||
uint32_t kSplitSize = (kLoops == 1) ? param.singleK : baseK;
|
||||
uint32_t kSplitTailSize = (param.singleK % baseK) ? (param.singleK % baseK) : kSplitSize;
|
||||
uint64_t L1Boffset = kSplitSize * param.singleN;
|
||||
uint64_t L0Coffset = mSplitSize * param.singleN;
|
||||
|
||||
for (uint32_t k = 0; k < kLoops; k++) {
|
||||
kSplitSize = (k == (kLoops - 1)) ? kSplitTailSize : kSplitSize;
|
||||
for (uint32_t m = 0; m < mLoops; m++){
|
||||
mSplitSize = (m == (mLoops - 1)) ? mplitTailSize : mSplitSize;
|
||||
Buffer<BufferType::L0A> l0aBuffer = aL0BuffsDb.Get();
|
||||
l0aBuffer.Wait<HardEvent::M_MTE1>(); // 占用
|
||||
LocalTensor<A> L0ATensor = l0aBuffer.GetTensor<A>();
|
||||
LoadDataToL0A(L0ATensor, aL1Tensor, param,
|
||||
k * param.singleM * kSplitSize + m * kSplitSize * mSplitSize,
|
||||
kSplitSize, mSplitSize);
|
||||
l0aBuffer.Set<HardEvent::MTE1_M>(); // 通知
|
||||
l0aBuffer.Wait<HardEvent::MTE1_M>(); // 等待L0A
|
||||
|
||||
Buffer<BufferType::L0B> l0bBuffer = bL0BuffsDb.Get();
|
||||
l0bBuffer.Wait<HardEvent::M_MTE1>(); // 占用
|
||||
LocalTensor<B> L0BTensor = l0bBuffer.GetTensor<B>();
|
||||
LoadDataToL0B(L0BTensor, bL1Tensor, param, k * L1Boffset, kSplitSize, param.singleN);
|
||||
l0bBuffer.Set<HardEvent::MTE1_M>(); // 通知
|
||||
l0bBuffer.Wait<HardEvent::MTE1_M>(); // 等待L0B
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = mSplitSize;
|
||||
mmadParams.n = param.singleN;
|
||||
mmadParams.k = kSplitSize;
|
||||
mmadParams.cmatrixInitVal = param.isOutKFisrt && (k == 0); // 配置C矩阵初始值是否为0。默认值ture
|
||||
mmadParams.cmatrixSource = false; // 来源于CO1
|
||||
Mmad(cL0Tensor[m * L0Coffset], L0ATensor, L0BTensor, mmadParams);
|
||||
|
||||
l0aBuffer.Set<HardEvent::M_MTE1>(); // 释放L0A
|
||||
l0bBuffer.Set<HardEvent::M_MTE1>(); // 释放L0B
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL, typename L0AType, typename L0BType>
|
||||
__aicore__ inline void MatmulBase(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
L0AType &aL0BuffsDb,
|
||||
L0BType &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
struct MMParam ¶m)
|
||||
{
|
||||
if ((param.singleK + baseK - 1) / baseK > 1) {
|
||||
MatmulK<A, B, C, baseM, baseN, baseK, AL, BL>(aL1Tensor, bL1Tensor, aL0BuffsDb, bL0BuffsDb, cL0Tensor, param);
|
||||
} else if ((param.singleN + baseN - 1) / baseN > 1) {
|
||||
MatmulN<A, B, C, baseM, baseN, baseK, AL, BL>(aL1Tensor, bL1Tensor, aL0BuffsDb, bL0BuffsDb, cL0Tensor, param);
|
||||
} else {
|
||||
MatmulFull<A, B, C, baseM, baseN, baseK, AL, BL>(aL1Tensor, bL1Tensor, aL0BuffsDb, bL0BuffsDb, cL0Tensor, param);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename A, typename B, typename C, uint32_t baseM, uint32_t baseN, uint32_t baseK, ABLayout AL, ABLayout BL>
|
||||
__aicore__ inline void MatmulKPP(const LocalTensor<A> &aL1Tensor,
|
||||
const LocalTensor<B> &bL1Tensor,
|
||||
BuffersPolicyDB<BufferType::L0A> &aL0BuffsDb,
|
||||
BuffersPolicyDB<BufferType::L0B> &bL0BuffsDb,
|
||||
const LocalTensor<C> &cL0Tensor,
|
||||
const MMParam ¶m)
|
||||
{
|
||||
uint32_t kLoops = (param.singleK + baseK - 1) / baseK;
|
||||
uint32_t kSplitSize = (kLoops == 1) ? param.singleK : baseK;
|
||||
uint32_t kSplitSizeAlign = AlignUp(kSplitSize, FP16_ONE_FRACTAL_ELEMENT);
|
||||
uint64_t L1Aoffset = AlignUp(param.singleM, FP16_ONE_FRACTAL_ELEMENT) * kSplitSize;
|
||||
uint64_t L1Boffset = AlignUp(param.singleN, FP16_ONE_FRACTAL_ELEMENT) * kSplitSize;
|
||||
for (uint32_t k = 0; k < kLoops; k++) {
|
||||
if (k == kLoops - 1) {
|
||||
kSplitSize = (param.singleK % baseK) ? (param.singleK % baseK) : kSplitSize;
|
||||
kSplitSizeAlign = AlignUp(kSplitSize, FP16_ONE_FRACTAL_ELEMENT);
|
||||
}
|
||||
Buffer<BufferType::L0A> l0aBuffer = aL0BuffsDb.Get();
|
||||
l0aBuffer.Wait<HardEvent::M_MTE1>(); // mte1等Matmul:上一轮matmul完成后才能搬运新数据到L0A
|
||||
LocalTensor<A> L0ATensor = l0aBuffer.GetTensor<A>();
|
||||
LoadDataToL0A<A, AL>(L0ATensor, aL1Tensor, param, k * L1Aoffset, kSplitSizeAlign, param.singleM);
|
||||
Buffer<BufferType::L0B> l0bBuffer = bL0BuffsDb.Get();
|
||||
LocalTensor<B> L0BTensor = l0bBuffer.GetTensor<B>();
|
||||
LoadDataToL0B<B, BL>(L0BTensor, bL1Tensor, param, k * L1Boffset, kSplitSizeAlign, param.singleN);
|
||||
|
||||
l0aBuffer.Set<HardEvent::MTE1_M>(); // mte1搬运完后,通知可以开始matmul
|
||||
l0aBuffer.Wait<HardEvent::MTE1_M>(); // matmul等mte1:L0A数据搬运完成后才能开始matmul
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = param.singleM;
|
||||
mmadParams.n = param.singleN;
|
||||
mmadParams.k = kSplitSize;
|
||||
if (mmadParams.m == 1) { //m等于1会默认开GEMV模式,且不可关闭GEMV,所以规避当作矩阵计算
|
||||
mmadParams.m = FP16_ONE_FRACTAL_ELEMENT;
|
||||
}
|
||||
mmadParams.cmatrixInitVal = (param.isOutKFisrt == true) && (k == 0);
|
||||
mmadParams.cmatrixSource = false;
|
||||
if (param.unitFlag != 0) {
|
||||
mmadParams.unitFlag = (param.unitFlag == UNITFLAG_EN_OUTER_LAST) && (k == kLoops - 1) ?
|
||||
UNITFLAG_EN_OUTER_LAST : UNITFLAG_ENABLE;
|
||||
}
|
||||
|
||||
Mmad(cL0Tensor, L0ATensor, L0BTensor, mmadParams);
|
||||
#if (__CCE_AICORE__ != 310) && (!(defined __DAV_310R6__))
|
||||
if ((mmadParams.m / FP16_ONE_FRACTAL_ELEMENT) * (mmadParams.n / FP16_ONE_FRACTAL_ELEMENT) < MMAD_MN_SIZE_10) {
|
||||
AscendC::PipeBarrier<PIPE_M>();
|
||||
}
|
||||
#endif
|
||||
l0aBuffer.Set<HardEvent::M_MTE1>(); // matmul完成后,通知mte1可以开始搬运新数据到L0A
|
||||
}
|
||||
}
|
||||
template <typename T, ABLayout AL>
|
||||
__aicore__ inline void LoadDataToL0A(LocalTensor<T>& aL0Tensor, const LocalTensor<T>& aL1Tensor,
|
||||
uint32_t rowSize, uint32_t kSplitSize, uint32_t mSplitSize)
|
||||
{
|
||||
uint32_t blockElementCnt = ONE_FRACTAL_W_BYTE / sizeof(T);
|
||||
if constexpr (IsSameType<T, int4b_t>::value) {
|
||||
blockElementCnt = INT4_ONE_FRACTAL_ELEMENT;
|
||||
}
|
||||
if constexpr (AL == ABLayout::MK) {
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0; // 分型矩阵ID,表明搬运起始位置为源操作数中第0个分型
|
||||
loadData2DParams.srcStride = 1; // 相邻迭代间,源操作数前一个分型和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.dstGap = kSplitSize / blockElementCnt - 1; // 相邻迭代间,目的操作数前一个分型的结束地址和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.repeatTimes = mSplitSize / ONE_FRACTAL_H_ELEMENT; // 迭代次数,每个迭代可以处理512B数据
|
||||
loadData2DParams.ifTranspose = false;
|
||||
uint32_t loopTimes = kSplitSize / blockElementCnt;
|
||||
uint64_t l1Offset = rowSize * blockElementCnt;
|
||||
uint64_t l0Offset = ONE_FRACTAL_H_ELEMENT * blockElementCnt;
|
||||
for(uint32_t loop = 0; loop < loopTimes; loop++) {
|
||||
LoadData(aL0Tensor[loop * l0Offset], aL1Tensor[loop * l1Offset], loadData2DParams);
|
||||
}
|
||||
} else if constexpr (AL == ABLayout::KM) {
|
||||
LoadData2dTransposeParams loadData2dTransposeParams;
|
||||
loadData2dTransposeParams.startIndex = 0;
|
||||
loadData2dTransposeParams.srcStride = 1;
|
||||
loadData2dTransposeParams.dstFracGap = (kSplitSize + blockElementCnt -1) / blockElementCnt;
|
||||
loadData2dTransposeParams.dstGap = mSplitSize / ONE_FRACTAL_H_ELEMENT - 1;
|
||||
if(rowSize == kSplitSize) {
|
||||
loadData2dTransposeParams.repeatTimes = (kSplitSize + blockElementCnt - 1) / blockElementCnt;
|
||||
uint32_t loopTimes = mSplitSize / blockElementCnt;
|
||||
uint64_t l1Offset = rowSize * blockElementCnt;
|
||||
uint64_t l0Offset = kSplitSize * blockElementCnt;
|
||||
for(uint32_t loop = 0; loop < loopTimes; loop++) {
|
||||
LoadDataWithTranspose(aL0Tensor[loop * l0Offset], aL1Tensor[loop * l1Offset], loadData2dTransposeParams);
|
||||
}
|
||||
} else {
|
||||
loadData2dTransposeParams.repeatTimes = ((kSplitSize + blockElementCnt - 1) / blockElementCnt) * (mSplitSize / blockElementCnt);
|
||||
LoadDataWithTranspose(aL0Tensor, aL1Tensor, loadData2dTransposeParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, ABLayout BL>
|
||||
__aicore__ inline void LoadDataToL0B(LocalTensor<T>& bL0Tensor, const LocalTensor<T>& bL1Tensor,
|
||||
uint32_t rowSize, uint32_t kSplitSize, uint32_t nSplitSize)
|
||||
{
|
||||
uint32_t blockElementCnt = ONE_FRACTAL_W_BYTE / sizeof(T);
|
||||
if constexpr (IsSameType<T, int4b_t>::value) {
|
||||
blockElementCnt = INT4_ONE_FRACTAL_ELEMENT;
|
||||
}
|
||||
if constexpr (BL == ABLayout::KN) {
|
||||
LoadData2dTransposeParams loadData2dTransposeParams;
|
||||
|
||||
loadData2dTransposeParams.startIndex = 0;
|
||||
loadData2dTransposeParams.srcStride = 1;
|
||||
loadData2dTransposeParams.dstFracGap = 0;
|
||||
loadData2dTransposeParams.dstGap = nSplitSize / ONE_FRACTAL_H_ELEMENT - 1;
|
||||
loadData2dTransposeParams.repeatTimes = (kSplitSize + blockElementCnt - 1) / blockElementCnt;
|
||||
|
||||
uint32_t loopTimes = nSplitSize / blockElementCnt;
|
||||
uint64_t l1Offset = rowSize * blockElementCnt;
|
||||
uint64_t l0Offset = blockElementCnt * blockElementCnt;
|
||||
for(uint32_t loop = 0; loop < loopTimes; loop++) {
|
||||
LoadDataWithTranspose(bL0Tensor[loop * l0Offset], bL1Tensor[loop * l1Offset], loadData2dTransposeParams);
|
||||
}
|
||||
} else if constexpr (BL == ABLayout::NK) {
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0; // 分型矩阵ID,表明搬运起始位置为源操作数中第0个分型
|
||||
loadData2DParams.srcStride = 1; // 相邻迭代间,源操作数前一个分型和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.dstGap = 0; // 相邻迭代间,目的操作数前一个分型的结束地址和后一个分型起始地址的间隔(单位512B)
|
||||
loadData2DParams.ifTranspose = false;
|
||||
if(rowSize == kSplitSize) {
|
||||
loadData2DParams.repeatTimes = ((nSplitSize + ONE_FRACTAL_H_ELEMENT - 1) / ONE_FRACTAL_H_ELEMENT) * (kSplitSize / blockElementCnt);// 迭代次数,每个迭代可以处理512B数据
|
||||
LoadData(bL0Tensor, bL1Tensor, loadData2DParams);
|
||||
} else {
|
||||
loadData2DParams.repeatTimes = (nSplitSize + ONE_FRACTAL_H_ELEMENT - 1) / ONE_FRACTAL_H_ELEMENT;// 迭代次数,每个迭代可以处理512B数据
|
||||
uint32_t loopTimes = kSplitSize / blockElementCnt;
|
||||
uint64_t l1Offset = nSplitSize * blockElementCnt;
|
||||
uint64_t l0Offset = rowSize * blockElementCnt;
|
||||
for (uint32_t loop = 0; loop < loopTimes; loop++) {
|
||||
LoadData(bL0Tensor[loop * l0Offset], bL1Tensor[loop * l1Offset], loadData2DParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,882 @@
|
||||
/**
|
||||
* 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 offset_calculator.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef OFFSET_CALCULATOR_H
|
||||
#define OFFSET_CALCULATOR_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using AscendC::GlobalTensor;
|
||||
|
||||
static constexpr uint32_t SHAPE_AXIS_DIM_0 = 0U;
|
||||
static constexpr uint32_t SHAPE_AXIS_DIM_1 = 1U;
|
||||
static constexpr uint32_t SHAPE_AXIS_DIM_2 = 2U;
|
||||
static constexpr uint32_t SHAPE_AXIS_DIM_3 = 3U;
|
||||
static constexpr uint32_t SHAPE_AXIS_DIM_4 = 4U;
|
||||
|
||||
// ----------------------------------------------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,
|
||||
SBNGD = 12,
|
||||
SBND = 13
|
||||
};
|
||||
|
||||
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::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::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);
|
||||
}
|
||||
};
|
||||
|
||||
// ----------------------------------------------ActualSeqLensParser--------------------------------
|
||||
enum class ActualSeqLensMode
|
||||
{
|
||||
BY_BATCH = 0,
|
||||
ACCUM = 1,
|
||||
};
|
||||
|
||||
template <ActualSeqLensMode MODE>
|
||||
class ActualSeqLensParser {
|
||||
};
|
||||
|
||||
template <>
|
||||
class ActualSeqLensParser<ActualSeqLensMode::ACCUM> {
|
||||
public:
|
||||
__aicore__ inline ActualSeqLensParser() = default;
|
||||
|
||||
__aicore__ inline void Init(GlobalTensor<int32_t> actualSeqLengthsGm, uint32_t actualLenDims)
|
||||
{
|
||||
this->actualSeqLengthsGm = actualSeqLengthsGm;
|
||||
this->actualLenDims = actualLenDims;
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetTBase(uint32_t bIdx) const
|
||||
{
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetActualSeqLength(uint32_t bIdx) const
|
||||
{
|
||||
return (actualSeqLengthsGm.GetValue(bIdx + 1) - actualSeqLengthsGm.GetValue(bIdx));
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetTSize() const
|
||||
{
|
||||
return actualSeqLengthsGm.GetValue(actualLenDims - 1);
|
||||
}
|
||||
private:
|
||||
GlobalTensor<int32_t> actualSeqLengthsGm;
|
||||
uint32_t actualLenDims;
|
||||
};
|
||||
|
||||
template <>
|
||||
class ActualSeqLensParser<ActualSeqLensMode::BY_BATCH> {
|
||||
public:
|
||||
__aicore__ inline ActualSeqLensParser() = default;
|
||||
|
||||
__aicore__ inline void Init(GlobalTensor<int32_t> actualSeqLengthsGm, uint32_t actualLenDims, int64_t defaultVal)
|
||||
{
|
||||
this->actualSeqLengthsGm = actualSeqLengthsGm;
|
||||
this->actualLenDims = actualLenDims;
|
||||
this->defaultVal = defaultVal;
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetActualSeqLength(uint32_t bIdx) const
|
||||
{
|
||||
if (actualLenDims == 0) {
|
||||
return defaultVal;
|
||||
}
|
||||
if (actualLenDims == 1) {
|
||||
return actualSeqLengthsGm.GetValue(0);
|
||||
}
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
private:
|
||||
GlobalTensor<int32_t> actualSeqLengthsGm;
|
||||
uint32_t actualLenDims;
|
||||
int64_t defaultVal;
|
||||
};
|
||||
|
||||
// ----------------------------------------------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;
|
||||
};
|
||||
|
||||
// ----------------------------------------------GmLayoutParams--------------------------------
|
||||
enum class FormatCategory
|
||||
{
|
||||
GM_Q_OUT_BNGSD = 0,
|
||||
GM_Q_OUT_TND = 1,
|
||||
GM_KV_BNSD = 2,
|
||||
GM_KV_TND = 3,
|
||||
GM_KV_PA_BNBD = 4,
|
||||
GM_KV_PA_NZ = 5,
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct GmLayoutParams {};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::BSNGD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_BNGSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::BNGSD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_BNGSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::NGBSD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_BNGSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::TNGD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_TND;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::NGTD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_TND;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::BSND> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_BNSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::BNSD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_BNSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::TND> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_TND;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::NTD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_TND;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::PA_BNBSND> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_PA_BNBD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::PA_BNNBSD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_PA_BNBD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::PA_NZ> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_PA_NZ;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::SBNGD> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_Q_OUT_BNGSD;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct GmLayoutParams<GmFormat::SBND> {
|
||||
static constexpr FormatCategory CATEGORY = FormatCategory::GM_KV_BNSD;
|
||||
};
|
||||
|
||||
// ----------------------------------------------OffsetCalculator--------------------------------
|
||||
template <GmFormat FORMAT, FormatCategory CATEGORY>
|
||||
struct OffsetCalculatorImpl {};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_Q_OUT_BNGSD> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t b, uint32_t n2, uint32_t g, uint32_t s1, uint32_t d)
|
||||
{
|
||||
gmLayout.MakeLayout(b, n2, g, s1, d);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t gIdx, uint32_t s1Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t offset = bIdx * GetStrideB() + n2Idx * GetStrideN2() + gIdx * GetStrideG() + s1Idx * GetStrideS1() +
|
||||
dIdx * GetStrideD();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideB()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideG()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideS1()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_4>(gmLayout.stride);
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetDimB()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimG()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimS1()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_4>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_Q_OUT_TND> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
ActualSeqLensParser<ActualSeqLensMode::ACCUM> actualSeqLensQParser;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t n2, uint32_t g, uint32_t d, GlobalTensor<int32_t> actualSeqLengthsGmQ,
|
||||
uint32_t actualLenQDims)
|
||||
{
|
||||
actualSeqLensQParser.Init(actualSeqLengthsGmQ, actualLenQDims);
|
||||
gmLayout.MakeLayout(actualSeqLensQParser.GetTSize(), n2, g, d);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t gIdx, uint32_t s1Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t tIdx = actualSeqLensQParser.GetTBase(bIdx) + s1Idx;
|
||||
uint64_t offset = tIdx * GetStrideT() + n2Idx * GetStrideN2() + gIdx * GetStrideG() + dIdx * GetStrideD();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideT()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideG()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideS1()
|
||||
{
|
||||
return GetStrideT();
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetDimT()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimG()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_KV_BNSD> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t b, uint32_t n2, uint32_t s2, uint32_t d)
|
||||
{
|
||||
gmLayout.MakeLayout(b, n2, s2, d);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t s2Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t offset = bIdx * GetStrideB() + n2Idx * GetStrideN2() + s2Idx * GetStrideS2() + dIdx * GetStrideD();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideB()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideS2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.stride);
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetDimB()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimS2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_KV_TND> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
ActualSeqLensParser<ActualSeqLensMode::ACCUM> actualSeqLensKVParser;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t n2, uint32_t d, GlobalTensor<int32_t> actualSeqLengthsGmKV,
|
||||
uint32_t actualLenKVDims)
|
||||
{
|
||||
actualSeqLensKVParser.Init(actualSeqLengthsGmKV, actualLenKVDims);
|
||||
gmLayout.MakeLayout(actualSeqLensKVParser.GetTSize(), n2, d);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t s2Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t tIdx = actualSeqLensKVParser.GetTBase(bIdx) + s2Idx;
|
||||
uint64_t offset = tIdx * GetStrideT() + n2Idx * GetStrideN2() + dIdx * GetStrideD();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideT()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideS2()
|
||||
{
|
||||
return GetStrideT();
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetDimT()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetDimD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_KV_PA_BNBD> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
BlockTableParser blockTableParser;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t n2, uint32_t blockSize, uint32_t d, GlobalTensor<int32_t> blockTableGm,
|
||||
uint32_t maxblockNumPerBatch)
|
||||
{
|
||||
blockTableParser.Init(blockTableGm, maxblockNumPerBatch);
|
||||
gmLayout.MakeLayout(n2, blockSize, d);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t s2Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t blockIdxInBatch = s2Idx / GetBlockSize(); // 获取block table上的索引
|
||||
uint64_t bsIdx = s2Idx % GetBlockSize(); // 获取在单个块上超出的行数
|
||||
int32_t blockIdx = blockTableParser.GetBlockIdx(bIdx, blockIdxInBatch);
|
||||
uint64_t offset =
|
||||
blockIdx * GetStrideBlockNum() + n2Idx * GetStrideN2() + bsIdx * GetStrideBlockSize() + dIdx * GetStrideD();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideBlockNum()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideBlockSize()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.stride);
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetBlockSize()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetD()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculatorImpl<FORMAT, FormatCategory::GM_KV_PA_NZ> {
|
||||
GmLayout<FORMAT> gmLayout;
|
||||
BlockTableParser blockTableParser;
|
||||
|
||||
__aicore__ inline OffsetCalculatorImpl() = default;
|
||||
|
||||
__aicore__ inline void Init(uint32_t n2, uint32_t blockSize, uint32_t d1, uint32_t d0,
|
||||
GlobalTensor<int32_t> blockTableGm, uint32_t maxblockNumPerBatch)
|
||||
{
|
||||
blockTableParser.Init(blockTableGm, maxblockNumPerBatch);
|
||||
gmLayout.MakeLayout(n2, blockSize, d1, d0);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetOffset(uint32_t bIdx, uint32_t n2Idx, uint32_t s2Idx, uint32_t dIdx)
|
||||
{
|
||||
uint64_t blockIdxInBatch = s2Idx / GetBlockSize(); // 获取block table上的索引
|
||||
uint64_t bsIdx = s2Idx % GetBlockSize(); // 获取在单个块上超出的行数
|
||||
int32_t blockIdx = blockTableParser.GetBlockIdx(bIdx, blockIdxInBatch);
|
||||
|
||||
uint32_t d1Idx = dIdx / GetD0();
|
||||
uint32_t d0Idx = dIdx % GetD0();
|
||||
uint64_t offset = blockIdx * GetStrideBlockNum() + n2Idx * GetStrideN2() +
|
||||
d1Idx * GetStrideD1() + bsIdx * GetStrideBlockSize() + d0Idx * GetStrideD0();
|
||||
return offset;
|
||||
}
|
||||
|
||||
// Get Stride
|
||||
__aicore__ inline uint64_t GetStrideBlockNum()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD1()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideBlockSize()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.stride);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetStrideD0()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_4>(gmLayout.stride);
|
||||
}
|
||||
|
||||
// Get Dim
|
||||
__aicore__ inline uint64_t GetN2()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_0>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetD1()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_1>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetBlockSize()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_2>(gmLayout.shape);
|
||||
}
|
||||
|
||||
__aicore__ inline uint64_t GetD0()
|
||||
{
|
||||
return AscendC::Std::get<SHAPE_AXIS_DIM_3>(gmLayout.shape);
|
||||
}
|
||||
};
|
||||
|
||||
template <GmFormat FORMAT>
|
||||
struct OffsetCalculator : public OffsetCalculatorImpl<FORMAT, GmLayoutParams<FORMAT>::CATEGORY> {
|
||||
};
|
||||
|
||||
// ----------------------------------------------CopyQueryGmToL1--------------------------------
|
||||
template <typename Q_T, GmFormat FORMAT>
|
||||
struct FaGmTensor {
|
||||
GlobalTensor<Q_T> gmTensor;
|
||||
OffsetCalculator<FORMAT> offsetCalculator;
|
||||
};
|
||||
|
||||
enum class L1Format
|
||||
{
|
||||
NZ = 0
|
||||
};
|
||||
|
||||
template <typename Q_T, L1Format FORMAT>
|
||||
struct FaL1Tensor {
|
||||
LocalTensor<Q_T> tensor;
|
||||
uint32_t rowCount;
|
||||
};
|
||||
|
||||
struct GmCoord {
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx;
|
||||
uint32_t gS1Idx;
|
||||
uint32_t dIdx;
|
||||
uint32_t gS1DealSize;
|
||||
uint32_t dDealSize;
|
||||
};
|
||||
#endif
|
||||
@@ -0,0 +1,95 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_common_arch35.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_AHSREDKV_COMMON_ARCH35_H
|
||||
#define KV_QUANT_SPARSE_ATTN_AHSREDKV_COMMON_ARCH35_H
|
||||
#include <type_traits>
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common.h"
|
||||
|
||||
constexpr uint64_t BLOCK_BYTE = 32;
|
||||
constexpr uint32_t NEGATIVE_MIN_VAULE_FP32 = 0xFF7FFFFF;
|
||||
|
||||
constexpr uint32_t L0AB_SHARED_SIZE_64K = 65536; // 65536表示64*1024
|
||||
constexpr uint32_t L0C_SHARED_SIZE_256K = 262144; // 262144表示256 * 1024
|
||||
|
||||
constexpr uint32_t BUFFER_SIZE_16K = 16384; // 16384表示16 * 1024
|
||||
constexpr uint32_t BUFFER_SIZE_32K = 32768; // 32768表示32 * 1024
|
||||
constexpr uint32_t BUFFER_SIZE_128K = 131072; // 131072表示128 * 1024
|
||||
|
||||
constexpr uint32_t CV_RATIO = 2;
|
||||
constexpr uint64_t SYNC_MODE = 4;
|
||||
|
||||
namespace BaseApi {
|
||||
__aicore__ constexpr uint64_t Align2Func(uint64_t data) {
|
||||
return (data + 1UL) >> 1UL << 1UL; // 向上2对齐, +1移位2
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align8Func(uint64_t data) {
|
||||
return (data + 7UL) >> 3UL << 3UL; // 向上8对齐, +7移位3
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align16Func(uint64_t data) {
|
||||
return (data + 15UL) >> 4UL << 4UL; // 向上16对齐, +15移位4
|
||||
}
|
||||
|
||||
__aicore__ constexpr uint64_t Align64Func(uint64_t data) {
|
||||
return (data + 63UL) >> 6UL << 6UL; // 向上64对齐, +63移位6
|
||||
}
|
||||
}
|
||||
|
||||
#define TEMPLATE_INTF \
|
||||
template <typename Q_T, typename KV_T, typename T, typename OUTPUT_T, bool isFd, bool isPa, SAS_LAYOUT LAYOUT_T, \
|
||||
SAS_LAYOUT KV_LAYOUT_T, SASTemplateMode TEMPLATE_MODE, bool IS_SPLIT_G>
|
||||
|
||||
#define TEMPLATE_INTF_ARGS \
|
||||
Q_T, KV_T, T, OUTPUT_T, isFd, isPa, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G
|
||||
|
||||
#define CUBE_BLOCK_TRAITS_TYPE_FIELDS(X) \
|
||||
X(Q_T) \
|
||||
X(KV_T) \
|
||||
X(T) \
|
||||
X(OUTPUT_T) \
|
||||
|
||||
#define CUBE_BLOCK_TRAITS_CONST_FIELDS(X) \
|
||||
X(isFd, bool, false) \
|
||||
X(isPa, bool, true) \
|
||||
X(LAYOUT_T, SAS_LAYOUT, SAS_LAYOUT::BSND) \
|
||||
X(KV_LAYOUT_T, SAS_LAYOUT, SAS_LAYOUT::PA_ND) \
|
||||
X(TEMPLATE_MODE, SASTemplateMode, SASTemplateMode::SCFA_TEMPLATE_MODE) \
|
||||
X(IS_SPLIT_G, bool, false) \
|
||||
|
||||
|
||||
/* 1. 生成带默认值的模版Template */
|
||||
#define GEN_TYPE_PARAM(name) typename name,
|
||||
#define GEN_CONST_PARAM(name, type, default_val) type name = default_val,
|
||||
|
||||
#define TEMPLATES_DEF \
|
||||
template <CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TYPE_PARAM) \
|
||||
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_CONST_PARAM) bool end = true>
|
||||
|
||||
/* 2. 生成不带带默认值的模版Template */
|
||||
#define GEN_TEMPLATE_TYPE_NODEF(name) typename name,
|
||||
#define GEN_TEMPLATE_CONST_NODEF(name, type, default_val) type name,
|
||||
#define TEMPLATES_DEF_NO_DEFAULT \
|
||||
template <CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TEMPLATE_TYPE_NODEF) \
|
||||
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TEMPLATE_CONST_NODEF) bool end>
|
||||
|
||||
/* 3. 生成有默认值的Args */
|
||||
#define GEN_ARG_NAME(name, ...) name,
|
||||
#define TEMPLATE_ARGS \
|
||||
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARG_NAME) \
|
||||
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARG_NAME) end
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,254 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_kvcache.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_KVCACHE_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_KVCACHE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
#include "util_regbase.h"
|
||||
|
||||
using namespace matmul;
|
||||
using namespace regbaseutil;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void GetSingleCoreParam(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
__gm__ int32_t *cuSeqlensQAddr, __gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t * actualSeqKvlenAddr)
|
||||
{
|
||||
int32_t actualS1Size = 0;
|
||||
int32_t actualS2Size = 0;
|
||||
int32_t actualSeqMin = 1;
|
||||
int32_t actualSeqKVMin = 1;
|
||||
int32_t sIdx = runParam.boIdx;
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
// actual seq length first
|
||||
actualS1Size = (actualSeqQlenAddr == nullptr) ? (cuSeqlensQAddr[sIdx + 1] - cuSeqlensQAddr[sIdx]) :
|
||||
actualSeqQlenAddr[sIdx];
|
||||
} else {
|
||||
actualS1Size = (actualSeqQlenAddr == nullptr) ? constInfo.s1Size :
|
||||
actualSeqQlenAddr[sIdx];
|
||||
}
|
||||
|
||||
if (constInfo.isActualLenDimsKVNull) {
|
||||
actualS2Size = constInfo.s2Size;
|
||||
} else {
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
actualS2Size = actualSeqKvlenAddr[sIdx];
|
||||
if ((sIdx > 0) && (!isPa)) {
|
||||
actualS2Size -= actualSeqKvlenAddr[sIdx - 1];
|
||||
}
|
||||
} else {
|
||||
actualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
|
||||
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
|
||||
}
|
||||
}
|
||||
|
||||
runParam.actualS1Size = actualS1Size;
|
||||
runParam.actualS2Size = actualS2Size;
|
||||
runParam.nextTokensPerBatch = runParam.actualS2Size - runParam.actualS1Size;
|
||||
if (constInfo.oriWinLeft == -1) {
|
||||
runParam.preTokensPerBatch = runParam.actualS1Size;
|
||||
} else {
|
||||
runParam.preTokensPerBatch = -(runParam.actualS2Size - runParam.actualS1Size - constInfo.oriWinLeft);
|
||||
}
|
||||
runParam.preTokensPerBatch = Min(runParam.preTokensPerBatch, runParam.actualS1Size);
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeParamBatch(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
__gm__ int32_t *cuSeqlensQAddr, __gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t *actualSeqKvlenAddr)
|
||||
{
|
||||
GetSingleCoreParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, cuSeqlensQAddr, actualSeqQlenAddr, actualSeqKvlenAddr);
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeS1LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo, bool lastBN,
|
||||
int64_t nextGs1Idx, int64_t gS1StartIdx)
|
||||
{
|
||||
// 计算每个基本快可以拷贝多少行s
|
||||
runParam.qSNumInOneBlock = (constInfo.gSize <= 64) ? \
|
||||
(constInfo.s1BaseSize / constInfo.gSize) : 1;
|
||||
runParam.gs1LoopStartIdx = gS1StartIdx;
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
int64_t gs1LoopStartIdx = runParam.nextTokensPerBatch * (-1) / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock;
|
||||
if (gs1LoopStartIdx > gS1StartIdx) {
|
||||
runParam.gs1LoopStartIdx = gs1LoopStartIdx;
|
||||
}
|
||||
}
|
||||
|
||||
int32_t gs1LoopEndIdx = 0;
|
||||
if constexpr (TEMPLATE_MODE == SASTemplateMode::SCFA_TEMPLATE_MODE) {
|
||||
gs1LoopEndIdx = runParam.actualS1Size; // 对于SCFA, 不切G轴, 每次拷贝一行的topk,只算一行的qs
|
||||
} else { // SWA/CFA
|
||||
// 不需要取topk, 每次计算gSize行, 循环qs次
|
||||
gs1LoopEndIdx = (runParam.actualS1Size + runParam.qSNumInOneBlock - 1) / runParam.qSNumInOneBlock;
|
||||
}
|
||||
// 不是最后一个bn, 赋值souterBlockNum
|
||||
if (!lastBN) {
|
||||
runParam.gs1LoopEndIdx = gs1LoopEndIdx;
|
||||
} else { // 最后一个bn, 从数组下一个元素取值
|
||||
runParam.gs1LoopEndIdx = nextGs1Idx == 0 ? gs1LoopEndIdx : nextGs1Idx;
|
||||
}
|
||||
|
||||
if (runParam.gs1LoopStartIdx > runParam.gs1LoopEndIdx) {
|
||||
runParam.gs1LoopStartIdx = runParam.gs1LoopEndIdx;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void ComputeSouterParam(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
uint32_t sOuterLoopIdx)
|
||||
{
|
||||
int64_t cubeSOuterOffset = sOuterLoopIdx * runParam.qSNumInOneBlock;
|
||||
if (runParam.actualS1Size == 0) {
|
||||
runParam.s1RealSize = 0;
|
||||
runParam.mRealSize = 0;
|
||||
} else {
|
||||
runParam.s1RealSize = Min(runParam.qSNumInOneBlock, runParam.actualS1Size - cubeSOuterOffset);
|
||||
runParam.mRealSize = runParam.s1RealSize * constInfo.gSize;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
runParam.mRealSize = runParam.mRealSize >> 1;
|
||||
}
|
||||
}
|
||||
|
||||
runParam.cubeMOuterOffset = cubeSOuterOffset * constInfo.gSize;
|
||||
runParam.halfMRealSize = (runParam.mRealSize + 1) >> 1;
|
||||
runParam.firstHalfMRealSize = runParam.halfMRealSize;
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.halfMRealSize = runParam.mRealSize - runParam.halfMRealSize;
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset + runParam.firstHalfMRealSize;
|
||||
} else {
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset;
|
||||
}
|
||||
|
||||
runParam.halfS1RealSize = (runParam.s1RealSize + 1) >> 1;
|
||||
runParam.firstHalfS1RealSize = runParam.halfS1RealSize;
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.halfS1RealSize = runParam.s1RealSize - runParam.halfS1RealSize;
|
||||
runParam.sOuterOffset = cubeSOuterOffset + runParam.halfMRealSize / constInfo.gSize;
|
||||
} else {
|
||||
runParam.sOuterOffset = cubeSOuterOffset;
|
||||
}
|
||||
runParam.cubeSOuterOffset = cubeSOuterOffset;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void LoopSOuterOffsetInit(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
int32_t sIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
int64_t seqOffset = 0;
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
seqOffset = cuSeqlensQAddr[sIdx];
|
||||
} else {
|
||||
seqOffset = sIdx * constInfo.s1Size;
|
||||
}
|
||||
|
||||
int64_t attentionOutSeqOffset = seqOffset * constInfo.n2GDv;
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::BSND || LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
runParam.attentionOutOffset = attentionOutSeqOffset +
|
||||
runParam.sOuterOffset * constInfo.n2GDv + runParam.n2oIdx * constInfo.gDv +
|
||||
runParam.goIdx * constInfo.dSizeV;
|
||||
}
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.attentionOutOffset += runParam.halfMRealSize * constInfo.dSizeV;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeParamS1(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
uint32_t sOuterLoopIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
if (runParam.s1oIdx < (runParam.nextTokensPerBatch * (-1)) / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
ComputeSouterParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, sOuterLoopIdx);
|
||||
|
||||
LoopSOuterOffsetInit<TEMPLATE_INTF_ARGS>(runParam, constInfo, runParam.boIdx, cuSeqlensQAddr);
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeLastBN(RunParamStr& runParam, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
// TND格式下 相邻Batch中当actualSeqQlen相等时则返回true
|
||||
if (runParam.boIdx > 0 && cuSeqlensQAddr[runParam.boIdx + 1] - cuSeqlensQAddr[runParam.boIdx] == 0) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline int64_t ClipSInnerTokenCube(int64_t sInnerToken, int64_t minValue, int64_t maxValue)
|
||||
{
|
||||
sInnerToken = sInnerToken > minValue ? sInnerToken : minValue;
|
||||
sInnerToken = sInnerToken < maxValue ? sInnerToken : maxValue;
|
||||
return sInnerToken;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeS2LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo)
|
||||
{
|
||||
if (runParam.actualS2Size == 0) {
|
||||
runParam.oriKvLoopEndIdx = 0;
|
||||
runParam.cmpKvLoopEndIdx = 0;
|
||||
runParam.s2LoopEndIdx = 0;
|
||||
return true;
|
||||
}
|
||||
uint32_t s2BaseSize = constInfo.s2BaseSize;
|
||||
|
||||
runParam.s2LineStartIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(runParam.cubeSOuterOffset - runParam.preTokensPerBatch,
|
||||
0, runParam.actualS2Size);
|
||||
runParam.s2LineEndIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(runParam.cubeSOuterOffset + runParam.nextTokensPerBatch +
|
||||
runParam.s1RealSize, 0, runParam.actualS2Size);
|
||||
runParam.oriKvLoopEndIdx = (runParam.s2LineEndIdx - runParam.s2LineStartIdx + s2BaseSize - 1) / s2BaseSize;
|
||||
if constexpr (TEMPLATE_MODE == SASTemplateMode::SWA_TEMPLATE_MODE) {
|
||||
runParam.cmpKvLoopEndIdx = 0;
|
||||
runParam.s2CmpLineEndIdx = 0;
|
||||
} else if constexpr (TEMPLATE_MODE == SASTemplateMode::CFA_TEMPLATE_MODE) {
|
||||
runParam.s2CmpLineEndIdx = runParam.s2LineEndIdx / constInfo.cmpRatio;
|
||||
runParam.cmpKvLoopEndIdx = (runParam.s2CmpLineEndIdx + s2BaseSize - 1) / s2BaseSize;
|
||||
} else { // SCFA_TEMPLATE_MODE
|
||||
runParam.s2CmpLineEndIdx = Min(runParam.s2LineEndIdx / constInfo.cmpRatio, constInfo.sparseBlockCount); // 当前LI输出的block size只可能是1
|
||||
runParam.cmpKvLoopEndIdx = (runParam.s2CmpLineEndIdx + s2BaseSize - 1) / s2BaseSize;
|
||||
}
|
||||
runParam.s2LoopEndIdx = runParam.oriKvLoopEndIdx + runParam.cmpKvLoopEndIdx;
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void InitTaskParamByRun(const RunParamStr& runParam, RunInfo &runInfo)
|
||||
{
|
||||
runInfo.boIdx = runParam.boIdx;
|
||||
runInfo.preTokensPerBatch = runParam.preTokensPerBatch;
|
||||
runInfo.nextTokensPerBatch = runParam.nextTokensPerBatch;
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.softmaxLseOffset = runParam.softmaxLseOffset;
|
||||
runInfo.qSNumInOneBlock = runParam.qSNumInOneBlock;
|
||||
runInfo.oriKvLoopEndIdx = runParam.oriKvLoopEndIdx;
|
||||
runInfo.cmpKvLoopEndIdx = runParam.cmpKvLoopEndIdx;
|
||||
runInfo.isCmp = runInfo.s2LoopCount >= runInfo.oriKvLoopEndIdx;
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHAREDKV_KVCACHE_H
|
||||
@@ -0,0 +1,358 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_scfa_block_cube.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H
|
||||
#include "common/offset_calculator.h"
|
||||
#include "common/matmul.h"
|
||||
#include "common/FixpipeOut.h"
|
||||
#include "common/CopyInL1.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
|
||||
#include "util_regbase.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
using namespace fa_base_matmul;
|
||||
namespace BaseApi {
|
||||
struct CubeCoordInfo {
|
||||
uint32_t curBIdx;
|
||||
uint32_t s1Coord;
|
||||
uint32_t s2Coord;
|
||||
};
|
||||
|
||||
template <SAS_LAYOUT LAYOUT>
|
||||
__aicore__ inline constexpr GmFormat GetQueryGmFormat()
|
||||
{
|
||||
if constexpr (LAYOUT == SAS_LAYOUT::BSND) {
|
||||
return GmFormat::BSNGD;
|
||||
} else {
|
||||
return GmFormat::TNGD;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF
|
||||
class SCFABlockCube {
|
||||
public:
|
||||
/* =================编译期常量的基本块信息================= */
|
||||
static constexpr uint32_t s1BaseSize = 64;
|
||||
static constexpr uint32_t s2BaseSize = 128;
|
||||
static constexpr uint32_t dBaseSize = 512;
|
||||
static constexpr uint32_t dBaseMatmulSize = 128;
|
||||
|
||||
__aicore__ inline SCFABlockCube() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *l1BufferManagerPtr, __gm__ uint8_t *query);
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
|
||||
__aicore__ inline void IterateBmm1(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &output,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
__aicore__ inline void IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
|
||||
private:
|
||||
__aicore__ inline void InitLocalBuffer();
|
||||
__aicore__ inline void InitGmTensor(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
|
||||
__aicore__ inline void CalcS1Coord(RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
__aicore__ inline void IterateBmm1SCFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
// --------------------Bmm2--------------------------
|
||||
__aicore__ inline void IterateBmm2SCFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
TPipe *tPipe;
|
||||
/* =====================GM变量==================== */
|
||||
static constexpr GmFormat Q_FORMAT = GetQueryGmFormat<LAYOUT_T>();
|
||||
FaGmTensor<Q_T, Q_FORMAT> queryGm;
|
||||
|
||||
/* =====================运行时变量==================== */
|
||||
CubeCoordInfo coordInfo[3];
|
||||
TEventID mte1ToMte2Id[3];
|
||||
TEventID mte2ToMte1Id[3];
|
||||
/* =====================LocalBuffer变量==================== */
|
||||
BufferManager<BufferType::L1> *l1BufferManagerPtr;
|
||||
BufferManager<BufferType::L0A> l0aBufferManager;
|
||||
BufferManager<BufferType::L0B> l0bBufferManager;
|
||||
BufferManager<BufferType::L0C> l0cBufferManager;
|
||||
|
||||
// D小于等于256 mm1左矩阵Q,GS1循环内左矩阵复用, GS1循环间开pingpong;D大于256使用单块Buffer,S1循环间驻留;fp32场景单块不驻留
|
||||
BuffersPolicySingleBuffer<BufferType::L1> l1QBuffers;
|
||||
|
||||
// mm1右矩阵K
|
||||
BuffersPolicy3buff<BufferType::L1> l1KBuffers;
|
||||
|
||||
// L0A
|
||||
BuffersPolicyDB<BufferType::L0A> mmL0ABuffers;
|
||||
// L0B
|
||||
BuffersPolicyDB<BufferType::L0B> mmL0BBuffers;
|
||||
// L0C
|
||||
BuffersPolicyDB<BufferType::L0C> mmL0CBuffers;
|
||||
};
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::InitCubeBlock(
|
||||
TPipe *pipe, BufferManager<BufferType::L1> *l1BuffMgr, __gm__ uint8_t *query)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
tPipe = pipe;
|
||||
l1BufferManagerPtr = l1BuffMgr;
|
||||
this->queryGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)query);
|
||||
InitLocalBuffer();
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::InitCubeInput(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
InitGmTensor(cuSeqlensQ, constInfo);
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
mte1ToMte2Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte1ToMte2Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte1ToMte2Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
|
||||
mte2ToMte1Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
mte2ToMte1Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
mte2ToMte1Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::InitLocalBuffer()
|
||||
{
|
||||
constexpr uint32_t mm1LeftSize = s1BaseSize * dBaseSize * sizeof(Q_T);
|
||||
constexpr uint32_t mm1RightSize = dBaseSize * s2BaseSize * sizeof(Q_T);
|
||||
l1QBuffers.Init((*l1BufferManagerPtr), mm1LeftSize);
|
||||
l1KBuffers.Init((*l1BufferManagerPtr), mm1RightSize);
|
||||
|
||||
// L0A B C 当前写死,能否通过基础api获取
|
||||
l0aBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
|
||||
l0bBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
|
||||
l0cBufferManager.Init(tPipe, L0C_SHARED_SIZE_256K);
|
||||
|
||||
mmL0ABuffers.Init(l0aBufferManager, BUFFER_SIZE_16K); // db类型,填入数值是总大小的一半
|
||||
mmL0BBuffers.Init(l0bBufferManager, BUFFER_SIZE_32K);
|
||||
mmL0CBuffers.Init(l0cBufferManager, BUFFER_SIZE_128K);
|
||||
}
|
||||
|
||||
/* 初始化GmTensor,设置shape信息并计算strides */
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::InitGmTensor(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::BSND) {
|
||||
this->queryGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize,
|
||||
constInfo.s1Size, constInfo.dSize);
|
||||
} else { // SAS_LAYOUT::TND
|
||||
GlobalTensor<int32_t> actualSeqQLen;
|
||||
actualSeqQLen.SetGlobalBuffer((__gm__ int32_t *)cuSeqlensQ);
|
||||
this->queryGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSize,
|
||||
actualSeqQLen, constInfo.actualSeqLenSize);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::CalcS1Coord(RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
// 计算s1方向偏移
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord = runInfo.s1oIdx * runInfo.qSNumInOneBlock;
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::IterateBmm1(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
CalcS1Coord(runInfo, constInfo);
|
||||
|
||||
IterateBmm1SCFA(outputBuf, inputRightBuf, v0ResGm, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
IterateBmm2SCFA(outputBuf, inputLeftBuffers, inputRightBuf, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::IterateBmm1SCFA(
|
||||
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
Buffer<BufferType::L1> inputLeftBuf;
|
||||
// 左矩阵复用,S2的第一次循环加载左矩阵
|
||||
// 加载左矩阵到L1, 全载
|
||||
// query对ori_kv, cmp_kv都一样,无需区分
|
||||
if (unlikely(runInfo.s2LoopCount == 0)) { // sOuter循环第一个基本块:搬运Q
|
||||
inputLeftBuf = l1QBuffers.Get();
|
||||
inputLeftBuf.Wait<HardEvent::MTE1_MTE2>(); // 占用L1A
|
||||
LocalTensor<Q_T> inputLeftTensor = inputLeftBuf.GetTensor<Q_T>();
|
||||
|
||||
uint64_t gmOffset = this->queryGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx, runInfo.goIdx,
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord, 0);
|
||||
CopyToL1Nd2Nz<Q_T>(inputLeftTensor, this->queryGm.gmTensor[gmOffset], runInfo.mRealSize, constInfo.dSize,
|
||||
constInfo.mm1Ka);
|
||||
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
} else { // 非S2的第一次循环直接复用Q
|
||||
inputLeftBuf = l1QBuffers.GetPre();
|
||||
// 左矩阵复用时,sinner循环内不需要MTE2同步等待
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
}
|
||||
|
||||
// 加载当前轮的右矩阵到L1
|
||||
inputRightBuf.WaitCrossCore(); // 核间同步,这里需要根据V0操作处理同步,确保取tensor时,数据已经准备好
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
|
||||
|
||||
LocalTensor<Q_T> dst = inputRightBuf.GetTensor<Q_T>();
|
||||
v0ResGm.WaitCrossCore();
|
||||
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
|
||||
DataCopy(dst, v0ResGmTensor, Align16Func(runInfo.s2RealSize) * constInfo.dSize);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
|
||||
}
|
||||
|
||||
inputLeftBuf.Wait<HardEvent::MTE2_MTE1>(); // 等待L1A
|
||||
Buffer<BufferType::L0C> mm1ResL0C = mmL0CBuffers.Get();
|
||||
mm1ResL0C.Wait<HardEvent::FIX_M>(); // 占用
|
||||
MMParam param = {static_cast<uint32_t>(runInfo.mRealSize), // singleM
|
||||
static_cast<uint32_t>(runInfo.s2RealSize), // singleN
|
||||
static_cast<uint32_t>(constInfo.dSize), // singleK
|
||||
0, // isLeftTranspose
|
||||
1 // isRightTranspose
|
||||
};
|
||||
MatmulK<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>( // m,n不切,k切128
|
||||
inputLeftBuf.GetTensor<Q_T>(), inputRightBuf.GetTensor<Q_T>(), // mm1B直接用tensor的数据
|
||||
mmL0ABuffers, mmL0BBuffers,
|
||||
mm1ResL0C.GetTensor<T>(),
|
||||
param);
|
||||
if (unlikely(runInfo.s2LoopCount == runInfo.s2LoopLimit)) {
|
||||
inputLeftBuf.Set<HardEvent::MTE1_MTE2>(); // 释放L1A
|
||||
}
|
||||
|
||||
mm1ResL0C.Set<HardEvent::M_FIX>(); // 通知
|
||||
mm1ResL0C.Wait<HardEvent::M_FIX>(); // 等待L0C
|
||||
|
||||
outputBuf.WaitCrossCore();
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB
|
||||
fixpipeParams.nSize = Align8Func(runInfo.s2RealSize); // L0C上的bmm1结果矩阵N方向的size大小; 同mmadParams.n; 为什么要8个元素对齐(32B对齐) // 128
|
||||
fixpipeParams.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小(必须为偶数) // 128
|
||||
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔), 单位为16*sizeof(T) // 源Nz矩阵中相邻大Z排布的起始地址偏移
|
||||
fixpipeParams.dstStride = s2BaseSize; // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到, ND方案(S1*S2)时脏数据用mask剔除
|
||||
fixpipeParams.dualDstCtl = 1; // 双目标模式,按M维度拆分,M / 2 * N写入每个UB, M必须为2的倍数
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
|
||||
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(), mm1ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
mm1ResL0C.Set<HardEvent::FIX_M>(); // 释放L0C
|
||||
outputBuf.SetCrossCore();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void SCFABlockCube<TEMPLATE_ARGS>::IterateBmm2SCFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> inputLeftBuf = inputLeftBuffers.Get(); // P直接用无需搬运
|
||||
inputLeftBuf.WaitCrossCore();
|
||||
|
||||
Buffer<BufferType::L0C> mm2ResL0C = mmL0CBuffers.Get();
|
||||
mm2ResL0C.Wait<HardEvent::FIX_M>(); // 占用
|
||||
MMParam param = {static_cast<uint32_t>(s1BaseSize), // singleM 64
|
||||
static_cast<uint32_t>(constInfo.dSizeV), // singleN 512
|
||||
static_cast<uint32_t>(runInfo.s2RealSize), // singleK 128
|
||||
0, // isLeftTranspose
|
||||
0 // isRightTranspose
|
||||
};
|
||||
MatmulN<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>(
|
||||
inputLeftBuf.GetTensor<Q_T>(),
|
||||
inputRightBuf.GetTensor<Q_T>(),
|
||||
mmL0ABuffers,
|
||||
mmL0BBuffers,
|
||||
mm2ResL0C.GetTensor<T>(),
|
||||
param);
|
||||
|
||||
mm2ResL0C.Set<HardEvent::M_FIX>(); // 通知
|
||||
mm2ResL0C.Wait<HardEvent::M_FIX>(); // 等待
|
||||
|
||||
outputBuf.WaitCrossCore(); //占用
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB;FixpipeParamsM300:L0C→UB
|
||||
fixpipeParams.nSize = Align8Func(constInfo.dSizeV); // L0C上的bmm1结果矩阵N方向的size大小, 分档计算且vector2中通过mask筛选出实际有效值
|
||||
fixpipeParams.mSize = s1BaseSize; // 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小; 同mmadParams.m
|
||||
fixpipeParams.srcStride = Align16Func(s1BaseSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔)
|
||||
fixpipeParams.dstStride = Align16Func(constInfo.dSizeV);
|
||||
fixpipeParams.dualDstCtl = 1;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(), mm2ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
mm2ResL0C.Set<HardEvent::FIX_M>(); // 释放
|
||||
|
||||
outputBuf.SetCrossCore();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF
|
||||
class SCFABlockCubeDummy {
|
||||
public:
|
||||
__aicore__ inline SCFABlockCubeDummy() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *l1BufferManagerPtr, __gm__ uint8_t *query) {}
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo) {}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct CubeBlockTraits; // 声明
|
||||
|
||||
/* 生成CubeBlockTraits */
|
||||
#define GEN_TRAIT_TYPE(name, ...) using name##_TRAITS = name;
|
||||
#define GEN_TRAIT_CONST(name, type, ...) static constexpr type name##Traits = name;
|
||||
|
||||
#define DEFINE_CUBE_BLOCK_TRAITS(CUBE_BLOCK_CLASS) \
|
||||
TEMPLATES_DEF_NO_DEFAULT \
|
||||
struct CubeBlockTraits<CUBE_BLOCK_CLASS<TEMPLATE_ARGS>> { \
|
||||
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TRAIT_TYPE) \
|
||||
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TRAIT_CONST) \
|
||||
}
|
||||
|
||||
DEFINE_CUBE_BLOCK_TRAITS(SCFABlockCube);
|
||||
DEFINE_CUBE_BLOCK_TRAITS(SCFABlockCubeDummy);
|
||||
|
||||
// /* 生成Arg Traits, kernel中只需要调用ARGS_TRAITS就可以获取所有CubeBlock中的模板参数 */
|
||||
#define GEN_ARGS_TYPE(name, ...) using name = typename CubeBlockTraits<CubeBlockType>::name##_TRAITS;
|
||||
#define GEN_ARGS_CONST(name, type, ...) static constexpr type name = CubeBlockTraits<CubeBlockType>::name##Traits;
|
||||
#define ARGS_TRAITS \
|
||||
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARGS_TYPE) \
|
||||
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARGS_CONST)
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,514 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_scfa_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H
|
||||
#include "kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_kvcache.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_scfa_block_cube.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_scfa_block_vector.h"
|
||||
#include "kernel_operator.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
#include "common/matmul.h"
|
||||
#include "common/FixpipeOut.h"
|
||||
#include "common/CopyInL1.h"
|
||||
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
|
||||
using matmul::MatmulType;
|
||||
using namespace AscendC;
|
||||
using namespace optiling;
|
||||
using namespace optiling::detail;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace BaseApi {
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
class KvQuantSparseAttnSharedkvScfa {
|
||||
public:
|
||||
ARGS_TRAITS;
|
||||
__aicore__ inline KvQuantSparseAttnSharedkvScfa() {};
|
||||
|
||||
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV,
|
||||
__gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable,
|
||||
__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedKv, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseAttnSharedkvTilingData *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void Process();
|
||||
private:
|
||||
__aicore__ inline void ProcessMainLoop();
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *cmpSparseIndices,
|
||||
__gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *cuSeqlensQ,
|
||||
__gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedKv, __gm__ uint8_t *sinks, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseAttnSharedkvTilingData *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void InitLocalBuffer();
|
||||
__aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace);
|
||||
__aicore__ inline void ComputeConstexpr();
|
||||
__aicore__ inline void SetRunInfo(RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount,
|
||||
int64_t s2LoopLimit, int64_t multiCoreInnerIdx);
|
||||
__aicore__ inline void ComputeBmm1Tail(RunInfo &runInfo, RunParamStr &runParam);
|
||||
__aicore__ inline void InitUniqueConstInfo();
|
||||
__aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam);
|
||||
__aicore__ inline void InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo);
|
||||
TPipe *pipe;
|
||||
|
||||
const KvQuantSparseAttnSharedkvTilingData *__restrict tilingData;
|
||||
static constexpr uint64_t SYNC_MODE = 4;
|
||||
static constexpr uint32_t PRELOAD_NUM = 2;
|
||||
/* 核间通道 */
|
||||
BufferManager<BufferType::GM> v0ResGmBufferManager;
|
||||
|
||||
BufferManager<BufferType::UB> ubBufferManager;
|
||||
BuffersPolicyDB<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm1Buffers;
|
||||
BuffersPolicySingleBuffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm2Buffers;
|
||||
|
||||
// mm2左矩阵P
|
||||
BufferManager<BufferType::L1> l1BufferManager;
|
||||
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> l1PBuffers;
|
||||
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> l1RightBuffers;
|
||||
CVSharedParams sharedParams;
|
||||
/* GM信息 */
|
||||
GlobalTensor<uint32_t> metadataGm;
|
||||
__gm__ int32_t *cuSeqlensQAddr = nullptr;
|
||||
__gm__ int32_t *actualSeqKvlenAddr = nullptr;
|
||||
__gm__ int32_t *actualSeqQlenAddr = nullptr;
|
||||
/* workspace 空间 */
|
||||
BuffersPolicy3buff<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> v0ResGmBuffers;
|
||||
/* 核Index信息 */
|
||||
int32_t aicIdx;
|
||||
|
||||
/* 初始化后不变的信息 */
|
||||
ConstInfo constInfo;
|
||||
|
||||
/* 模板库Block */
|
||||
CubeBlockType cubeBlock;
|
||||
VecBlockType vecBlock;
|
||||
};
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::Init(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV,
|
||||
__gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable,
|
||||
__gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedKv, __gm__ uint8_t *sinks, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseAttnSharedkvTilingData *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
fa_base_matmul::idCounterNum = 0;
|
||||
constInfo.subBlockIdx = GetSubBlockIdx();
|
||||
if ASCEND_IS_AIC {
|
||||
this->aicIdx = GetBlockIdx();
|
||||
constInfo.aivIdx = 0;
|
||||
} else {
|
||||
constInfo.aivIdx = GetBlockIdx();
|
||||
this->aicIdx = constInfo.aivIdx >> 1;
|
||||
this->tilingData = tiling;
|
||||
}
|
||||
|
||||
if (metadata == nullptr) {
|
||||
return;
|
||||
}
|
||||
this->metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata);
|
||||
|
||||
constInfo.s1BaseSize = 64;
|
||||
constInfo.s2BaseSize = 128;
|
||||
|
||||
this->pipe = tPipe;
|
||||
vecBlock.InitVecBlock(tPipe, this->tilingData, this->sharedParams, this->aicIdx, constInfo.subBlockIdx, cuSeqlensQ, sequsedKv);
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.bSize = this->sharedParams.bSize;
|
||||
constInfo.gSize = this->sharedParams.gSize;
|
||||
constInfo.s1Size = this->sharedParams.s1Size;
|
||||
constInfo.dSizeV = this->sharedParams.dSize;
|
||||
constInfo.needInit = this->sharedParams.needInit;
|
||||
}
|
||||
vecBlock.CleanOutput(attentionOut, constInfo);
|
||||
/* cube侧不依赖sharedParams的scalar前置 */
|
||||
InitMMResBuf(workspace);
|
||||
if ASCEND_IS_AIC {
|
||||
cubeBlock.InitCubeBlock(pipe, &l1BufferManager, query);
|
||||
/* wait kfc message */
|
||||
CrossCoreWaitFlag<SYNC_MODE, PIPE_S>(15);
|
||||
auto tempTilingSSbuf = reinterpret_cast<__ssbuf__ uint32_t*>(0); // 从ssbuf的0地址开始拷贝
|
||||
auto tempTiling = reinterpret_cast<uint32_t *>(&sharedParams);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < sizeof(CVSharedParams) / sizeof(uint32_t); ++i, ++tempTilingSSbuf, ++tempTiling) {
|
||||
*tempTiling = *tempTilingSSbuf;
|
||||
}
|
||||
}
|
||||
this->ComputeConstexpr();
|
||||
this->InitGlobalBuffer(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, sequsedQ, sequsedKv, sinks,
|
||||
workspace, tiling, tPipe); // gm设置
|
||||
this->InitLocalBuffer();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::InitGlobalBuffer(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV, __gm__ uint8_t *cmpSparseIndices,
|
||||
__gm__ uint8_t *oriBlockTable, __gm__ uint8_t *cmpBlockTable, __gm__ uint8_t *cuSeqlensQ,
|
||||
__gm__ uint8_t *sequsedQ, __gm__ uint8_t *sequsedKv, __gm__ uint8_t *sinks, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseAttnSharedkvTilingData *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
if (cuSeqlensQ != nullptr) {
|
||||
cuSeqlensQAddr = (__gm__ int32_t *)cuSeqlensQ;
|
||||
}
|
||||
if (sequsedKv != nullptr) {
|
||||
actualSeqKvlenAddr = (__gm__ int32_t *)sequsedKv;
|
||||
}
|
||||
|
||||
if (sequsedQ != nullptr) {
|
||||
actualSeqQlenAddr = (__gm__ int32_t *)sequsedQ;
|
||||
}
|
||||
|
||||
vecBlock.InitGlobalBuffer(oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, sequsedQ, sinks);
|
||||
cubeBlock.InitCubeInput(cuSeqlensQ, constInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::InitMMResBuf(__gm__ uint8_t *workspace)
|
||||
{
|
||||
uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize * sizeof(T);
|
||||
uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512 * sizeof(T);
|
||||
uint32_t mm2LeftSize = constInfo.s1BaseSize * constInfo.s2BaseSize * sizeof(Q_T);
|
||||
uint32_t mm1RightSize = constInfo.s2BaseSize * 512 * sizeof(Q_T);
|
||||
l1BufferManager.Init(pipe, 524288); // 512 * 1024
|
||||
// 保存p结果的L1内存必须放在第一个L1 policy上,保证和vec申请的地址相同
|
||||
l1PBuffers.Init(l1BufferManager, mm2LeftSize);
|
||||
l1RightBuffers.Init(l1BufferManager, mm1RightSize);
|
||||
|
||||
ubBufferManager.Init(pipe, mm1ResultSize * 2 + mm2ResultSize);
|
||||
bmm2Buffers.Init(ubBufferManager, mm2ResultSize);
|
||||
if ASCEND_IS_AIV {
|
||||
bmm2Buffers.Get().SetCrossCore();
|
||||
}
|
||||
bmm1Buffers.Init(ubBufferManager, mm1ResultSize);
|
||||
if ASCEND_IS_AIV {
|
||||
bmm1Buffers.Get().SetCrossCore();
|
||||
bmm1Buffers.Get().SetCrossCore();
|
||||
}
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
uint32_t v0ResSize = constInfo.s2BaseSize * 512U * sizeof(Q_T);
|
||||
int64_t totalOffset = v0ResSize * 3 * (aicIdx >> 1U);
|
||||
v0ResGmBufferManager.Init(workspace + totalOffset);
|
||||
v0ResGmBuffers.Init(v0ResGmBufferManager, v0ResSize);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::InitLocalBuffer()
|
||||
{
|
||||
vecBlock.InitLocalBuffer(pipe, constInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::ComputeConstexpr()
|
||||
{
|
||||
// 计算轴的乘积
|
||||
if ASCEND_IS_AIC {
|
||||
constInfo.bSize = this->sharedParams.bSize;
|
||||
constInfo.gSize = this->sharedParams.gSize;
|
||||
constInfo.s1Size = this->sharedParams.s1Size;
|
||||
constInfo.dSizeV = this->sharedParams.dSize;
|
||||
constInfo.needInit = this->sharedParams.needInit;
|
||||
}
|
||||
constInfo.n2Size = sharedParams.n2Size;
|
||||
constInfo.s2Size = sharedParams.s2Size;
|
||||
constInfo.dSize = sharedParams.dSize;
|
||||
constInfo.dSizeVInput = sharedParams.dSizeVInput;
|
||||
constInfo.dSizeRope = sharedParams.dSizeRope;
|
||||
constInfo.dSizeNope = constInfo.dSize - constInfo.dSizeRope;
|
||||
constInfo.tileSize = sharedParams.tileSize;
|
||||
constInfo.sparseBlockCount = sharedParams.sparseBlockCount;
|
||||
constInfo.sparseBlockSize = 1;
|
||||
constInfo.cmpRatio = sharedParams.cmpRatio;
|
||||
constInfo.oriWinLeft = sharedParams.oriWinLeft;
|
||||
constInfo.oriWinRight = sharedParams.oriWinRight;
|
||||
constInfo.s1S2 = constInfo.s1Size * constInfo.s2Size;
|
||||
constInfo.gS1 = constInfo.gSize * constInfo.s1Size;
|
||||
constInfo.n2G = constInfo.n2Size * constInfo.gSize;
|
||||
|
||||
constInfo.s1Dv = constInfo.s1Size * constInfo.dSizeV;
|
||||
constInfo.s2Dv = constInfo.s2Size * constInfo.dSizeV;
|
||||
constInfo.n2Dv = constInfo.n2Size * constInfo.dSizeV;
|
||||
constInfo.gDv = constInfo.gSize * constInfo.dSizeV;
|
||||
constInfo.gS1Dv = constInfo.gSize * constInfo.s1Dv;
|
||||
constInfo.n2S2Dv = constInfo.n2Size * constInfo.s2Dv;
|
||||
constInfo.n2GDv = constInfo.n2Size * constInfo.gDv;
|
||||
constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv;
|
||||
constInfo.n2GS1Dv = constInfo.n2Size * constInfo.gS1Dv;
|
||||
|
||||
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
|
||||
// (BS)ND
|
||||
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
|
||||
|
||||
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
|
||||
}
|
||||
} else if constexpr (LAYOUT_T == SAS_LAYOUT::BSND) {
|
||||
// BSH/BSNGD
|
||||
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
|
||||
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
|
||||
}
|
||||
}
|
||||
if ASCEND_IS_AIV {
|
||||
constInfo.softmaxScale = sharedParams.softmaxScale;
|
||||
constInfo.oriKvStride = sharedParams.oriKvStride;
|
||||
constInfo.cmpKvStride = sharedParams.cmpKvStride;
|
||||
constInfo.oriBlockSize = sharedParams.oriBlockSize;
|
||||
constInfo.cmpBlockSize = sharedParams.cmpBlockSize;
|
||||
constInfo.oriMaxBlockNumPerBatch = sharedParams.oriMaxBlockNumPerBatch;
|
||||
constInfo.cmpMaxBlockNumPerBatch = sharedParams.cmpMaxBlockNumPerBatch;
|
||||
}
|
||||
|
||||
InitUniqueConstInfo();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::InitUniqueConstInfo()
|
||||
{
|
||||
this->constInfo.actualSeqLenSize = this->sharedParams.bSize + 1;
|
||||
this->constInfo.actualSeqLenKVSize = this->sharedParams.bSize;
|
||||
this->constInfo.isActualLenDimsKVNull = static_cast<bool>(this->sharedParams.isActualSeqLengthsKVNull);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::Process()
|
||||
{
|
||||
// SyncAll Cube和Vector都需要调用
|
||||
if (this->sharedParams.needInit) {
|
||||
SyncAll<false>();
|
||||
}
|
||||
ProcessMainLoop();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::ProcessMainLoop()
|
||||
{
|
||||
uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_CORE_ENABLE_INDEX, false));
|
||||
int64_t maxS2LoopCnt = 0;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt = static_cast<int64_t>(metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_MAX_NUM, false)));
|
||||
}
|
||||
if (hasLoad == 0) {
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) {
|
||||
CrossCoreSetFlag<0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// 从meta data解析分核信息
|
||||
uint32_t bN2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_START_INDEX, false));
|
||||
uint32_t gS1StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_START_INDEX, false));
|
||||
uint32_t s2StartIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_START_INDEX, false));
|
||||
uint32_t bN2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_BN2_END_INDEX, false));
|
||||
uint32_t nextGs1Idx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_M_END_INDEX, false));
|
||||
uint32_t s2EndIdx = metadataGm.GetValue(GetAttrAbsIndex(aicIdx, FA_S2_END_INDEX, false));
|
||||
uint32_t s2LoopLimit = 0;
|
||||
|
||||
if (nextGs1Idx != 0) {
|
||||
bN2EndIdx++;
|
||||
}
|
||||
|
||||
int64_t taskId = 0;
|
||||
bool notLast = true;
|
||||
RunInfo runInfo[3];
|
||||
RunParamStr runParam;
|
||||
int64_t multiCoreInnerIdx = 1;
|
||||
for (int64_t bnIdx = bN2StartIdx; bnIdx < bN2EndIdx; bnIdx++) {
|
||||
bool lastBN = (bnIdx == bN2EndIdx - 1);
|
||||
runParam.boIdx = bnIdx;
|
||||
runParam.n2oIdx = 0;
|
||||
ComputeParamBatch<TEMPLATE_INTF_ARGS>(runParam, this->constInfo,
|
||||
this->cuSeqlensQAddr, this->actualSeqQlenAddr, this->actualSeqKvlenAddr);
|
||||
ComputeS1LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo, lastBN, nextGs1Idx, gS1StartIdx);
|
||||
|
||||
int64_t gS1LoopEnd = lastBN ? (runParam.gs1LoopEndIdx + PRELOAD_NUM) : runParam.gs1LoopEndIdx;
|
||||
for (int64_t gS1Index = runParam.gs1LoopStartIdx; gS1Index < gS1LoopEnd; gS1Index++) {
|
||||
bool notLastTwoLoop = true;
|
||||
if (lastBN) {
|
||||
int32_t extraGS1 = gS1Index - runParam.gs1LoopEndIdx;
|
||||
switch (extraGS1) {
|
||||
case 0:
|
||||
notLastTwoLoop = false;
|
||||
break;
|
||||
case 1:
|
||||
notLast = false;
|
||||
notLastTwoLoop = false;
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (notLastTwoLoop) {
|
||||
this->ComputeAxisIdxByBnAndGs1(bnIdx, gS1Index, runParam);
|
||||
bool s1NoNeedCalc = ComputeParamS1<TEMPLATE_INTF_ARGS>(
|
||||
runParam, this->constInfo, gS1Index, this->cuSeqlensQAddr);
|
||||
bool s2NoNeedCalc =
|
||||
ComputeS2LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo);
|
||||
// s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分
|
||||
if (s1NoNeedCalc || s2NoNeedCalc) {
|
||||
continue;
|
||||
}
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt -= runParam.s2LoopEndIdx;
|
||||
}
|
||||
s2LoopLimit = runParam.s2LoopEndIdx - 1;
|
||||
} else {
|
||||
s2LoopLimit = 0;
|
||||
}
|
||||
for (int64_t s2LoopCount = 0; s2LoopCount <= s2LoopLimit; ++s2LoopCount) {
|
||||
if (notLastTwoLoop) {
|
||||
RunInfo &runInfo1 = runInfo[taskId % 3];
|
||||
this->SetRunInfo(runInfo1, runParam, taskId, s2LoopCount, s2LoopLimit, multiCoreInnerIdx);
|
||||
if ASCEND_IS_AIC {
|
||||
this->cubeBlock.IterateBmm1(this->bmm1Buffers.Get(), this->l1RightBuffers.Get(), v0ResGmBuffers.Get(), runInfo1,
|
||||
this->constInfo);
|
||||
} else {
|
||||
this->vecBlock.ProcessVec0(this->l1RightBuffers.Get(), v0ResGmBuffers.Get(), runInfo1, this->constInfo);
|
||||
}
|
||||
} else {
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
if (maxS2LoopCnt > 0) {
|
||||
maxS2LoopCnt--;
|
||||
CrossCoreSetFlag<0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if (taskId > 0 && notLast) {
|
||||
auto &runInfo2 = runInfo[(taskId + 2) % 3];
|
||||
if ASCEND_IS_AIV {
|
||||
this->vecBlock.ProcessVec1(this->l1PBuffers.Get(), this->bmm1Buffers.Get(), runInfo2,
|
||||
this->constInfo);
|
||||
} else {
|
||||
RunInfo &runInfo2 = runInfo[(taskId + 2) % 3];
|
||||
this->cubeBlock.IterateBmm2(this->bmm2Buffers.Get(), this->l1PBuffers, this->l1RightBuffers.GetReused(), runInfo2,
|
||||
this->constInfo);
|
||||
}
|
||||
}
|
||||
if (taskId > 1) {
|
||||
if ASCEND_IS_AIV {
|
||||
RunInfo &runInfo3 = runInfo[(taskId + 1) % 3];
|
||||
this->vecBlock.ProcessVec2(this->bmm2Buffers.Get(), runInfo3, this->constInfo);
|
||||
}
|
||||
}
|
||||
++taskId;
|
||||
}
|
||||
++multiCoreInnerIdx;
|
||||
}
|
||||
gS1StartIdx = 0;
|
||||
}
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) {
|
||||
CrossCoreSetFlag<0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::ComputeAxisIdxByBnAndGs1(
|
||||
int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam)
|
||||
{
|
||||
// GS1合轴, 不切G, 只切S1
|
||||
runParam.s1oIdx = gS1Index * runParam.qSNumInOneBlock;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
runParam.goIdx = (aicIdx % 2 == 0) ? 0 : 64;
|
||||
} else {
|
||||
runParam.goIdx = 0;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::SetRunInfo(
|
||||
RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount, int64_t s2LoopLimit, int64_t multiCoreInnerIdx)
|
||||
{
|
||||
if (s2LoopCount < runParam.oriKvLoopEndIdx) {
|
||||
runInfo.s2StartIdx = runParam.s2LineStartIdx;
|
||||
runInfo.s2EndIdx = runParam.s2LineEndIdx;
|
||||
} else {
|
||||
runInfo.s2StartIdx = 0;
|
||||
runInfo.s2EndIdx = runParam.s2CmpLineEndIdx;
|
||||
}
|
||||
runInfo.s2LoopCount = s2LoopCount;
|
||||
if (runInfo.multiCoreInnerIdx != multiCoreInnerIdx) {
|
||||
runInfo.s1oIdx = runParam.s1oIdx;
|
||||
runInfo.boIdx = runParam.boIdx;
|
||||
runInfo.n2oIdx = runParam.n2oIdx;
|
||||
runInfo.goIdx = runParam.goIdx;
|
||||
runInfo.multiCoreInnerIdx = multiCoreInnerIdx;
|
||||
runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1;
|
||||
runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3;
|
||||
}
|
||||
|
||||
runInfo.taskId = taskId;
|
||||
runInfo.taskIdMod2 = taskId & 1;
|
||||
runInfo.taskIdMod3 = taskId % 3;
|
||||
runInfo.s2LoopLimit = s2LoopLimit;
|
||||
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.attentionOutOffset = runParam.attentionOutOffset;
|
||||
runInfo.sOuterOffset = runParam.sOuterOffset;
|
||||
this->ComputeBmm1Tail(runInfo, runParam);
|
||||
InitUniqueRunInfo(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void
|
||||
KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::InitUniqueRunInfo(
|
||||
const RunParamStr &runParam, RunInfo &runInfo)
|
||||
{
|
||||
InitTaskParamByRun<TEMPLATE_INTF_ARGS>(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseAttnSharedkvScfa<CubeBlockType, VecBlockType>::ComputeBmm1Tail(
|
||||
RunInfo &runInfo, RunParamStr &runParam)
|
||||
{
|
||||
// ------------------------S1 Base Related---------------------------
|
||||
runInfo.s1RealSize = runParam.s1RealSize;
|
||||
runInfo.halfS1RealSize = runParam.halfS1RealSize;
|
||||
runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize;
|
||||
runInfo.mRealSize = runParam.mRealSize;
|
||||
runInfo.halfMRealSize = runParam.halfMRealSize;
|
||||
runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize;
|
||||
|
||||
runInfo.vec2S1BaseSize = runInfo.halfS1RealSize; // D>128 这里需要适配
|
||||
runInfo.vec2MBaseSize = runInfo.halfMRealSize;
|
||||
|
||||
// ------------------------S2 Base Related----------------------------
|
||||
runInfo.s2RealSize = constInfo.s2BaseSize;
|
||||
runInfo.s2AlignedSize = runInfo.s2RealSize;
|
||||
int64_t curS2LoopCnt = (runInfo.s2LoopCount >= runParam.oriKvLoopEndIdx) ? (runInfo.s2LoopCount - runParam.oriKvLoopEndIdx) : runInfo.s2LoopCount;
|
||||
if (runInfo.s2StartIdx + (curS2LoopCnt + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) {
|
||||
runInfo.s2RealSize = runInfo.s2EndIdx - curS2LoopCnt * runInfo.s2RealSize - runInfo.s2StartIdx;
|
||||
runInfo.s2AlignedSize = Align(runInfo.s2RealSize);
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H
|
||||
@@ -0,0 +1,256 @@
|
||||
/**
|
||||
* 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 util_regbase.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SAS_UTIL_REGBASE_H
|
||||
#define KV_QUANT_SAS_UTIL_REGBASE_H
|
||||
|
||||
#include "util.h"
|
||||
|
||||
using AscendC::TQue;
|
||||
using AscendC::QuePosition;
|
||||
|
||||
namespace regbaseutil {
|
||||
constexpr int64_t MAX_PRE_NEXT_TOKENS = 0x7FFFFFFF;
|
||||
|
||||
#define COMMON_RUN_PARAM \
|
||||
int64_t boIdx; \
|
||||
int64_t s1oIdx; \
|
||||
int64_t n2oIdx; \
|
||||
int64_t goIdx; \
|
||||
int64_t s2LoopEndIdx; /* S2方向的循环控制信息 souter层确定 */ \
|
||||
int64_t s2LineStartIdx = 0; /* S2方向按行的起始位置 */ \
|
||||
int64_t s2LineEndIdx; /* S2方向按行的结束位置 */ \
|
||||
int64_t s2CmpLineEndIdx; \
|
||||
/* cube视角的sOuter,在SAMEAB场景中cubeSOuterSize为两倍的 halfS1RealSize souter层确定 */ \
|
||||
uint32_t s1RealSize; \
|
||||
uint32_t halfS1RealSize; \
|
||||
uint32_t firstHalfS1RealSize; \
|
||||
uint32_t mRealSize; \
|
||||
uint32_t halfMRealSize; \
|
||||
uint32_t firstHalfMRealSize; \
|
||||
int64_t attentionOutOffset; /* attentionOut的offset souter层确定 */ \
|
||||
int32_t actualS1Size; /* Q的actualSeqLength */ \
|
||||
int32_t actualS2Size /* KV的actualSeqLength */ \
|
||||
|
||||
struct RunParamStr { // 分核与切块需要使用到参数
|
||||
COMMON_RUN_PARAM;
|
||||
/* 推理新增 */
|
||||
int64_t gs1LoopStartIdx;
|
||||
int64_t gs1LoopEndIdx;
|
||||
// BN循环生产的数据
|
||||
int64_t preTokensPerBatch = MAX_PRE_NEXT_TOKENS; // 左上顶点的pretoken
|
||||
int64_t nextTokensPerBatch = MAX_PRE_NEXT_TOKENS; // 左上顶点的nexttoken
|
||||
|
||||
// NBS1循环生产的数据
|
||||
int64_t sOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定
|
||||
int64_t cubeSOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定
|
||||
int64_t mOuterOffset;
|
||||
int64_t cubeMOuterOffset;
|
||||
|
||||
// lse 输出offset
|
||||
int64_t softmaxLseOffset; // souter层确定
|
||||
|
||||
int64_t qSNumInOneBlock;
|
||||
int64_t oriKvLoopEndIdx;
|
||||
int64_t cmpKvLoopEndIdx;
|
||||
};
|
||||
|
||||
#define COMMON_RUN_INFO \
|
||||
int64_t s2StartIdx; /* s2的起始位置,sparse场景下可能不是0 */ \
|
||||
int64_t s2EndIdx; \
|
||||
int64_t s2LoopCount; /* s2循环当前的循环index */ \
|
||||
int64_t s2LoopLimit; \
|
||||
int64_t s1oIdx = 0; /* s1轴的index */ \
|
||||
int64_t loop = 0; /* for v0 perload loop */ \
|
||||
int64_t boIdx = 0; /* b轴的index */ \
|
||||
int64_t n2oIdx = 0; /* n2轴的index */ \
|
||||
int64_t goIdx = 0; /* g轴的index */ \
|
||||
int32_t s1RealSize; \
|
||||
int32_t halfS1RealSize; /* vector侧实际的s1基本块大小,如果Cube基本块=128,那么halfS1RealSize=64 */ \
|
||||
int32_t firstHalfS1RealSize; /* 当s1RealSize不是2的整数倍时,v0比v1少计算一行,计算subblock偏移的时候需要使用v0的s1 size */ \
|
||||
int32_t mRealSize; \
|
||||
int32_t halfMRealSize; \
|
||||
int32_t firstHalfMRealSize; \
|
||||
int32_t s2RealSize; /* s2方向基本块的真实长度 */ \
|
||||
int64_t s2AlignedSize; /* s2方向基本块对齐到16之后的长度 */ \
|
||||
int32_t vec2S1BaseSize; /* vector2侧开循环之后,经过切分的S1大小,例如把64切分成两份32 */ \
|
||||
int32_t vec2S1RealSize; /* vector2侧开循环之后,经过切分的S1的尾块大小,例如把63切分成两份32和31,第二份的实际大小是31 */ \
|
||||
int32_t vec2MBaseSize; \
|
||||
int32_t vec2MRealSize; \
|
||||
int64_t taskId; \
|
||||
int64_t multiCoreInnerIdx = 0; \
|
||||
int64_t attentionOutOffset; \
|
||||
int32_t actualS1Size; /* 非TND场景=总s1Size, Tnd场景下当前batch对应的s1 */ \
|
||||
int32_t actualS2Size; /* 非TND场景=总s2Size, Tnd场景下当前batch对应的s2 */ \
|
||||
int64_t preTokensPerBatch; /* vector2 左上顶点的pretoken */ \
|
||||
int64_t nextTokensPerBatch; /* vector2 左上顶点的nexttoken */ \
|
||||
uint8_t taskIdMod2; \
|
||||
uint8_t taskIdMod3; \
|
||||
uint8_t multiCoreIdxMod2 = 0; \
|
||||
uint8_t multiCoreIdxMod3 = 0; \
|
||||
int64_t sOuterOffset; \
|
||||
int64_t mOuterOffset; \
|
||||
bool isCmp; \
|
||||
|
||||
struct RunInfo {
|
||||
COMMON_RUN_INFO;
|
||||
// 推理新增
|
||||
// lse 输出offset
|
||||
int64_t softmaxLseOffset;
|
||||
|
||||
int64_t qSNumInOneBlock;
|
||||
int64_t oriKvLoopEndIdx;
|
||||
int64_t cmpKvLoopEndIdx;
|
||||
};
|
||||
|
||||
#define COMMON_CONST_INFO \
|
||||
/* 全局的基本块信息 */ \
|
||||
uint32_t bSize; \
|
||||
uint32_t needInit; \
|
||||
uint32_t s1BaseSize; \
|
||||
uint32_t s2BaseSize; \
|
||||
int64_t dSize; /* query d 512 */ \
|
||||
int64_t dSizeV; /* key d 512 */ \
|
||||
int64_t dSizeVInput; /* key inpue d 640 = rope + nope + scale + pad */ \
|
||||
int64_t dSizeNope; /* key nope d 448 */ \
|
||||
int64_t dSizeRope; /* key rope d 64 */ \
|
||||
int64_t tileSize; /* 64 */ \
|
||||
int64_t sparseMode = 3; \
|
||||
int64_t gSize; /* g轴的大小 */ \
|
||||
int64_t n2Size; \
|
||||
int64_t s1Size; /* s1总大小 */ \
|
||||
int64_t s2Size; /* s2总大小 */ \
|
||||
/* 轴的乘积 */ \
|
||||
int64_t s1D; \
|
||||
int64_t gS1D; \
|
||||
int64_t n2GS1D; \
|
||||
int64_t s2D; \
|
||||
int64_t n2S2D; \
|
||||
int64_t s1Dv; \
|
||||
int64_t gS1Dv; \
|
||||
int64_t n2GS1Dv; \
|
||||
int64_t s2Dv; \
|
||||
int64_t n2S2Dv; \
|
||||
int64_t s1S2; \
|
||||
int64_t gS1; \
|
||||
int64_t gD; \
|
||||
int64_t n2D; \
|
||||
int64_t bN2D; \
|
||||
int64_t gDv; \
|
||||
int64_t n2Dv; \
|
||||
int64_t bN2Dv; \
|
||||
int64_t n2G; \
|
||||
int64_t n2GD; \
|
||||
int64_t bN2GD; \
|
||||
int64_t n2GDv; \
|
||||
int64_t bN2GDv; \
|
||||
int64_t gS2; \
|
||||
int64_t s1Dr; \
|
||||
int64_t gS1Dr; \
|
||||
int64_t n2GS1Dr; \
|
||||
int64_t s2Dr; \
|
||||
int64_t n2S2Dr; \
|
||||
int64_t gDr; \
|
||||
int64_t n2Dr; \
|
||||
int64_t bN2Dr; \
|
||||
int64_t n2GDr; \
|
||||
int64_t bN2GDr; \
|
||||
int32_t s2BaseN2D; \
|
||||
int32_t s1BaseN2GD; \
|
||||
int64_t s2BaseBN2D; \
|
||||
int64_t s1BaseBN2GD; \
|
||||
int32_t s1BaseD; \
|
||||
int32_t s2BaseD; \
|
||||
int64_t s2BaseN2Dv; \
|
||||
int64_t s2BaseBN2Dv; \
|
||||
int64_t s1BaseN2GDv; \
|
||||
int64_t s1BaseBN2GDv; \
|
||||
int32_t s1BaseDv; \
|
||||
int32_t s2BaseDv; \
|
||||
/* matmul跳读参数 */ \
|
||||
int64_t mm1Ka; \
|
||||
/* dq 或者attentionOut的Stride */ \
|
||||
int64_t attentionOutStride; \
|
||||
uint32_t aivIdx; \
|
||||
uint8_t subBlockIdx;\
|
||||
|
||||
#define INFER_CONST_INFO \
|
||||
/* 推理 */ \
|
||||
bool isActualLenDimsNull; /* 判断是否有actualseq */ \
|
||||
bool isActualLenDimsKVNull; /* 判断是否有actualseq_kv */ \
|
||||
bool isSoftmaxLseEnable; \
|
||||
bool rsvd1; \
|
||||
uint32_t sparseBlockCount; \
|
||||
uint32_t actualSeqLenSize; /* 用户输入的actualseq的长度 */ \
|
||||
uint32_t actualSeqLenKVSize; /* 用户输入的actualseq_kv的长度 */ \
|
||||
/* service mm1 mm2 pageAttention */ \
|
||||
uint32_t oriBlockSize; \
|
||||
uint32_t cmpBlockSize; \
|
||||
uint32_t paLayoutType; \
|
||||
uint32_t oriMaxBlockNumPerBatch; \
|
||||
uint32_t cmpMaxBlockNumPerBatch; \
|
||||
int32_t oriWinLeft; \
|
||||
int32_t oriWinRight; \
|
||||
uint32_t sparseBlockSize; \
|
||||
uint32_t cmpRatio; \
|
||||
float softmaxScale; \
|
||||
uint32_t oriKvStride; \
|
||||
uint32_t cmpKvStride; \
|
||||
|
||||
#define CV_SHARED_PARAMS \
|
||||
/* base params */ \
|
||||
uint32_t s1BaseSize; \
|
||||
uint32_t s2BaseSize; \
|
||||
uint32_t bSize; \
|
||||
uint32_t n2Size; \
|
||||
uint32_t gSize; \
|
||||
uint32_t s1Size; \
|
||||
uint32_t s2Size; \
|
||||
uint32_t dSize : 10; \
|
||||
uint32_t dSizeVInput : 12; \
|
||||
uint32_t needInit : 4; \
|
||||
uint32_t layoutType : 4; \
|
||||
uint32_t isActualSeqLengthsNull : 1; \
|
||||
uint32_t isActualSeqLengthsKVNull : 1; \
|
||||
uint32_t sparseBlockCount; \
|
||||
float softmaxScale; \
|
||||
uint32_t cmpRatio : 9; \
|
||||
uint32_t dSizeRope : 11; \
|
||||
uint32_t oriMaskMode : 6;\
|
||||
uint32_t cmpMaskMode : 6; \
|
||||
int32_t oriWinLeft; \
|
||||
int32_t oriWinRight; \
|
||||
uint32_t tileSize : 8; \
|
||||
/* pa params */ \
|
||||
uint32_t oriBlockSize : 12; \
|
||||
uint32_t cmpBlockSize : 12; \
|
||||
uint32_t oriMaxBlockNumPerBatch; \
|
||||
uint32_t cmpMaxBlockNumPerBatch; \
|
||||
uint32_t oriKvStride; \
|
||||
uint32_t cmpKvStride; \
|
||||
|
||||
|
||||
struct ConstInfo{
|
||||
COMMON_CONST_INFO;
|
||||
INFER_CONST_INFO;
|
||||
};
|
||||
|
||||
/* only support b32 or b64 */
|
||||
struct CVSharedParams {
|
||||
CV_SHARED_PARAMS;
|
||||
};
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SAS_UTIL_REGBASE_H
|
||||
@@ -0,0 +1,126 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_NO_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
// 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;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_SCFA_H
|
||||
@@ -0,0 +1,131 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED128_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED128_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
// 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;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_SCFA_H
|
||||
@@ -0,0 +1,134 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_NO_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
|
||||
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;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_SCFA_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_update_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED128_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
|
||||
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;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_pse_bf16_src;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_pse_bf16;
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> vreg_pse_bf16_unroll;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_SCFA_H
|
||||
@@ -0,0 +1,116 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_UNALIGNED64_NO_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
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;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_H
|
||||
@@ -0,0 +1,127 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_ALIGNED64_UPDATE_SCFA_H
|
||||
#define VF_BASIC_BLOCK_ALIGNED64_UPDATE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
// 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;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
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_SCFA_H
|
||||
@@ -0,0 +1,86 @@
|
||||
/**
|
||||
* 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_utils.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef VF_BASIC_BLOCK_UTILS_H
|
||||
#define VF_BASIC_BLOCK_UTILS_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
constexpr uint32_t floatRepSize = 64;
|
||||
constexpr uint32_t blockBytesU8 = 32;
|
||||
constexpr float fp8e4m3MaxValue = 448.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,
|
||||
};
|
||||
}
|
||||
#endif // VF_BASIC_BLOCK_UTILS_H
|
||||
@@ -0,0 +1,168 @@
|
||||
/**
|
||||
* 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_flashupdate_new_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef FLASH_UPDATE_NEW_INTERFACE_SCFA_H
|
||||
#define FLASH_UPDATE_NEW_INTERFACE_SCFA_H
|
||||
|
||||
#include "vf_basic_block_utils.h"
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
constexpr uint16_t REDUCE_SIZE = 1;
|
||||
/* **************************************************************************************************
|
||||
* FlashUpdate, fp32
|
||||
* ************************************************************************************************* */
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize>
|
||||
__simd_vf__ inline void FlashUpdateBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, const uint16_t m)
|
||||
{
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_pre;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_cur;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_mul;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_add;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// dstTensor = preTensor * expMaxTensor + curTensor
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize); // [m,8]
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_pre, preUb + i * srcD + j * floatRepSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_cur, curUb + i * srcD + j * floatRepSize);
|
||||
AscendC::MicroAPI::MulDstAdd(vreg_input_pre, vreg_exp_max, vreg_input_cur, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * srcD + j * floatRepSize, vreg_input_pre, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
* @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] srcD, input columns, should be 32 bytes aligned
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD>
|
||||
__aicore__ inline void FlashUpdateNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const uint16_t m)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdate, T must be float");
|
||||
|
||||
__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();
|
||||
FlashUpdateBasicVF<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE>(dstUb, curUb, preUb, expMaxUb, m);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD, uint16_t reduceSize>
|
||||
__simd_vf__ inline void FlashUpdateLastBasicVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * preUb,
|
||||
__ubuf__ float * expMaxUb, __ubuf__ float * expSumUb, const uint16_t m)
|
||||
{
|
||||
constexpr uint16_t dLoops = srcD / floatRepSize;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_max;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_pre;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_cur;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_mul;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_add;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_div;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
|
||||
AscendC::MicroAPI::MaskReg preg_all = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_max, expMaxUb + i * reduceSize);
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * reduceSize);
|
||||
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_pre, preUb + i * srcD + j * floatRepSize);
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_cur, curUb + i * srcD + j * floatRepSize);
|
||||
AscendC::MicroAPI::MulDstAdd(vreg_input_pre, vreg_exp_max, vreg_input_cur, preg_all);
|
||||
AscendC::MicroAPI::Div(vreg_div, vreg_input_pre, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * srcD + j * floatRepSize, vreg_div, preg_all);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/*
|
||||
* @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] srcD, input columns, 32 bytes align
|
||||
*/
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint16_t srcD>
|
||||
__aicore__ inline void FlashUpdateLastNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& preTensor, const LocalTensor<T>& expMaxTensor, const LocalTensor<T>& expSumTensor, uint16_t m)
|
||||
{
|
||||
static_assert(IsSameType<T, float>::value, "VF FlashUpdateLast, T must be float");
|
||||
__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();
|
||||
|
||||
FlashUpdateLastBasicVF<T, INPUT_T, OUTPUT_T, srcD, REDUCE_SIZE>(dstUb, curUb, preUb, expMaxUb, expSumUb, m);
|
||||
}
|
||||
|
||||
template <typename T, typename INPUT_T, typename OUTPUT_T, uint32_t srcD>
|
||||
__simd_vf__ inline void LastDivNewVF(__ubuf__ float * dstUb, __ubuf__ float * curUb, __ubuf__ float * expSumUb,
|
||||
const uint16_t m)
|
||||
{
|
||||
const uint16_t dLoops = srcD >> 6;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_input_cur;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_div;
|
||||
AscendC::MicroAPI::RegTensor<float> vreg_exp_sum;
|
||||
AscendC::MicroAPI::MaskReg preg_all = CreateMask<float, MaskPattern::ALL>();
|
||||
uint32_t sreg_init = srcD;
|
||||
AscendC::MicroAPI::MaskReg preg_update = UpdateMask<float>(sreg_init);
|
||||
|
||||
for (uint16_t i = 0; i < m; ++i) {
|
||||
AscendC::MicroAPI::LoadAlign<T, MicroAPI::LoadDist::DIST_BRC_B32>(vreg_exp_sum, expSumUb + i * REDUCE_SIZE);
|
||||
for (uint16_t j = 0; j < dLoops; ++j) {
|
||||
AscendC::MicroAPI::LoadAlign(vreg_input_cur, curUb + i * srcD + j * floatRepSize);
|
||||
AscendC::MicroAPI::Div(vreg_div, vreg_input_cur, vreg_exp_sum, preg_all);
|
||||
AscendC::MicroAPI::StoreAlign<T, MicroAPI::StoreDist::DIST_NORM_B32>(
|
||||
(__ubuf__ T *&)dstUb + i * srcD + 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>
|
||||
__aicore__ inline void LastDivNew(const LocalTensor<T>& dstTensor, const LocalTensor<T>& curTensor,
|
||||
const LocalTensor<T>& expSumTensor, const uint16_t m)
|
||||
{
|
||||
__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>(dstUb, curUb, expSumUb, m);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
#endif // FLASH_UPDATE_NEW_INTERFACE_SCFA_H
|
||||
@@ -0,0 +1,162 @@
|
||||
/**
|
||||
* 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_scfa.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SCFA_INTERFACE_H
|
||||
#define MUL_SEL_SOFTMAX_FLASH_V2_CAST_NZ_SCFA_INTERFACE_H
|
||||
|
||||
#include "../util_regbase.h"
|
||||
#include "../kv_quant_sparse_attn_sharedkv_common_arch35.h"
|
||||
#include "vf_basic_block_aligned128_no_update_scfa.h"
|
||||
#include "vf_basic_block_aligned128_update_scfa.h"
|
||||
#include "vf_basic_block_unaligned64_update_scfa.h"
|
||||
#include "vf_basic_block_unaligned64_no_update_scfa.h"
|
||||
#include "vf_basic_block_unaligned128_no_update_scfa.h"
|
||||
#include "vf_basic_block_unaligned128_update_scfa.h"
|
||||
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace SCFaVectorApi {
|
||||
/* **************************************************************************************************
|
||||
* Muls + Select(optional) + SoftmaxFlashV2 + Cast(fp32->fp16/bf16) + ND2NZ
|
||||
* ************************************************************************************************* */
|
||||
using AscendC::LocalTensor;
|
||||
|
||||
enum class OriginNRange {
|
||||
EQ_128_SCFA = 0, // originN == 128, better performance than GT_64_AND_LTE_128 (s2BaseSize=128)
|
||||
GT_0_AND_LTE_64_SCFA, // 0 < originN <= 64 (s2BaseSize <= 64 or tail s2)
|
||||
GT_64_AND_LTE_128_SCFA, // 64 < originN <= 128, support for non-alignment (s2BaseSize=128)
|
||||
N_INVALID_SCFA
|
||||
};
|
||||
template <typename T, typename T2, uint32_t s1BaseSize = 64, uint32_t s2BaseSize = 128, OriginNRange oriNRange = OriginNRange::EQ_128_SCFA>
|
||||
__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_SCFA) {
|
||||
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_SCFA){
|
||||
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_SCFA){
|
||||
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_SCFA>
|
||||
__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_SCFA) {
|
||||
ProcessVec1UpdateImpl128<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_0_AND_LTE_64_SCFA) {
|
||||
ProcessVec1UpdateImpl64<T, T2, s1BaseSize, s2BaseSize>(
|
||||
dstTensor, srcTensor, inMaxTensor, sharedTmpBuffer, m, originN, scale, minValue);
|
||||
} else if constexpr (oriNRange == OriginNRange::GT_64_AND_LTE_128_SCFA){
|
||||
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_SCFA>
|
||||
__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, bfloat16_t>::value, "VF mul_sel_softmaxflashv2_cast_nz, T2 must be 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 SCFAUpdateExpSumAndExpMax(
|
||||
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_SCFA_INTERFACE_H
|
||||
@@ -0,0 +1,69 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_template_tiling_key.h"
|
||||
#include "arch35/kv_quant_sparse_attn_sharedkv_scfa_kernel.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_common.h"
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
#if defined(__DAV_C310_CUBE__)
|
||||
#define SAS_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::SCFABlockCube<__VA_ARGS__>, BaseApi::SCFABlockCubeDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::SCFABlockVecDummy<__VA_ARGS__>, BaseApi::SCFABlockVec<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
op.Init(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, \
|
||||
seqUsedQ, seqUsedKV, sinks, metadata, attentionOut, user, nullptr, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#else
|
||||
#define SAS_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::SCFABlockCube<__VA_ARGS__>, BaseApi::SCFABlockCubeDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::SCFABlockVecDummy<__VA_ARGS__>, BaseApi::SCFABlockVec<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tilingDataIn, tiling); \
|
||||
const tilingdataClass *__restrict tilingData = &tilingDataIn; \
|
||||
op.Init(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, \
|
||||
seqUsedQ, seqUsedKV, sinks, metadata, attentionOut, user, tilingData, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
template<int FLASH_DECODE, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE, int SPLIT_G>
|
||||
__global__ __aicore__ void
|
||||
kv_quant_sparse_attn_sharedkv(__gm__ uint8_t *query, __gm__ uint8_t *oriKV, __gm__ uint8_t *cmpKV,
|
||||
__gm__ uint8_t *oriSparseIndices, __gm__ uint8_t *cmpSparseIndices, __gm__ uint8_t* oriBlockTable,
|
||||
__gm__ uint8_t* cmpBlockTable, __gm__ uint8_t *cuSeqlensQ, __gm__ uint8_t *cuSeqlensOriKv,
|
||||
__gm__ uint8_t *cuSeqlensCmpKv, __gm__ uint8_t *seqUsedQ, __gm__ uint8_t *seqUsedKV,
|
||||
__gm__ uint8_t *sinks, __gm__ uint8_t *metadata, __gm__ uint8_t *attentionOut, __gm__ uint8_t *softmax_lse,
|
||||
__gm__ uint8_t *workspace, __gm__ uint8_t *tiling)
|
||||
{
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
|
||||
TPipe tPipe;
|
||||
__gm__ uint8_t *user = GetUserWorkspace(workspace);
|
||||
|
||||
SAS_OP_IMPL(BaseApi::KvQuantSparseAttnSharedkvScfa, KvQuantSparseAttnSharedkvTilingData, bfloat16_t,
|
||||
fp8_e4m3fn_t, float, bfloat16_t, FLASH_DECODE, true, static_cast<SAS_LAYOUT>(LAYOUT_T),
|
||||
static_cast<SAS_LAYOUT>(KV_LAYOUT_T), static_cast<SASTemplateMode>(TEMPLATE_MODE), SPLIT_G);
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_common.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
enum class SAS_LAYOUT {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_ND = 2
|
||||
};
|
||||
|
||||
enum class SASTemplateMode {
|
||||
SWA_TEMPLATE_MODE = 0,
|
||||
CFA_TEMPLATE_MODE = 1,
|
||||
SCFA_TEMPLATE_MODE = 2
|
||||
};
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
@@ -0,0 +1,80 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_metadata.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_H
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
// Constants
|
||||
constexpr uint32_t AIC_CORE_NUM = 36;
|
||||
constexpr uint32_t AIV_CORE_NUM = 72;
|
||||
constexpr uint32_t SAS_META_SIZE = 1024;
|
||||
using SAS_METADATA_T = int32_t;
|
||||
|
||||
constexpr uint32_t FA_METADATA_SIZE = 9;
|
||||
constexpr uint32_t FD_METADATA_SIZE = 8;
|
||||
|
||||
// FA Metadata Index Definitions
|
||||
constexpr uint32_t FA_CORE_ENABLE_INDEX = 0;
|
||||
constexpr uint32_t FA_BN2_START_INDEX = 1;
|
||||
constexpr uint32_t FA_M_START_INDEX = 2;
|
||||
constexpr uint32_t FA_S2_START_INDEX = 3;
|
||||
constexpr uint32_t FA_BN2_END_INDEX = 4;
|
||||
constexpr uint32_t FA_M_END_INDEX = 5;
|
||||
constexpr uint32_t FA_S2_END_INDEX = 6;
|
||||
constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7;
|
||||
constexpr uint32_t FA_S2_MAX_NUM = 8;
|
||||
|
||||
// FD Metadata Index Definitions
|
||||
constexpr uint32_t FD_CORE_ENABLE_INDEX = 0;
|
||||
constexpr uint32_t FD_BN2_IDX_INDEX = 1;
|
||||
constexpr uint32_t FD_M_IDX_INDEX = 2;
|
||||
constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3;
|
||||
constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4;
|
||||
constexpr uint32_t FD_M_START_INDEX = 5;
|
||||
constexpr uint32_t FD_M_NUM_INDEX = 6;
|
||||
|
||||
/**
|
||||
* @brief 获取属性的绝对索引
|
||||
* @param coreIdx 核索引
|
||||
* @param metaIdx 元数据索引
|
||||
* @param isAIV 是否为AIV数据,默认为false
|
||||
* @return 返回属性的绝对索引
|
||||
*/
|
||||
#ifdef __CCE_AICORE__
|
||||
__aicore__ inline uint32_t GetAttrAbsIndex(uint32_t coreIdx, uint32_t metaIdx, bool isAIV=false)
|
||||
{
|
||||
if (isAIV) {
|
||||
return FA_METADATA_SIZE * AIC_CORE_NUM + FD_METADATA_SIZE * coreIdx + metaIdx;
|
||||
} else {
|
||||
return FA_METADATA_SIZE * coreIdx + metaIdx;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
namespace detail {
|
||||
struct SasMetaData {
|
||||
uint32_t faMetadata[AIC_CORE_NUM][FA_METADATA_SIZE];
|
||||
uint32_t fdMetadata[AIV_CORE_NUM][FD_METADATA_SIZE];
|
||||
};
|
||||
};
|
||||
|
||||
static_assert(SAS_META_SIZE * sizeof(SAS_METADATA_T) >= sizeof(detail::SasMetaData));
|
||||
};
|
||||
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_H
|
||||
@@ -0,0 +1,68 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_attn_sharedkv_template_tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H
|
||||
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define SAS_LAYOUT_BSND 0
|
||||
#define SAS_LAYOUT_TND 1
|
||||
#define SAS_LAYOUT_PA_ND 2
|
||||
|
||||
#define ASCENDC_TPL_4_BW 4
|
||||
|
||||
#define SWA_TEMPLATE 0
|
||||
#define CFA_TEMPLATE 1
|
||||
#define SCFA_TEMPLATE 2
|
||||
// 模板参数支持的范围定义
|
||||
ASCENDC_TPL_ARGS_DECL(SparseAttnSharedkv, // 算子OpType
|
||||
ASCENDC_TPL_BOOL_DECL(FLASH_DECODE, 0, 1),
|
||||
ASCENDC_TPL_UINT_DECL(LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_DECL(KV_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_PA_ND),
|
||||
ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE, CFA_TEMPLATE, SCFA_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_DECL(SPLIT_G, 0, 1),
|
||||
);
|
||||
|
||||
// 支持的模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_PA_ND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_PA_ND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CFA_TEMPLATE), // V模板不支持非PA
|
||||
ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, SAS_LAYOUT_PA_ND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SCFA_TEMPLATE), // V模板不支持非PA
|
||||
ASCENDC_TPL_BOOL_SEL(SPLIT_G, 0, 1),
|
||||
),
|
||||
);
|
||||
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H
|
||||
Reference in New Issue
Block a user