@@ -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
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user