* Copyright (c) 2025 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
* \file rain_fusion_attention_kernel.h
* \brief Rain Fusion Attention Kernel Implementation
*/
#ifndef RAIN_FUSION_ATTENTION_KERNEL_H
#define RAIN_FUSION_ATTENTION_KERNEL_H
#include "rain_fusion_attention_kernel_common.hpp"
using namespace NpuArch;
using namespace RfaKenelCommon;
namespace RainFusion {
* @brief Rain Fusion Attention Inference Kernel
*
* This kernel implements rain fusion attention where attention is computed only on
* selected KV blocks specified by selectIdx. This reduces computation for long sequences
* by focusing on relevant tokens.
*
* @tparam BlockMmadQK Block-level QK matmul module
* @tparam BlockMmadPV Block-level PV matmul module
* @tparam EpilogueOnlineSoftmax Online softmax epilogue
* @tparam EpilogueRescaleO Output rescaling epilogue
* @tparam PAGED_CACHE_FLAG Whether to use paged KV cache
* @tparam QUERY_LAYOUT Query tensor layout (0=TND, 1=BNSD)
* @tparam KV_CACHE_LAYOUT KV cache layout (0=TND, 1=BNSD)
*/
template <
class BlockMmadQK,
class BlockMmadPV,
class EpilogueOnlineSoftmax,
class EpilogueRescaleO,
bool PAGED_CACHE_FLAG,
uint32_t QUERY_LAYOUT,
uint32_t KV_CACHE_LAYOUT>
class RainFusionAttentionKernel {
public:
using ArchTag = typename BlockMmadQK::ArchTag;
using L1TileShape = typename BlockMmadQK::L1TileShape;
using ElementQ = typename BlockMmadQK::ElementA;
using LayoutQ = typename BlockMmadQK::LayoutA;
using ElementK = typename BlockMmadQK::ElementB;
using LayoutK = typename BlockMmadQK::LayoutB;
using ElementS = typename BlockMmadQK::ElementC;
using LayoutS = typename BlockMmadQK::LayoutC;
using ElementP = typename BlockMmadPV::ElementA;
using LayoutP = typename BlockMmadPV::LayoutA;
using ElementV = typename BlockMmadPV::ElementB;
using LayoutV = typename BlockMmadPV::LayoutB;
using ElementMask = typename EpilogueOnlineSoftmax::ElementMask;
using ElementO = typename EpilogueRescaleO::ElementOutput;
using LayoutO = typename EpilogueRescaleO::LayoutOutput;
using ElementOTmp = typename EpilogueRescaleO::ElementInput;
using LayoutOTmp = typename EpilogueRescaleO::LayoutInput;
using ElementLse = typename EpilogueRescaleO::ElementLse;
using LayoutLse = typename EpilogueRescaleO::LayoutLse;
using ElementUpdate = typename EpilogueRescaleO::ElementUpdate;
using LayoutUpdate = typename EpilogueRescaleO::LayoutUpdate;
static constexpr Epilogue::LseMode LSE_MODE = EpilogueRescaleO::LSE_MODE;
__aicore__ inline
RainFusionAttentionKernel() {}
__aicore__ inline void operator()(RainFusionAttentionKernelParams const ¶ms)
{
__gm__ RainFusionAttentionTilingData *rainFusionAttentionTilingData = reinterpret_cast<__gm__ RainFusionAttentionTilingData *>(params.tiling);
uint64_t mm1OutSize = rainFusionAttentionTilingData->mm1OutSize;
uint64_t smOnlineOutSize = rainFusionAttentionTilingData->smOnlineOutSize;
uint64_t mm2OutSize = rainFusionAttentionTilingData->mm2OutSize;
uint32_t batch = rainFusionAttentionTilingData->batch;
uint32_t qHeads = rainFusionAttentionTilingData->numHeads;
uint32_t kvHeads = rainFusionAttentionTilingData->kvHeads;
uint32_t embed = rainFusionAttentionTilingData->embeddingSize;
uint32_t pagedBlockSize = rainFusionAttentionTilingData->blockSize;
uint32_t maxNumBlocksPerBatch = rainFusionAttentionTilingData->maxNumBlocksPerBatch;
uint32_t firstBatchTaskNum = rainFusionAttentionTilingData->firstBatchTaskNum;
uint32_t totalTaskNum = rainFusionAttentionTilingData->totalTaskNum;
uint32_t maskType = rainFusionAttentionTilingData->maskType;
ElementS scaleValue = static_cast<ElementS>(rainFusionAttentionTilingData->scaleValue);
uint32_t totalQBlocks = rainFusionAttentionTilingData->totalQBlocks;
uint32_t maxKvBlockNum = rainFusionAttentionTilingData->maxKvBlockNum;
uint32_t qBlockX = rainFusionAttentionTilingData->blockShapeX;
uint32_t qBlockY = rainFusionAttentionTilingData->blockShapeY;
uint32_t qBlockNum = totalQBlocks / qBlockX;
uint32_t qBlockInX = (qBlockX + BASIC_BLOCK_SIZE - 1) / BASIC_BLOCK_SIZE;
uint32_t firstQBlockNum = rainFusionAttentionTilingData->firstQBlockNum;
uint32_t maxQSeqlen = rainFusionAttentionTilingData->maxQSeqlen;
uint32_t maxKvSeqlen = rainFusionAttentionTilingData->maxKvSeqlen;
uint32_t useUniformQSeqlen = rainFusionAttentionTilingData->useUniformQSeqlen;
uint32_t useUniformKvSeqlen = rainFusionAttentionTilingData->useUniformKvSeqlen;
AscendC::GlobalTensor<ElementQ> gQ;
gQ.SetGlobalBuffer((__gm__ ElementQ *)params.q);
AscendC::GlobalTensor<ElementK> gK;
gK.SetGlobalBuffer((__gm__ ElementK *)params.k);
AscendC::GlobalTensor<ElementK> gV;
gV.SetGlobalBuffer((__gm__ ElementK *)params.v);
AscendC::GlobalTensor<int32_t> gBlockTable;
gBlockTable.SetGlobalBuffer((__gm__ int32_t *)(params.blockTables));
AscendC::GlobalTensor<int64_t> gActualQseqlen;
gActualQseqlen.SetGlobalBuffer((__gm__ int64_t *)params.actualQseqlen);
AscendC::GlobalTensor<int64_t> gActualKvseqlen;
gActualKvseqlen.SetGlobalBuffer((__gm__ int64_t *)params.actualKvseqlen);
AscendC::GlobalTensor<int64_t> gSelectIdx;
gSelectIdx.SetGlobalBuffer((__gm__ int64_t *)params.selectIdx);
AscendC::GlobalTensor<int64_t> gSelectNumIdx;
gSelectNumIdx.SetGlobalBuffer((__gm__ int64_t *)params.selectNumIdx);
AscendC::GlobalTensor<ElementO> gO;
gO.SetGlobalBuffer((__gm__ ElementO *)params.o);
AscendC::GlobalTensor<ElementLse> gLse;
gLse.SetGlobalBuffer((__gm__ ElementLse *)params.lse);
AscendC::GlobalTensor<ElementS> gS;
gS.SetGlobalBuffer((__gm__ ElementS *)params.workspace);
AscendC::GlobalTensor<ElementP> gP;
gP.SetGlobalBuffer((__gm__ ElementP *)(params.workspace + mm1OutSize));
AscendC::GlobalTensor<ElementOTmp> gOTmp;
gOTmp.SetGlobalBuffer((__gm__ ElementOTmp *)(params.workspace + mm1OutSize + smOnlineOutSize));
AscendC::GlobalTensor<ElementOTmp> gOUpdate;
gOUpdate.SetGlobalBuffer((__gm__ ElementOTmp *)(params.workspace + mm1OutSize + smOnlineOutSize + mm2OutSize));
uint32_t coreIdx = AscendC::GetBlockIdx();
uint32_t coreNum = AscendC::GetBlockNum();
#ifdef __DAV_C220_CUBE__
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID0);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID1);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID2);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID3);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID4);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID5);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID6);
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID7);
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(EVENT_ID0);
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(EVENT_ID1);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID0);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID1);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID2);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID3);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID4);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID5);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID6);
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID7);
static constexpr uint32_t L1_QK_SIZE =
BlockMmadQK::L1TileShape::M * BlockMmadQK::L1TileShape::K * sizeof(ElementQ) +
BlockMmadQK::L1TileShape::N * BlockMmadQK::L1TileShape::K * sizeof(ElementK) * 2;
BlockMmadQK blockMmadQK(resource);
BlockMmadPV blockMmadPV(resource, L1_QK_SIZE);
#endif
#ifdef __DAV_C220_VEC__
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID0);
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID1);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID2);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID3);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID4);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID5);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID6);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID3);
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID2);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2);
EpilogueOnlineSoftmax epilogueOnlineSoftmax(resource, scaleValue);
EpilogueRescaleO epilogueRescaleO(resource);
coreIdx = AscendC::GetBlockIdx() / AscendC::GetSubBlockNum();
#endif
uint64_t strideQO = 0;
uint64_t strideKV = 0;
uint64_t strideQOB = 0;
uint64_t strideQON = 0;
uint64_t strideQOS = 0;
uint64_t strideKVB = 0;
uint64_t strideKVN = 0;
uint64_t strideKVS = 0;
if constexpr (QUERY_LAYOUT == 1) {
strideQOB = static_cast<uint64_t>(qHeads) * maxQSeqlen * embed;
strideQON = static_cast<uint64_t>(maxQSeqlen) * embed;
strideQOS = embed;
} else {
strideQO = static_cast<uint64_t>(qHeads) * embed;
}
if constexpr (KV_CACHE_LAYOUT == 1) {
strideKVB = static_cast<uint64_t>(kvHeads) * maxKvSeqlen * embed;
strideKVN = static_cast<uint64_t>(maxKvSeqlen) * embed;
strideKVS = embed;
} else {
strideKV = static_cast<uint64_t>(kvHeads) * embed;
}
uint32_t embedRound = AlignUp<uint32_t>(embed, BLOCK_SIZE);
uint32_t groupSize = qHeads / kvHeads;
uint64_t qBOffset = 0;
uint64_t kBOffset = 0;
uint64_t vBOffset = 0;
uint64_t oBOffset = 0;
uint64_t blockBOffset = 0;
uint32_t preTotalTaskNum = 0;
uint32_t preTotalQBlockNum = 0;
uint32_t curBatch = 0;
uint32_t qSeqlen = useUniformQSeqlen ? maxQSeqlen :
static_cast<uint32_t>(static_cast<int64_t>(gActualQseqlen.GetValue(curBatch)));
uint32_t kvSeqlen = useUniformKvSeqlen ? maxKvSeqlen :
static_cast<uint32_t>(static_cast<int64_t>(gActualKvseqlen.GetValue(curBatch)));
uint32_t curQNBlockTile = GetQNBlockTile(qSeqlen, groupSize);
uint32_t qNBlockNumPerGroup = curQNBlockTile == 0 ? 1 : (groupSize + curQNBlockTile - 1) / curQNBlockTile;
uint32_t curQNBlockNum = qNBlockNumPerGroup * kvHeads;
uint32_t curQSBlockTile = GetQSBlockTile(kvSeqlen);
uint32_t curQSBlockNum = GetQBlocks(qSeqlen, qBlockX);
uint32_t curTotalTaskNum = firstBatchTaskNum;
uint32_t curQXBlockNum = (qSeqlen + qBlockX - 1) / qBlockX;
uint32_t curTotalQBlockNum = firstQBlockNum;
for (uint32_t taskIdx = coreIdx; taskIdx < totalTaskNum; taskIdx += uint32_t(coreNum)) {
while (taskIdx >= curTotalTaskNum) {
++curBatch;
preTotalTaskNum = curTotalTaskNum;
preTotalQBlockNum = curTotalQBlockNum;
if constexpr (QUERY_LAYOUT == 1) {
qBOffset = static_cast<uint64_t>(curBatch) * strideQOB;
oBOffset = static_cast<uint64_t>(curBatch) * strideQOB;
} else {
qBOffset += static_cast<uint64_t>(qSeqlen) * strideQO;
oBOffset += static_cast<uint64_t>(qSeqlen) * strideQO;
}
if constexpr (!PAGED_CACHE_FLAG) {
if constexpr (KV_CACHE_LAYOUT == 1) {
kBOffset = static_cast<uint64_t>(curBatch) * strideKVB;
vBOffset = static_cast<uint64_t>(curBatch) * strideKVB;
} else {
kBOffset += static_cast<uint64_t>(kvSeqlen) * strideKV;
vBOffset += static_cast<uint64_t>(kvSeqlen) * strideKV;
}
} else {
blockBOffset += maxNumBlocksPerBatch;
}
qSeqlen = useUniformQSeqlen ? maxQSeqlen :
static_cast<uint32_t>(static_cast<int64_t>(gActualQseqlen.GetValue(curBatch)));
kvSeqlen = useUniformKvSeqlen ? maxKvSeqlen :
static_cast<uint32_t>(static_cast<int64_t>(gActualKvseqlen.GetValue(curBatch)));
curQNBlockTile = GetQNBlockTile(qSeqlen, groupSize);
qNBlockNumPerGroup = curQNBlockTile == 0 ? 1 : (groupSize + curQNBlockTile - 1) / curQNBlockTile;
curQNBlockNum = qNBlockNumPerGroup * kvHeads;
curQSBlockTile = GetQSBlockTile(kvSeqlen);
curQSBlockNum = GetQBlocks(qSeqlen, qBlockX);
curTotalTaskNum += curQNBlockNum * curQSBlockNum;
curQXBlockNum = (qSeqlen + qBlockX - 1) / qBlockX;
curTotalQBlockNum += qHeads * curQXBlockNum;
}
uint32_t taskIdxCurBatch = taskIdx - preTotalTaskNum;
uint32_t qSBlockIdx = taskIdxCurBatch / curQNBlockNum;
uint32_t qXIdx = qSBlockIdx / qBlockInX;
uint32_t qXInnerIdx = qSBlockIdx - qXIdx * qBlockInX;
uint32_t qNBlockIdx = taskIdxCurBatch - qSBlockIdx * curQNBlockNum;
uint32_t qNBlockIdxCurGroup = qNBlockIdx % qNBlockNumPerGroup;
uint32_t xBlockNum = qSeqlen / qBlockX;
uint32_t xTailNum = qSeqlen - xBlockNum * qBlockX;
uint32_t kvHeadIdx = qNBlockIdx / qNBlockNumPerGroup;
uint32_t qHeadIdx = kvHeadIdx * groupSize + qNBlockIdxCurGroup * curQNBlockTile;
uint32_t curSelectIdx = preTotalQBlockNum + qXIdx * qHeads + qHeadIdx;
uint32_t curSelectNum = static_cast<uint32_t>(gSelectNumIdx.GetValue(curSelectIdx));
if (curSelectNum == 0) {
continue;
}
uint32_t lastSelectIdx = static_cast<int32_t>(
gSelectIdx.GetValue(curSelectIdx * maxKvBlockNum + curSelectNum - 1));
uint32_t kvYBlockNum = (kvSeqlen + qBlockY - 1) / qBlockY;
uint32_t curKvSeqLen = (lastSelectIdx == kvYBlockNum - 1 && kvSeqlen % qBlockY != 0) ?
qBlockY * (curSelectNum - 1) + kvSeqlen % qBlockY : qBlockY * curSelectNum;
uint64_t gmOffsetQ = 0;
uint64_t gmOffsetK = 0;
uint64_t gmOffsetV = 0;
uint64_t gmOffsetO = 0;
if constexpr (QUERY_LAYOUT == 1) {
uint32_t qSeqOffset = qXIdx * qBlockX + qXInnerIdx * BASIC_BLOCK_SIZE;
gmOffsetQ = qBOffset + qHeadIdx * strideQON + qSeqOffset * strideQOS;
gmOffsetO = oBOffset + qHeadIdx * strideQON + qSeqOffset * strideQOS;
} else {
uint32_t qSeqOffset = qXIdx * qBlockX + qXInnerIdx * BASIC_BLOCK_SIZE;
gmOffsetQ = qBOffset + qSeqOffset * strideQO + qHeadIdx * embed;
gmOffsetO = oBOffset + qSeqOffset * strideQO + qHeadIdx * embed;
}
if constexpr (KV_CACHE_LAYOUT == 1) {
gmOffsetK = kBOffset + kvHeadIdx * strideKVN;
gmOffsetV = vBOffset + kvHeadIdx * strideKVN;
} else {
gmOffsetK = kBOffset + kvHeadIdx * embed;
gmOffsetV = vBOffset + kvHeadIdx * embed;
}
uint32_t qSBlockSize = (qXIdx == xBlockNum) ?
(qXInnerIdx == xTailNum / curQSBlockTile ?
xTailNum - qXInnerIdx * curQSBlockTile : curQSBlockTile) :
((qXInnerIdx == qBlockInX - 1) ? qBlockX - qXInnerIdx * curQSBlockTile : curQSBlockTile);
uint32_t qNBlockSize = (qNBlockIdxCurGroup == (qNBlockNumPerGroup - 1)) ?
(groupSize - qNBlockIdxCurGroup * curQNBlockTile) : curQNBlockTile;
uint32_t rowNum = qSBlockSize * qNBlockSize;
uint32_t rowNumRound = AlignUp<uint32_t>(rowNum, BLOCK_SIZE);
uint32_t noSkipKvS = curKvSeqLen;
uint32_t kvSLoopNumTotal = (noSkipKvS + pagedBlockSize - 1) / pagedBlockSize;
uint32_t blockStackNum = MAX_KV_STACK_LEN / pagedBlockSize;
uint32_t stackSeqTile;
uint32_t stackSeqTilePad = blockStackNum * pagedBlockSize;
uint32_t preKVNum = PRE_LAUNCH * blockStackNum;
int32_t stackSeqCount = 0;
#ifdef __DAV_C220_CUBE__
LayoutQ layoutQTemp(rowNum, embed);
uint64_t actualStrideKV = 0;
if constexpr (KV_CACHE_LAYOUT == 1) {
actualStrideKV = strideKVS;
} else {
actualStrideKV = strideKV;
}
LayoutK layoutKTemp(actualStrideKV, blockStackNum * pagedBlockSize);
LayoutV layoutVTemp(blockStackNum * pagedBlockSize, actualStrideKV);
uint64_t qGmStride = 0;
if constexpr (QUERY_LAYOUT == 1) {
qGmStride = strideQOS;
} else {
qGmStride = strideQO;
}
blockMmadQK.loadQGM(gQ[gmOffsetQ], layoutQTemp, rowNum, qNBlockSize, qGmStride);
#endif
for (uint32_t kvSIdx = 0; kvSIdx < kvSLoopNumTotal + preKVNum; kvSIdx += blockStackNum) {
if (kvSIdx < kvSLoopNumTotal) {
stackSeqTile = noSkipKvS - kvSIdx * pagedBlockSize;
if (stackSeqTile >= pagedBlockSize * blockStackNum) {
stackSeqTile = pagedBlockSize * blockStackNum;
}
uint32_t curStackTileMod = stackSeqCount % (PRE_LAUNCH + 1);
uint64_t gmOffsetS = coreIdx * WORKSPACE_BLOCK_SIZE_DB * (PRE_LAUNCH + 1) +
curStackTileMod * WORKSPACE_BLOCK_SIZE_DB;
GemmCoord actualBlockShapeQK{rowNum, stackSeqTile, embed};
LayoutS layOutS(rowNum, stackSeqTile, stackSeqTilePad);
#ifdef __DAV_C220_CUBE__
uint64_t actualStrideKVForQK = 0;
if constexpr (KV_CACHE_LAYOUT == 1) {
actualStrideKVForQK = strideKVS;
} else {
actualStrideKVForQK = strideKV;
}
blockMmadQK(gQ[gmOffsetQ],
gK[gmOffsetK],
gS[gmOffsetS],
gBlockTable[blockBOffset],
gSelectIdx[curSelectIdx * maxKvBlockNum],
layoutQTemp,
layoutKTemp,
layOutS,
actualBlockShapeQK,
kvSIdx,
kvSLoopNumTotal,
pagedBlockSize,
actualStrideKVForQK,
qBlockY,
curSelectNum,
kvYBlockNum,
kvSeqlen);
NpuArch::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(qkReady);
#endif
#ifdef __DAV_C220_VEC__
LayoutP layOutP(rowNum, stackSeqTile, stackSeqTilePad);
uint64_t gmOffsetP = gmOffsetS;
NpuArch::Arch::CrossCoreWaitFlag(qkReady);
epilogueOnlineSoftmax(gP[gmOffsetP],
gS[gmOffsetS],
layOutP,
layOutS,
actualBlockShapeQK,
(stackSeqCount == 0),
0,
qSBlockSize,
qNBlockSize,
curStackTileMod);
NpuArch::Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(softmaxReady);
#endif
}
if (kvSIdx >= preKVNum) {
uint32_t nowkvSIdx = kvSIdx - preKVNum;
stackSeqTile = noSkipKvS - nowkvSIdx * pagedBlockSize;
if (stackSeqTile >= pagedBlockSize * blockStackNum) {
stackSeqTile = pagedBlockSize * blockStackNum;
}
uint32_t curStackTileMod = (stackSeqCount - PRE_LAUNCH) % (PRE_LAUNCH + 1);
uint64_t gmOffsetOTmp = coreIdx * WORKSPACE_BLOCK_SIZE_DB * (PRE_LAUNCH + 1) +
curStackTileMod * WORKSPACE_BLOCK_SIZE_DB;
GemmCoord actualBlockShapePV{rowNum, embed, stackSeqTile};
LayoutOTmp layoutOTmp(rowNum, embed, embedRound);
#ifdef __DAV_C220_CUBE__
LayoutP layoutPTemp(rowNum, stackSeqTile, stackSeqTilePad);
uint64_t gmOffsetP = coreIdx * WORKSPACE_BLOCK_SIZE_DB * (PRE_LAUNCH + 1) +
curStackTileMod * WORKSPACE_BLOCK_SIZE_DB;
uint64_t actualStrideKVForPV = 0;
if constexpr (KV_CACHE_LAYOUT == 1) {
actualStrideKVForPV = strideKVS;
} else {
actualStrideKVForPV = strideKV;
}
blockMmadPV(gP[gmOffsetP],
gV[gmOffsetV],
gOTmp[gmOffsetOTmp],
gBlockTable[blockBOffset],
gSelectIdx[curSelectIdx * maxKvBlockNum],
layoutPTemp,
layoutVTemp,
layoutOTmp,
actualBlockShapePV,
nowkvSIdx,
kvSLoopNumTotal,
pagedBlockSize,
kvSeqlen,
actualStrideKVForPV,
blockStackNum,
softmaxReady,
qBlockY,
curSelectNum,
kvYBlockNum);
NpuArch::Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(pvReady);
#endif
#ifdef __DAV_C220_VEC__
LayoutO layoutO;
if constexpr (QUERY_LAYOUT == 1) {
layoutO = LayoutO(qSeqlen, embed);
} else {
layoutO = LayoutO(qSeqlen, qHeads * embed);
}
LayoutUpdate layoutUpdate(rowNum, embed, embedRound);
uint64_t gmOffsetUpdate = (uint64_t)(coreIdx * WORKSPACE_BLOCK_SIZE_DB);
NpuArch::Arch::CrossCoreWaitFlag(pvReady);
epilogueRescaleO(
gO[gmOffsetO],
gOTmp[gmOffsetOTmp],
gOUpdate[gmOffsetUpdate],
layoutO,
layoutOTmp,
layoutUpdate,
actualBlockShapePV,
qSBlockSize,
qNBlockSize,
(stackSeqCount - PRE_LAUNCH == 0),
nowkvSIdx + blockStackNum >= kvSLoopNumTotal,
curStackTileMod);
#endif
}
stackSeqCount++;
}
}
#ifdef __DAV_C220_CUBE__
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID1);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID2);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID3);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID4);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID5);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID6);
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(EVENT_ID7);
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(EVENT_ID1);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID1);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID2);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID3);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID4);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID5);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID6);
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(EVENT_ID7);
#endif
#ifdef __DAV_C220_VEC__
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID2);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID3);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID4);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID5);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID6);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID1);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(EVENT_ID2);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID3);
#endif
AscendC::PipeBarrier<PIPE_ALL>();
}
private:
NpuArch::Arch::Resource<ArchTag> resource;
NpuArch::Arch::CrossCoreFlag qkReady{QK_READY_ID};
NpuArch::Arch::CrossCoreFlag softmaxReady{SOFTMAX_READY_ID};
NpuArch::Arch::CrossCoreFlag pvReady{PV_READY_ID};
};
}
#endif