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,853 @@
/**
 * 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 sparse_attn_sharedkv_scfa_block_cube.h
* \brief use 7 buffer for matmul l1, better pipeline
*/
#ifndef SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H
#define SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "kernel_tiling/kernel_tiling.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#include "../sparse_attn_sharedkv_common.h"
namespace SASKernel {
template <typename SAST>
class SASCubeBlock {
public:
// 中间计算数据类型为float, 高精度模式
using T = float;
using Q_T = typename SAST::queryType;
using KV_T = typename SAST::kvType;
using OUT_T = typename SAST::outputType;
using MM_OUT_T = T;
__aicore__ inline SASCubeBlock(){};
__aicore__ inline void InitParams(const ConstInfo &constInfo);
__aicore__ inline void InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<KV_T> oriKvGm,
GlobalTensor<KV_T> cmpKV, GlobalTensor<MM_OUT_T> mm1ResGm);
__aicore__ inline void InitMm2GlobalTensor(GlobalTensor<KV_T> vec1ResGm, GlobalTensor<MM_OUT_T> mm2ResGm,
GlobalTensor<OUT_T> attentionOutGm);
__aicore__ inline void InitPageAttentionInfo(GlobalTensor<KV_T> oriKvGm, const GlobalTensor<KV_T> &kvMergeGm,
GlobalTensor<int32_t> oriBlockTableGm,
GlobalTensor<int32_t> cmpBlockTableGm);
__aicore__ inline void InitBuffers(TPipe *pipe);
__aicore__ inline void AllocEventID();
__aicore__ inline void FreeEventID();
__aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo);
__aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo);
private:
static constexpr bool PAGE_ATTENTION = SAST::pageAttention;
static constexpr int TEMPLATE_MODE = SAST::templateMode;
static constexpr bool FLASH_DECODE = SAST::flashDecode;
static constexpr SAS_LAYOUT LAYOUT_T = SAST::layout;
static constexpr SAS_LAYOUT KV_LAYOUT_T = SAST::kvLayout;
static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分
static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分
static constexpr uint32_t K_L0_SPLIT_SIZE = 128; // k方向L0切分
static constexpr uint32_t K_L1_SPLIT_SIZE = 256; // k方向L1切分
static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分
static constexpr uint32_t D_SPLIT_SIZE = 256; // d轴切分
static constexpr uint32_t L1_BLOCK_SIZE = (64 * 512 * sizeof(Q_T));
static constexpr uint32_t L1_BLOCK_OFFSET = 64 * 512;
static constexpr uint32_t L0A_PP_SIZE = (32 * 1024);
static constexpr uint32_t L0B_PP_SIZE = (32 * 1024);
static constexpr uint32_t L0C_PP_SIZE = (64 * 1024);
// mte2 <> mte1 EventID
// L1 3buf, 使用3个eventId
static constexpr uint32_t L1_EVENT0 = EVENT_ID2;
static constexpr uint32_t L1_EVENT1 = EVENT_ID3;
static constexpr uint32_t L1_EVENT2 = EVENT_ID4;
static constexpr uint32_t L1_EVENT3 = EVENT_ID5;
static constexpr uint32_t L1_EVENT4 = EVENT_ID6;
static constexpr uint32_t L1_EVENT5 = EVENT_ID7;
static constexpr uint32_t L1_EVENT6 = EVENT_ID1;
// m <> mte1 EventID
static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3;
static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4;
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding
static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用
static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6};
ConstInfo constInfo{};
// L1分成3块buf, 用于记录
uint32_t qpL1BufIter = 0;
uint32_t kvL1BufIter = -1;
uint32_t abL0BufIter = 0;
uint32_t cL0BufIter = 0;
// mm1
GlobalTensor<Q_T> queryGm;
GlobalTensor<KV_T> keyGm;
GlobalTensor<MM_OUT_T> mm1ResGm;
GlobalTensor<KV_T> oriKvGm;
GlobalTensor<KV_T> kvMergeGm_;
GlobalTensor<KV_T> cmpKvGm;
// mm2
GlobalTensor<KV_T> vec1ResGm;
GlobalTensor<KV_T> valueGm;
GlobalTensor<MM_OUT_T> mm2ResGm;
GlobalTensor<OUT_T> attentionOutGm;
// block_table
GlobalTensor<int32_t> oriBlockTableGm;
GlobalTensor<int32_t> cmpBlockTableGm;
TBuf<TPosition::A1> bufQPL1;
TBuf<TPosition::A1> bufKVL1;
TBuf<TPosition::A2> tmpBufL0A;
TBuf<TPosition::B2> tmpBufL0B;
TBuf<TPosition::CO1> tmpBufL0C;
LocalTensor<Q_T> l1QPTensor;
LocalTensor<Q_T> l1KVTensor;
LocalTensor<KV_T> aL0TensorPingPong;
LocalTensor<KV_T> bL0TensorPingPong;
LocalTensor<MM_OUT_T> cL0TensorPingPong;
// L0AB m <> mte1 EventID
__aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx)
{
return (L0AB_EVENT0 + idx);
}
__aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx)
{
uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连
return idxMap[mIdx % 2] + k1Idx;
}
__aicore__ inline void CopyGmToL1(LocalTensor<KV_T> &l1Tensor, GlobalTensor<KV_T> &gmSrcTensor, uint32_t srcN,
uint32_t srcD, uint32_t srcDstride);
__aicore__ inline void CopyInMm1AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset);
__aicore__ inline void CopyInMm2AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset);
__aicore__ inline void LoadDataMm1A(LocalTensor<KV_T> &aL0Tensor, LocalTensor<KV_T> &aL1Tensor, uint32_t idx,
uint32_t kSplitSize, uint32_t mSize, uint32_t kSize);
__aicore__ inline void LoadDataMm1B(LocalTensor<KV_T> &bL0Tensor, LocalTensor<KV_T> &bL1Tensor, uint32_t idx,
uint32_t kSplitSize, uint32_t kSize, uint32_t nSize);
};
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::InitParams(const ConstInfo &constInfo)
{
this->constInfo = constInfo;
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<KV_T> oriKvGm,
GlobalTensor<KV_T> cmpKvGm,
GlobalTensor<MM_OUT_T> mm1ResGm)
{
// mm1
this->queryGm = queryGm;
this->oriKvGm = oriKvGm;
this->cmpKvGm = cmpKvGm;
this->mm1ResGm = mm1ResGm;
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::InitMm2GlobalTensor(GlobalTensor<KV_T> vec1ResGm,
GlobalTensor<MM_OUT_T> mm2ResGm,
GlobalTensor<OUT_T> attentionOutGm)
{
// mm2
this->vec1ResGm = vec1ResGm;
this->mm2ResGm = mm2ResGm;
this->attentionOutGm = attentionOutGm;
}
template <typename SAST>
__aicore__ inline void
SASCubeBlock<SAST>::InitPageAttentionInfo(GlobalTensor<KV_T> oriKvGm, const GlobalTensor<KV_T> &kvMergeGm,
GlobalTensor<int32_t> oriBlockTableGm, GlobalTensor<int32_t> cmpBlockTableGm)
{
this->oriKvGm = oriKvGm;
this->kvMergeGm_ = kvMergeGm;
this->oriBlockTableGm = oriBlockTableGm;
this->cmpBlockTableGm = cmpBlockTableGm;
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::InitBuffers(TPipe *pipe)
{
pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4);
l1QPTensor = bufQPL1.Get<Q_T>();
pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3);
l1KVTensor = bufKVL1.Get<KV_T>();
// L0A
pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K
aL0TensorPingPong = tmpBufL0A.Get<KV_T>();
// L0B
pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K
bL0TensorPingPong = tmpBufL0B.Get<KV_T>();
// L0C
pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K
cL0TensorPingPong = tmpBufL0C.Get<MM_OUT_T>();
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::AllocEventID()
{
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::FreeEventID()
{
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::CopyGmToL1(LocalTensor<KV_T> &l1Tensor, GlobalTensor<KV_T> &gmSrcTensor,
uint32_t srcN, uint32_t srcD, uint32_t srcDstride)
{
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = srcN; // 行数
nd2nzPara.dValue = srcD;
nd2nzPara.srcDValue = srcDstride;
nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(l1Tensor, gmSrcTensor, nd2nzPara);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::CopyInMm1AToL1(LocalTensor<KV_T> &l1Tensor, const RunInfo &info,
uint32_t mSeqIdx, uint32_t mSizeAct, uint32_t headSize,
uint32_t headOffset)
{
auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.headDim + headOffset];
CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, constInfo.headDim);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::LoadDataMm1A(LocalTensor<KV_T> &aL0Tensor, LocalTensor<KV_T> &aL1Tensor,
uint32_t idx, uint32_t kSplitSize, uint32_t mSize,
uint32_t kSize)
{
LocalTensor<KV_T> srcTensor = aL1Tensor[mSize * kSplitSize * idx];
LoadData3DParamsV2<KV_T> loadData3DParams;
// SetFmatrixParams
loadData3DParams.l1H = mSize / 16; // Hin=M1=8
loadData3DParams.l1W = 16; // Win=M0
loadData3DParams.padList[0] = 0;
loadData3DParams.padList[1] = 0;
loadData3DParams.padList[2] = 0;
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
// SetLoadToA0Params
loadData3DParams.mExtension = mSize; // M
loadData3DParams.kExtension = kSize; // K
loadData3DParams.mStartPt = 0;
loadData3DParams.kStartPt = 0;
loadData3DParams.strideW = 1;
loadData3DParams.strideH = 1;
loadData3DParams.filterW = 1;
loadData3DParams.filterSizeW = (1 >> 8) & 255;
loadData3DParams.filterH = 1;
loadData3DParams.filterSizeH = (1 >> 8) & 255;
loadData3DParams.dilationFilterW = 1;
loadData3DParams.dilationFilterH = 1;
loadData3DParams.enTranspose = 0;
loadData3DParams.fMatrixCtrl = 0;
loadData3DParams.channelSize = kSize; // Cin=K
LoadData<KV_T, LOAD3DV2_CONFIG>(aL0Tensor, srcTensor, loadData3DParams);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::LoadDataMm1B(LocalTensor<KV_T> &l0Tensor, LocalTensor<KV_T> &l1Tensor,
uint32_t idx, uint32_t kSplitSize, uint32_t kSize,
uint32_t nSize)
{
// N 方向全载
LocalTensor<KV_T> srcTensor = l1Tensor[nSize * kSplitSize * idx];
LoadData2DParams loadData2DParams;
loadData2DParams.startIndex = 0;
loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(KV_T));
loadData2DParams.srcStride = 1;
loadData2DParams.dstGap = 0;
loadData2DParams.ifTranspose = false;
LoadData(l0Tensor, srcTensor, loadData2DParams);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::CopyInMm2AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info,
uint32_t mSeqIdx, uint32_t subMSizeAct, uint32_t nSize,
uint32_t nOffset)
{
auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize +
mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset];
CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign);
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo)
{
uint32_t mSize = mSplitInfo.nBufferDealM;
uint32_t mL1Size = M_SPLIT_SIZE;
uint32_t mL1SizeAlign = SASAlign(M_SPLIT_SIZE, 16);
uint32_t mL1Loops = CeilDiv(mSize, M_SPLIT_SIZE);
uint32_t nSize = info.actualSingleProcessSInnerSize;
uint32_t nL1Size = N_SPLIT_SIZE;
uint32_t nL1SizeAlign = SASAlign(N_SPLIT_SIZE, 16);
uint32_t nL1Loops = CeilDiv(nSize, N_SPLIT_SIZE);
uint32_t kSize = 512;
uint32_t kL1Size = 256;
uint32_t kL1Loops = 2;
uint32_t kL0Size = 128;
uint32_t kL0Loops = CeilDiv(kL1Size, kL0Size);
LocalTensor<KV_T> bL1Tensor;
uint32_t ka = 0, kb = 0;
// L1 切n切k
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) {
if (nL1 == (nL1Loops - 1)) {
// 尾块重新计算size
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
nL1SizeAlign = SASAlign(nL1Size, 16);
}
for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) {
kvL1BufIter++;
uint32_t kb = kvL1BufIter % 3;
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
// 从k当中取当前的块
bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
uint32_t curSeqIdx = info.s2BatchOffset + nL1 * N_SPLIT_SIZE;
if (info.isOri) {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
uint32_t copyFinishRowCnt = 0;
LocalTensor<KV_T> kTensor;
uint32_t copyRowCnt = 0;
while (copyFinishRowCnt < nL1Size) {
// 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位
copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize;
if (copyFinishRowCnt + copyRowCnt > nL1Size) {
copyRowCnt = nL1Size - copyFinishRowCnt;
}
PAShape shape;
shape.blockSize = constInfo.paOriBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.oriKvStride;
shape.actHeadDim = D_SPLIT_SIZE;
shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = nL1SizeAlign;
kTensor = bL1Tensor[copyFinishRowCnt * 16];
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
startPos.dIdx = kL1 * D_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
DataCopyPA<KV_T>(kTensor, oriKvGm, oriBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size; // 行数
nd2nzPara.dValue = D_SPLIT_SIZE; // 256
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize * seqStride;
uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE;
DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
if (kL1 == 0) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor, oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim +
nL1 * N_SPLIT_SIZE * constInfo.headDim], nd2nzPara);
} else {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor,
oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + (constInfo.headDim >> 1) +
nL1 * N_SPLIT_SIZE * constInfo.headDim],
nd2nzPara);
}
}
} else {
if (kL1 == 0) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor,
kvMergeGm_[info.cmpLoop % 4 * N_WORKSPACE_SIZE * kSize +
nL1 * N_SPLIT_SIZE * constInfo.headDim],
nd2nzPara);
} else {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor,
kvMergeGm_[info.cmpLoop % 4 * N_WORKSPACE_SIZE * kSize + (constInfo.headDim >> 1) +
nL1 * N_SPLIT_SIZE * constInfo.headDim],
nd2nzPara);
}
}
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
mL1Size = M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(M_SPLIT_SIZE, 16U);
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
uint32_t aL1PaddingSize = 0; // 用于使左矩阵对齐到尾部, 以保证两块32K内存连续
if (mL1 == (mL1Loops - 1)) {
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(mL1Size, 16U);
aL1PaddingSize = (M_SPLIT_SIZE - mL1SizeAlign) * 256;
}
uint32_t mIdx = qpL1BufIter + mL1;
ka = GetQPL1RealIdx(mIdx, kL1);
LocalTensor<Q_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + (1 - kL1) * aL1PaddingSize];
if (nL1 == 0) {
if (kL1 == 0) {
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka + 1]);
CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, 0);
} else {
LocalTensor<Q_T> qTmpTensor = aL1Tensor;
CopyInMm1AToL1(qTmpTensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256,
256);
}
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
}
// 使用unitflag同步
LocalTensor cL0Tensor =
cL0TensorPingPong[(cL0BufIter % 2) *
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
LocalTensor<KV_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))];
LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size);
LocalTensor<KV_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))];
LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign);
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
MmadParams mmadParams;
mmadParams.m = mL1SizeAlign;
mmadParams.n = nL1SizeAlign;
mmadParams.k = kL0Size;
mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0);
mmadParams.cmatrixSource = false;
mmadParams.unitFlag =
(kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
PipeBarrier<PIPE_M>();
}
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
abL0BufIter++;
}
if (nL1 == (nL1Loops - 1)) {
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
}
if (kL1 == 1) { // 最后一轮kL1循环
FixpipeParamsV220 fixParams;
fixParams.nSize = nL1SizeAlign;
fixParams.mSize = mL1SizeAlign;
fixParams.srcStride = mL1SizeAlign;
// 改成nSizeAlign
fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔
fixParams.unitFlag = 0b11;
fixParams.ndNum = 1; // 输出ND
Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE +
(mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) *
info.actualSingleProcessSInnerSizeAlign],
cL0Tensor, fixParams);
}
if (mL1Loops == 2) {
cL0BufIter++;
}
}
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
}
if (mL1Loops == 1) {
cL0BufIter++;
}
}
qpL1BufIter += mL1Loops;
}
template <typename SAST>
__aicore__ inline void SASCubeBlock<SAST>::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo)
{
uint32_t mSize = mSplitInfo.nBufferDealM;
uint32_t mSizeAlign = (mSize + 16 - 1) / 16;
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐
uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小
uint32_t nSize = BlockAlign<KV_T>(constInfo.headDim);
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐
uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小
uint32_t kSize = info.actualSingleProcessSInnerSize;
uint32_t kL1Size = 256;
uint32_t kL1SizeAlign = SASAlign(kL1Size, 16U);
uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size;
uint32_t kL0Size = 128;
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
uint32_t kL0SizeAlign = kL0Size;
LocalTensor<KV_T> bL1Tensor;
LocalTensor<KV_T> subvTensor;
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
uint32_t ka = 0, kb = 0;
uint32_t mBaseIdx = qpL1BufIter;
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1 -> D
if (nL1 == (nL1Loops - 1)) {
// 尾块
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
nL1SizeAlign = SASAlign(nL1Size, 16U);
}
// k l1写成一个循环, 和mm1保持一致
kL1Size = 256;
kL1SizeAlign = SASAlign(kL1Size, 16U);
uint32_t copyRowCnt = 0;
for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作 -> S2,每次256
if (k1 == (kL1Loops - 1)) {
// 尾块
kL1Size = kSize - (kL1Loops - 1) * 256;
kL1SizeAlign = SASAlign(kL1Size, 16U);
}
kvL1BufIter++;
uint32_t kb = kvL1BufIter % 3;
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
uint32_t kOffset = k1 * kL0Loops;
kL0Size = 128;
// 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误
kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
kL0SizeAlign = kL0Size;
for (uint32_t kL1 = kOffset; kL1 < kL0Loops + kOffset; kL1++) { // 128 循环搬pa,每次128
if (kL1 == kOffset + kL0Loops - 1) {
// 尾块
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
kL0SizeAlign = SASAlign(kL0Size, 16U);
}
uint32_t curSeqIdx = info.s2BatchOffset + (kL1 - kOffset) * 128 + k1 * 256;
if (info.isOri) {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t copyFinishRowCnt = 0;
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
while (copyFinishRowCnt < kL0Size) {
copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize;
if (copyFinishRowCnt + copyRowCnt > kL0Size) {
copyRowCnt = kL0Size - copyFinishRowCnt;
}
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
startPos.dIdx = nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
PAShape shape;
shape.blockSize = constInfo.paOriBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.oriKvStride;
shape.actHeadDim = nL1Size;
shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = kL0SizeAlign;
subvTensor = bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE + copyFinishRowCnt * 16];
DataCopyPA<KV_T>(subvTensor, oriKvGm, oriBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize * seqStride;
uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE;
subvTensor = bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE];
DataCopy(subvTensor, oriKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE],
oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + kL1 * 128 * constInfo.headDim +
nL1 * N_SPLIT_SIZE], nd2nzPara);
}
} else {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor[(kL1 - kOffset) * 128 * N_SPLIT_SIZE],
kvMergeGm_[info.cmpLoop % 4 * N_WORKSPACE_SIZE * 512 + kL1 * 128 * constInfo.headDim +
nL1 * N_SPLIT_SIZE],
nd2nzPara);
}
}
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
mL1SizeAlign = M_SPLIT_SIZE;
mL1Size = M_SPLIT_SIZE; // m的实际大小
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
if (mL1 == (mL1Loops - 1)) {
// 尾块
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(mL1Size, 16U);
}
uint32_t mIdx = mBaseIdx + mL1;
ka = GetQPL1RealIdx(mIdx, k1);
LocalTensor<KV_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
if (nL1 == 0) {
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, kL1Size,
256 * k1);
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
}
LocalTensor cL0Tensor =
cL0TensorPingPong[(cL0BufIter % 2) *
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
uint32_t baseK = 128;
uint32_t baseN = 128;
kL0Size = 128;
kL0SizeAlign = kL0Size;
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
if (kL0 + 1 == kL0Loops) {
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
kL0SizeAlign = SASAlign(kL0Size, 16U);
}
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
LocalTensor<KV_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))];
LoadData3DParamsV2<KV_T> loadData3DParamsForB;
loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height
loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W
loadData3DParamsForB.padList[0] = 0;
loadData3DParamsForB.padList[1] = 0;
loadData3DParamsForB.padList[2] = 0;
loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果
loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度
loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度
loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点
loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点
loadData3DParamsForB.strideW = 1;
loadData3DParamsForB.strideH = 1;
loadData3DParamsForB.filterW = 1;
loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
loadData3DParamsForB.filterH = 1;
loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数
loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数
loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能
loadData3DParamsForB.fMatrixCtrl =
0; // 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1
loadData3DParamsForB.channelSize =
nL1SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
LoadData<KV_T, LOAD3DV2_CONFIG>(bL0Tensor, bL1Tensor[kL0 * baseK * baseN], loadData3DParamsForB);
LocalTensor<KV_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))];
LoadData3DParamsV2<KV_T> loadData3DParamsForA;
loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height
loadData3DParamsForA.l1W = 16; // 源操作数weight
loadData3DParamsForA.padList[0] = 0;
loadData3DParamsForA.padList[1] = 0;
loadData3DParamsForA.padList[2] = 0;
loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果
loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度
loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度
loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点
loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点
loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长
loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长
loadData3DParamsForA.filterW = 1; // 卷积核width
loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
loadData3DParamsForA.filterH = 1; // 卷积核height
loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数
loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数
loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置
loadData3DParamsForA.fMatrixCtrl = 0;
loadData3DParamsForA.channelSize =
kL0SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
LoadData<KV_T, LOAD3DV2_CONFIG>(aL0Tensor, aL1Tensor[kL0 * baseK * mL1SizeAlign],
loadData3DParamsForA);
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
MmadParams mmadParams;
mmadParams.m = mL1SizeAlign;
mmadParams.n = nL1SizeAlign;
mmadParams.k = kL0Size;
mmadParams.cmatrixInitVal = (kL0 == 0 && k1 == 0);
mmadParams.cmatrixSource = false;
mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (kL0 == (kL0Loops - 1))) ? 0b11 : 0b10;
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
PipeBarrier<PIPE_M>();
}
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
abL0BufIter++;
}
if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算?
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
}
if (k1 == (kL1Loops - 1)) {
// ND
FixpipeParamsV220 fixParams;
fixParams.nSize = nL1SizeAlign;
fixParams.mSize = mL1SizeAlign;
fixParams.srcStride = mL1SizeAlign;
fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔
fixParams.ndNum = 1; // 输出ND
fixParams.unitFlag = 0b11;
uint64_t mm2Offset = (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * nSize + nL1 * N_SPLIT_SIZE;
Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.bmm2ResUbSize + mm2Offset],
cL0Tensor, fixParams);
}
if (mL1Loops == 2) {
cL0BufIter++;
}
}
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
}
// cL0BufIter已经不在使用
if (mL1Loops == 1) {
cL0BufIter++;
}
}
qpL1BufIter += mL1Loops;
}
} // namespace SASKernel
#endif // SPARSE_ATTN_SHAREDKV_SCFA_BLOCK_CUBE_H

View File

@@ -0,0 +1,830 @@
/**
 * 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 sparse_attn_sharedkv_scfa_kernel.h
* \brief
*/
#ifndef SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H
#define SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "kernel_tiling/kernel_tiling.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#include "../sparse_attn_sharedkv_common.h"
#include "sparse_attn_sharedkv_scfa_block_cube.h"
#include "sparse_attn_sharedkv_scfa_block_vector.h"
#include "../sparse_attn_sharedkv_metadata.h"
namespace SASKernel {
using namespace matmul;
using namespace optiling;
using AscendC::CrossCoreSetFlag;
using AscendC::CrossCoreWaitFlag;
// 由于S2循环前,RunInfo还没有赋值,使用Bngs1Param临时存放B、N、S1轴相关的信息;同时减少重复计算
struct TempLoopInfo {
uint32_t bn2IdxInCurCore = 0;
uint32_t bIdx = 0U;
uint32_t n2Idx = 0U;
uint64_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小
uint32_t s2LoopTimes = 0U; // S2方向循环的总次数,无论TND还是BXXD都是等于实际次数,不用减1
int32_t actS1Size = 0; // TND场景下当前Batch循环处理的S1轴的大小
int32_t actOriS2Size = 0;
int32_t actCmpS2Size = 0;
bool curActSeqLenIsZero = false;
uint32_t tndCoreStartKVSplitPos = 0;
bool tndIsS2SplitCore = false;
uint32_t gS1Idx = 0U;
uint32_t s1StartIdx = 0;
uint32_t s1EndIdx = 0;
uint64_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小
uint32_t cmpLoopTimes = 0;
uint32_t oriLoopTimes = 0;
uint32_t v0OriSize = 0;
uint32_t v0CmpSize = 0;
// sparsemode = 4
int32_t oriMaskRight = 0;
int32_t oriMaskLeft = 0;
// sparsemode = 3
int32_t cmpMaskRight = 0;
uint64_t actualSeqQPrefixSum = 0;
uint64_t actualSeqKVPrefixSum = 0;
uint64_t actualSeqCmpKVPrefixSum = 0;
};
template <typename SAST>
class SparseAttnSharedkvScfa {
public:
// 中间计算数据类型为float,高精度模式
using T = float;
using Q_T = typename SAST::queryType;
using KV_T = typename SAST::kvType;
using OUT_T = typename SAST::outputType;
using SINKS_T = float;
using UPDATE_T = T;
using MM1_OUT_T = T;
using MM2_OUT_T = T;
__aicore__ inline SparseAttnSharedkvScfa(){};
__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* cuSeqlensKV, __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 *softmaxLse, __gm__ uint8_t *workspace,
const SparseAttnSharedkvTilingData *__restrict tiling, __gm__ uint8_t *gmTiling,
TPipe *tPipe);
__aicore__ inline void Process();
private:
static constexpr bool PAGE_ATTENTION = SAST::pageAttention;
static constexpr bool FLASH_DECODE = SAST::flashDecode;
static constexpr SAS_LAYOUT LAYOUT_T = SAST::layout;
static constexpr SAS_LAYOUT KV_LAYOUT_T = SAST::kvLayout;
static constexpr uint32_t PRELOAD_NUM = 2;
static constexpr uint32_t N_BUFFER_M_BASIC_SIZE = 256;
static constexpr uint32_t SAS_PRELOAD_TASK_CACHE_SIZE = 3;
static constexpr uint32_t SYNC_V0_C1_FLAG = 6;
static constexpr uint32_t SYNC_C1_V1_FLAG = 7;
static constexpr uint32_t SYNC_V1_C2_FLAG = 8;
static constexpr uint32_t SYNC_C2_V2_FLAG = 9;
static constexpr uint64_t kvHeadNum = 1ULL;
static constexpr uint64_t headDim = 512ULL;
static constexpr uint32_t dbWorkspaceRatio = PRELOAD_NUM;
const SparseAttnSharedkvTilingData *__restrict tilingData = nullptr;
TPipe *pipe = nullptr;
GlobalTensor<uint32_t> metadataGm;
uint64_t mSizeVStart = 0ULL;
uint64_t topKBaseOffset = 0ULL;
uint64_t tensorACoreOffset = 0ULL;
uint64_t tensorBCoreOffset = 0ULL;
uint64_t tensorCmpBCoreOffset = 0ULL;
uint32_t tmpBlockIdx = 0U;
uint32_t aiCoreIdx = 0U;
ConstInfo constInfo{};
TempLoopInfo tempLoopInfo{};
SASCubeBlock<SAST> cubeBlock;
SASVectorBlock<SAST> vectorBlock;
GlobalTensor<Q_T> queryGm;
GlobalTensor<KV_T> oriKvGm;
GlobalTensor<KV_T> cmpKvGm;
GlobalTensor<SINKS_T> sinksGm;
GlobalTensor<OUT_T> attentionOutGm;
GlobalTensor<T> softmaxLseGm;
GlobalTensor<int32_t> oriBlockTableGm;
GlobalTensor<int32_t> cmpBlockTableGm;
GlobalTensor<int32_t> topKGm;
GlobalTensor<int32_t> actualSeqLengthsQGm;
GlobalTensor<int32_t> actualSeqLengthsKVGm;
GlobalTensor<int32_t> actualSeqLengthsCmpKVGm;
// workspace
GlobalTensor<MM1_OUT_T> mm1ResGm;
GlobalTensor<KV_T> vec1ResGm;
GlobalTensor<MM2_OUT_T> mm2ResGm;
GlobalTensor<KV_T> kvMergeGm_;
GlobalTensor<int32_t> kvValidSizeGm_;
GlobalTensor<UPDATE_T> vec2ResGm;
GlobalTensor<T> accumOutGm;
// ================================Init functions==================================
__aicore__ inline void InitTilingData();
__aicore__ inline void InitCalcParamsEach();
__aicore__ inline void InitBuffers();
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV);
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV,
__gm__ uint8_t *actualSeqLengthsCmpKV);
__aicore__ inline void InitOutputSingleCore();
// ================================Process functions================================
__aicore__ inline void ProcessBalance();
__aicore__ inline void PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint64_t s2LoopIdx,
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE]);
// ================================Offset Calc=====================================
__aicore__ inline void GetSparseActualSeqLen();
__aicore__ inline void UpdateInnerLoopCond();
__aicore__ inline void CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint32_t s2LoopIdx,
RunInfo &info);
__aicore__ inline int32_t GetActualSeqLenQ(uint32_t bIdx);
__aicore__ inline int32_t GetActualSeqLenKV(uint32_t bIdx);
__aicore__ inline void GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx);
// ================================Mm1==============================================
__aicore__ inline void ComputeMm1(const RunInfo &info);
// ================================Mm2==============================================
__aicore__ inline void ComputeMm2(const RunInfo &info);
__aicore__ inline void InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx);
};
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitTilingData()
{
// singleCoreParams
// singleCoreTensorSize
constInfo.mmResUbSize = tilingData->baseParams.mmResUbSize;
constInfo.bmm2ResUbSize = tilingData->baseParams.bmm2ResUbSize;
// baseParams
constInfo.batchSize = tilingData->baseParams.batchSize;
constInfo.qHeadNum = constInfo.gSize = tilingData->baseParams.nNumOfQInOneGroup;
constInfo.kvSeqSize = tilingData->baseParams.kvSeqSize;
constInfo.qSeqSize = tilingData->baseParams.qSeqSize;
constInfo.oriMaxBlockNumPerBatch = tilingData->baseParams.oriMaxBlockNumPerBatch;
constInfo.cmpMaxBlockNumPerBatch = tilingData->cmpParams.cmpMaxBlockNumPerBatch;
constInfo.kvCacheBlockSize = tilingData->baseParams.paBlockSize;
constInfo.paOriBlockSize = tilingData->baseParams.oriBlockSize;
constInfo.paCmpBlockSize = tilingData->baseParams.cmpBlockSize;
constInfo.outputLayout = static_cast<SAS_LAYOUT>(tilingData->baseParams.outputLayout);
constInfo.kvHeadNum = kvHeadNum;
constInfo.headDim = headDim;
constInfo.oriMaskMode = tilingData->baseParams.oriMaskMode;
constInfo.oriKvStride = tilingData->baseParams.oriKvStride;
constInfo.oriWinLeft = tilingData->baseParams.oriWinLeft;
constInfo.oriWinRight = tilingData->baseParams.oriWinRight;
constInfo.actualLenDimsQ = tilingData->baseParams.actualLenDimsQ;
constInfo.actualLenDimsKV = tilingData->baseParams.actualLenDimsKV;
constInfo.returnSoftmaxLse = tilingData->baseParams.returnSoftmaxLse;
// innerSplitParams
constInfo.mBaseSize = constInfo.gSize;
constInfo.s2BaseSize = tilingData->baseParams.s2BaseSize;
constInfo.preLoadNum = PRELOAD_NUM;
constInfo.nBufferMBaseSize = N_BUFFER_M_BASIC_SIZE;
constInfo.syncV0C1 = SYNC_V0_C1_FLAG;
constInfo.syncC1V1 = SYNC_C1_V1_FLAG;
constInfo.syncV1C2 = SYNC_V1_C2_FLAG;
constInfo.syncC2V2 = SYNC_C2_V2_FLAG;
// cmp
constInfo.cmpRatio = tilingData->cmpParams.cmpRatio;
constInfo.sparseBlockCount = tilingData->cmpParams.sparseBlockCount;
constInfo.sparseBlockSize = 1; // sparseBlockSize 固定为1
constInfo.cmpMaskMode = tilingData->cmpParams.cmpMaskMode;
constInfo.cmpKvStride = tilingData->cmpParams.cmpKvStride;
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitBuffers()
{
if ASCEND_IS_AIV {
vectorBlock.InitBuffers(pipe);
} else {
cubeBlock.InitBuffers(pipe);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ,
__gm__ uint8_t *actualSeqLengthsKV)
{
if (constInfo.actualLenDimsKV != 0) {
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV);
}
if (constInfo.actualLenDimsQ != 0) {
actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ);
}
}
template <typename SAST>
__aicore__ inline void
SparseAttnSharedkvScfa<SAST>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV,
__gm__ uint8_t *actualSeqLengthsCmpKV)
{
if (constInfo.actualLenDimsKV != 0) {
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV);
actualSeqLengthsCmpKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsCmpKV, constInfo.actualLenDimsKV);
}
if (constInfo.actualLenDimsQ != 0) {
actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx)
{
if (constInfo.outputLayout == SAS_LAYOUT::TND) {
if (tempLoopInfo.actS1Size == 0) {
return;
}
uint32_t tBase = actualSeqLengthsQGm.GetValue(bIdx);
uint32_t s1Count = tempLoopInfo.actS1Size;
uint64_t attenOutOffset = (tBase + s1Idx) * kvHeadNum * constInfo.gSize * headDim + // T轴、s1轴偏移
n2Idx * constInfo.gSize * headDim; // N2轴偏移
uint64_t lseOffset = (tBase + s1Idx) * constInfo.gSize + // T轴、s1轴偏移
n2Idx * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移
matmul::InitOutput<OUT_T>(attentionOutGm[attenOutOffset], constInfo.gSize * headDim, 0);
if (constInfo.returnSoftmaxLse) {
matmul::InitOutput<T>(softmaxLseGm[lseOffset], constInfo.gSize, 0);
}
} else if (constInfo.outputLayout == SAS_LAYOUT::BSND) {
uint64_t attenOutOffset = bIdx * constInfo.qSeqSize * kvHeadNum * constInfo.gSize * headDim +
s1Idx * kvHeadNum * constInfo.gSize * headDim + // B轴、S1轴偏移
n2Idx * constInfo.gSize * headDim; // N2轴偏移
uint64_t lseOffset = bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移
n2Idx * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移
s1Idx * constInfo.gSize; // S1轴偏移
matmul::InitOutput<OUT_T>(attentionOutGm[attenOutOffset], constInfo.gSize * headDim, 0);
if (constInfo.returnSoftmaxLse) {
matmul::InitOutput<T>(softmaxLseGm[lseOffset], constInfo.gSize, 0);
}
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitOutputSingleCore()
{
uint32_t coreNum = GetBlockNum();
if (coreNum != 0) {
uint64_t totalOutputSize = constInfo.batchSize * constInfo.qHeadNum * constInfo.qSeqSize * constInfo.headDim;
uint64_t singleCoreSize = (totalOutputSize + (2 * coreNum) - 1) / (2 * coreNum); // 2 means c:v = 1:2
uint64_t tailSize = totalOutputSize - tmpBlockIdx * singleCoreSize;
uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize;
if (singleInitOutputSize > 0) {
matmul::InitOutput<OUT_T>(attentionOutGm[tmpBlockIdx * singleCoreSize], singleInitOutputSize, 0);
}
SyncAll();
}
}
template <typename SAST>
__aicore__ inline int32_t SparseAttnSharedkvScfa<SAST>::GetActualSeqLenQ(uint32_t bIdx)
{
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
int32_t actualSeqQPrefixSum = actualSeqLengthsQGm.GetValue(bIdx);
int32_t actualSeqQNextSum = actualSeqLengthsQGm.GetValue(bIdx + 1);
tempLoopInfo.actualSeqQPrefixSum = static_cast<uint64_t>(actualSeqQPrefixSum);
return actualSeqQNextSum - actualSeqQPrefixSum;
} else {
tempLoopInfo.actualSeqQPrefixSum = static_cast<uint64_t>(bIdx * constInfo.qSeqSize);
if (constInfo.actualLenDimsQ == 0) {
return static_cast<int32_t>(constInfo.qSeqSize);
} else {
return actualSeqLengthsQGm.GetValue(bIdx);
}
}
}
template <typename SAST>
__aicore__ inline int32_t SparseAttnSharedkvScfa<SAST>::GetActualSeqLenKV(uint32_t bIdx)
{
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
tempLoopInfo.actualSeqKVPrefixSum = static_cast<uint64_t>(bIdx * constInfo.kvSeqSize);
if (constInfo.actualLenDimsKV == 0) {
return static_cast<int32_t>(constInfo.kvSeqSize);
}
return actualSeqLengthsKVGm.GetValue(bIdx);
} else if constexpr(KV_LAYOUT_T == SAS_LAYOUT::BSND) {
return static_cast<int32_t>(constInfo.kvSeqSize);
} else if constexpr(KV_LAYOUT_T == SAS_LAYOUT::TND) {
int32_t actualSeqKVPrefixSum = actualSeqLengthsKVGm.GetValue(bIdx);
int32_t actualSeqKVNextSum = actualSeqLengthsKVGm.GetValue(bIdx + 1);
tempLoopInfo.actualSeqCmpKVPrefixSum = actualSeqLengthsCmpKVGm.GetValue(bIdx);
tempLoopInfo.actualSeqKVPrefixSum = actualSeqKVPrefixSum;
return actualSeqKVNextSum - actualSeqKVPrefixSum;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::GetSparseActualSeqLen()
{
// 行无效通过ori部分判断, ori部分如果有行无效那么ori和cmp都有
if (static_cast<int32_t>(tempLoopInfo.s1EndIdx) < -(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size)) {
tempLoopInfo.actOriS2Size = 0;
tempLoopInfo.actCmpS2Size = 0;
return;
}
// 对于cmp部分还有top k, tempLoopInfo.actS2Size只针对cmp
int32_t thresHold = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio;
tempLoopInfo.actCmpS2Size = Min(constInfo.sparseBlockCount * constInfo.sparseBlockSize, thresHold);
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::UpdateInnerLoopCond()
{
if ((tempLoopInfo.actCmpS2Size == 0 && tempLoopInfo.actOriS2Size == 0) || (tempLoopInfo.actS1Size == 0)) {
tempLoopInfo.curActSeqLenIsZero = true;
return;
}
tempLoopInfo.curActSeqLenIsZero = false;
tempLoopInfo.mBasicSizeTail = (tempLoopInfo.actS1Size * constInfo.gSize) % constInfo.mBaseSize;
tempLoopInfo.mBasicSizeTail =
(tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail;
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::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* cuSeqlensKV, __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 *softmaxLse,
__gm__ uint8_t *workspace, const SparseAttnSharedkvTilingData *__restrict tiling, __gm__ uint8_t *gmTiling,
TPipe *tPipe)
{
if ASCEND_IS_AIV {
tmpBlockIdx = GetBlockIdx(); // vec:0-47
aiCoreIdx = tmpBlockIdx / 2;
} else {
tmpBlockIdx = GetBlockIdx(); // cube:0-23
aiCoreIdx = tmpBlockIdx;
}
// init tiling data
tilingData = tiling;
InitTilingData();
if (KV_LAYOUT_T == SAS_LAYOUT::TND && LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(cuSeqlensQ, cuSeqlensKV, cuSeqlensCmpKV);
} else if (KV_LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(seqUsedQ, cuSeqlensKV, cuSeqlensCmpKV);
} else if ((KV_LAYOUT_T == SAS_LAYOUT::PA_ND || KV_LAYOUT_T == SAS_LAYOUT::BSND) && LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(cuSeqlensQ, seqUsedKV);
} else if ((KV_LAYOUT_T == SAS_LAYOUT::PA_ND || KV_LAYOUT_T == SAS_LAYOUT::BSND)) {
InitActualSeqLen(seqUsedQ, seqUsedKV);
}
metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata);
InitCalcParamsEach();
pipe = tPipe;
// init global buffer
queryGm.SetGlobalBuffer((__gm__ Q_T *)query);
oriKvGm.SetGlobalBuffer((__gm__ KV_T *)oriKV);
cmpKvGm.SetGlobalBuffer((__gm__ KV_T *)cmpKV);
if (sinks != nullptr) {
sinksGm.SetGlobalBuffer((__gm__ SINKS_T *)sinks);
}
attentionOutGm.SetGlobalBuffer((__gm__ OUT_T *)attentionOut);
softmaxLseGm.SetGlobalBuffer((__gm__ T *)softmaxLse);
if ASCEND_IS_AIV {
if (LAYOUT_T != SAS_LAYOUT::TND) {
if (constInfo.needInit) {
InitOutputSingleCore();
}
}
}
if constexpr (PAGE_ATTENTION) {
oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable);
cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable);
}
topKGm.SetGlobalBuffer((__gm__ int32_t *)cmpSparseIndices);
// workspace 内存排布
// |Q--|mm1ResGm|vec1ResGm|mm2ResGm|vec2ResGm
// |Core0_Q1-Core0_Q2-Core1_Q1-Core1_Q2....Core32_Q1-Core32_Q2|Core0_mmRes
uint64_t offset = 0;
mm1ResGm.SetGlobalBuffer(
(__gm__ MM1_OUT_T *)(workspace + offset +
aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T);
vec1ResGm.SetGlobalBuffer(
(__gm__ Q_T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T);
mm2ResGm.SetGlobalBuffer(
(__gm__ MM2_OUT_T *)(workspace + offset +
aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T);
vec2ResGm.SetGlobalBuffer(
(__gm__ T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T);
kvMergeGm_.SetGlobalBuffer((__gm__ KV_T *)(workspace + offset + aiCoreIdx * 512 * 512 * 4 * sizeof(KV_T)));
offset += GetBlockNum() * 512 * 512 * 4 * sizeof(KV_T);
kvValidSizeGm_.SetGlobalBuffer(
(__gm__ int32_t *)(workspace + offset + (aiCoreIdx * 2) * 128 * 4 * sizeof(int32_t)));
if ASCEND_IS_AIV {
vectorBlock.InitParams(constInfo, tilingData);
vectorBlock.InitVec0GlobalTensor(kvValidSizeGm_, kvMergeGm_, oriKvGm, cmpKvGm, oriBlockTableGm,
cmpBlockTableGm);
vectorBlock.InitVec1GlobalTensor(mm1ResGm, vec1ResGm, actualSeqLengthsQGm, actualSeqLengthsKVGm, topKGm,
sinksGm, softmaxLseGm);
vectorBlock.InitVec2GlobalTensor(accumOutGm, vec2ResGm, mm2ResGm, attentionOutGm);
}
if ASCEND_IS_AIC {
cubeBlock.InitParams(constInfo);
cubeBlock.InitMm1GlobalTensor(queryGm, oriKvGm, cmpKvGm, mm1ResGm);
cubeBlock.InitMm2GlobalTensor(vec1ResGm, mm2ResGm, attentionOutGm);
cubeBlock.InitPageAttentionInfo(oriKvGm, kvMergeGm_, oriBlockTableGm, cmpBlockTableGm);
}
// 要在InitParams之后执行
if (pipe != nullptr) {
InitBuffers();
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::InitCalcParamsEach()
{
if (aiCoreIdx != 0) {
constInfo.bN2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_START_INDEX, false));
constInfo.gS1Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_START_INDEX, false));
constInfo.s2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_START_INDEX, false));
}
constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_END_INDEX, false));
constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_END_INDEX, false));
constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_END_INDEX, false));
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start,
uint32_t s2LoopIdx, RunInfo &info)
{
info.isValid = s2LoopIdx < tempLoopInfo.s2LoopTimes;
info.loop = loop;
info.cmpLoop = cmpLoop;
info.bIdx = tempLoopInfo.bIdx;
info.n2IdxReal = tempLoopInfo.n2Idx;
info.gS1Idx = tempLoopInfo.gS1Idx;
info.s1Idx = tempLoopInfo.gS1Idx / constInfo.gSize;
info.s2Idx = s2LoopIdx;
info.curSInnerLoopTimes = tempLoopInfo.s2LoopTimes;
info.tndIsS2SplitCore = tempLoopInfo.tndIsS2SplitCore;
info.tndCoreStartKVSplitPos = tempLoopInfo.tndCoreStartKVSplitPos;
info.isBmm2Output = false;
info.actS1Size = tempLoopInfo.actS1Size;
// M方向的尾块
info.actMBaseSize = tempLoopInfo.mBasicSizeTail;
if ASCEND_IS_AIV {
info.mSize = info.actMBaseSize;
info.mSizeV = (info.mSize <= 16) ? info.mSize : ((CeilDiv(info.mSize, 16) + 1) / 2 * 16);
info.mSizeVStart = 0;
if (tmpBlockIdx % 2 == 1) {
info.mSizeVStart = info.mSizeV;
info.mSizeV = info.mSize - info.mSizeV;
}
}
info.isFirstSInnerLoop = s2LoopIdx == s2Start;
if (info.isFirstSInnerLoop) {
tempLoopInfo.bn2IdxInCurCore++;
}
info.isLastS2Loop = (s2LoopIdx == (tempLoopInfo.s2LoopTimes - 1));
info.bn2IdxInCurCore = tempLoopInfo.bn2IdxInCurCore - 1;
uint64_t tndBIdxOffsetForQ = tempLoopInfo.actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim;
uint64_t tndBIdxOffsetForKV = tempLoopInfo.actualSeqKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim;
uint64_t tndBIdxOffsetForCmpKV = tempLoopInfo.actualSeqCmpKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim;
if (info.isFirstSInnerLoop) {
tensorACoreOffset = tndBIdxOffsetForQ + info.gS1Idx * constInfo.headDim;
tensorBCoreOffset = tndBIdxOffsetForKV + info.n2Idx * constInfo.headDim; // 当前为PA场景,该变量失效
tensorCmpBCoreOffset = tndBIdxOffsetForCmpKV + info.n2Idx * constInfo.headDim;
if constexpr (LAYOUT_T == SAS_LAYOUT::BSND) { // B,S1,N2 K
topKBaseOffset = (info.bIdx * constInfo.qSeqSize + tempLoopInfo.s1StartIdx) * constInfo.kvHeadNum *
constInfo.sparseBlockCount +
info.n2Idx * constInfo.sparseBlockCount;
} else if (LAYOUT_T == SAS_LAYOUT::TND) { // T N2 K
topKBaseOffset = (tempLoopInfo.actualSeqQPrefixSum + tempLoopInfo.s1StartIdx) * constInfo.kvHeadNum *
constInfo.sparseBlockCount +
info.n2Idx * constInfo.sparseBlockCount;
}
}
info.tensorAOffset = tensorACoreOffset;
info.tensorBOffset = tensorBCoreOffset;
info.tensorCmpBOffset = tensorCmpBCoreOffset;
info.attenOutOffset = tensorACoreOffset;
info.topKBaseOffset = topKBaseOffset;
if (s2LoopIdx < tempLoopInfo.oriLoopTimes) {
// S2首次循环只能在ori_kv
info.isOri = true;
info.relativeS2Idx = 0;
uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize;
if (s2LoopIdx + 1 == tempLoopInfo.oriLoopTimes) {
info.actualSingleProcessSInnerSize = (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset;
} else {
info.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
}
info.s2StartPoint = tempLoopInfo.oriMaskLeft;
info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio;
} else {
info.isOri = false;
info.relativeS2Idx = info.s2Idx - tempLoopInfo.oriLoopTimes;
uint64_t s2Offset = (info.s2Idx - tempLoopInfo.oriLoopTimes) * constInfo.s2BaseSize;
if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes) {
info.actualSingleProcessSInnerSize = tempLoopInfo.actCmpS2Size - s2Offset;
} else {
info.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
}
info.s2StartPoint = 0;
info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio;
}
info.actualSingleProcessSInnerSizeAlign = SASAlign(info.actualSingleProcessSInnerSize, SASVectorBlock<SAST>::BYTE_BLOCK);
if (info.isOri) {
info.v0S2Start = 0;
info.v0S2DealSize = 0;
} else {
info.v0S2Start = 0;
if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes && s2LoopIdx == 2) { // tail
info.v0S2Start = 512;
}
info.v0S2DealSize = 512;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::ComputeMm1(const RunInfo &info)
{
uint32_t nBufferLoopTimes = CeilDiv(info.actMBaseSize, constInfo.nBufferMBaseSize);
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
cubeBlock.ComputeMm1(info, mSplitInfo);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V1);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::ComputeMm2(const RunInfo &info)
{
uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize;
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
CrossCoreWaitFlag(constInfo.syncV1C2);
cubeBlock.ComputeMm2(info, mSplitInfo);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_FIX>(constInfo.syncC2V2);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::Process()
{
uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_CORE_ENABLE_INDEX, false));
if (hasLoad == 0) {
return;
}
if ASCEND_IS_AIV {
vectorBlock.AllocEventID();
vectorBlock.InitSoftmaxDefaultBuffer();
} else {
cubeBlock.AllocEventID();
}
ProcessBalance();
if ASCEND_IS_AIV {
vectorBlock.FreeEventID();
} else {
cubeBlock.FreeEventID();
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx)
{
bIdx = bN2Idx / kvHeadNum;
n2Idx = bN2Idx % kvHeadNum;
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::ProcessBalance()
{
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE];
uint32_t gloop = 0;
uint32_t cmpLoop = 0;
uint32_t gS1LoopEnd = 0;
bool globalLoopStart = true;
if ASCEND_IS_AIC {
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE2>(3);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE2>(3);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE2>(3);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE2>(3);
}
// 适配左闭右开
if (constInfo.bN2Start == constInfo.bN2End) {
if (constInfo.gS1Start != constInfo.gS1End || constInfo.s2Start != constInfo.s2End) {
constInfo.bN2End += 1;
}
} else if ((constInfo.gS1End != 0) || (constInfo.s2End != 0)) {
constInfo.bN2End += 1;
}
for (uint32_t bN2LoopIdx = constInfo.bN2Start; bN2LoopIdx < constInfo.bN2End; bN2LoopIdx++) {
GetBN2Idx(bN2LoopIdx, tempLoopInfo.bIdx, tempLoopInfo.n2Idx);
tempLoopInfo.actS1Size = GetActualSeqLenQ(tempLoopInfo.bIdx); // 获取actualSeqLength
bool isS1ZeroAndLastBatch = (tempLoopInfo.actS1Size == 0) &&
((constInfo.outputLayout == SAS_LAYOUT::BSND) || (bN2LoopIdx + 1 == constInfo.bN2End));
uint32_t gS1SplitNum = CeilDiv(tempLoopInfo.actS1Size * constInfo.gSize, constInfo.mBaseSize);
// 当处于最后一个BN2时, 且gS1End为0时, 说明当前BN2里的所有数据都在当前核处理
gS1LoopEnd = (bN2LoopIdx + 1 == constInfo.bN2End && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum;
// 当处于最后一个BN2且当前S1为0时,需要进入循环计算preload导致的未完成的部分
gS1LoopEnd = isS1ZeroAndLastBatch ? gS1LoopEnd + 1 : gS1LoopEnd;
for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < gS1LoopEnd; gS1LoopIdx++) {
tempLoopInfo.actOriS2Size = GetActualSeqLenKV(tempLoopInfo.bIdx);
// 计算需要的数据, 避免重复计算
tempLoopInfo.gS1Idx = gS1LoopIdx * constInfo.mBaseSize;
tempLoopInfo.s1StartIdx = tempLoopInfo.gS1Idx / constInfo.gSize;
tempLoopInfo.s1EndIdx =
Min((tempLoopInfo.s1StartIdx + constInfo.mBaseSize / constInfo.gSize - 1), tempLoopInfo.actS1Size - 1);
// 此处均为闭区间
tempLoopInfo.oriMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size +
static_cast<int32_t>(tempLoopInfo.s1EndIdx) + constInfo.oriWinRight;
tempLoopInfo.oriMaskLeft = Max(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size +
static_cast<int32_t>(tempLoopInfo.s1EndIdx) - constInfo.oriWinLeft,
0);
tempLoopInfo.cmpMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size;
GetSparseActualSeqLen();
UpdateInnerLoopCond();
uint32_t oriS2Size = tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1;
uint32_t oriSplitNum = 0;
uint32_t cmpSplitNum = 0;
uint32_t cmpS2Size = 0;
bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == gS1LoopEnd);
if (tempLoopInfo.curActSeqLenIsZero) {
if ASCEND_IS_AIV {
InitAllZeroOutput(tempLoopInfo.bIdx, tempLoopInfo.s1StartIdx, tempLoopInfo.n2Idx);
}
if (!isEnd) {
continue;
}
} else {
oriSplitNum = CeilDiv(oriS2Size, constInfo.s2BaseSize);
cmpS2Size = tempLoopInfo.actCmpS2Size;
cmpSplitNum = CeilDiv(cmpS2Size, constInfo.s2BaseSize);
}
uint32_t s2SplitNum = oriSplitNum + cmpSplitNum;
constexpr uint32_t V0_SPLIT = 32; // align to 32
uint32_t v0OriSize = CeilDiv(oriS2Size * cmpS2Size, oriS2Size + cmpS2Size);
if (cmpS2Size > V0_SPLIT * oriSplitNum) {
v0OriSize = SASAlign(v0OriSize, V0_SPLIT * oriSplitNum);
}
uint32_t v0CmpSize = cmpS2Size - v0OriSize;
tempLoopInfo.oriLoopTimes = oriSplitNum;
tempLoopInfo.cmpLoopTimes = cmpSplitNum;
tempLoopInfo.s2LoopTimes = s2SplitNum;
tempLoopInfo.v0OriSize = v0OriSize;
tempLoopInfo.v0CmpSize = v0CmpSize;
uint32_t s2LoopEnd = (isEnd && constInfo.s2End != 0) ? constInfo.s2End : tempLoopInfo.s2LoopTimes;
tempLoopInfo.s2LoopTimes = s2LoopEnd;
// 分核修改后需要打开
// 当前s2是否被切,决定了输出是否要写到attenOut上
tempLoopInfo.tndIsS2SplitCore = ((constInfo.s2Start == 0) && (s2LoopEnd == s2SplitNum)) ? false : true;
tempLoopInfo.tndCoreStartKVSplitPos = globalLoopStart ? constInfo.coreStartKVSplitPos : 0;
uint32_t extraLoop = isEnd ? 2 : 0;
uint32_t curTopKIdx = 0;
for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (s2LoopEnd + extraLoop); s2LoopIdx++) {
PreloadPipeline(gloop, cmpLoop, constInfo.s2Start, s2LoopIdx, extraInfo);
++gloop;
if (s2LoopIdx >= tempLoopInfo.oriLoopTimes && s2LoopIdx < s2LoopEnd) { // 用于判断v0使用的循环GM的id
++cmpLoop;
}
}
globalLoopStart = false;
constInfo.s2Start = 0;
}
constInfo.gS1Start = 0;
}
if ASCEND_IS_AIV {
CrossCoreWaitFlag(3);
CrossCoreWaitFlag(3);
CrossCoreWaitFlag(3);
CrossCoreWaitFlag(3);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvScfa<SAST>::PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start,
uint64_t s2LoopIdx,
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE])
{
RunInfo &extraInfo0 = extraInfo[loop % SAS_PRELOAD_TASK_CACHE_SIZE]; // 本轮任务
RunInfo &extraInfo2 = extraInfo[(loop + 2) % SAS_PRELOAD_TASK_CACHE_SIZE]; // 上一轮任务
RunInfo &extraInfo1 = extraInfo[(loop + 1) % SAS_PRELOAD_TASK_CACHE_SIZE]; // 上两轮任务
CalcParams(loop, cmpLoop, s2Start, s2LoopIdx, extraInfo0);
if (extraInfo0.isValid) {
if ASCEND_IS_AIC {
if (!extraInfo0.isOri) {
CrossCoreWaitFlag(constInfo.syncV0C1);
}
ComputeMm1(extraInfo0);
} else {
if (extraInfo0.isFirstSInnerLoop) {
CrossCoreWaitFlag(3);
}
vectorBlock.ProcessVec0L(extraInfo0);
if (!extraInfo0.isOri) {
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE3>(constInfo.syncV0C1);
}
}
}
if (extraInfo2.isValid) {
if ASCEND_IS_AIV {
vectorBlock.ProcessVec1L(extraInfo2);
}
if ASCEND_IS_AIC {
ComputeMm2(extraInfo2);
if (extraInfo2.isLastS2Loop) {
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE2>(3);
}
}
}
if (extraInfo1.isValid) {
if ASCEND_IS_AIV {
vectorBlock.ProcessVec2L(extraInfo1);
}
extraInfo1.isValid = false;
}
}
} // namespace SASKernel
#endif // SPARSE_ATTN_SHAREDKV_SCFA_KERNEL_H

View File

@@ -0,0 +1,953 @@
/**
 * 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 sparse_attn_sharedkv_swa_block_cube.h
* \brief use 7 buffer for matmul l1, better pipeline
*/
#ifndef SPARSE_ATTN_SHAREDKV_SWA_BLOCK_CUBE_H
#define SPARSE_ATTN_SHAREDKV_SWA_BLOCK_CUBE_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "kernel_tiling/kernel_tiling.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#include "../sparse_attn_sharedkv_common.h"
namespace SASKernel {
template <typename SAST>
class SWACubeBlock {
public:
// 中间计算数据类型为float, 高精度模式
using T = float;
using Q_T = typename SAST::queryType;
using KV_T = typename SAST::kvType;
using OUT_T = typename SAST::outputType;
using MM_OUT_T = T;
__aicore__ inline SWACubeBlock(){};
__aicore__ inline void InitParams(const ConstInfo &constInfo);
__aicore__ inline void InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<KV_T> oriKvGm,
GlobalTensor<KV_T> cmpKV, GlobalTensor<MM_OUT_T> mm1ResGm);
__aicore__ inline void InitMm2GlobalTensor(GlobalTensor<KV_T> vec1ResGm, GlobalTensor<MM_OUT_T> mm2ResGm,
GlobalTensor<OUT_T> attentionOutGm);
__aicore__ inline void InitPageAttentionInfo(GlobalTensor<KV_T> oriKvGm, // const GlobalTensor<KV_T>& kvMergeGm,
GlobalTensor<int32_t> oriBlockTableGm,
GlobalTensor<int32_t> cmpBlockTableGm);
__aicore__ inline void InitBuffers(TPipe *pipe);
__aicore__ inline void AllocEventID();
__aicore__ inline void FreeEventID();
__aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo);
__aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo);
private:
static constexpr bool PAGE_ATTENTION = SAST::pageAttention;
// static constexpr int TEMPLATE_MODE = SAST::templateMode;
static constexpr bool FLASH_DECODE = SAST::flashDecode;
static constexpr SAS_LAYOUT LAYOUT_T = SAST::layout;
static constexpr SAS_LAYOUT KV_LAYOUT_T = SAST::kvLayout;
static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分
static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分
static constexpr uint32_t K_L0_SPLIT_SIZE = 128; // k方向L0切分
static constexpr uint32_t K_L1_SPLIT_SIZE = 256; // k方向L1切分
static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分
static constexpr uint32_t D_SPLIT_SIZE = 256; // d轴切分
static constexpr uint32_t L1_BLOCK_SIZE = (64 * 512 * sizeof(Q_T));
static constexpr uint32_t L1_BLOCK_OFFSET = 64 * 512;
static constexpr uint32_t L0A_PP_SIZE = (32 * 1024);
static constexpr uint32_t L0B_PP_SIZE = (32 * 1024);
static constexpr uint32_t L0C_PP_SIZE = (64 * 1024);
// mte2 <> mte1 EventID
// L1 3buf, 使用3个eventId
static constexpr uint32_t L1_EVENT0 = EVENT_ID2;
static constexpr uint32_t L1_EVENT1 = EVENT_ID3;
static constexpr uint32_t L1_EVENT2 = EVENT_ID4;
static constexpr uint32_t L1_EVENT3 = EVENT_ID5;
static constexpr uint32_t L1_EVENT4 = EVENT_ID6;
static constexpr uint32_t L1_EVENT5 = EVENT_ID7;
static constexpr uint32_t L1_EVENT6 = EVENT_ID1;
// m <> mte1 EventID
static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3;
static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4;
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用
static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6};
ConstInfo constInfo{};
// L1分成3块buf, 用于记录
uint32_t qpL1BufIter = 0;
uint32_t kvL1BufIter = -1;
uint32_t abL0BufIter = 0;
uint32_t cL0BufIter = 0;
// mm1
GlobalTensor<Q_T> queryGm;
GlobalTensor<KV_T> keyGm;
GlobalTensor<MM_OUT_T> mm1ResGm;
// GlobalTensor<KV_T> kvMergeGm_;
GlobalTensor<KV_T> oriKvGm;
GlobalTensor<KV_T> cmpKvGm;
// mm2
GlobalTensor<KV_T> vec1ResGm;
GlobalTensor<KV_T> valueGm;
GlobalTensor<MM_OUT_T> mm2ResGm;
GlobalTensor<OUT_T> attentionOutGm;
// block_table
GlobalTensor<int32_t> oriBlockTableGm;
GlobalTensor<int32_t> cmpBlockTableGm;
TBuf<TPosition::A1> bufQPL1;
TBuf<TPosition::A1> bufKVL1;
TBuf<TPosition::A2> tmpBufL0A;
TBuf<TPosition::B2> tmpBufL0B;
TBuf<TPosition::CO1> tmpBufL0C;
LocalTensor<Q_T> l1QPTensor;
LocalTensor<Q_T> l1KVTensor;
LocalTensor<KV_T> aL0TensorPingPong;
LocalTensor<KV_T> bL0TensorPingPong;
LocalTensor<MM_OUT_T> cL0TensorPingPong;
// L0AB m <> mte1 EventID
__aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx)
{
return (L0AB_EVENT0 + idx);
}
__aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx)
{
uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连
return idxMap[mIdx % 2] + k1Idx;
}
__aicore__ inline void CopyGmToL1(LocalTensor<KV_T> &l1Tensor, GlobalTensor<KV_T> &gmSrcTensor, uint32_t srcN,
uint32_t srcD, uint32_t srcDstride);
__aicore__ inline void CopyInMm1AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset);
__aicore__ inline void CopyInMm2AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset);
__aicore__ inline void LoadDataMm1A(LocalTensor<KV_T> &aL0Tensor, LocalTensor<KV_T> &aL1Tensor, uint32_t idx,
uint32_t kSplitSize, uint32_t mSize, uint32_t kSize);
__aicore__ inline void LoadDataMm1B(LocalTensor<KV_T> &bL0Tensor, LocalTensor<KV_T> &bL1Tensor, uint32_t idx,
uint32_t kSplitSize, uint32_t kSize, uint32_t nSize);
};
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::InitParams(const ConstInfo &constInfo)
{
this->constInfo = constInfo;
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<KV_T> oriKvGm,
GlobalTensor<KV_T> cmpKvGm,
GlobalTensor<MM_OUT_T> mm1ResGm)
{
// mm1
this->queryGm = queryGm;
this->oriKvGm = oriKvGm;
if (constInfo.templateMode == CFA_TEMPLATE) {
this->cmpKvGm = cmpKvGm;
}
this->mm1ResGm = mm1ResGm;
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::InitMm2GlobalTensor(GlobalTensor<KV_T> vec1ResGm,
GlobalTensor<MM_OUT_T> mm2ResGm,
GlobalTensor<OUT_T> attentionOutGm)
{
// mm2
this->vec1ResGm = vec1ResGm;
this->mm2ResGm = mm2ResGm;
this->attentionOutGm = attentionOutGm;
}
template <typename SAST>
__aicore__ inline void
SWACubeBlock<SAST>::InitPageAttentionInfo(GlobalTensor<KV_T> oriKvGm, // const GlobalTensor<KV_T>& kvMergeGm,
GlobalTensor<int32_t> oriBlockTableGm, GlobalTensor<int32_t> cmpBlockTableGm)
{
this->oriKvGm = oriKvGm;
this->oriBlockTableGm = oriBlockTableGm;
if (constInfo.templateMode == CFA_TEMPLATE) {
this->cmpBlockTableGm = cmpBlockTableGm;
}
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::InitBuffers(TPipe *pipe)
{
pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4);
l1QPTensor = bufQPL1.Get<Q_T>();
pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3);
l1KVTensor = bufKVL1.Get<KV_T>();
// L0A
pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K
aL0TensorPingPong = tmpBufL0A.Get<KV_T>();
// L0B
pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K
bL0TensorPingPong = tmpBufL0B.Get<KV_T>();
// L0C
pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K
cL0TensorPingPong = tmpBufL0C.Get<MM_OUT_T>();
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::AllocEventID()
{
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::FreeEventID()
{
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::CopyGmToL1(LocalTensor<KV_T> &l1Tensor, GlobalTensor<KV_T> &gmSrcTensor,
uint32_t srcN, uint32_t srcD, uint32_t srcDstride)
{
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = srcN; // 行数
nd2nzPara.dValue = srcD;
nd2nzPara.srcDValue = srcDstride;
nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(l1Tensor, gmSrcTensor, nd2nzPara);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::CopyInMm1AToL1(LocalTensor<KV_T> &l1Tensor, const RunInfo &info,
uint32_t mSeqIdx, uint32_t mSizeAct, uint32_t headSize,
uint32_t headOffset)
{
auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.headDim + headOffset];
CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, constInfo.headDim);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::LoadDataMm1A(LocalTensor<KV_T> &aL0Tensor, LocalTensor<KV_T> &aL1Tensor,
uint32_t idx, uint32_t kSplitSize, uint32_t mSize,
uint32_t kSize)
{
LocalTensor<KV_T> srcTensor = aL1Tensor[mSize * kSplitSize * idx];
LoadData3DParamsV2<KV_T> loadData3DParams;
// SetFmatrixParams
loadData3DParams.l1H = mSize / 16; // Hin=M1=8
loadData3DParams.l1W = 16; // Win=M0
loadData3DParams.padList[0] = 0;
loadData3DParams.padList[1] = 0;
loadData3DParams.padList[2] = 0;
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
// SetLoadToA0Params
loadData3DParams.mExtension = mSize; // M
loadData3DParams.kExtension = kSize; // K
loadData3DParams.mStartPt = 0;
loadData3DParams.kStartPt = 0;
loadData3DParams.strideW = 1;
loadData3DParams.strideH = 1;
loadData3DParams.filterW = 1;
loadData3DParams.filterSizeW = (1 >> 8) & 255;
loadData3DParams.filterH = 1;
loadData3DParams.filterSizeH = (1 >> 8) & 255;
loadData3DParams.dilationFilterW = 1;
loadData3DParams.dilationFilterH = 1;
loadData3DParams.enTranspose = 0;
loadData3DParams.fMatrixCtrl = 0;
loadData3DParams.channelSize = kSize; // Cin=K
LoadData<KV_T, LOAD3DV2_CONFIG>(aL0Tensor, srcTensor, loadData3DParams);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::LoadDataMm1B(LocalTensor<KV_T> &l0Tensor, LocalTensor<KV_T> &l1Tensor,
uint32_t idx, uint32_t kSplitSize, uint32_t kSize,
uint32_t nSize)
{
// N 方向全载
LocalTensor<KV_T> srcTensor = l1Tensor[nSize * kSplitSize * idx];
LoadData2DParams loadData2DParams;
loadData2DParams.startIndex = 0;
loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(KV_T));
loadData2DParams.srcStride = 1;
loadData2DParams.dstGap = 0;
loadData2DParams.ifTranspose = false;
LoadData(l0Tensor, srcTensor, loadData2DParams);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::CopyInMm2AToL1(LocalTensor<KV_T> &aL1Tensor, const RunInfo &info,
uint32_t mSeqIdx, uint32_t subMSizeAct, uint32_t nSize,
uint32_t nOffset)
{
auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize +
mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset];
CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign);
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo)
{
uint32_t mSize = mSplitInfo.nBufferDealM;
uint32_t mL1Size = M_SPLIT_SIZE;
uint32_t mL1SizeAlign = SASAlign(M_SPLIT_SIZE, 16);
uint32_t mL1Loops = CeilDiv(mSize, M_SPLIT_SIZE);
uint32_t nSize = info.actualSingleProcessSInnerSize;
uint32_t nL1Size = N_SPLIT_SIZE;
uint32_t nL1SizeAlign = SASAlign(N_SPLIT_SIZE, 16);
uint32_t nL1Loops = CeilDiv(nSize, N_SPLIT_SIZE);
uint32_t kSize = 512;
uint32_t kL1Size = 256;
uint32_t kL1Loops = 2;
uint32_t kL0Size = 128;
uint32_t kL0Loops = CeilDiv(kL1Size, kL0Size);
LocalTensor<KV_T> bL1Tensor;
LocalTensor<KV_T> kTensor;
uint32_t ka = 0, kb = 0;
uint32_t copyRowCnt = 0;
uint32_t copyRowCntTmp = 0;
// L1 切n切k
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // L1切n, 512/128=4
if (nL1 == (nL1Loops - 1)) {
// 尾块重新计算size
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
nL1SizeAlign = SASAlign(nL1Size, 16);
}
for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) {
kvL1BufIter++;
uint32_t kb = kvL1BufIter % 3;
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
// 从k当中取当前的块
bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
uint32_t curSeqIdx = info.s2BatchOffset + nL1 * N_SPLIT_SIZE;
uint32_t copyFinishRowCnt = 0;
if (info.isOri) {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + nL1 * N_SPLIT_SIZE;
uint32_t copyFinishRowCnt = 0;
LocalTensor<KV_T> kTensor;
uint32_t copyRowCnt = 0;
while (copyFinishRowCnt < nL1Size) {
// 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位
copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize;
if (copyFinishRowCnt + copyRowCnt > nL1Size) {
copyRowCnt = nL1Size - copyFinishRowCnt;
}
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
// 256、32等待7buf命名更改
startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
PAShape shape;
shape.blockSize = constInfo.paOriBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.oriKvStride;
shape.actHeadDim = 256;
shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = nL1SizeAlign;
kTensor = bL1Tensor[copyFinishRowCnt * 16];
DataCopyPA<KV_T>(kTensor, oriKvGm, oriBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size; // 行数
nd2nzPara.dValue = D_SPLIT_SIZE; // 256
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize * seqStride;
uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE;
DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + nL1 * N_SPLIT_SIZE;
if (kL1 == 0) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor, oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim +
nL1 * N_SPLIT_SIZE * constInfo.headDim], nd2nzPara);
} else {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor,
oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + (constInfo.headDim >> 1) +
nL1 * N_SPLIT_SIZE * constInfo.headDim],
nd2nzPara);
}
}
} else {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t curS2Offset = info.relativeS2Idx * constInfo.s2BaseSize + nL1 * N_SPLIT_SIZE;
while (copyFinishRowCnt < nL1Size) {
// 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位
copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize;
if (copyFinishRowCnt + copyRowCnt > nL1Size) {
copyRowCnt = nL1Size - copyFinishRowCnt;
}
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
// 256、32等待7buf命名更改
startPos.dIdx = kL1 * 256; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
PAShape shape;
shape.blockSize = constInfo.paCmpBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.cmpKvStride;
shape.actHeadDim = 256;
shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = nL1SizeAlign;
kTensor = bL1Tensor[copyFinishRowCnt * 16];
DataCopyPA<KV_T>(kTensor, cmpKvGm, cmpBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size; // 行数
nd2nzPara.dValue = D_SPLIT_SIZE; // 256
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize / constInfo.cmpRatio * seqStride;
uint32_t curS2 = info.relativeS2Idx * constInfo.s2BaseSize + nL1 * N_SPLIT_SIZE;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE;
DataCopy(bL1Tensor, cmpKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.relativeS2Idx * constInfo.s2BaseSize + nL1 * N_SPLIT_SIZE;
if (kL1 == 0) {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor, cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim], nd2nzPara);
} else {
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = nL1Size;
nd2nzPara.dValue = constInfo.headDim >> 1;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = nL1SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor,
cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim + (constInfo.headDim >> 1)],
nd2nzPara);
}
}
}
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
mL1Size = M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(M_SPLIT_SIZE, 16U);
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
uint32_t aL1PaddingSize = 0; // 用于使左矩阵对齐到尾部, 以保证两块32K内存连续
if (mL1 == (mL1Loops - 1)) {
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(mL1Size, 16U);
aL1PaddingSize = (M_SPLIT_SIZE - mL1SizeAlign) * 256;
}
uint32_t mIdx = qpL1BufIter + mL1;
ka = GetQPL1RealIdx(mIdx, kL1);
LocalTensor<Q_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + (1 - kL1) * aL1PaddingSize];
if (nL1 == 0) {
if (kL1 == 0) {
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka + 1]);
CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256, 0);
} else {
LocalTensor<Q_T> qTmpTensor = aL1Tensor;
CopyInMm1AToL1(qTmpTensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 256,
256);
}
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
}
// 使用unitflag同步
LocalTensor cL0Tensor =
cL0TensorPingPong[(cL0BufIter % 2) *
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
LocalTensor<KV_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))];
LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size);
LocalTensor<KV_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))];
LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign);
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
MmadParams mmadParams;
mmadParams.m = mL1SizeAlign;
mmadParams.n = nL1SizeAlign;
mmadParams.k = kL0Size;
mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0);
mmadParams.cmatrixSource = false;
mmadParams.unitFlag =
(kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
PipeBarrier<PIPE_M>();
}
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
abL0BufIter++;
}
if (nL1 == (nL1Loops - 1)) {
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
}
if (kL1 == 1) { // 最后一轮kL1循环
FixpipeParamsV220 fixParams;
fixParams.nSize = nL1SizeAlign;
fixParams.mSize = mL1SizeAlign;
fixParams.srcStride = mL1SizeAlign;
// 改成nSizeAlign
fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔
fixParams.unitFlag = 0b11;
fixParams.ndNum = 1; // 输出ND
Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE +
(mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) *
info.actualSingleProcessSInnerSizeAlign],
cL0Tensor, fixParams);
}
if (mL1Loops == 2) {
cL0BufIter++;
}
}
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
}
if (mL1Loops == 1) {
cL0BufIter++;
}
}
qpL1BufIter += mL1Loops;
}
template <typename SAST>
__aicore__ inline void SWACubeBlock<SAST>::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo)
{
uint32_t mSize = mSplitInfo.nBufferDealM;
uint32_t mSizeAlign = (mSize + 16 - 1) / 16;
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐
uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小
uint32_t nSize = BlockAlign<KV_T>(constInfo.headDim);
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐
uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小
uint32_t kSize = info.actualSingleProcessSInnerSize;
uint32_t kL1Size = 256;
uint32_t kL1SizeAlign = SASAlign(kL1Size, 16U);
uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size;
uint32_t kL0Size = 128;
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
uint32_t kL0SizeAlign = kL0Size;
LocalTensor<KV_T> bL1Tensor;
LocalTensor<KV_T> subvTensor;
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
uint32_t ka = 0, kb = 0;
uint32_t mBaseIdx = qpL1BufIter;
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1
if (nL1 == (nL1Loops - 1)) {
// 尾块
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
nL1SizeAlign = SASAlign(nL1Size, 16U);
}
// k l1写成一个循环, 和mm1保持一致
kL1Size = 256;
kL1SizeAlign = SASAlign(kL1Size, 16U);
uint32_t copyRowCnt = 0;
for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作
if (k1 == (kL1Loops - 1)) {
// 尾块
kL1Size = kSize - (kL1Loops - 1) * K_L1_SPLIT_SIZE;
kL1SizeAlign = SASAlign(kL1Size, 16U);
}
kvL1BufIter++;
uint32_t kb = kvL1BufIter % 3;
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
uint32_t kOffset = k1 * kL0Loops;
kL0Size = 128;
// 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误
kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
kL0SizeAlign = kL0Size;
for (uint32_t kL1 = kOffset; kL1 < kL0Loops + kOffset; kL1++) { // 128 循环搬pa
if (kL1 == kOffset + kL0Loops - 1) {
// 尾块
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
kL0SizeAlign = SASAlign(kL0Size, 16U);
}
uint32_t copyFinishRowCnt = 0;
if (info.isOri) {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE;
while (copyFinishRowCnt < kL0Size) {
copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize;
if (copyFinishRowCnt + copyRowCnt > kL0Size) {
copyRowCnt = kL0Size - copyFinishRowCnt;
}
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
startPos.dIdx = nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
PAShape shape;
shape.blockSize = constInfo.paOriBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.oriKvStride;
shape.actHeadDim = nL1Size;
shape.maxblockNumPerBatch = constInfo.oriMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = kL0SizeAlign;
subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + copyFinishRowCnt * 16];
DataCopyPA<KV_T>(subvTensor, oriKvGm, oriBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE];
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = nL1Size; // 256
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize * seqStride;
uint32_t curS2 = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE;
DataCopy(subvTensor, oriKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE;
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE],
oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim +
nL1 * N_SPLIT_SIZE], nd2nzPara);
}
} else {
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
uint32_t curS2Offset = info.relativeS2Idx * constInfo.s2BaseSize + K_L0_SPLIT_SIZE * kL1;
while (copyFinishRowCnt < kL0Size) {
copyRowCnt = constInfo.paCmpBlockSize - curS2Offset % constInfo.paCmpBlockSize;
if (copyFinishRowCnt + copyRowCnt > kL0Size) {
copyRowCnt = kL0Size - copyFinishRowCnt;
}
Position startPos;
startPos.bIdx = info.bIdx;
startPos.n2Idx = info.n2Idx;
startPos.s2Idx = curS2Offset;
// 256、32等待7buf命名更改
startPos.dIdx = nL1 * N_SPLIT_SIZE; // mm1 右矩阵 bn2s2d, d为k轴不切; mm2 右矩阵, s2为k轴, d轴切分
PAShape shape;
shape.blockSize = constInfo.paCmpBlockSize;
shape.headNum = constInfo.kvHeadNum;
shape.headDim = constInfo.headDim;
shape.kvStride = constInfo.cmpKvStride;
shape.actHeadDim = nL1Size;
shape.maxblockNumPerBatch = constInfo.cmpMaxBlockNumPerBatch;
shape.copyRowNum = copyRowCnt;
shape.copyRowNumAlign = kL0SizeAlign;
subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE + copyFinishRowCnt * 16];
DataCopyPA<KV_T>(subvTensor, cmpKvGm, cmpBlockTableGm, shape, startPos);
// 更新循环变量
copyFinishRowCnt += copyRowCnt;
curS2Offset += copyRowCnt;
}
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::BSND) {
subvTensor = bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE];
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = nL1Size; // 256
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
uint32_t headStride = constInfo.headDim;
uint32_t seqStride = constInfo.kvHeadNum * constInfo.headDim;
uint32_t batchStride = constInfo.kvSeqSize / constInfo.cmpRatio * seqStride;
uint32_t curS2 = info.relativeS2Idx * constInfo.s2BaseSize + K_L0_SPLIT_SIZE * kL1;
uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE;
DataCopy(subvTensor, cmpKvGm[offset], nd2nzPara);
} else if constexpr (KV_LAYOUT_T == SAS_LAYOUT::TND) {
uint32_t curS2Offset = info.relativeS2Idx * constInfo.s2BaseSize + K_L0_SPLIT_SIZE * kL1;
Nd2NzParams nd2nzPara;
nd2nzPara.ndNum = 1;
nd2nzPara.nValue = kL0Size; // 行数
nd2nzPara.dValue = N_SPLIT_SIZE; // constInfo.headDim;
nd2nzPara.srcDValue = constInfo.headDim;
nd2nzPara.dstNzC0Stride = kL0SizeAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE],
cmpKvGm[info.tensorCmpBOffset + curS2Offset * constInfo.headDim +
nL1 * N_SPLIT_SIZE], nd2nzPara);
}
}
}
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
mL1SizeAlign = M_SPLIT_SIZE;
mL1Size = M_SPLIT_SIZE; // m的实际大小
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
if (mL1 == (mL1Loops - 1)) {
// 尾块
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
mL1SizeAlign = SASAlign(mL1Size, 16U);
}
uint32_t mIdx = mBaseIdx + mL1;
ka = GetQPL1RealIdx(mIdx, k1);
LocalTensor<KV_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
if (nL1 == 0) {
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, kL1Size,
256 * k1);
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
}
LocalTensor cL0Tensor =
cL0TensorPingPong[(cL0BufIter % 2) *
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
uint32_t baseK = 128;
uint32_t baseN = 128;
kL0Size = 128;
kL0SizeAlign = kL0Size;
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
if (kL0 + 1 == kL0Loops) {
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
kL0SizeAlign = SASAlign(kL0Size, 16U);
}
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
LocalTensor<KV_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE / sizeof(KV_T))];
LoadData3DParamsV2<KV_T> loadData3DParamsForB;
loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height
loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W
loadData3DParamsForB.padList[0] = 0;
loadData3DParamsForB.padList[1] = 0;
loadData3DParamsForB.padList[2] = 0;
loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果
loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度
loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度
loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点
loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点
loadData3DParamsForB.strideW = 1;
loadData3DParamsForB.strideH = 1;
loadData3DParamsForB.filterW = 1;
loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
loadData3DParamsForB.filterH = 1;
loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数
loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数
loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能
loadData3DParamsForB.fMatrixCtrl =
0; // 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1
loadData3DParamsForB.channelSize =
nL1SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
LoadData<KV_T, LOAD3DV2_CONFIG>(bL0Tensor, bL1Tensor[kL0 * baseK * baseN], loadData3DParamsForB);
LocalTensor<KV_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE / sizeof(KV_T))];
LoadData3DParamsV2<KV_T> loadData3DParamsForA;
loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height
loadData3DParamsForA.l1W = 16; // 源操作数weight
loadData3DParamsForA.padList[0] = 0;
loadData3DParamsForA.padList[1] = 0;
loadData3DParamsForA.padList[2] = 0;
loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果
loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度
loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度
loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点
loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点
loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长
loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长
loadData3DParamsForA.filterW = 1; // 卷积核width
loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
loadData3DParamsForA.filterH = 1; // 卷积核height
loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数
loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数
loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置
loadData3DParamsForA.fMatrixCtrl = 0;
loadData3DParamsForA.channelSize =
kL0SizeAlign; // 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
LoadData<KV_T, LOAD3DV2_CONFIG>(aL0Tensor, aL1Tensor[kL0 * baseK * mL1SizeAlign],
loadData3DParamsForA);
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
MmadParams mmadParams;
mmadParams.m = mL1SizeAlign;
mmadParams.n = nL1SizeAlign;
mmadParams.k = kL0Size;
mmadParams.cmatrixInitVal = (kL0 == 0 && k1 == 0);
mmadParams.cmatrixSource = false;
mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (kL0 == (kL0Loops - 1))) ? 0b11 : 0b10;
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
PipeBarrier<PIPE_M>();
}
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
abL0BufIter++;
}
if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算?
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
}
if (k1 == (kL1Loops - 1)) {
// ND
FixpipeParamsV220 fixParams;
fixParams.nSize = nL1SizeAlign;
fixParams.mSize = mL1SizeAlign;
fixParams.srcStride = mL1SizeAlign;
fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔
fixParams.ndNum = 1; // 输出ND
fixParams.unitFlag = 0b11;
uint64_t mm2Offset = (mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) * nSize + nL1 * N_SPLIT_SIZE;
Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.bmm2ResUbSize + mm2Offset],
cL0Tensor, fixParams);
}
if (mL1Loops == 2) {
cL0BufIter++;
}
}
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
}
// cL0BufIter已经不在使用
if (mL1Loops == 1) {
cL0BufIter++;
}
}
qpL1BufIter += mL1Loops;
}
} // namespace SASKernel
#endif

View File

@@ -0,0 +1,850 @@
/**
 * 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 sparse_attn_sharedkv_swa_block_vector.h
* \brief
*/
#ifndef SPARSE_ATTN_SHAREDKV_SWA_BLOCK_VECTOR_H
#define SPARSE_ATTN_SHAREDKV_SWA_BLOCK_VECTOR_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "kernel_tiling/kernel_tiling.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#include "../sparse_attn_sharedkv_common.h"
namespace SASKernel {
using AscendC::CrossCoreSetFlag;
using AscendC::CrossCoreWaitFlag;
template <typename SAST>
class SWAVectorBlock {
public:
// 中间计算数据类型为float,高精度模式
using T = float;
using KV_T = typename SAST::kvType;
using OUT_T = typename SAST::outputType;
using UPDATE_T = T;
using SINKS_T = T;
using MM1_OUT_T = float;
using MM2_OUT_T = float;
__aicore__ inline SWAVectorBlock(){};
__aicore__ inline void ProcessVec1L(const RunInfo &info);
__aicore__ inline void ProcessVec2L(const RunInfo &info);
__aicore__ inline void InitBuffers(TPipe *pipe);
__aicore__ inline void InitParams(const struct ConstInfo &constInfo,
const SparseAttnSharedkvTilingData *__restrict tilingData);
__aicore__ inline void InitVec1GlobalTensor(GlobalTensor<MM1_OUT_T> mm1ResGm, GlobalTensor<KV_T> vec1ResGm,
GlobalTensor<int32_t> actualSeqLengthsQGm,
GlobalTensor<int32_t> actualSeqLengthsKVGm, GlobalTensor<T> sinksGm, GlobalTensor<T> softmaxLseGm);
__aicore__ inline void InitVec2GlobalTensor(GlobalTensor<T> accumOutGm, GlobalTensor<UPDATE_T> vec2ResGm,
GlobalTensor<MM2_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm);
__aicore__ inline void AllocEventID();
__aicore__ inline void FreeEventID();
__aicore__ inline void CopySinksIn();
__aicore__ inline void SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount);
__aicore__ inline void InitSoftmaxDefaultBuffer();
// ================================Base Vector==========================================
__aicore__ inline void RowDivs(LocalTensor<float> dstUb, LocalTensor<float> src0Ub, LocalTensor<float> src1Ub,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount);
__aicore__ inline void RowMuls(LocalTensor<T> dstUb, LocalTensor<T> src0Ub, LocalTensor<T> src1Ub,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount);
// ================================Vector1==========================================
__aicore__ inline void ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo);
__aicore__ inline void DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow,
uint32_t dealRowCount, uint32_t columnCount, uint32_t loopId);
__aicore__ inline void SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo,
LocalTensor<T> &mmResUb, LocalTensor<uint8_t> &softmaxTmpUb,
uint32_t startRow, uint32_t dealRowCount, uint32_t columnCount,
uint32_t actualColumnCount);
__aicore__ inline void ElewiseCompute(const RunInfo &info, const LocalTensor<T> &mmResUb, uint32_t dealRowCount,
uint32_t columnCount);
__aicore__ inline void ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo);
// ================================Vecotr2==========================================
__aicore__ inline void ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo);
__aicore__ inline void DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t startRow,
uint32_t dealRowCount, uint32_t columnCount,
uint32_t actualColumnCount);
__aicore__ inline void ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo, uint32_t mStartRow,
uint32_t mDealSize);
__aicore__ inline void Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor<OUT_T> &attenOutUb, uint32_t wsMStart,
uint32_t dealRowCount, uint32_t columnCount,
uint32_t actualColumnCount);
__aicore__ inline void Bmm2ResCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb, uint32_t wsMStart,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount);
__aicore__ inline void Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb, uint32_t wsMStart,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount);
__aicore__ inline void Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb, uint32_t wsMStart,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount);
__aicore__ inline uint64_t CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx);
// BLOCK和REPEAT的字节数
static constexpr uint64_t BYTE_BLOCK = 32UL;
static constexpr uint32_t REPEAT_BLOCK_BYTE = 256U;
// BLOCK和REPEAT的FP32元素数
static constexpr uint32_t FP32_BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(float);
static constexpr uint32_t FP32_REPEAT_ELEMENT_NUM = REPEAT_BLOCK_BYTE / sizeof(float);
// repeat stride不能超过256
static constexpr uint32_t REPEATE_STRIDE_UP_BOUND = 256;
private:
static constexpr bool PAGE_ATTENTION = SAST::pageAttention;
static constexpr bool FLASH_DECODE = SAST::flashDecode;
static constexpr SAS_LAYOUT LAYOUT_T = SAST::layout;
static constexpr SAS_LAYOUT KV_LAYOUT_T = SAST::kvLayout;
static constexpr uint64_t SYNC_INPUT_BUF1_FLAG = 2;
static constexpr uint64_t SYNC_INPUT_BUF1_PONG_FLAG = 3;
static constexpr uint64_t SYNC_INPUT_BUF2_FLAG = 4;
static constexpr uint64_t SYNC_INPUT_BUF2_PONG_FLAG = 5;
static constexpr uint64_t SYNC_OUTPUT_BUF1_FLAG = 4;
static constexpr uint64_t SYNC_OUTPUT_BUF2_FLAG = 5;
static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6;
static constexpr uint32_t INPUT1_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_32K;
static constexpr uint32_t SOFTMAX_TMP_BUFFER_OFFSET = ConstInfo::BUFFER_SIZE_BYTE_1K;
static constexpr uint32_t BASE_BLOCK_MAX_ELEMENT_NUM = ConstInfo::BUFFER_SIZE_BYTE_32K / sizeof(T); // 32768/4=8096
static constexpr uint32_t BLOCK_ELEMENT_NUM = BYTE_BLOCK / sizeof(T); // 32/4=8
static constexpr uint32_t MAX_N1_SIZE = 128U;
static constexpr T SOFTMAX_MIN_NUM = -2e38;
static constexpr SINKS_T R0 = 1.0f;
const SparseAttnSharedkvTilingData *__restrict tilingData;
uint32_t pingpongFlag = 0U;
ConstInfo constInfo = {};
GlobalTensor<MM1_OUT_T> mm1ResGm;
GlobalTensor<KV_T> vec1ResGm;
GlobalTensor<T> softmaxMaxGm;
GlobalTensor<T> softmaxSumGm;
GlobalTensor<T> sinksGm;
GlobalTensor<int32_t> actualSeqLengthsQGm;
GlobalTensor<int32_t> actualSeqLengthsKVGm;
GlobalTensor<UPDATE_T> vec2ResGm;
GlobalTensor<MM2_OUT_T> mm2ResGm;
GlobalTensor<T> accumOutGm;
GlobalTensor<OUT_T> attentionOutGm;
GlobalTensor<int32_t> blkTableGm_;
GlobalTensor<KV_T> keyGm_;
GlobalTensor<int32_t> kvValidSizeGm_;
GlobalTensor<KV_T> oriKvGm_;
GlobalTensor<KV_T> cmpKvGm_;
GlobalTensor<int32_t> oriBlockTableGm_;
GlobalTensor<int32_t> cmpBlockTableGm_;
GlobalTensor<T> softmaxLseGm;
// ================================Local Buffer区====================================
TBuf<> inputBuff1; // 32K
TBuf<> inputBuff2; // 16K
TBuf<> outputBuff1; // 32K
TBuf<> outputBuff2; // 4K
TBuf<> tmpBuff1; // 32K
TBuf<> v0ValidSizeBuff; // 8K
TBuf<> sinksBuff; // 1K
TBuf<> sinksBrcbBuff; // 12K
TBuf<> softmaxMaxBuff; // PRE_LOAD_NUM * 2K
TBuf<> softmaxExpBuff; // PRE_LOAD_NUM * 2K
TBuf<> softmaxSumBuff; // PRE_LOAD_NUM * 2K
TBuf<> softmaxMaxDefaultBuff; // 2K
TBuf<> softmaxSumDefaultBuff; // 2K
LocalTensor<T> softmaxMaxDefaultUb;
LocalTensor<T> softmaxSumDefaultUb;
LocalTensor<T> softmaxMaxUb;
LocalTensor<T> softmaxSumUb;
LocalTensor<T> softmaxExpUb;
LocalTensor<SINKS_T> sinksUb;
LocalTensor<SINKS_T> sinksBrcbUb;
};
// ============================== init ==============================================
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::InitBuffers(TPipe *pipe)
{
pipe->InitBuffer(inputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K * 2); // 2:pingpong
pipe->InitBuffer(inputBuff2, ConstInfo::BUFFER_SIZE_BYTE_8K * 2); // 2:pingpong
pipe->InitBuffer(outputBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K);
pipe->InitBuffer(outputBuff2, ConstInfo::BUFFER_SIZE_BYTE_4K);
pipe->InitBuffer(tmpBuff1, ConstInfo::BUFFER_SIZE_BYTE_32K);
pipe->InitBuffer(v0ValidSizeBuff, ConstInfo::BUFFER_SIZE_BYTE_8K);
// M_MAX = 512/2vector = 256, 256 * sizeof(T) * N_Buffer
pipe->InitBuffer(softmaxMaxBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum);
pipe->InitBuffer(softmaxExpBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum);
pipe->InitBuffer(softmaxSumBuff, ConstInfo::BUFFER_SIZE_BYTE_1K * constInfo.preLoadNum);
pipe->InitBuffer(softmaxMaxDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K);
pipe->InitBuffer(softmaxSumDefaultBuff, ConstInfo::BUFFER_SIZE_BYTE_1K);
pipe->InitBuffer(sinksBuff, MAX_N1_SIZE * sizeof(SINKS_T));
// 分配256+N1大小内存,其中256是m轴VEC最大切块
pipe->InitBuffer(sinksBrcbBuff, MAX_N1_SIZE * sizeof(SINKS_T) * BLOCK_ELEMENT_NUM * 3U);
softmaxMaxUb = softmaxMaxBuff.Get<T>();
softmaxSumUb = softmaxSumBuff.Get<T>();
softmaxExpUb = softmaxExpBuff.Get<T>();
softmaxMaxDefaultUb = softmaxMaxDefaultBuff.Get<T>();
softmaxSumDefaultUb = softmaxSumDefaultBuff.Get<T>();
sinksUb = sinksBuff.Get<SINKS_T>();
sinksBrcbUb = sinksBrcbBuff.Get<SINKS_T>();
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::InitParams(const struct ConstInfo &constInfo,
const SparseAttnSharedkvTilingData *__restrict tilingData)
{
this->constInfo = constInfo;
this->tilingData = tilingData;
}
template <typename SAST>
__aicore__ inline void
SWAVectorBlock<SAST>::InitVec1GlobalTensor(GlobalTensor<MM1_OUT_T> mm1ResGm, GlobalTensor<KV_T> vec1ResGm,
GlobalTensor<int32_t> actualSeqLengthsQGm,
GlobalTensor<int32_t> actualSeqLengthsKVGm,
GlobalTensor<SINKS_T> sinksGm, GlobalTensor<T> softmaxLseGm)
{
this->mm1ResGm = mm1ResGm;
this->vec1ResGm = vec1ResGm;
this->actualSeqLengthsQGm = actualSeqLengthsQGm;
this->actualSeqLengthsKVGm = actualSeqLengthsKVGm;
this->sinksGm = sinksGm;
this->softmaxLseGm = softmaxLseGm;
}
template <typename SAST>
__aicore__ inline void
SWAVectorBlock<SAST>::InitVec2GlobalTensor(GlobalTensor<T> accumOutGm, GlobalTensor<UPDATE_T> vec2ResGm,
GlobalTensor<MM2_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm)
{
this->accumOutGm = accumOutGm;
this->vec2ResGm = vec2ResGm;
this->mm2ResGm = mm2ResGm;
this->attentionOutGm = attentionOutGm;
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::AllocEventID()
{
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG);
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_PONG_FLAG);
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_FLAG);
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_PONG_FLAG);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF2_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::FreeEventID()
{
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_PONG_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_PONG_FLAG);
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF2_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::CopySinksIn()
{
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = 1U;
dataCopyParams.blockLen = constInfo.qHeadNum * sizeof(T);
dataCopyParams.srcStride = 0U;
dataCopyParams.dstStride = 0U;
DataCopyPadExtParams<T> padParams;
DataCopyPad(sinksUb, sinksGm, dataCopyParams, padParams);
SetFlag<AscendC::HardEvent::MTE2_V>(SYNC_SINKS_BUF_FLAG);
WaitFlag<AscendC::HardEvent::MTE2_V>(SYNC_SINKS_BUF_FLAG);
uint32_t repeatTimes = (constInfo.qHeadNum + BLOCK_ELEMENT_NUM - 1U) / BLOCK_ELEMENT_NUM; // 每次处理 8 datablocks
Brcb(sinksBrcbUb, sinksUb, repeatTimes, {1, BLOCK_ELEMENT_NUM});
PipeBarrier<PIPE_V>();
DataCopyParams repeatParams;
repeatParams.blockCount = 1; // 搬到有一个块超过单个vec核减分核M轴大小即可,核间切分每个vec256
repeatParams.blockLen = constInfo.qHeadNum;
repeatParams.srcStride = 0U;
repeatParams.dstStride = 0U;
for (uint32_t i = 1U; i <= 256U / constInfo.qHeadNum; i++) {
DataCopy(sinksBrcbUb[constInfo.qHeadNum * BLOCK_ELEMENT_NUM * i], sinksBrcbUb, repeatParams);
}
PipeBarrier<PIPE_V>();
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::SliceAndContactSinksValue(uint32_t nIdx, uint32_t dealRowCount)
{
// 由于WholeReduceMax接口中repeatTimes支持范围(0,255),因此需要分多次调用WholeReduceMax,这里就使用每次repeatTime=128
uint32_t repeatTimesOnce = 128;
uint32_t loopTimes = (dealRowCount + repeatTimesOnce - 1) / repeatTimesOnce;
uint32_t repeatTimes = repeatTimesOnce;
for (uint32_t loop = 0; loop < loopTimes; ++loop) {
if (loop == loopTimes - 1) {
repeatTimes = dealRowCount - loop * repeatTimesOnce;
}
WholeReduceMax(softmaxMaxDefaultUb[loop * repeatTimesOnce],
sinksBrcbUb[(nIdx + loop * repeatTimesOnce) * BLOCK_ELEMENT_NUM],
BLOCK_ELEMENT_NUM * BLOCK_ELEMENT_NUM, repeatTimes, 1, 0, 1, ReduceOrder::ORDER_ONLY_VALUE);
PipeBarrier<PIPE_V>();
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::InitSoftmaxDefaultBuffer()
{
CopySinksIn();
Duplicate(softmaxMaxDefaultUb, SOFTMAX_MIN_NUM, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T));
Duplicate(softmaxSumDefaultUb, R0, SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T));
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ElewiseCompute(const RunInfo &info, const LocalTensor<T> &mmResUb,
uint32_t dealRowCount, uint32_t columnCount)
{
Muls(mmResUb, mmResUb, static_cast<T>(tilingData->baseParams.softmaxScale), dealRowCount * columnCount);
}
template <typename SAST>
__aicore__ inline void
SWAVectorBlock<SAST>::SoftmaxFlashV2Compute(const RunInfo &info, const MSplitInfo &mSplitInfo, LocalTensor<T> &mmResUb,
LocalTensor<uint8_t> &softmaxTmpUb, uint32_t startRow,
uint32_t dealRowCount, uint32_t columnCount, uint32_t actualColumnCount)
{
LocalTensor<T> inSumTensor;
LocalTensor<T> inMaxTensor;
uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow;
uint32_t outIdx = info.loop % (constInfo.preLoadNum);
uint32_t softmaxOutOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset;
if (info.isFirstSInnerLoop) {
inMaxTensor = softmaxMaxDefaultUb[startRow];
inSumTensor = softmaxSumDefaultUb;
} else {
uint32_t inIdx = (info.loop - 1) % (constInfo.preLoadNum);
inMaxTensor = softmaxMaxUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset];
inSumTensor = softmaxSumUb[inIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset];
}
if (actualColumnCount != 0) {
SoftMaxShapeInfo srcShape{dealRowCount, columnCount, dealRowCount, actualColumnCount};
SoftMaxTiling newTiling =
SoftMaxFlashV2TilingFunc(srcShape, sizeof(T), sizeof(T), softmaxTmpUb.GetSize(), true, false);
SoftmaxFlashV2<T, true, true, false, false, SAS_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC>(
mmResUb, softmaxSumUb[softmaxOutOffset], softmaxMaxUb[softmaxOutOffset], mmResUb,
softmaxExpUb[softmaxOutOffset], inSumTensor, inMaxTensor, softmaxTmpUb, newTiling, srcShape);
} else {
uint32_t dealRowCountAlign = SASAlign(dealRowCount, FP32_BLOCK_ELEMENT_NUM);
DataCopy(softmaxSumUb[softmaxOutOffset], inSumTensor, dealRowCountAlign);
PipeBarrier<PIPE_V>();
DataCopy(softmaxMaxUb[softmaxOutOffset], inMaxTensor, dealRowCountAlign);
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessLse(const RunInfo &info, const MSplitInfo &mSplitInfo)
{
if (mSplitInfo.vecDealM == 0) {
return;
}
uint64_t lseOffset;
if (constInfo.outputLayout == SAS_LAYOUT::TND) {
uint32_t tBase = actualSeqLengthsQGm.GetValue(info.bIdx);
lseOffset = (tBase + info.s1Idx) * constInfo.gSize + // T轴、s1轴偏移
info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移
} else if (constInfo.outputLayout == SAS_LAYOUT::BSND) {
lseOffset = info.bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移
info.n2IdxReal * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移
info.s1Idx * constInfo.gSize; // S1轴偏移
}
lseOffset = lseOffset + mSplitInfo.nBufferStartM + mSplitInfo.vecStartM;
uint32_t baseOffset = mSplitInfo.nBufferStartM / 2;
uint32_t outIdx = info.loop % (constInfo.preLoadNum);
uint32_t softmaxOffset = outIdx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset;
auto sumTensor = softmaxSumUb[softmaxOffset];
auto maxTensor = softmaxMaxUb[softmaxOffset];
auto outLSETensor = outputBuff2.Get<T>();
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = 1;
dataCopyParams.blockLen = mSplitInfo.vecDealM * sizeof(T);
dataCopyParams.srcStride = 0;
dataCopyParams.dstStride = 0;
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF2_FLAG);
PipeBarrier<PIPE_V>();
Log(outLSETensor, sumTensor, mSplitInfo.vecDealM);
PipeBarrier<PIPE_V>();
Add(outLSETensor, outLSETensor, maxTensor, mSplitInfo.vecDealM);
SetFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF2_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF2_FLAG);
DataCopyPad(softmaxLseGm[lseOffset], outLSETensor, dataCopyParams);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF2_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::DealBmm1ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo,
uint32_t startRow, uint32_t dealRowCount,
uint32_t columnCount, uint32_t loopId)
{
uint32_t computeSize = dealRowCount * columnCount;
uint64_t inOutGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize +
(mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow) * columnCount;
LocalTensor<MM1_OUT_T> mmResUb = inputBuff1.Get<MM1_OUT_T>();
mmResUb = mmResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM1_OUT_T)];
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG + pingpongFlag);
DataCopy(mmResUb, mm1ResGm[inOutGmOffset], computeSize);
SetFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF1_FLAG);
ElewiseCompute(info, mmResUb, dealRowCount, columnCount);
PipeBarrier<PIPE_V>();
LocalTensor<T> tmpAFloorUb = tmpBuff1.Get<T>();
LocalTensor<uint8_t> softmaxTmpUb = tmpAFloorUb.template ReinterpretCast<uint8_t>();
SoftmaxFlashV2Compute(info, mSplitInfo, mmResUb, softmaxTmpUb, startRow, dealRowCount, columnCount,
info.actualSingleProcessSInnerSize);
PipeBarrier<PIPE_V>();
LocalTensor<KV_T> tmpMMResCastTensor = outputBuff1.Get<KV_T>();
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
Cast(tmpMMResCastTensor, mmResUb, AscendC::RoundMode::CAST_ROUND, computeSize);
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG + pingpongFlag);
SetFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
DataCopy(vec1ResGm[inOutGmOffset], tmpMMResCastTensor, computeSize);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessVec1SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo)
{
if (mSplitInfo.vecDealM == 0) {
return;
}
uint32_t mSplitSize = info.actualSingleProcessSInnerSize == 0 ?
16 :
BASE_BLOCK_MAX_ELEMENT_NUM / info.actualSingleProcessSInnerSizeAlign;
// 1. 向下8对齐是因为UB操作至少32B
// 2. info.actualSingleProcessSInnerSizeAlign最大512, mSplitSize可以确保最小为16
mSplitSize = mSplitSize / 8 * 8;
if (mSplitSize > mSplitInfo.vecDealM) {
mSplitSize = mSplitInfo.vecDealM;
}
uint32_t loopCount = (mSplitInfo.vecDealM + mSplitSize - 1) / mSplitSize;
uint32_t tailSplitSize = mSplitInfo.vecDealM - (loopCount - 1) * mSplitSize;
SliceAndContactSinksValue((mSplitInfo.nBufferStartM + mSplitInfo.vecStartM) % constInfo.qHeadNum,
mSplitInfo.vecDealM);
for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) {
if (i == (loopCount - 1)) {
dealSize = tailSplitSize;
}
DealBmm1ResBaseBlock(info, mSplitInfo, i * mSplitSize, dealSize, info.actualSingleProcessSInnerSizeAlign, i);
pingpongFlag ^= 1; // pingpong 0 1切换
}
}
// =======================vec1=============================
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessVec1L(const RunInfo &info)
{
uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize;
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferIdx = i;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM :
(((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16);
mSplitInfo.vecStartM = 0;
if (GetBlockIdx() % 2 == 1) {
mSplitInfo.vecStartM = mSplitInfo.vecDealM;
mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM;
}
CrossCoreWaitFlag(constInfo.syncC1V1);
// vec1 compute
ProcessVec1SingleBuf(info, mSplitInfo);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_MTE3>(constInfo.syncV1C2);
// move lse for flash decode or FA
if (constInfo.returnSoftmaxLse && info.s2Idx == info.curSInnerLoopTimes - 1) {
ProcessLse(info, mSplitInfo);
}
}
}
// =======================vec2=============================
template <typename SAST>
__aicore__ inline uint64_t SWAVectorBlock<SAST>::CalcAccumOffset(uint32_t bN2Idx, uint32_t gS1Idx)
{
return 0;
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessVec2SingleBuf(const RunInfo &info, const MSplitInfo &mSplitInfo)
{
if (mSplitInfo.vecDealM == 0) {
return;
}
ProcessVec2Inner(info, mSplitInfo, 0, mSplitInfo.vecDealM);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessVec2L(const RunInfo &info)
{
uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize;
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferIdx = i;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
mSplitInfo.vecDealM = (mSplitInfo.nBufferDealM <= 16) ? mSplitInfo.nBufferDealM :
(((mSplitInfo.nBufferDealM + 15) / 16 + 1) / 2 * 16);
mSplitInfo.vecStartM = 0;
if (GetBlockIdx() % 2 == 1) {
mSplitInfo.vecStartM = mSplitInfo.vecDealM;
mSplitInfo.vecDealM = mSplitInfo.nBufferDealM - mSplitInfo.vecDealM;
}
CrossCoreWaitFlag(constInfo.syncC2V2);
ProcessVec2SingleBuf(info, mSplitInfo);
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::ProcessVec2Inner(const RunInfo &info, const MSplitInfo &mSplitInfo,
uint32_t mStartRow, uint32_t mDealSize)
{
uint32_t mSplitSize = BASE_BLOCK_MAX_ELEMENT_NUM / constInfo.headDim;
if (mSplitSize > mDealSize) {
mSplitSize = mDealSize;
}
uint32_t loopCount = (mDealSize + mSplitSize - 1) / mSplitSize;
uint32_t tailSplitSize = mDealSize - (loopCount - 1) * mSplitSize;
for (uint32_t i = 0, dealSize = mSplitSize; i < loopCount; i++) {
if (i == (loopCount - 1)) {
dealSize = tailSplitSize;
}
DealBmm2ResBaseBlock(info, mSplitInfo, i * mSplitSize + mStartRow, dealSize, constInfo.headDim,
constInfo.headDim);
pingpongFlag ^= 1; // pingpong 0 1切换
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::Bmm2FDDataCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb,
uint32_t wsMStart, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
LocalTensor<T> tmp = outputBuff1.Get<T>();
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
DataCopy(tmp, bmm2ResUb, columnCount * dealRowCount);
SetFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
uint64_t accumTmpOutNum = CalcAccumOffset(info.bIdx, info.gS1Idx);
uint64_t offset =
accumTmpOutNum * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // taskoffset
info.tndCoreStartKVSplitPos * constInfo.kvHeadNum * constInfo.mBaseSize * constInfo.headDim + // 份数offset
wsMStart * actualColumnCount; // m轴offset
GlobalTensor<T> dst = accumOutGm[offset];
if (info.actualSingleProcessSInnerSize == 0) {
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = dealRowCount;
dataCopyParams.blockLen = actualColumnCount * sizeof(T);
dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(T));
dataCopyParams.dstStride = 0;
DataCopyPad(dst, tmp, dataCopyParams);
} else {
matmul::InitOutput<T>(dst, dealRowCount * actualColumnCount, ConstInfo::FLOAT_ZERO);
}
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::Bmm2DataCopyOutTrans(const RunInfo &info, LocalTensor<OUT_T> &attenOutUb,
uint32_t wsMStart, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = dealRowCount;
dataCopyParams.blockLen = actualColumnCount * sizeof(OUT_T);
dataCopyParams.srcStride = (columnCount - actualColumnCount) / (BYTE_BLOCK / sizeof(OUT_T));
dataCopyParams.dstStride = 0;
DataCopyPad(attentionOutGm[info.attenOutOffset + wsMStart * actualColumnCount], attenOutUb, dataCopyParams);
return;
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::Bmm2CastAndCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb,
uint32_t wsMStart, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
LocalTensor<OUT_T> tmpBmm2ResCastTensor = outputBuff1.Get<OUT_T>();
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
if constexpr (IsSameType<OUT_T, bfloat16_t>::value) { // bf16 采取四舍六入五成双模式
Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_RINT, dealRowCount * columnCount);
} else {
Cast(tmpBmm2ResCastTensor, bmm2ResUb, AscendC::RoundMode::CAST_ROUND, dealRowCount * columnCount);
}
SetFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
Bmm2DataCopyOutTrans(info, tmpBmm2ResCastTensor, wsMStart, dealRowCount, columnCount, actualColumnCount);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::Bmm2ResCopyOut(const RunInfo &info, LocalTensor<T> &bmm2ResUb,
uint32_t wsMStart, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
if constexpr (FLASH_DECODE) {
if (info.tndIsS2SplitCore) {
Bmm2FDDataCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount);
} else {
Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount);
}
} else {
Bmm2CastAndCopyOut(info, bmm2ResUb, wsMStart, dealRowCount, columnCount, actualColumnCount);
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::DealBmm2ResBaseBlock(const RunInfo &info, const MSplitInfo &mSplitInfo,
uint32_t startRow, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
uint32_t vec2ComputeSize = dealRowCount * columnCount;
uint32_t mStart = mSplitInfo.nBufferStartM + mSplitInfo.vecStartM + startRow;
uint64_t srcGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + mStart * columnCount;
LocalTensor<MM2_OUT_T> tmpBmm2ResUb = inputBuff1.Get<MM2_OUT_T>();
tmpBmm2ResUb = tmpBmm2ResUb[pingpongFlag * INPUT1_BUFFER_OFFSET / sizeof(MM2_OUT_T)];
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG + pingpongFlag);
DataCopy(tmpBmm2ResUb, mm2ResGm[srcGmOffset], vec2ComputeSize);
SetFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF1_FLAG);
LocalTensor<T> bmm2ResUb = tmpBuff1.Get<T>();
bmm2ResUb.SetSize(vec2ComputeSize);
DataCopy(bmm2ResUb, tmpBmm2ResUb, vec2ComputeSize);
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF1_FLAG + pingpongFlag);
uint32_t inOutBaseOffset = mStart * columnCount;
uint32_t baseOffset = mSplitInfo.nBufferStartM / 2 + startRow;
// 除第一个循环外,均需要更新中间计算结果
if (!info.isFirstSInnerLoop) {
event_t eventIdMte2WaitMte3 = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::MTE3_MTE2));
SetFlag<HardEvent::MTE3_MTE2>(eventIdMte2WaitMte3);
WaitFlag<HardEvent::MTE3_MTE2>(eventIdMte2WaitMte3);
LocalTensor<MM2_OUT_T> bmm2ResPreUb = inputBuff2.Get<MM2_OUT_T>();
WaitFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_FLAG);
uint64_t vec2ResGmOffset = ((info.loop - 1) % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset;
DataCopy(bmm2ResPreUb, vec2ResGm[vec2ResGmOffset], vec2ComputeSize);
SetFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF2_FLAG);
WaitFlag<AscendC::HardEvent::MTE2_V>(SYNC_INPUT_BUF2_FLAG);
uint32_t idx = info.loop % (constInfo.preLoadNum);
LocalTensor<T> expUb = v0ValidSizeBuff.Get<T>()[384]; // sumUb用临时内存 16 * 32B = 512B
Brcb(expUb, softmaxExpUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8,
{1, 8});
PipeBarrier<PIPE_V>();
RowMuls(bmm2ResPreUb, bmm2ResPreUb, expUb, dealRowCount, columnCount, actualColumnCount);
AscendC::PipeBarrier<PIPE_V>();
Add(bmm2ResUb, bmm2ResUb, bmm2ResPreUb, vec2ComputeSize);
AscendC::PipeBarrier<PIPE_V>();
SetFlag<AscendC::HardEvent::V_MTE2>(SYNC_INPUT_BUF2_FLAG);
}
// 最后一次输出计算结果,否则将中间结果暂存至workspace
if (info.isLastS2Loop) {
uint32_t idx = info.loop % (constInfo.preLoadNum);
LocalTensor<T> tmpSumUb = v0ValidSizeBuff.Get<T>()[384]; // sumUb用临时内存 16 * 32B = 512B
Brcb(tmpSumUb, softmaxSumUb[idx * SOFTMAX_TMP_BUFFER_OFFSET / sizeof(T) + baseOffset], (dealRowCount + 7) / 8,
{1, 8});
PipeBarrier<PIPE_V>();
RowDivs(bmm2ResUb, bmm2ResUb, tmpSumUb, dealRowCount, columnCount, actualColumnCount);
PipeBarrier<PIPE_V>();
Bmm2ResCopyOut(info, bmm2ResUb, mStart, dealRowCount, columnCount, actualColumnCount);
} else {
LocalTensor<T> outUb = outputBuff1.Get<T>();
WaitFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
DataCopy(outUb, bmm2ResUb, dealRowCount * columnCount);
SetFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
WaitFlag<AscendC::HardEvent::V_MTE3>(SYNC_OUTPUT_BUF1_FLAG);
uint64_t vec2ResGmOffset = (info.loop % constInfo.preLoadNum) * constInfo.bmm2ResUbSize + inOutBaseOffset;
DataCopy(vec2ResGm[vec2ResGmOffset], outUb, vec2ComputeSize);
SetFlag<AscendC::HardEvent::MTE3_V>(SYNC_OUTPUT_BUF1_FLAG);
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::RowDivs(LocalTensor<float> dstUb, LocalTensor<float> src0Ub,
LocalTensor<float> src1Ub, uint32_t dealRowCount,
uint32_t columnCount, uint32_t actualColumnCount)
{
// divs by row, 每行的元素除以相同的元素
// dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] / src1Ub[i, 0 : 7]
// src0Ub:[dealRowCount, columnCount], src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount,
// columnCount]
uint32_t dtypeMask = FP32_REPEAT_ELEMENT_NUM;
uint32_t dLoop = actualColumnCount / dtypeMask;
uint32_t dRemain = actualColumnCount % dtypeMask;
BinaryRepeatParams repeatParamsDiv;
repeatParamsDiv.src0BlkStride = 1;
repeatParamsDiv.src1BlkStride = 0;
repeatParamsDiv.dstBlkStride = 1;
repeatParamsDiv.src0RepStride = columnCount / FP32_BLOCK_ELEMENT_NUM;
repeatParamsDiv.src1RepStride = 1;
repeatParamsDiv.dstRepStride = columnCount / FP32_BLOCK_ELEMENT_NUM;
uint32_t columnRepeatCount = dLoop;
if (columnRepeatCount <= dealRowCount) {
uint32_t offset = 0;
for (uint32_t i = 0; i < dLoop; i++) {
Div(dstUb[offset], src0Ub[offset], src1Ub, dtypeMask, dealRowCount, repeatParamsDiv);
offset += dtypeMask;
}
} else {
BinaryRepeatParams columnRepeatParams;
columnRepeatParams.src0BlkStride = 1;
columnRepeatParams.src1BlkStride = 0;
columnRepeatParams.dstBlkStride = 1;
columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block
columnRepeatParams.src1RepStride = 0;
columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block
uint32_t offset = 0;
for (uint32_t i = 0; i < dealRowCount; i++) {
Div(dstUb[offset], src0Ub[offset], src1Ub[i * FP32_BLOCK_ELEMENT_NUM], dtypeMask, columnRepeatCount,
columnRepeatParams);
offset += columnCount;
}
}
if (dRemain > 0) {
Div(dstUb[dLoop * dtypeMask], src0Ub[dLoop * dtypeMask], src1Ub, dRemain, dealRowCount, repeatParamsDiv);
}
}
template <typename SAST>
__aicore__ inline void SWAVectorBlock<SAST>::RowMuls(LocalTensor<T> dstUb, LocalTensor<T> src0Ub, LocalTensor<T> src1Ub,
uint32_t dealRowCount, uint32_t columnCount,
uint32_t actualColumnCount)
{
// muls by row, 每行的元素乘以相同的元素
// dstUb[i, (j * 8) : (j * 8 + 7)] = src0Ub[i, (j * 8) : (j * 8 + 7)] * src1Ub[i, 0 : 7]
// src0Ub:[dealRowCount, columnCount] src1Ub:[dealRowCount, FP32_BLOCK_ELEMENT_NUM] dstUb:[dealRowCount,
// columnCount]
// dealRowCount is repeat times, must be less 256
uint32_t repeatElementNum = FP32_REPEAT_ELEMENT_NUM;
uint32_t blockElementNum = FP32_BLOCK_ELEMENT_NUM;
if constexpr (std::is_same<T, half>::value) {
// 此限制由于每个repeat至多连续读取256B数据
repeatElementNum = FP32_REPEAT_ELEMENT_NUM * 2; // 256/4 * 2=128
blockElementNum = FP32_BLOCK_ELEMENT_NUM * 2; // 32/4 * 2 = 16
}
// 每次只能连续读取256B的数据进行计算,故每次只能处理256B/sizeof(dType)=
// 列方向分dLoop次,每次处理8列数据
uint32_t dLoop = actualColumnCount / repeatElementNum;
uint32_t dRemain = actualColumnCount % repeatElementNum;
// REPEATE_STRIDE_UP_BOUND=256, 此限制由于src0RepStride数据类型为uint8之多256个datablock间距
if (columnCount < REPEATE_STRIDE_UP_BOUND * blockElementNum) {
BinaryRepeatParams repeatParams;
repeatParams.src0BlkStride = 1;
repeatParams.src1BlkStride = 0;
repeatParams.dstBlkStride = 1;
repeatParams.src0RepStride = columnCount / blockElementNum;
repeatParams.src1RepStride = 1;
repeatParams.dstRepStride = columnCount / blockElementNum;
// 如果以列为repeat所处理的次数小于行处理次数,则以列方式处理。反之则以行进行repeat处理
if (dLoop <= dealRowCount) {
uint32_t offset = 0;
for (uint32_t i = 0; i < dLoop; i++) {
Mul(dstUb[offset], src0Ub[offset], src1Ub, repeatElementNum, dealRowCount, repeatParams);
offset += repeatElementNum;
}
} else {
BinaryRepeatParams columnRepeatParams;
columnRepeatParams.src0BlkStride = 1;
columnRepeatParams.src1BlkStride = 0;
columnRepeatParams.dstBlkStride = 1;
columnRepeatParams.src0RepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block
columnRepeatParams.src1RepStride = 0;
columnRepeatParams.dstRepStride = 8; // 列方向上两次repeat起始地址间隔dtypeMask=64个元素,即8个block
for (uint32_t i = 0; i < dealRowCount; i++) {
Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum,
dLoop, columnRepeatParams);
}
}
// 最后一次完成[dealRowCount, dRemain] * [dealRowCount, blockElementNum] 只计算有效部分
if (dRemain > 0) {
Mul(dstUb[dLoop * repeatElementNum], src0Ub[dLoop * repeatElementNum], src1Ub, dRemain, dealRowCount,
repeatParams);
}
} else {
BinaryRepeatParams repeatParams;
repeatParams.src0RepStride = 8; // 每个repeat为256B数据,正好8个datablock
repeatParams.src0BlkStride = 1;
repeatParams.src1RepStride = 0;
repeatParams.src1BlkStride = 0;
repeatParams.dstRepStride = 8;
repeatParams.dstBlkStride = 1;
// 每次计算一行,共计算dealRowCount行
for (uint32_t i = 0; i < dealRowCount; i++) {
// 计算一行中的dLoop个repeat, 每个repeat计算256/block_size 个data_block
Mul(dstUb[i * columnCount], src0Ub[i * columnCount], src1Ub[i * blockElementNum], repeatElementNum, dLoop,
repeatParams);
// 计算一行中的尾块
if (dRemain > 0) {
Mul(dstUb[i * columnCount + dLoop * repeatElementNum],
src0Ub[i * columnCount + dLoop * repeatElementNum], src1Ub[i * blockElementNum], dRemain, 1,
repeatParams);
}
}
}
}
} // namespace SASKernel
#endif

View File

@@ -0,0 +1,771 @@
/**
 * 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 sparse_attn_sharedkv_swa_kernel.h
* \brief
*/
#ifndef SPARSE_ATTN_SHAREDKV_SWA_KERNEL_H
#define SPARSE_ATTN_SHAREDKV_SWA_KERNEL_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "kernel_tiling/kernel_tiling.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#include "../sparse_attn_sharedkv_common.h"
#include "sparse_attn_sharedkv_swa_block_cube.h"
#include "sparse_attn_sharedkv_swa_block_vector.h"
#include "../sparse_attn_sharedkv_metadata.h"
namespace SASKernel {
using namespace matmul;
using namespace optiling;
using AscendC::CrossCoreSetFlag;
using AscendC::CrossCoreWaitFlag;
// 由于S2循环前,RunInfo还没有赋值,使用Bngs1Param临时存放B、N、S1轴相关的信息;同时减少重复计算
struct SwaTempLoopInfo {
uint32_t bn2IdxInCurCore = 0;
uint32_t bIdx = 0U;
uint32_t n2Idx = 0U;
uint64_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小
uint32_t s2LoopTimes = 0U; // S2方向循环的总次数,无论TND还是BXXD都是等于实际次数,不用减1
int32_t actS1Size = 0; // TND场景下当前Batch循环处理的S1轴的大小
int32_t actOriS2Size = 0;
int32_t actCmpS2Size = 0;
bool curActSeqLenIsZero = false;
uint32_t tndCoreStartKVSplitPos = 0;
bool tndIsS2SplitCore = false;
uint32_t gS1Idx = 0U;
uint32_t s1StartIdx = 0;
uint32_t s1EndIdx = 0;
uint64_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小
uint32_t cmpLoopTimes = 0;
uint32_t oriLoopTimes = 0;
int32_t oriMaskRight = 0;
int32_t oriMaskLeft = 0;
int32_t cmpMaskRight = 0;
uint64_t actualSeqQPrefixSum = 0;
uint64_t actualSeqKVPrefixSum = 0;
uint64_t actualSeqCmpKVPrefixSum = 0;
};
template <typename SAST>
class SparseAttnSharedkvSwa {
public:
// 中间计算数据类型为float,高精度模式
using T = float;
using Q_T = typename SAST::queryType;
using KV_T = typename SAST::kvType;
using OUT_T = typename SAST::outputType;
using SINKS_T = float;
using UPDATE_T = T;
using MM1_OUT_T = T;
using MM2_OUT_T = T;
__aicore__ inline SparseAttnSharedkvSwa(){};
__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* cuSeqlensKV, __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 *softmaxLse, __gm__ uint8_t *workspace,
const SparseAttnSharedkvTilingData *__restrict tiling, __gm__ uint8_t *gmTiling,
TPipe *tPipe);
__aicore__ inline void Process();
private:
static constexpr bool PAGE_ATTENTION = SAST::pageAttention;
static constexpr int TEMPLATE_MODE = SAST::templateMode;
static constexpr bool FLASH_DECODE = SAST::flashDecode;
static constexpr SAS_LAYOUT LAYOUT_T = SAST::layout;
static constexpr SAS_LAYOUT KV_LAYOUT_T = SAST::kvLayout;
static constexpr uint32_t PRELOAD_NUM = 2;
static constexpr uint32_t N_BUFFER_M_BASIC_SIZE = 256;
static constexpr uint32_t SAS_PRELOAD_TASK_CACHE_SIZE = 3;
static constexpr uint32_t SYNC_V0_C1_FLAG = 6;
static constexpr uint32_t SYNC_C1_V1_FLAG = 7;
static constexpr uint32_t SYNC_V1_C2_FLAG = 8;
static constexpr uint32_t SYNC_C2_V2_FLAG = 9;
static constexpr uint64_t SYNC_MM2RES_BUF1_FLAG = 10;
static constexpr uint64_t SYNC_MM2RES_BUF2_FLAG = 11;
static constexpr uint64_t SYNC_FDOUTPUT_BUF_FLAG = 12;
static constexpr uint64_t kvHeadNum = 1ULL;
static constexpr uint64_t headDim = 512ULL;
static constexpr uint64_t headDimAlign = 512ULL;
static constexpr uint32_t msdIterNum = 2U;
static constexpr uint32_t dbWorkspaceRatio = PRELOAD_NUM;
const SparseAttnSharedkvTilingData *__restrict tilingData = nullptr;
TPipe *pipe = nullptr;
GlobalTensor<uint32_t> metadataGm;
uint64_t mSizeVStart = 0ULL;
int64_t threshold = 0;
uint64_t s2BatchBaseOffset = 0;
uint64_t tensorACoreOffset = 0ULL;
uint64_t tensorBCoreOffset = 0ULL;
uint64_t tensorCmpBCoreOffset = 0ULL;
uint64_t attenOutOffset = 0ULL;
uint32_t tmpBlockIdx = 0U;
uint32_t aiCoreIdx = 0U;
ConstInfo constInfo{};
SwaTempLoopInfo tempLoopInfo{};
SWACubeBlock<SAST> cubeBlock;
SWAVectorBlock<SAST> vectorBlock;
GlobalTensor<Q_T> queryGm;
GlobalTensor<KV_T> oriKvGm;
GlobalTensor<KV_T> cmpKvGm;
GlobalTensor<SINKS_T> sinksGm;
GlobalTensor<OUT_T> attentionOutGm;
GlobalTensor<T> softmaxLseGm;
GlobalTensor<int32_t> oriBlockTableGm;
GlobalTensor<int32_t> cmpBlockTableGm;
GlobalTensor<int32_t> actualSeqLengthsQGm;
GlobalTensor<int32_t> actualSeqLengthsKVGm;
GlobalTensor<int32_t> actualSeqLengthsCmpKVGm;
// workspace
GlobalTensor<MM1_OUT_T> mm1ResGm;
GlobalTensor<KV_T> vec1ResGm;
GlobalTensor<MM2_OUT_T> mm2ResGm;
GlobalTensor<UPDATE_T> vec2ResGm;
GlobalTensor<T> accumOutGm;
// ================================Init functions==================================
__aicore__ inline void InitTilingData();
__aicore__ inline void InitCalcParamsEach();
__aicore__ inline void InitBuffers();
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKv);
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV,
__gm__ uint8_t *actualSeqLengthsCmpKV);
__aicore__ inline void InitOutputSingleCore();
// ================================Process functions================================
__aicore__ inline void ProcessBalance();
__aicore__ inline void PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint64_t s2LoopIdx,
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE]);
// ================================Offset Calc=====================================
__aicore__ inline void GetSparseActualSeqLen();
__aicore__ inline void UpdateInnerLoopCond();
__aicore__ inline void CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start, uint32_t s2LoopIdx,
RunInfo &info);
__aicore__ inline int32_t GetActualSeqLenQ(uint32_t bIdx);
__aicore__ inline int32_t GetActualSeqLenKV(uint32_t bIdx);
__aicore__ inline void GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx);
// ================================Mm1==============================================
__aicore__ inline void ComputeMm1(const RunInfo &info);
// ================================Mm2==============================================
__aicore__ inline void ComputeMm2(const RunInfo &info);
__aicore__ inline void InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx);
};
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitTilingData()
{
// singleCoreParams
// singleCoreTensorSize
constInfo.mmResUbSize = tilingData->baseParams.mmResUbSize;
constInfo.bmm2ResUbSize = tilingData->baseParams.bmm2ResUbSize;
// baseParams
constInfo.batchSize = tilingData->baseParams.batchSize;
constInfo.qHeadNum = constInfo.gSize = tilingData->baseParams.nNumOfQInOneGroup;
constInfo.kvSeqSize = tilingData->baseParams.kvSeqSize;
constInfo.qSeqSize = tilingData->baseParams.qSeqSize;
constInfo.oriMaxBlockNumPerBatch = tilingData->baseParams.oriMaxBlockNumPerBatch;
constInfo.kvCacheBlockSize = tilingData->baseParams.paBlockSize;
constInfo.paOriBlockSize = tilingData->baseParams.oriBlockSize;
constInfo.paCmpBlockSize = tilingData->baseParams.cmpBlockSize;
constInfo.outputLayout = static_cast<SAS_LAYOUT>(tilingData->baseParams.outputLayout);
constInfo.kvHeadNum = kvHeadNum;
constInfo.headDim = headDim;
constInfo.oriMaskMode = tilingData->baseParams.oriMaskMode;
constInfo.oriKvStride = tilingData->baseParams.oriKvStride;
constInfo.oriWinLeft = tilingData->baseParams.oriWinLeft;
constInfo.oriWinRight = tilingData->baseParams.oriWinRight;
constInfo.returnSoftmaxLse = tilingData->baseParams.returnSoftmaxLse;
constInfo.actualLenDimsQ = tilingData->baseParams.actualLenDimsQ;
constInfo.actualLenDimsKV = tilingData->baseParams.actualLenDimsKV;
// innerSplitParams
constInfo.mBaseSize = constInfo.gSize;;
constInfo.s2BaseSize = tilingData->baseParams.s2BaseSize;
// tilingData->baseParams.s2BaseSize
constInfo.preLoadNum = PRELOAD_NUM;
constInfo.nBufferMBaseSize = N_BUFFER_M_BASIC_SIZE;
constInfo.syncV0C1 = SYNC_V0_C1_FLAG;
constInfo.syncC1V1 = SYNC_C1_V1_FLAG;
constInfo.syncV1C2 = SYNC_V1_C2_FLAG;
constInfo.syncC2V2 = SYNC_C2_V2_FLAG;
constInfo.templateMode = TEMPLATE_MODE;
// cmp
if (constInfo.templateMode == CFA_TEMPLATE) {
constInfo.cmpRatio = tilingData->cmpParams.cmpRatio;
constInfo.cmpMaskMode = tilingData->cmpParams.cmpMaskMode;
constInfo.cmpKvStride = tilingData->cmpParams.cmpKvStride;
constInfo.cmpMaxBlockNumPerBatch = tilingData->cmpParams.cmpMaxBlockNumPerBatch;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitBuffers()
{
if ASCEND_IS_AIV {
vectorBlock.InitBuffers(pipe);
} else {
cubeBlock.InitBuffers(pipe);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ,
__gm__ uint8_t *actualSeqLengthsKv)
{
if (constInfo.actualLenDimsKV != 0) {
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKv, constInfo.actualLenDimsKV);
}
if (constInfo.actualLenDimsQ != 0) {
actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ);
}
}
template <typename SAST>
__aicore__ inline void
SparseAttnSharedkvSwa<SAST>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsKV,
__gm__ uint8_t *actualSeqLengthsCmpKV)
{
if (constInfo.actualLenDimsKV != 0) {
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsKV, constInfo.actualLenDimsKV);
if (constInfo.templateMode == CFA_TEMPLATE) {
actualSeqLengthsCmpKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsCmpKV, constInfo.actualLenDimsKV);
}
}
if (constInfo.actualLenDimsQ != 0) {
actualSeqLengthsQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ, constInfo.actualLenDimsQ);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitAllZeroOutput(uint32_t bIdx, uint32_t s1Idx, uint32_t n2Idx)
{
if (constInfo.outputLayout == SAS_LAYOUT::TND) {
if (tempLoopInfo.actS1Size == 0) {
return;
}
uint32_t tBase = actualSeqLengthsQGm.GetValue(bIdx);
uint64_t attenOutOffset = (tBase + s1Idx) * kvHeadNum * constInfo.gSize * headDim + // T轴、s1轴偏移
n2Idx * constInfo.gSize * headDim; // N2轴偏移
uint64_t lseOffset = (tBase + s1Idx) * constInfo.gSize + // T轴、s1轴偏移
n2Idx * constInfo.qSeqSize * constInfo.gSize; // N2轴偏移
matmul::InitOutput<OUT_T>(attentionOutGm[attenOutOffset], constInfo.gSize * headDim, 0);
if (constInfo.returnSoftmaxLse) {
matmul::InitOutput<T>(softmaxLseGm[lseOffset], constInfo.gSize, 0);
}
} else if (constInfo.outputLayout == SAS_LAYOUT::BSND) {
uint64_t attenOutOffset = bIdx * constInfo.qSeqSize * kvHeadNum * constInfo.gSize * headDim +
s1Idx * kvHeadNum * constInfo.gSize * headDim + // B轴、S1轴偏移
n2Idx * constInfo.gSize * headDim; // N2轴偏移
uint64_t lseOffset = bIdx * constInfo.qSeqSize * constInfo.kvHeadNum * constInfo.gSize + // B轴偏移
n2Idx * constInfo.qSeqSize * constInfo.gSize + // N2轴偏移
s1Idx * constInfo.gSize; // S1轴偏移
matmul::InitOutput<OUT_T>(attentionOutGm[attenOutOffset], constInfo.gSize * headDim, 0);
if (constInfo.returnSoftmaxLse) {
matmul::InitOutput<T>(softmaxLseGm[lseOffset], constInfo.gSize, 0);
}
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitOutputSingleCore()
{
uint32_t coreNum = GetBlockNum();
if (coreNum != 0) {
uint64_t totalOutputSize = constInfo.batchSize * constInfo.qHeadNum * constInfo.qSeqSize * constInfo.headDim;
uint64_t singleCoreSize = (totalOutputSize + (2 * coreNum) - 1) / (2 * coreNum); // 2 means c:v = 1:2
uint64_t tailSize = totalOutputSize - tmpBlockIdx * singleCoreSize;
uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize;
if (singleInitOutputSize > 0) {
matmul::InitOutput<OUT_T>(attentionOutGm[tmpBlockIdx * singleCoreSize], singleInitOutputSize, 0);
}
SyncAll();
}
}
template <typename SAST>
__aicore__ inline int32_t SparseAttnSharedkvSwa<SAST>::GetActualSeqLenQ(uint32_t bIdx)
{
if constexpr (LAYOUT_T == SAS_LAYOUT::TND) {
int32_t actualSeqQPrefixSum = actualSeqLengthsQGm.GetValue(bIdx);
int32_t actualSeqQNextSum = actualSeqLengthsQGm.GetValue(bIdx + 1);
tempLoopInfo.actualSeqQPrefixSum = static_cast<uint64_t>(actualSeqQPrefixSum);
return actualSeqQNextSum - actualSeqQPrefixSum;
} else {
tempLoopInfo.actualSeqQPrefixSum = static_cast<uint64_t>(bIdx * constInfo.qSeqSize);
if (constInfo.actualLenDimsQ == 0) {
return static_cast<int32_t>(constInfo.qSeqSize);
} else {
return actualSeqLengthsQGm.GetValue(bIdx);
}
}
}
template <typename SAST>
__aicore__ inline int32_t SparseAttnSharedkvSwa<SAST>::GetActualSeqLenKV(uint32_t bIdx)
{
if constexpr (KV_LAYOUT_T == SAS_LAYOUT::PA_ND) {
tempLoopInfo.actualSeqKVPrefixSum = static_cast<uint64_t>(bIdx * constInfo.kvSeqSize);
if (constInfo.actualLenDimsKV == 0) {
return static_cast<int32_t>(constInfo.kvSeqSize);
}
return actualSeqLengthsKVGm.GetValue(bIdx);
} else if constexpr(KV_LAYOUT_T == SAS_LAYOUT::BSND) {
return static_cast<int32_t>(constInfo.kvSeqSize);
} else if constexpr(KV_LAYOUT_T == SAS_LAYOUT::TND) {
int32_t actualSeqKVPrefixSum = actualSeqLengthsKVGm.GetValue(bIdx);
int32_t actualSeqKVNextSum = actualSeqLengthsKVGm.GetValue(bIdx + 1);
if (constInfo.templateMode == CFA_TEMPLATE) {
tempLoopInfo.actualSeqCmpKVPrefixSum = actualSeqLengthsCmpKVGm.GetValue(bIdx);
}
tempLoopInfo.actualSeqKVPrefixSum = actualSeqKVPrefixSum;
return actualSeqKVNextSum - actualSeqKVPrefixSum;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::GetSparseActualSeqLen()
{
// 行无效通过ori部分判断, ori部分如果有行无效那么ori和cmp都有
if (static_cast<int32_t>(tempLoopInfo.s1EndIdx) < -(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size)) {
tempLoopInfo.actOriS2Size = 0;
tempLoopInfo.actCmpS2Size = 0;
return;
}
// 对于cmp部分还有top k, tempLoopInfo.actS2Size只针对cmp
if (constInfo.templateMode == CFA_TEMPLATE) {
int32_t thresHold = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio;
tempLoopInfo.actCmpS2Size = thresHold;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::UpdateInnerLoopCond()
{
if ((tempLoopInfo.actCmpS2Size == 0 && tempLoopInfo.actOriS2Size == 0) || (tempLoopInfo.actS1Size == 0)) {
tempLoopInfo.curActSeqLenIsZero = true;
return;
}
tempLoopInfo.curActSeqLenIsZero = false;
tempLoopInfo.mBasicSizeTail = (tempLoopInfo.actS1Size * constInfo.gSize) % constInfo.mBaseSize;
tempLoopInfo.mBasicSizeTail =
(tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail;
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::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 *cuSeqlensKV, __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 *softmaxLse,
__gm__ uint8_t *workspace, const SparseAttnSharedkvTilingData *__restrict tiling, __gm__ uint8_t *gmTiling,
TPipe *tPipe)
{
if ASCEND_IS_AIV {
tmpBlockIdx = GetBlockIdx(); // vec:0-47
aiCoreIdx = tmpBlockIdx / 2;
} else {
tmpBlockIdx = GetBlockIdx(); // cube:0-23
aiCoreIdx = tmpBlockIdx;
}
// init tiling data
tilingData = tiling;
InitTilingData();
if (KV_LAYOUT_T == SAS_LAYOUT::TND && LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(cuSeqlensQ, cuSeqlensKV, cuSeqlensCmpKV);
} else if (KV_LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(seqUsedQ, cuSeqlensKV, cuSeqlensCmpKV);
} else if ((KV_LAYOUT_T == SAS_LAYOUT::PA_ND || KV_LAYOUT_T == SAS_LAYOUT::BSND) && LAYOUT_T == SAS_LAYOUT::TND) {
InitActualSeqLen(cuSeqlensQ, seqUsedKV);
} else if ((KV_LAYOUT_T == SAS_LAYOUT::PA_ND || KV_LAYOUT_T == SAS_LAYOUT::BSND)) {
InitActualSeqLen(seqUsedQ, seqUsedKV);
}
metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata);
InitCalcParamsEach();
pipe = tPipe;
// init global buffer
queryGm.SetGlobalBuffer((__gm__ Q_T *)query);
oriKvGm.SetGlobalBuffer((__gm__ KV_T *)oriKV);
if (constInfo.templateMode == CFA_TEMPLATE) {
cmpKvGm.SetGlobalBuffer((__gm__ KV_T *)cmpKV);
}
if (sinks != nullptr) {
sinksGm.SetGlobalBuffer((__gm__ SINKS_T *)sinks);
}
attentionOutGm.SetGlobalBuffer((__gm__ OUT_T *)attentionOut);
softmaxLseGm.SetGlobalBuffer((__gm__ T *)softmaxLse);
if ASCEND_IS_AIV {
if (LAYOUT_T != SAS_LAYOUT::TND) {
if (constInfo.needInit) {
InitOutputSingleCore();
}
}
}
if constexpr (PAGE_ATTENTION) {
oriBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)oriBlockTable);
if (constInfo.templateMode == CFA_TEMPLATE) {
cmpBlockTableGm.SetGlobalBuffer((__gm__ int32_t *)cmpBlockTable);
}
}
// workspace 内存排布
// |Q--|mm1ResGm|vec1ResGm|mm2ResGm|vec2ResGm
// |Core0_Q1-Core0_Q2-Core1_Q1-Core1_Q2....Core32_Q1-Core32_Q2|Core0_mmRes
uint64_t offset = 0;
mm1ResGm.SetGlobalBuffer(
(__gm__ MM1_OUT_T *)(workspace + offset +
aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(MM1_OUT_T);
vec1ResGm.SetGlobalBuffer(
(__gm__ Q_T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.mmResUbSize * sizeof(KV_T);
mm2ResGm.SetGlobalBuffer(
(__gm__ MM2_OUT_T *)(workspace + offset +
aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(MM2_OUT_T);
vec2ResGm.SetGlobalBuffer(
(__gm__ T *)(workspace + offset + aiCoreIdx * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T)));
offset += GetBlockNum() * dbWorkspaceRatio * constInfo.bmm2ResUbSize * sizeof(T);
if ASCEND_IS_AIV {
vectorBlock.InitParams(constInfo, tilingData);
vectorBlock.InitVec1GlobalTensor(mm1ResGm, vec1ResGm, actualSeqLengthsQGm, actualSeqLengthsKVGm, sinksGm, softmaxLseGm);
vectorBlock.InitVec2GlobalTensor(accumOutGm, vec2ResGm, mm2ResGm, attentionOutGm);
}
if ASCEND_IS_AIC {
cubeBlock.InitParams(constInfo);
cubeBlock.InitMm1GlobalTensor(queryGm, oriKvGm, cmpKvGm, mm1ResGm);
cubeBlock.InitMm2GlobalTensor(vec1ResGm, mm2ResGm, attentionOutGm);
cubeBlock.InitPageAttentionInfo(oriKvGm, oriBlockTableGm, cmpBlockTableGm);
}
// 要在InitParams之后执行
if (pipe != nullptr) {
InitBuffers();
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::InitCalcParamsEach()
{
if (aiCoreIdx != 0) {
constInfo.bN2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_START_INDEX, false));
constInfo.gS1Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_START_INDEX, false));
constInfo.s2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_START_INDEX, false));
}
constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_BN2_END_INDEX, false));
constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_M_END_INDEX, false));
constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_S2_END_INDEX, false));
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::CalcParams(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start,
uint32_t s2LoopIdx, RunInfo &info)
{
info.isValid = s2LoopIdx < tempLoopInfo.s2LoopTimes;
info.loop = loop;
info.cmpLoop = cmpLoop;
info.bIdx = tempLoopInfo.bIdx;
info.gS1Idx = tempLoopInfo.gS1Idx;
info.s1Idx = tempLoopInfo.gS1Idx / constInfo.gSize;
info.s2Idx = s2LoopIdx;
info.n2IdxReal = tempLoopInfo.n2Idx;
info.curSInnerLoopTimes = tempLoopInfo.s2LoopTimes;
info.tndIsS2SplitCore = tempLoopInfo.tndIsS2SplitCore;
info.tndCoreStartKVSplitPos = tempLoopInfo.tndCoreStartKVSplitPos;
info.isBmm2Output = false;
info.actS1Size = tempLoopInfo.actS1Size;
// M方向的尾块
info.actMBaseSize = tempLoopInfo.mBasicSizeTail;
if ASCEND_IS_AIV {
info.mSize = info.actMBaseSize;
info.mSizeV = (info.mSize <= 16) ? info.mSize : ((CeilDiv(info.mSize, 16) + 1) / 2 * 16);
info.mSizeVStart = 0;
if (tmpBlockIdx % 2 == 1) {
info.mSizeVStart = info.mSizeV;
info.mSizeV = info.mSize - info.mSizeV;
}
}
info.isFirstSInnerLoop = s2LoopIdx == s2Start;
if (info.isFirstSInnerLoop) {
tempLoopInfo.bn2IdxInCurCore++;
}
info.isLastS2Loop = (s2LoopIdx == (tempLoopInfo.s2LoopTimes - 1));
info.bn2IdxInCurCore = tempLoopInfo.bn2IdxInCurCore - 1;
uint64_t tndBIdxOffsetForQ = tempLoopInfo.actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim;
uint64_t tndBIdxOffsetForKV = tempLoopInfo.actualSeqKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim;
uint64_t tndBIdxOffsetForCmpKV = tempLoopInfo.actualSeqCmpKVPrefixSum * constInfo.kvHeadNum * constInfo.headDim;
if (info.isFirstSInnerLoop) {
tensorACoreOffset = tndBIdxOffsetForQ + info.gS1Idx * constInfo.headDim;
tensorBCoreOffset = tndBIdxOffsetForKV + info.n2Idx * constInfo.headDim;
tensorCmpBCoreOffset = tndBIdxOffsetForCmpKV + info.n2Idx * constInfo.headDim;
}
info.tensorAOffset = tensorACoreOffset;
info.tensorBOffset = tensorBCoreOffset;
info.tensorCmpBOffset = tensorCmpBCoreOffset;
info.attenOutOffset = tensorACoreOffset;
if (s2LoopIdx < tempLoopInfo.oriLoopTimes) {
// S2首次循环只能在ori_kv
info.isOri = true;
info.relativeS2Idx = 0;
uint64_t s2Offset = info.s2Idx * constInfo.s2BaseSize;
if (s2LoopIdx + 1 == tempLoopInfo.oriLoopTimes) {
info.actualSingleProcessSInnerSize = (tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1) - s2Offset;
} else {
info.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
}
info.s2StartPoint = tempLoopInfo.oriMaskLeft;
info.cmpS2IdLimit = 0;
} else {
if (constInfo.templateMode == CFA_TEMPLATE) {
info.isOri = false;
info.relativeS2Idx = info.s2Idx - tempLoopInfo.oriLoopTimes;
uint64_t s2Offset = (info.s2Idx - tempLoopInfo.oriLoopTimes) * constInfo.s2BaseSize;
if (s2LoopIdx + 1 == tempLoopInfo.s2LoopTimes) {
info.actualSingleProcessSInnerSize = tempLoopInfo.actCmpS2Size - s2Offset;
} else {
info.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
}
info.s2StartPoint = 0;
info.cmpS2IdLimit = (tempLoopInfo.cmpMaskRight + tempLoopInfo.s1EndIdx + 1) / constInfo.cmpRatio;
}
}
info.actualSingleProcessSInnerSizeAlign =
SASAlign(info.actualSingleProcessSInnerSize, SASVectorBlock<SAST>::BYTE_BLOCK);
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::ComputeMm1(const RunInfo &info)
{
uint32_t nBufferLoopTimes = CeilDiv(info.actMBaseSize, constInfo.nBufferMBaseSize);
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
cubeBlock.ComputeMm1(info, mSplitInfo);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V1);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::ComputeMm2(const RunInfo &info)
{
uint32_t nBufferLoopTimes = (info.actMBaseSize + constInfo.nBufferMBaseSize - 1) / constInfo.nBufferMBaseSize;
uint32_t nBufferTail = info.actMBaseSize - (nBufferLoopTimes - 1) * constInfo.nBufferMBaseSize;
for (uint32_t i = 0; i < nBufferLoopTimes; i++) {
MSplitInfo mSplitInfo;
mSplitInfo.nBufferStartM = i * constInfo.nBufferMBaseSize;
mSplitInfo.nBufferDealM = (i + 1 != nBufferLoopTimes) ? constInfo.nBufferMBaseSize : nBufferTail;
CrossCoreWaitFlag(constInfo.syncV1C2);
cubeBlock.ComputeMm2(info, mSplitInfo);
CrossCoreSetFlag<ConstInfo::SAS_SYNC_MODE2, PIPE_FIX>(constInfo.syncC2V2);
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::Process()
{
uint32_t hasLoad = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, FA_CORE_ENABLE_INDEX, false));
if (hasLoad == 0) {
return;
}
if ASCEND_IS_AIV {
vectorBlock.AllocEventID();
vectorBlock.InitSoftmaxDefaultBuffer();
} else {
cubeBlock.AllocEventID();
}
ProcessBalance();
if ASCEND_IS_AIV {
vectorBlock.FreeEventID();
} else {
cubeBlock.FreeEventID();
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::GetBN2Idx(uint32_t bN2Idx, uint32_t &bIdx, uint32_t &n2Idx)
{
bIdx = bN2Idx / kvHeadNum;
n2Idx = bN2Idx % kvHeadNum;
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::ProcessBalance()
{
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE];
uint32_t gloop = 0;
uint32_t cmpLoop = 0;
uint32_t gS1LoopEnd = 0;
bool globalLoopStart = true;
// 适配左闭右开
if (constInfo.bN2Start == constInfo.bN2End) {
if (constInfo.gS1Start != constInfo.gS1End || constInfo.s2Start != constInfo.s2End) {
constInfo.bN2End += 1;
}
} else if ((constInfo.gS1End != 0) || (constInfo.s2End != 0)) {
constInfo.bN2End += 1;
}
for (uint32_t bN2LoopIdx = constInfo.bN2Start; bN2LoopIdx < constInfo.bN2End; bN2LoopIdx++) {
GetBN2Idx(bN2LoopIdx, tempLoopInfo.bIdx, tempLoopInfo.n2Idx);
tempLoopInfo.actS1Size = GetActualSeqLenQ(tempLoopInfo.bIdx); // 获取actualSeqLength
bool isS1ZeroAndLastBatch = (tempLoopInfo.actS1Size == 0) && ((constInfo.outputLayout == SAS_LAYOUT::BSND) ||
(bN2LoopIdx + 1 == constInfo.bN2End));
uint32_t gS1SplitNum = CeilDiv(tempLoopInfo.actS1Size * constInfo.gSize, constInfo.mBaseSize);
// 当处于最后一个BN2时, 且gS1End为0时, 说明当前BN2里的所有数据都在当前核处理
gS1LoopEnd = (bN2LoopIdx == constInfo.bN2End - 1 && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum;
// 当处于最后一个BN2且当前S1为0时,需要进入循环计算preload导致的未完成的部分
gS1LoopEnd = isS1ZeroAndLastBatch ? gS1LoopEnd + 1 : gS1LoopEnd;
for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < gS1LoopEnd; gS1LoopIdx++) {
tempLoopInfo.actOriS2Size = GetActualSeqLenKV(tempLoopInfo.bIdx);
// 对于各轴上的真实的idx, 采用左闭右闭的方案
tempLoopInfo.gS1Idx = gS1LoopIdx * constInfo.mBaseSize;
tempLoopInfo.s1StartIdx = tempLoopInfo.gS1Idx / constInfo.gSize;
tempLoopInfo.s1EndIdx =
Min((tempLoopInfo.s1StartIdx + constInfo.mBaseSize / constInfo.gSize - 1), tempLoopInfo.actS1Size - 1);
// 此处均为闭区间
tempLoopInfo.oriMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size +
static_cast<int32_t>(tempLoopInfo.s1EndIdx) + constInfo.oriWinRight;
tempLoopInfo.oriMaskLeft = Max(tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size +
static_cast<int32_t>(tempLoopInfo.s1EndIdx) - constInfo.oriWinLeft,
0);
if (constInfo.templateMode == CFA_TEMPLATE) {
tempLoopInfo.cmpMaskRight = tempLoopInfo.actOriS2Size - tempLoopInfo.actS1Size;
}
GetSparseActualSeqLen();
UpdateInnerLoopCond();
uint32_t oriSplitNum = 0;
uint32_t s2SplitNum = 0;
bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == gS1LoopEnd);
if (tempLoopInfo.curActSeqLenIsZero) {
if ASCEND_IS_AIV {
InitAllZeroOutput(tempLoopInfo.bIdx, tempLoopInfo.s1StartIdx, tempLoopInfo.n2Idx);
}
if (!isEnd) {
continue;
}
} else {
oriSplitNum = CeilDiv(tempLoopInfo.oriMaskRight - tempLoopInfo.oriMaskLeft + 1, constInfo.s2BaseSize);
s2SplitNum = oriSplitNum;
if (constInfo.templateMode == CFA_TEMPLATE) {
uint32_t cmpSplitNum = CeilDiv(tempLoopInfo.actCmpS2Size, constInfo.s2BaseSize);
s2SplitNum = oriSplitNum + cmpSplitNum;
tempLoopInfo.cmpLoopTimes = cmpSplitNum;
}
}
tempLoopInfo.s2LoopTimes = s2SplitNum;
tempLoopInfo.oriLoopTimes = oriSplitNum;
uint32_t s2LoopEnd = (isEnd && constInfo.s2End != 0) ? constInfo.s2End : tempLoopInfo.s2LoopTimes;
tempLoopInfo.s2LoopTimes = s2LoopEnd;
// 分核修改后需要打开
// 当前s2是否被切,决定了输出是否要写到attenOut上
tempLoopInfo.tndIsS2SplitCore = ((constInfo.s2Start == 0) && (s2LoopEnd == s2SplitNum)) ? false : true;
tempLoopInfo.tndCoreStartKVSplitPos = globalLoopStart ? constInfo.coreStartKVSplitPos : 0;
uint32_t extraLoop = isEnd ? PRELOAD_NUM : 0;
for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (s2LoopEnd + extraLoop); s2LoopIdx++) {
// PreloadPipeline loop初始值要求为 PRELOAD_NUM
PreloadPipeline(gloop, cmpLoop, constInfo.s2Start, s2LoopIdx, extraInfo);
++gloop;
}
globalLoopStart = false;
constInfo.s2Start = 0;
}
constInfo.gS1Start = 0;
}
}
template <typename SAST>
__aicore__ inline void SparseAttnSharedkvSwa<SAST>::PreloadPipeline(uint32_t loop, uint32_t cmpLoop, uint64_t s2Start,
uint64_t s2LoopIdx,
RunInfo extraInfo[SAS_PRELOAD_TASK_CACHE_SIZE])
{
RunInfo &extraInfo0 = extraInfo[loop % SAS_PRELOAD_TASK_CACHE_SIZE]; // 本轮任务
RunInfo &extraInfo2 = extraInfo[(loop + 2) % SAS_PRELOAD_TASK_CACHE_SIZE]; // 上一轮任务
RunInfo &extraInfo1 = extraInfo[(loop + 1) % SAS_PRELOAD_TASK_CACHE_SIZE]; // 上两轮任务
CalcParams(loop, cmpLoop, s2Start, s2LoopIdx, extraInfo0);
if (extraInfo0.isValid) {
if ASCEND_IS_AIC {
ComputeMm1(extraInfo0);
}
}
if (extraInfo2.isValid) {
if ASCEND_IS_AIV {
vectorBlock.ProcessVec1L(extraInfo2);
}
if ASCEND_IS_AIC {
ComputeMm2(extraInfo2);
}
}
if (extraInfo1.isValid) {
if ASCEND_IS_AIV {
vectorBlock.ProcessVec2L(extraInfo1);
}
extraInfo1.isValid = false;
}
}
} // namespace SASKernel
#endif // SPARSE_ATTN_SHAREDKV_SWA_KERNEL_H

View File

@@ -0,0 +1,72 @@
/**
 * 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 sparse_attn_sharedkv.cpp
* \brief
*/
#include "kernel_operator.h"
#include "lib/matmul_intf.h"
#include "sparse_attn_sharedkv_template_tiling_key.h"
#include "arch32/sparse_attn_sharedkv_scfa_kernel.h"
#include "arch32/sparse_attn_sharedkv_swa_kernel.h"
#include "sparse_attn_sharedkv_metadata.h"
using namespace AscendC;
using namespace optiling::detail;
using namespace SASKernel;
#define SAS_OP_IMPL(templateClass, tilingdataClass, ...) \
do { \
templateClass<SASType<__VA_ARGS__>> op; \
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tiling_data_in, tiling); \
const tilingdataClass *__restrict tiling_data = &tiling_data_in; \
op.Init(query, oriKV, cmpKV, cmpSparseIndices, oriBlockTable, cmpBlockTable, cuSeqlensQ, \
cuSeqlensOriKv, cuSeqlensCmpKv, seqUsedQ, seqUsedKV, \
sinks, metadata, attentionOut, softmaxLse, user, tiling_data, tiling, &tPipe); \
op.Process(); \
} while (0)
template <int FLASH_DECODE, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE>
__global__ __aicore__ void
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 *softmaxLse, __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);
if constexpr (ORIG_DTYPE_Q == DT_FLOAT16 && ORIG_DTYPE_ORI_KV == DT_FLOAT16 && ORIG_DTYPE_ATTN_OUT == DT_FLOAT16) {
if constexpr (TEMPLATE_MODE == SCFA_TEMPLATE) {
SAS_OP_IMPL(SparseAttnSharedkvScfa, SparseAttnSharedkvTilingData, half, half, half, FLASH_DECODE,
static_cast<SAS_LAYOUT>(LAYOUT_T), static_cast<SAS_LAYOUT>(KV_LAYOUT_T), TEMPLATE_MODE);
} else {
SAS_OP_IMPL(SparseAttnSharedkvSwa, SparseAttnSharedkvTilingData, half, half, half, FLASH_DECODE,
static_cast<SAS_LAYOUT>(LAYOUT_T), static_cast<SAS_LAYOUT>(KV_LAYOUT_T), TEMPLATE_MODE);
}
}
if constexpr (ORIG_DTYPE_Q == DT_BF16 && ORIG_DTYPE_ORI_KV == DT_BF16 && ORIG_DTYPE_ATTN_OUT == DT_BF16) {
if constexpr (TEMPLATE_MODE == SCFA_TEMPLATE) {
SAS_OP_IMPL(SparseAttnSharedkvScfa, SparseAttnSharedkvTilingData, bfloat16_t, bfloat16_t, bfloat16_t,
FLASH_DECODE, static_cast<SAS_LAYOUT>(LAYOUT_T), static_cast<SAS_LAYOUT>(KV_LAYOUT_T),
TEMPLATE_MODE);
} else {
SAS_OP_IMPL(SparseAttnSharedkvSwa, SparseAttnSharedkvTilingData, bfloat16_t, bfloat16_t, bfloat16_t,
FLASH_DECODE, static_cast<SAS_LAYOUT>(LAYOUT_T), static_cast<SAS_LAYOUT>(KV_LAYOUT_T),
TEMPLATE_MODE);
}
}
}

View File

@@ -0,0 +1,325 @@
/**
 * 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 sparse_attn_sharedkv_common.h
* \brief
*/
#ifndef SPARSE_ATTN_SHAREDKV_COMMON_H
#define SPARSE_ATTN_SHAREDKV_COMMON_H
#include "kernel_operator.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
namespace SASKernel {
using namespace AscendC;
// 将isCheckTiling设置为false, 输入输出的max&sum&exp的shape为(m, 1)
constexpr SoftmaxConfig SAS_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC = {false, 0, 0, SoftmaxMode::SOFTMAX_OUTPUT_WITHOUT_BRC};
enum class SAS_RUN_MODE {
SWA_MODE = 0,
SCFA_MODE = 1,
CFA_MODE = 2,
};
enum class SAS_LAYOUT {
BSND = 0,
TND = 1,
PA_ND = 2
};
template <typename Q_T, typename KV_T, typename OUT_T, const bool FLASH_DECODE = false,
SAS_LAYOUT LAYOUT_T = SAS_LAYOUT::BSND, SAS_LAYOUT KV_LAYOUT_T = SAS_LAYOUT::PA_ND, int TEMPLATE_MODE = 0,
typename... Args>
struct SASType {
using queryType = Q_T;
using kvType = KV_T;
using outputType = OUT_T;
static constexpr bool flashDecode = FLASH_DECODE;
static constexpr SAS_LAYOUT layout = LAYOUT_T;
static constexpr SAS_LAYOUT kvLayout = KV_LAYOUT_T;
static constexpr bool pageAttention = (KV_LAYOUT_T == SAS_LAYOUT::PA_ND);
static constexpr int templateMode = TEMPLATE_MODE;
};
// ================================Util functions==================================
template <typename T1, typename T2>
__aicore__ inline T1 SASAlign(T1 num, T2 rnd)
{
return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd * rnd);
}
template <typename T1, typename T2>
__aicore__ inline T1 CeilDiv(T1 num, T2 rnd)
{
return (rnd == 0) ? 0 : ((num + rnd - 1) / rnd);
}
template <typename T1, typename T2>
__aicore__ inline T1 Min(T1 a, T2 b)
{
return (a > b) ? b : a;
}
template <typename T1, typename T2>
__aicore__ inline T1 Max(T1 a, T2 b)
{
return (a > b) ? a : b;
}
template <typename T>
__aicore__ inline size_t BlockAlign(size_t s)
{
if constexpr (IsSameType<T, int4b_t>::value) {
return (s + 63) / 64 * 64;
}
size_t n = (32 / sizeof(T));
return (s + n - 1) / n * n;
}
struct PAShape {
uint32_t blockSize;
uint32_t headNum; // 一般为kv的head num,对应n2
uint32_t headDim; // 512 对应d
uint32_t kvStride;
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 s2Idx;
uint32_t dIdx;
uint32_t s1Idx;
};
// 场景:query、key、value GM to L1
// GM按ND格式存储
// L1按NZ格式存储
// GM的行、列、列的stride
template <typename T>
__aicore__ inline void DataCopyGmNDToL1(LocalTensor<T> &l1Tensor, GlobalTensor<T> &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; // nd矩阵的行数
// T为int4场景下,dValue = col / 2,srcDValue = colStride / 2
nd2nzPara.dValue = col; // nd矩阵的列数
nd2nzPara.srcDValue = colStride; // 同一nd矩阵相邻行起始地址间的偏移
nd2nzPara.dstNzC0Stride = rowAlign;
nd2nzPara.dstNzNStride = 1;
nd2nzPara.srcNdMatrixStride = 0;
nd2nzPara.dstNzMatrixStride = 0;
DataCopy(l1Tensor, gmTensor, nd2nzPara);
}
/*
适用PA数据从GM拷贝到L1,支持ND、NZ数据;
PA的layout分 BNBD(blockNum,N,blockSize,D) BBH(blockNum,blockSize,N*D
BSH\BSND\TND 为BBH
shape.copyRowNumAlign 需要16字节对齐,如拷贝k矩阵,一次拷贝128*512,遇到尾块 10*512 需对齐到16*512
*/
template <typename T>
__aicore__ inline void DataCopyPA(LocalTensor<T> &dstTensor, //l1
GlobalTensor<T> &srcTensor, //gm
GlobalTensor<int32_t> &blockTableGm,
const PAShape &shape, // blockSize, headNum, headDim
const Position &startPos) // bacthIdx nIdx curSeqIdx
{
uint32_t copyFinishRowCnt = 0;
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
uint32_t curS2Idx = startPos.s2Idx;
uint32_t blockElementCnt = 32 / sizeof(T);
while (copyFinishRowCnt < shape.copyRowNum) {
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
uint64_t reaminRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
uint64_t idInBlockTable =
blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset); // 从block table上的获取编号
uint32_t copyRowCnt = shape.blockSize - reaminRowCnt; // 一次只能处理一个Block
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
}
// uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim; // PA的偏移
uint64_t offset = idInBlockTable * shape.kvStride; // PA的偏移
uint64_t dStride = shape.headDim;
offset += (uint64_t)(startPos.n2Idx * shape.headDim * shape.blockSize) +
reaminRowCnt * shape.headDim + startPos.dIdx;
uint32_t dValue = shape.actHeadDim;
uint32_t srcDValue = dStride;
LocalTensor<T> tmpDstTensor = dstTensor[copyFinishRowCnt * blockElementCnt];
GlobalTensor<T> tmpSrcTensor = srcTensor[offset];
DataCopyGmNDToL1<T>(tmpDstTensor, tmpSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
copyFinishRowCnt += copyRowCnt;
curS2Idx += copyRowCnt;
}
}
struct RunInfo {
uint32_t loop = 0;
uint32_t cmpLoop = 0; // 用于判断取 用于merge的4块GM 中的哪一块
uint32_t bIdx = 0;
uint32_t gIdx = 0;
uint32_t s1Idx = 0;
uint32_t s2Idx = 0;
uint32_t n2IdxReal = 0;
uint32_t relativeS2Idx = 0;
uint32_t bn2IdxInCurCore = 0;
uint32_t curSInnerLoopTimes = 0;
uint64_t tndBIdxOffsetForQ = 0;
uint64_t tndBIdxOffsetForKV = 0;
uint64_t tensorCmpBOffset = 0;
uint64_t tensorAOffset = 0;
uint64_t tensorBOffset = 0;
uint64_t attenOutOffset = 0;
uint64_t attenMaskOffset = 0;
uint64_t topKBaseOffset = 0;
uint32_t actualSingleProcessSInnerSize = 0;
uint32_t actualSingleProcessSInnerSizeAlign = 0;
bool isFirstSInnerLoop = false;
uint32_t s2BatchOffset = 0;
uint32_t gSize = 0;
uint32_t s1Size = 0;
uint32_t s2Size = 0;
uint32_t mSize = 0;
uint32_t mSizeV = 0;
uint32_t mSizeVStart = 0;
uint32_t tndIsS2SplitCore = 0;
uint32_t tndCoreStartKVSplitPos = 0;
bool isBmm2Output = false;
bool isValid = false;
static constexpr uint32_t n2Idx = 0;
uint64_t actS1Size = 1;
uint64_t actS2SizeOri = 0ULL;
uint32_t gS1Idx = 0;
uint64_t actS2Size = 1;
uint64_t actOriS2Size = 1;
uint32_t actMBaseSize = 0;
bool isLastS2Loop = 0;
int32_t nextTokensPerBatch = 0;
int64_t threshold = 0;
uint32_t curTopKIdx = 0;
uint64_t curOffsetInSparseBlock = 0;
bool isOri = true; // 判断当前块是在Ori部分还是Cmp部分
uint64_t s2StartPoint = 0;
int64_t cmpS2IdLimit = 0;
int32_t v0S2DealSize = 0;
int32_t v0S2Start = 0;
};
struct ConstInfo {
// CUBE与VEC核间同步的模式
static constexpr uint32_t SAS_SYNC_MODE2 = 2;
// BUFFER的字节数
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64;
static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256;
static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512;
static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024;
static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048;
static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096;
static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192;
static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384;
static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768;
// FP32的0值和极大值
static constexpr float FLOAT_ZERO = 0;
static constexpr float FLOAT_MAX = 3.402823466e+38F;
// preLoad的总次数
uint32_t preLoadNum = 0U;
uint32_t nBufferMBaseSize = 0U;
// CUBE和VEC的核间同步EventID
uint32_t syncV0C1 = 0U;
uint32_t syncC1V1 = 0U;
uint32_t syncV1C2 = 0U;
uint32_t syncC2V2 = 0U;
uint32_t mmResUbSize = 0U; // Matmul1输出结果GM上的大小
uint32_t vec1ResUbSize = 0U; // Vector1输出结果GM上的大小
uint32_t bmm2ResUbSize = 0U; // Matmul2输出结果GM上的大小
uint64_t batchSize = 0ULL;
uint64_t gSize = 0ULL;
uint64_t qHeadNum = 0ULL;
uint64_t kvHeadNum = 0;
uint64_t headDim = 0;
uint64_t kvSeqSize = 0ULL; // kv最大S长度
uint64_t qSeqSize = 1ULL; // q最大S长度
int64_t kvCacheBlockSize = 0; // PA场景的block size
uint64_t paCmpBlockSize = 0;
uint64_t paOriBlockSize = 0;
int64_t orikvCacheBlockSize = 0;
int64_t cmpkvCacheBlockSize = 0;
uint32_t oriMaxBlockNumPerBatch = 0; // PA场景的最大单batch block number
uint32_t cmpMaxBlockNumPerBatch = 0;
uint32_t splitKVNum = 0U; // S2核间切分的切分份数
SAS_LAYOUT outputLayout; // 输出的Transpose格式
uint32_t oriMaskMode = 0;
uint32_t cmpMaskMode = 0;
uint32_t oriKvStride = 0;
uint32_t cmpKvStride = 0;
bool needInit = false;
uint32_t templateMode = 0;
// FlashDecoding
uint32_t actualCombineLoopSize = 0U; // FlashDecoding场景, S2在核间切分的最大份数
uint64_t combineLseOffset = 0ULL;
uint64_t combineAccumOutOffset = 0ULL;
uint32_t actualLenDimsQ = 0U; // query的actualSeqLength 的维度
uint32_t actualLenDimsKV = 0U; // KV 的actualSeqLength 的维度
// TND
uint32_t s2Start = 0U; // TND场景下,S2的起始位置
uint32_t s2End = 0U; // 单核TND场景下S2循环index上限
uint32_t bN2Start = 0U;
uint32_t bN2End = 0U;
uint32_t gS1Start = 0U;
uint32_t gS1End = 0U;
uint32_t tndFDCoreArrLen = 0U; // TNDFlashDecoding相关分核信息array的长度
uint32_t coreStartKVSplitPos = 0U; // TNDFlashDecoding kv起始位置
uint32_t mBaseSize = 1ULL;
uint32_t s2BaseSize = 1ULL;
// sparse attr
int64_t sparseBlockSize = 0;
uint32_t sparseBlockCount = 0;
// cmp attr
int64_t cmpRatio = 0;
// win
int32_t oriWinRight = 0;
int32_t oriWinLeft = 128;
// 是否返回SoftmaxLse
bool returnSoftmaxLse = false;
};
struct MSplitInfo {
uint32_t nBufferIdx = 0U;
uint32_t nBufferStartM = 0U;
uint32_t nBufferDealM = 0U;
uint32_t vecStartM = 0U;
uint32_t vecDealM = 0U;
};
} // namespace SASKernel
#endif // SPARSE_ATTN_SHAREDKV_COMMON_H

View File

@@ -0,0 +1,79 @@
/**
 * 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 sparse_attn_sharedkv_metadata.h
* \brief
*/
#ifndef SPARSE_ATTN_SHAREDKV_METADATA_H
#define 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 = 8;
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;
// 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];
};
} // namespace detail
static_assert(SAS_META_SIZE * sizeof(SAS_METADATA_T) >= sizeof(detail::SasMetaData));
} // namespace optiling
#endif

View File

@@ -0,0 +1,59 @@
/**
 * 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 sparse_attn_sharedkv_template_tiling_key.h
* \brief
*/
#ifndef SPARSE_ATTN_SHARED_TEMPLATE_TILING_KEY_H
#define 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, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE,
CFA_TEMPLATE, SCFA_TEMPLATE), );
// 支持的模板参数组合
// 用于调用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, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SWA_TEMPLATE), ),
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, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, CFA_TEMPLATE),
),
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, SAS_LAYOUT_BSND, SAS_LAYOUT_TND),
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, SCFA_TEMPLATE), ), );
#endif // TEMPLATE_TILING_KEY