init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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 &param)
{
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 &param)
{
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 &param)
{
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 &param)
{
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 &param)
{
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 &param)
{
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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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);
}

View File

@@ -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

View File

@@ -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

View File

@@ -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