init v0.23.0

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

View File

@@ -0,0 +1,112 @@
/**
 * 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_flash_attention_common_arch35.h
* \brief
*/
#ifndef SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
#define SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
#include <type_traits>
#include "kernel_tiling/kernel_tiling.h"
constexpr uint64_t BLOCK_BYTE = 32;
constexpr uint32_t NEGATIVE_MIN_VALUE_FP32 = 0xFF7FFFFF;
constexpr uint32_t L0AB_SHARED_SIZE_64K = 65536; // 65536表示64*1024
constexpr uint32_t L0C_SHARED_SIZE_256K = 262144; // 262144表示256 * 1024
constexpr uint32_t BUFFER_SIZE_16K = 16384; // 16384表示16 * 1024
constexpr uint32_t BUFFER_SIZE_32K = 32768; // 32768表示32 * 1024
constexpr uint32_t BUFFER_SIZE_128K = 131072; // 131072表示128 * 1024
constexpr uint32_t CV_RATIO = 2;
constexpr uint64_t SYNC_MODE = 4;
static constexpr uint32_t SFA_SYNC_MODE0 = 0;
enum class SFA_LAYOUT {
BSND = 0,
TND = 1,
PA_BSND = 2,
};
enum class SFATemplateMode {
SWA_TEMPLATE_MODE = 0,
CFA_TEMPLATE_MODE = 1,
SCFA_TEMPLATE_MODE = 2
};
namespace BaseApi {
__aicore__ constexpr uint64_t Align2Func(uint64_t data)
{
return (data + 1UL) >> 1UL << 1UL; // 向上2对齐, +1移位2
}
__aicore__ constexpr uint64_t Align8Func(uint64_t data)
{
return (data + 7UL) >> 3UL << 3UL; // 向上8对齐, +7移位3
}
__aicore__ constexpr uint64_t Align16Func(uint64_t data)
{
return (data + 15UL) >> 4UL << 4UL; // 向上16对齐, +15移位4
}
__aicore__ constexpr uint64_t Align64Func(uint64_t data)
{
return (data + 63UL) >> 6UL << 6UL; // 向上64对齐, +63移位6
}
}
#define TEMPLATE_INTF \
template <typename Q_T, typename KV_T, typename T, typename OUTPUT_T, bool isFd, bool isPa, SFA_LAYOUT LAYOUT_T, \
SFA_LAYOUT KV_LAYOUT_T, SFATemplateMode TEMPLATE_MODE, bool IS_SPLIT_G>
#define TEMPLATE_INTF_ARGS \
Q_T, KV_T, T, OUTPUT_T, isFd, isPa, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G
#define CUBE_BLOCK_TRAITS_TYPE_FIELDS(X) \
X(Q_T) \
X(KV_T) \
X(T) \
X(OUTPUT_T) \
#define CUBE_BLOCK_TRAITS_CONST_FIELDS(X) \
X(isFd, bool, false) \
X(isPa, bool, true) \
X(LAYOUT_T, SFA_LAYOUT, SFA_LAYOUT::BSND) \
X(KV_LAYOUT_T, SFA_LAYOUT, SFA_LAYOUT::PA_BSND) \
X(TEMPLATE_MODE, SFATemplateMode, SFATemplateMode::SCFA_TEMPLATE_MODE) \
X(IS_SPLIT_G, bool, false)
/* 1. 生成带默认值的模版Template */
#define GEN_TYPE_PARAM(name) typename name,
#define GEN_CONST_PARAM(name, type, default_val) type name = default_val,
#define TEMPLATES_DEF \
template <CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TYPE_PARAM) \
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_CONST_PARAM) bool end = true>
/* 2. 生成不带默认值的模版Template */
#define GEN_TEMPLATE_TYPE_NODEF(name) typename name,
#define GEN_TEMPLATE_CONST_NODEF(name, type, default_val) type name,
#define TEMPLATES_DEF_NO_DEFAULT \
template <CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TEMPLATE_TYPE_NODEF) \
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TEMPLATE_CONST_NODEF) bool end>
/* 3. 生成有默认值的Args */
#define GEN_ARG_NAME(name, ...) name,
#define TEMPLATE_ARGS \
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARG_NAME) \
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARG_NAME) end
#endif // SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H

View File

@@ -0,0 +1,717 @@
/**
 * 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_flash_attention_kernel_mla.h
* \brief
*/
#ifndef SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
#define SPARSE_FLASH_ATTENTION_KERNEL_MLA_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_flash_attention_service_cube_mla.h"
#include "sparse_flash_attention_service_vector_mla.h"
#include "sparse_flash_attention_common_arch35.h"
#include "sparse_flash_attention_kvcache.h"
#if __has_include("../../common/op_kernel/matmul.h")
#include "../../common/op_kernel/matmul.h"
#else
#include "../common/matmul.h"
#endif
#if __has_include("../../common/op_kernel/FixpipeOut.h")
#include "../../common/op_kernel/FixpipeOut.h"
#else
#include "../common/FixpipeOut.h"
#endif
#if __has_include("../../common/op_kernel/CopyInL1.h")
#include "../../common/op_kernel/CopyInL1.h"
#else
#include "../common/CopyInL1.h"
#endif
using matmul::MatmulType;
using namespace AscendC;
using namespace AscendC::Impl::Detail;
using namespace regbaseutil;
namespace BaseApi {
template <typename CubeBlockType, typename VecBlockType> class SparseFlashAttentionKernelMla {
public:
ARGS_TRAITS;
__aicore__ inline SparseFlashAttentionKernelMla(){};
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *actualSeqLengthsQ,
__gm__ uint8_t *actualSeqLengths, __gm__ uint8_t *blockTable,
__gm__ uint8_t *queryRope, __gm__ uint8_t *keyRope, __gm__ uint8_t *attentionOut,
__gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum, __gm__ uint8_t *workspace,
const SparseFlashAttentionTilingDataMla *__restrict tiling,
__gm__ uint8_t *gmTiling, TPipe *tPipe);
__aicore__ inline void Process();
private:
__aicore__ inline void ProcessMainLoop();
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *queryRope, __gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
__gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum, __gm__ uint8_t *workspace,
const SparseFlashAttentionTilingDataMla *__restrict tiling,
TPipe *tPipe);
__aicore__ inline void InitLocalBuffer();
__aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace);
__aicore__ inline void ComputeConstexpr();
__aicore__ inline void SetRunInfo(RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount,
int64_t s2LoopLimit, int64_t multiCoreInnerIdx);
__aicore__ inline void ComputeBmm1Tail(RunInfo &runInfo, RunParamStr &runParam);
__aicore__ inline void InitUniqueConstInfo();
__aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam);
__aicore__ inline void InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo);
__aicore__ inline void InitCalcParamsEach();
__aicore__ inline uint64_t GetBalanceActualSeqLengths(GlobalTensor<int32_t> &actualSeqLengths, uint32_t bIdx);
__aicore__ inline void GetAxisStartIdx(uint32_t bN2EndPrev, uint32_t s1GEndPrev, uint32_t s2EndPrev);
TPipe *pipe;
const SparseFlashAttentionTilingDataMla *__restrict tilingData;
static constexpr uint64_t SYNC_MODE = 4;
static constexpr uint32_t PRELOAD_NUM = 2;
/* 核间通道 */
BufferManager<BufferType::GM> gmBufferManager;
BufferManager<BufferType::UB> ubBufferManager;
BuffersPolicyDB<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm1Buffers;
BuffersPolicySingleBuffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm2Buffers;
// mm2左矩阵P
BufferManager<BufferType::L1> l1BufferManager;
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> l1RightBuffers;
CVSharedParams sharedParams;
/* GM信息 */
__gm__ int32_t *actualSeqKvlenAddr = nullptr;
__gm__ int32_t *actualSeqQlenAddr = nullptr;
GlobalTensor<int32_t> actualSeqLengthsQGm;
uint32_t usedCoreNum = 0U;
GlobalTensor<int32_t> oriTopkLengthGm;
bool hasOriTopkLength = false;
/* workspace 空间 */
BuffersPolicy3buff<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> v0ResGmBuffers;
/* 核Index信息 */
int32_t aicIdx;
/* 切G时最大s2Loop */
int64_t maxS2LoopCnt;
/* 初始化后不变的信息 */
ConstInfo constInfo;
/* 模板库Block */
CubeBlockType cubeBlock;
VecBlockType vecBlock;
uint32_t crossCoreSyncBufId = 0;
};
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::Init(
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *actualSeqLengthsQ,
__gm__ uint8_t *actualSeqLengths, __gm__ uint8_t *blockTable,
__gm__ uint8_t *queryRope, __gm__ uint8_t *keyRope,
__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum,
__gm__ uint8_t *workspace, const SparseFlashAttentionTilingDataMla *__restrict tiling,
__gm__ uint8_t *gmTiling, TPipe *tPipe)
{
fa_base_matmul::idCounterNum = 0;
constInfo.subBlockIdx = GetSubBlockIdx();
if ASCEND_IS_AIC {
this->aicIdx = GetBlockIdx();
constInfo.aivIdx = 0;
} else {
constInfo.aivIdx = GetBlockIdx();
this->aicIdx = constInfo.aivIdx >> 1;
this->tilingData = tiling;
}
constInfo.s1BaseSize = 64;
constInfo.s2BaseSize = 128;
this->pipe = tPipe;
vecBlock.InitVecBlock(tPipe, this->tilingData, this->sharedParams, \
this->aicIdx, constInfo.subBlockIdx, actualSeqLengthsQ, actualSeqLengths);
if ASCEND_IS_AIV {
constInfo.bSize = this->sharedParams.bSize;
constInfo.gSize = this->sharedParams.gSize;
constInfo.s1Size = this->sharedParams.s1Size;
constInfo.dSizeV = 512;
constInfo.needInit = this->sharedParams.needInit;
constInfo.returnSoftmaxLse = this->sharedParams.returnSoftmaxLse;
}
vecBlock.CleanOutput(attentionOut, softmaxMax, softmaxSum, constInfo);
/* cube侧不依赖sharedParams的scalar前置 */
InitMMResBuf(workspace);
if ASCEND_IS_AIC {
cubeBlock.InitCubeBlock(pipe, l1BufferManager, query, queryRope);
/* wait kfc message */
CrossCoreWaitFlag<SYNC_MODE, PIPE_S>(15);
auto tempTilingSSbuf = reinterpret_cast<__ssbuf__ uint32_t*>(0); // 从ssbuf的0地址开始拷贝
auto tempTiling = reinterpret_cast<uint32_t *>(&sharedParams);
#pragma unroll
for (int i = 0; i < sizeof(CVSharedParams) / sizeof(uint32_t); ++i, ++tempTilingSSbuf, ++tempTiling) {
*tempTiling = *tempTilingSSbuf;
}
}
this->ComputeConstexpr();
this->InitGlobalBuffer(query, key, value, queryRope, keyRope, sparseIndices, \
blockTable, actualSeqLengthsQ, actualSeqLengths, softmaxMax, softmaxSum, \
workspace, tiling, tPipe); // gm设置
this->InitCalcParamsEach();
this->InitLocalBuffer();
}
template <typename CubeBlockType, typename VecBlockType> __aicore__ inline
void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitCalcParamsEach()
{
// 计算总的基本块
maxS2LoopCnt = 0; // 所有核中最大累计s2Loop
uint32_t totalBaseNum = 0;
uint32_t s1GBaseSize = constInfo.gSize;
uint32_t actBatchS2 = 1;
uint32_t coreNum = GetBlockNum(); // G128时相邻两个cube核处理一个s1,coreNum减半
uint32_t currCoreIdx = aicIdx;
if constexpr (IS_SPLIT_G) {
coreNum = coreNum >> 1;
currCoreIdx = currCoreIdx >> 1;
}
uint32_t actBatchS1 = 1;
for (uint32_t bIdx = 0; bIdx < constInfo.bSize; bIdx++) {
actBatchS1 = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx); // 不切S2,只关注S1
if (actBatchS1 < constInfo.s1Size) {
constInfo.needInit = true;
}
totalBaseNum += actBatchS1 * actBatchS2;
}
uint32_t avgBaseNum = 1;
if (totalBaseNum > coreNum) {
avgBaseNum = (totalBaseNum + coreNum - 1) / coreNum;
if constexpr (IS_SPLIT_G) {
usedCoreNum = (totalBaseNum + avgBaseNum - 1) / avgBaseNum << 1;
}
} else {
if constexpr (IS_SPLIT_G) {
usedCoreNum = totalBaseNum << 1;
} else {
usedCoreNum = totalBaseNum;
}
}
if constexpr (IS_SPLIT_G) {
maxS2LoopCnt = avgBaseNum * (constInfo.sparseBlockCount + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
}
if (aicIdx >= usedCoreNum) {
return;
}
// 计算当前核的基本块
uint32_t accumBaseNum = 0; // 当前累积的基本块数
uint32_t targetBaseNum = 0;
uint32_t lastValidBIdx = 0;
uint32_t lastValidactBatchS1 = 0;
bool setStart = false;
targetBaseNum = (currCoreIdx + 1) * avgBaseNum; // 计算当前的目标权重
uint32_t targetStartBaseNum = targetBaseNum - avgBaseNum;
for (uint32_t bN2Idx = 0; bN2Idx < constInfo.bSize * constInfo.n2Size; bN2Idx++) {
uint32_t bIdx = bN2Idx / constInfo.n2Size;
actBatchS1 = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx);
for (uint32_t s1GIdx = 0; s1GIdx < actBatchS1; s1GIdx++) {
accumBaseNum += 1;
if (!setStart && accumBaseNum >= targetStartBaseNum) {
constInfo.bN2Start = bN2Idx;
constInfo.gS1Start = s1GIdx;
setStart = true;
}
if (accumBaseNum >= targetBaseNum) {
// 更新当前核的End分核信息
constInfo.bN2End = bN2Idx;
constInfo.gS1End = s1GIdx;
constInfo.s2End = 0;
if (currCoreIdx != 0) {
GetAxisStartIdx(constInfo.bN2Start, constInfo.gS1Start, 0);
}
return;
}
}
if ((actBatchS1 > 0) && (actBatchS2 > 0)) {
lastValidBIdx = bIdx;
lastValidactBatchS1 = actBatchS1;
}
}
if (!setStart) {
constInfo.bN2Start = lastValidBIdx;
constInfo.gS1Start = lastValidactBatchS1 - 1;
}
if (accumBaseNum < targetBaseNum) {
// 更新最后一个核的End分核信息
constInfo.bN2End = lastValidBIdx;
constInfo.gS1End = lastValidactBatchS1 - 1;
constInfo.s2End = 0;
if (currCoreIdx != 0) {
GetAxisStartIdx(constInfo.bN2Start, constInfo.gS1Start, 0);
}
return;
}
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline uint64_t
SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::GetBalanceActualSeqLengths(
GlobalTensor<int32_t> &actualSeqLengths, uint32_t bIdx)
{
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
if (bIdx > 0) {
return actualSeqQlenAddr[bIdx] - actualSeqQlenAddr[bIdx - 1];
} else if (bIdx == 0) {
return actualSeqQlenAddr[0];
} else {
return 0;
}
} else {
if (constInfo.isActualLenDimsNull == 1) {
return constInfo.s1Size;
} else {
return actualSeqQlenAddr[bIdx];
}
}
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::GetAxisStartIdx(uint32_t bN2EndPrev,
uint32_t s1GEndPrev,
uint32_t s2EndPrev)
{
uint32_t bEndPrev = bN2EndPrev / constInfo.n2Size;
uint32_t actualSeqQPrev = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bEndPrev);
uint32_t s1GPrevBaseNum = actualSeqQPrev;
constInfo.bN2Start = bN2EndPrev;
constInfo.gS1Start = s1GEndPrev;
constInfo.s2Start = 0;
if (s1GEndPrev >= s1GPrevBaseNum - 1) { // 上个核把S1G处理完了
constInfo.gS1Start = 0;
constInfo.bN2Start++;
} else {
constInfo.gS1Start++;
}
}
template <typename CubeBlockType, typename VecBlockType> __aicore__ inline
void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitGlobalBuffer(
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *queryRope, __gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
__gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum,
__gm__ uint8_t *workspace, const SparseFlashAttentionTilingDataMla *__restrict tiling, TPipe *tPipe)
{
if (actualSeqLengthsQ != nullptr) {
actualSeqQlenAddr = (__gm__ int32_t *)actualSeqLengthsQ;
}
if (actualSeqLengths != nullptr) {
actualSeqKvlenAddr = (__gm__ int32_t *)actualSeqLengths;
}
vecBlock.InitGlobalBuffer(key, value, keyRope, sparseIndices, blockTable, softmaxMax, softmaxSum);
cubeBlock.InitCubeInput(key, keyRope, sparseIndices, blockTable, actualSeqLengthsQ, constInfo);
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void
SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitMMResBuf(__gm__ uint8_t *workspace)
{
uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize * sizeof(T);
uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512 * sizeof(T);
uint32_t mm2LeftSize = constInfo.s1BaseSize * constInfo.s2BaseSize * sizeof(Q_T);
uint32_t mm1RightSize = constInfo.s2BaseSize * 576 * sizeof(Q_T);
l1BufferManager.Init(pipe, 524288); // 512 * 1024
// 保存p结果的L1内存必须放在第一个L1 policy上,保证和vec申请的地址相同
l1RightBuffers.Init(l1BufferManager, mm1RightSize);
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
crossCoreSyncBufId++;
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
crossCoreSyncBufId++;
l1RightBuffers.Get().SetCrossCoreID(crossCoreSyncBufId, INVALID_CROSS_CORE_EVENT_ID);
crossCoreSyncBufId++;
if ASCEND_IS_AIC {
l1RightBuffers.Get().SetCrossCore();
l1RightBuffers.Get().SetCrossCore();
l1RightBuffers.Get().SetCrossCore();
}
ubBufferManager.Init(pipe, mm1ResultSize * 2 + mm2ResultSize);
bmm2Buffers.Init(ubBufferManager, mm2ResultSize);
bmm2Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
crossCoreSyncBufId++;
if ASCEND_IS_AIV {
bmm2Buffers.Get().SetCrossCore();
}
bmm1Buffers.Init(ubBufferManager, mm1ResultSize);
bmm1Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
crossCoreSyncBufId++;
bmm1Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
crossCoreSyncBufId++;
if ASCEND_IS_AIV {
bmm1Buffers.Get().SetCrossCore();
bmm1Buffers.Get().SetCrossCore();
}
uint32_t v0ResSize = constInfo.s2BaseSize * 576U * sizeof(Q_T);
int64_t totalOffset;
if constexpr (IS_SPLIT_G) {
totalOffset = v0ResSize * 3 * (aicIdx >> 1U);
} else {
totalOffset = v0ResSize * 3 * aicIdx;
}
gmBufferManager.Init(workspace + totalOffset);
v0ResGmBuffers.Init(gmBufferManager, v0ResSize);
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
crossCoreSyncBufId++;
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
crossCoreSyncBufId++;
v0ResGmBuffers.Get().SetCrossCoreID(INVALID_CROSS_CORE_EVENT_ID, crossCoreSyncBufId);
crossCoreSyncBufId++;
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitLocalBuffer()
{
vecBlock.InitLocalBuffer(pipe, constInfo);
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::ComputeConstexpr()
{
// 计算轴的乘积
usedCoreNum = sharedParams.usedCoreNum;
if ASCEND_IS_AIC {
constInfo.bSize = this->sharedParams.bSize;
constInfo.gSize = this->sharedParams.gSize;
constInfo.s1Size = this->sharedParams.s1Size;
constInfo.dSizeV = 512;
constInfo.needInit = this->sharedParams.needInit;
}
constInfo.n2Size = sharedParams.n2Size;
constInfo.s2Size = sharedParams.s2Size;
constInfo.dSize = sharedParams.dSize;
constInfo.dSizeVInput = sharedParams.dSizeVInput;
constInfo.dSizeRope = 64;
constInfo.dSizeNope = 512;
constInfo.tileSize = sharedParams.tileSize;
constInfo.sparseBlockCount = sharedParams.sparseBlockCount;
constInfo.sparseBlockSize = 1;
constInfo.cmpRatio = sharedParams.cmpRatio;
constInfo.oriWinLeft = sharedParams.oriWinLeft;
constInfo.oriWinRight = sharedParams.oriWinRight;
constInfo.sparseMode = sharedParams.oriMaskMode;
constInfo.s1S2 = constInfo.s1Size * constInfo.s2Size;
constInfo.gS1 = constInfo.gSize * constInfo.s1Size;
constInfo.n2G = constInfo.n2Size * constInfo.gSize;
constInfo.gD = constInfo.gSize * constInfo.dSize;
constInfo.n2GD = constInfo.n2Size * constInfo.gD;
constInfo.s1Dv = constInfo.s1Size * constInfo.dSizeV;
constInfo.s2Dv = constInfo.s2Size * constInfo.dSizeV;
constInfo.n2Dv = constInfo.n2Size * constInfo.dSizeV;
constInfo.gDv = constInfo.gSize * constInfo.dSizeV;
constInfo.gS1Dv = constInfo.gSize * constInfo.s1Dv;
constInfo.isActualLenDimsNull = sharedParams.isActualSeqLengthsNull;
constInfo.isActualLenDimsKVNull = sharedParams.isActualSeqLengthsKVNull;
constInfo.n2S2Dv = constInfo.n2Size * constInfo.s2Dv;
constInfo.n2GDv = constInfo.n2Size * constInfo.gDv;
constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv;
constInfo.n2GS1Dv = constInfo.n2Size * constInfo.gS1Dv;
constInfo.layoutType = sharedParams.layoutType;
constInfo.isActualLenDimsNull = sharedParams.isActualSeqLengthsNull;
constInfo.isActualLenDimsKVNull = sharedParams.isActualSeqLengthsKVNull;
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
// (BS)ND
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
if ASCEND_IS_AIV {
constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
}
} else if constexpr (LAYOUT_T == SFA_LAYOUT::BSND) {
// BSH/BSNGD
constInfo.s1BaseN2GDv = constInfo.s1BaseSize * constInfo.n2GDv;
constInfo.mm1Ka = constInfo.n2Size * constInfo.dSize;
if ASCEND_IS_AIV {
constInfo.attentionOutStride = (constInfo.n2G - constInfo.gSize) * constInfo.dSizeV * sizeof(OUTPUT_T);
}
}
if ASCEND_IS_AIV {
constInfo.softmaxScale = sharedParams.softmaxScale;
constInfo.oriBlockSize = sharedParams.oriBlockSize;
constInfo.oriMaxBlockNumPerBatch = sharedParams.oriMaxBlockNumPerBatch;
}
InitUniqueConstInfo();
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitUniqueConstInfo()
{
// bsize + 1-> bsize
this->constInfo.actualSeqLenSize = this->sharedParams.bSize;
this->constInfo.actualSeqLenKVSize = this->sharedParams.bSize;
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::Process()
{
// SyncAll Cube和Vector都需要调用
if (this->sharedParams.needInit) {
SyncAll<false>();
}
ProcessMainLoop();
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::ProcessMainLoop()
{
bool hasLoad = aicIdx < usedCoreNum;
if (!hasLoad) {
if ASCEND_IS_AIV {
if constexpr (IS_SPLIT_G) {
for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) {
CrossCoreSetFlag<SFA_SYNC_MODE0, PIPE_MTE3>(15);
CrossCoreWaitFlag<SFA_SYNC_MODE0, PIPE_MTE3>(15);
}
}
}
return;
}
// 适配分核左闭右开
uint32_t bIdx = constInfo.bN2End / constInfo.n2Size;
uint32_t actS1Size = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx);
uint32_t gS1max = actS1Size;
if (constInfo.gS1End + 1 < gS1max) {
/* constInfo.gS1End != gS1max时,gS1End需要往后加一格, bN2End不变 */
constInfo.gS1End = constInfo.gS1End + 1;
} else {
/* constInfo.gS1End == gS1max,bN2End需要往后加一格,bN2End变为0,以代表末尾 */
constInfo.bN2End = constInfo.bN2End + 1;
constInfo.gS1End = 0;
}
// 分核信息
uint32_t bN2StartIdx = constInfo.bN2Start;
uint32_t bN2EndIdx = constInfo.bN2End;
uint32_t gS1StartIdx = constInfo.gS1Start;
uint32_t nextGs1Idx = constInfo.gS1End;
uint32_t s2StartIdx = 0;
uint32_t s2EndIdx = 0;
uint32_t s2LoopLimit = 0;
if (nextGs1Idx != 0) {
bN2EndIdx++;
}
int64_t taskId = 0;
bool notLast = true;
RunInfo runInfo[3];
RunParamStr runParam;
int64_t multiCoreInnerIdx = 1;
for (int64_t bnIdx = bN2StartIdx; bnIdx < bN2EndIdx; bnIdx++) {
bool lastBN = (bnIdx == bN2EndIdx - 1);
runParam.boIdx = bnIdx;
runParam.n2oIdx = 0;
ComputeParamBatch<TEMPLATE_INTF_ARGS>(runParam, this->constInfo,
this->actualSeqQlenAddr, this->actualSeqKvlenAddr);
ComputeS1LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo, lastBN, nextGs1Idx, gS1StartIdx);
int64_t gS1LoopEnd = lastBN ? (runParam.gs1LoopEndIdx + PRELOAD_NUM) : runParam.gs1LoopEndIdx;
for (int64_t gS1Index = runParam.gs1LoopStartIdx; gS1Index < gS1LoopEnd; gS1Index++) {
bool notLastTwoLoop = true;
if (lastBN) {
int32_t extraGS1 = gS1Index - runParam.gs1LoopEndIdx;
switch (extraGS1) {
case 0:
notLastTwoLoop = false;
break;
case 1:
notLast = false;
notLastTwoLoop = false;
break;
default:
break;
}
}
if (notLastTwoLoop) {
this->ComputeAxisIdxByBnAndGs1(bnIdx, gS1Index, runParam);
bool s1NoNeedCalc = ComputeParamS1<TEMPLATE_INTF_ARGS>(
runParam, this->constInfo, gS1Index, this->actualSeqQlenAddr);
// s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分
bool s2NoNeedCalc =
ComputeS2LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo);
if (s1NoNeedCalc || s2NoNeedCalc) {
continue;
}
s2LoopLimit = runParam.s2LoopEndIdx - 1;
if constexpr (IS_SPLIT_G) {
maxS2LoopCnt -= (s2LoopLimit + 1);
}
} else {
s2LoopLimit = 0;
}
for (int64_t s2LoopCount = 0; s2LoopCount <= s2LoopLimit; ++s2LoopCount) {
if (notLastTwoLoop) {
RunInfo &runInfo1 = runInfo[taskId % 3];
this->SetRunInfo(runInfo1, runParam, taskId, s2LoopCount, s2LoopLimit, multiCoreInnerIdx);
if ASCEND_IS_AIC {
this->cubeBlock.IterateBmm1(this->bmm1Buffers.Get(),
this->l1RightBuffers.Get(), v0ResGmBuffers.Get(), runInfo1, this->constInfo);
} else {
this->vecBlock.ProcessVec0(this->l1RightBuffers.Get(), v0ResGmBuffers.Get(),
runInfo1, this->constInfo, 0);
}
} else {
if ASCEND_IS_AIV {
if constexpr (IS_SPLIT_G) {
if (maxS2LoopCnt > 0) {
maxS2LoopCnt--;
CrossCoreSetFlag<0, PIPE_MTE3>(15);
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
}
}
}
}
if (taskId > 0 && notLast) {
auto &runInfo2 = runInfo[(taskId + 2) % 3];
if ASCEND_IS_AIV {
this->vecBlock.ProcessVec1(this->l1RightBuffers.GetReused(),
this->bmm1Buffers.Get(), runInfo2, this->constInfo);
} else {
RunInfo &runInfo2 = runInfo[(taskId + 2) % 3];
this->cubeBlock.IterateBmm2(this->bmm2Buffers.Get(), this->l1RightBuffers,
this->l1RightBuffers.GetReused(), runInfo2, this->constInfo);
}
}
if (taskId > 1) {
if ASCEND_IS_AIV {
RunInfo &runInfo3 = runInfo[(taskId + 1) % 3];
this->vecBlock.ProcessVec2(this->bmm2Buffers.Get(), runInfo3, this->constInfo);
}
}
++taskId;
}
++multiCoreInnerIdx;
}
gS1StartIdx = 0;
}
if ASCEND_IS_AIV {
if constexpr (IS_SPLIT_G) {
for (int64_t loopCnt = 0; loopCnt < maxS2LoopCnt; loopCnt++) {
CrossCoreSetFlag<0, PIPE_MTE3>(15);
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
}
}
}
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::ComputeAxisIdxByBnAndGs1(
int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam)
{
// GS1合轴, 不切G, 只切S1
runParam.s1oIdx = gS1Index * runParam.qSNumInOneBlock;
if constexpr (IS_SPLIT_G) {
runParam.goIdx = (aicIdx % 2 == 0) ? 0 : 64; // N1=128场景,相邻cube核处理一个s1,第一个cube核承担0-63行g,第二个cube核承担后64行g
} else {
runParam.goIdx = 0;
}
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::SetRunInfo(
RunInfo &runInfo, RunParamStr &runParam, int64_t taskId, int64_t s2LoopCount,
int64_t s2LoopLimit, int64_t multiCoreInnerIdx)
{
if (s2LoopCount < runParam.kvLoopEndIdx) {
runInfo.s2StartIdx = runParam.s2LineStartIdx;
runInfo.s2EndIdx = runParam.s2LineEndIdx;
}
runInfo.s2LoopCount = s2LoopCount;
if (runInfo.multiCoreInnerIdx != multiCoreInnerIdx) {
runInfo.s1oIdx = runParam.s1oIdx;
runInfo.boIdx = runParam.boIdx;
runInfo.n2oIdx = runParam.n2oIdx;
runInfo.goIdx = runParam.goIdx;
runInfo.multiCoreInnerIdx = multiCoreInnerIdx;
runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1;
runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3;
}
runInfo.taskId = taskId;
runInfo.taskIdMod2 = taskId & 1;
runInfo.taskIdMod3 = taskId % 3;
runInfo.s2LoopLimit = s2LoopLimit;
runInfo.actualS1Size = runParam.actualS1Size;
runInfo.actualS2Size = runParam.actualS2Size;
runInfo.attentionOutOffset = runParam.attentionOutOffset;
runInfo.sOuterOffset = runParam.sOuterOffset;
runInfo.queryOffset = runParam.tensorQOffset;
this->ComputeBmm1Tail(runInfo, runParam);
InitUniqueRunInfo(runParam, runInfo);
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::InitUniqueRunInfo(
const RunParamStr &runParam, RunInfo &runInfo)
{
InitTaskParamByRun<TEMPLATE_INTF_ARGS>(runParam, runInfo);
}
template <typename CubeBlockType, typename VecBlockType>
__aicore__ inline void SparseFlashAttentionKernelMla<CubeBlockType, VecBlockType>::ComputeBmm1Tail(
RunInfo &runInfo, RunParamStr &runParam)
{
// ------------------------S1 Base Related---------------------------
runInfo.s1RealSize = runParam.s1RealSize;
runInfo.halfS1RealSize = runParam.halfS1RealSize;
runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize;
runInfo.mRealSize = runParam.mRealSize;
runInfo.halfMRealSize = runParam.halfMRealSize;
runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize;
runInfo.vec2S1BaseSize = runInfo.halfS1RealSize;
runInfo.vec2MBaseSize = runInfo.halfMRealSize;
// ------------------------S2 Base Related----------------------------
runInfo.s2RealSize = constInfo.s2BaseSize;
runInfo.s2AlignedSize = runInfo.s2RealSize;
int64_t curS2LoopCnt = runInfo.s2LoopCount;
if (runInfo.s2StartIdx + (curS2LoopCnt + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) {
runInfo.s2RealSize = runInfo.s2EndIdx - curS2LoopCnt * runInfo.s2RealSize - runInfo.s2StartIdx;
runInfo.s2AlignedSize = Align(runInfo.s2RealSize);
}
}
}
#endif // SPARSE_FLASH_ATTENTION_KERNEL_MLA_H

View File

@@ -0,0 +1,300 @@
/**
 * 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_kvcache.h
* \brief
*/
#ifndef SPARSE_FLASH_ATTENTION_KVCACHE_H
#define SPARSE_FLASH_ATTENTION_KVCACHE_H
#include "kernel_operator.h"
#include "kernel_operator_list_tensor_intf.h"
#include "sparse_flash_attention_common_arch35.h"
#include "util_regbase.h"
using namespace matmul;
using namespace regbaseutil;
using namespace AscendC;
using namespace AscendC::Impl::Detail;
static constexpr uint32_t sparseModeZero = 0;
static constexpr uint32_t sparseModeThree = 3;
TEMPLATE_INTF
__aicore__ inline void CalculateQueryOffset(RunParamStr& runParam,
const ConstInfo &constInfo, int32_t bIdx,
__gm__ int32_t* actualSeqQlenAddr)
{
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
runParam.qBOffset = (bIdx == 0) ? 0 : actualSeqQlenAddr[bIdx - 1] * constInfo.gSize * 512;
runParam.qRopeBOffset = (bIdx == 0) ? 0 : actualSeqQlenAddr[bIdx - 1] * constInfo.gSize * 64;
}
}
TEMPLATE_INTF
__aicore__ inline void GetSingleCoreParam(RunParamStr& runParam, const ConstInfo &constInfo,
__gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t * actualSeqKvlenAddr)
{
int32_t actualS1Size = 0;
int32_t actualS2Size = 0;
int32_t actualSeqMin = 1;
int32_t actualSeqKVMin = 1;
int32_t sIdx = runParam.boIdx;
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
// actual seq length first
if (actualSeqQlenAddr != nullptr) {
actualS1Size = (sIdx == 0) ? actualSeqQlenAddr[0] :
actualSeqQlenAddr[sIdx] - actualSeqQlenAddr[sIdx - 1];
} else {
actualS1Size = constInfo.s1Size;
}
} else {
actualS1Size = (actualSeqQlenAddr == nullptr) ? constInfo.s1Size :
actualSeqQlenAddr[sIdx];
}
if (constInfo.isActualLenDimsKVNull) {
actualS2Size = constInfo.s2Size;
} else {
if constexpr (isPa) {
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
actualS2Size = actualSeqKvlenAddr[sIdx];
} else {
actualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
}
} else {
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
actualS2Size = (sIdx == 0) ? actualSeqKvlenAddr[0] :
actualSeqKvlenAddr[sIdx] - actualSeqKvlenAddr[sIdx - 1];
} else {
actualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
}
}
}
runParam.actualS1Size = actualS1Size;
runParam.actualS2Size = actualS2Size;
if (constInfo.sparseMode == sparseModeZero) {
runParam.nextTokensPerBatch = MAX_PRE_NEXT_TOKENS;
} else {
runParam.nextTokensPerBatch = runParam.actualS2Size - runParam.actualS1Size;
}
runParam.preTokensPerBatch = runParam.actualS1Size;
}
TEMPLATE_INTF
__aicore__ inline void ComputeParamBatch(RunParamStr& runParam, const ConstInfo &constInfo,
__gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t *actualSeqKvlenAddr)
{
GetSingleCoreParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, actualSeqQlenAddr, actualSeqKvlenAddr);
}
TEMPLATE_INTF
__aicore__ inline void ComputeS1LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo, bool lastBN,
int64_t nextGs1Idx, int64_t gS1StartIdx)
{
runParam.qSNumInOneBlock = 1; // 不切G轴, 计算每个基本块可以拷贝多少行s
runParam.gs1LoopStartIdx = gS1StartIdx;
if (runParam.nextTokensPerBatch < 0) {
int64_t gs1LoopStartIdx = runParam.nextTokensPerBatch * (-1) / runParam.qSNumInOneBlock
* runParam.qSNumInOneBlock;
if (gs1LoopStartIdx > gS1StartIdx) {
runParam.gs1LoopStartIdx = gs1LoopStartIdx;
}
}
int32_t gs1LoopEndIdx = runParam.actualS1Size; // 不切G轴, 每次拷贝一行的topk,只算一行的qs
// 不是最后一个bn, 赋值souterBlockNum
if (!lastBN) {
runParam.gs1LoopEndIdx = gs1LoopEndIdx;
} else { // 最后一个bn, 从数组下一个元素取值
runParam.gs1LoopEndIdx = nextGs1Idx == 0 ? gs1LoopEndIdx : nextGs1Idx;
}
if (runParam.gs1LoopStartIdx > runParam.gs1LoopEndIdx) {
runParam.gs1LoopStartIdx = runParam.gs1LoopEndIdx;
}
}
TEMPLATE_INTF
__aicore__ inline void ComputeSouterParam(RunParamStr& runParam, const ConstInfo &constInfo,
uint32_t sOuterLoopIdx)
{
int64_t cubeSOuterOffset = sOuterLoopIdx * runParam.qSNumInOneBlock;
if (runParam.actualS1Size == 0) {
runParam.s1RealSize = 0;
runParam.mRealSize = 0;
} else {
runParam.s1RealSize = Min(runParam.qSNumInOneBlock, runParam.actualS1Size - cubeSOuterOffset);
runParam.mRealSize = runParam.s1RealSize * constInfo.gSize;
if constexpr (IS_SPLIT_G) {
runParam.mRealSize = runParam.mRealSize >> 1;
}
}
runParam.cubeMOuterOffset = cubeSOuterOffset * constInfo.gSize;
runParam.halfMRealSize = (runParam.mRealSize + 1) >> 1;
runParam.firstHalfMRealSize = runParam.halfMRealSize;
if (constInfo.subBlockIdx == 1) {
runParam.halfMRealSize = runParam.mRealSize - runParam.halfMRealSize;
runParam.mOuterOffset = runParam.cubeMOuterOffset + runParam.firstHalfMRealSize;
} else {
runParam.mOuterOffset = runParam.cubeMOuterOffset;
}
runParam.halfS1RealSize = (runParam.s1RealSize + 1) >> 1;
runParam.firstHalfS1RealSize = runParam.halfS1RealSize;
if (constInfo.subBlockIdx == 1) {
runParam.halfS1RealSize = runParam.s1RealSize - runParam.halfS1RealSize;
runParam.sOuterOffset = cubeSOuterOffset + runParam.halfMRealSize / constInfo.gSize;
} else {
runParam.sOuterOffset = cubeSOuterOffset;
}
runParam.cubeSOuterOffset = cubeSOuterOffset;
}
TEMPLATE_INTF
__aicore__ inline void LoopSOuterOffsetInit(RunParamStr& runParam, const ConstInfo &constInfo,
int32_t sIdx, __gm__ int32_t *cuSeqlensQAddr)
{
if ASCEND_IS_AIV {
int64_t seqOffset = 0;
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
seqOffset = sIdx == 0 ? 0 : cuSeqlensQAddr[sIdx - 1];
} else {
seqOffset = sIdx * constInfo.s1Size;
}
int64_t attentionOutSeqOffset = seqOffset * constInfo.n2GDv;
if constexpr (LAYOUT_T == SFA_LAYOUT::BSND || LAYOUT_T == SFA_LAYOUT::TND) {
runParam.attentionOutOffset = attentionOutSeqOffset +
runParam.sOuterOffset * constInfo.n2GDv + runParam.n2oIdx * constInfo.gDv +
runParam.goIdx * constInfo.dSizeV;
}
if (constInfo.subBlockIdx == 1) {
runParam.attentionOutOffset += runParam.firstHalfMRealSize * constInfo.dSizeV;
}
if (constInfo.returnSoftmaxLse) {
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
// [N2, T, G] (TND)
runParam.softmaxLseOffset = runParam.n2oIdx * constInfo.s1Size * constInfo.gSize +
(seqOffset + runParam.sOuterOffset) * constInfo.gSize;
} else {
// [B, N2, S1, G] (BSND)
runParam.softmaxLseOffset = sIdx * constInfo.n2Size * constInfo.s1Size * constInfo.gSize +
runParam.n2oIdx * constInfo.s1Size * constInfo.gSize +
runParam.sOuterOffset * constInfo.gSize;
}
uint32_t aicIdx = constInfo.aivIdx >> 1U;
if (IS_SPLIT_G && aicIdx % 2U != 0) {
runParam.softmaxLseOffset += 64; // splitG时,需要偏移64
}
if (constInfo.subBlockIdx == 1) {
runParam.softmaxLseOffset += runParam.firstHalfMRealSize;
}
}
} else {
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
runParam.tensorQOffset = runParam.qBOffset + runParam.cubeSOuterOffset * constInfo.n2GD +
runParam.n2oIdx * constInfo.gD + runParam.goIdx * constInfo.dSize;
runParam.tensorQRopeOffset = runParam.qRopeBOffset + runParam.cubeSOuterOffset * constInfo.n2GD +
runParam.n2oIdx * constInfo.gD + runParam.goIdx * constInfo.dSizeRope;
} else {
runParam.tensorQOffset = runParam.qBOffset + runParam.n2oIdx * constInfo.gS1D +
runParam.goIdx * constInfo.s1D + runParam.cubeSOuterOffset * constInfo.dSize;
runParam.tensorQRopeOffset = runParam.qRopeBOffset + runParam.n2oIdx * constInfo.gS1D +
runParam.goIdx * constInfo.s1D + runParam.cubeSOuterOffset * constInfo.dSizeRope;
}
}
}
TEMPLATE_INTF
__aicore__ inline bool ComputeParamS1(RunParamStr& runParam, const ConstInfo &constInfo,
uint32_t sOuterLoopIdx, __gm__ int32_t *cuSeqlensQAddr)
{
if (runParam.nextTokensPerBatch < 0) {
if (runParam.s1oIdx < (runParam.nextTokensPerBatch * (-1)) \
/ runParam.qSNumInOneBlock * runParam.qSNumInOneBlock) {
return true;
}
}
ComputeSouterParam<TEMPLATE_INTF_ARGS>(runParam, constInfo, sOuterLoopIdx);
LoopSOuterOffsetInit<TEMPLATE_INTF_ARGS>(runParam, constInfo, runParam.boIdx, cuSeqlensQAddr);
return false;
}
TEMPLATE_INTF
__aicore__ inline bool ComputeLastBN(RunParamStr& runParam, __gm__ int32_t *cuSeqlensQAddr)
{
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
// TND格式下 相邻Batch中当actualSeqQlen相等时则返回true
if (runParam.boIdx > 0 && ((runParam.boIdx == 0 && cuSeqlensQAddr[runParam.boIdx] == 0) ||
(cuSeqlensQAddr[runParam.boIdx] - cuSeqlensQAddr[runParam.boIdx - 1] == 0))) {
return true;
}
}
return false;
}
TEMPLATE_INTF
__aicore__ inline int64_t ClipSInnerTokenCube(int64_t sInnerToken, int64_t minValue, int64_t maxValue)
{
sInnerToken = sInnerToken > minValue ? sInnerToken : minValue;
sInnerToken = sInnerToken < maxValue ? sInnerToken : maxValue;
return sInnerToken;
}
TEMPLATE_INTF
__aicore__ inline bool ComputeS2LoopInfo(RunParamStr& runParam, const ConstInfo &constInfo)
{
if (runParam.actualS2Size == 0) {
runParam.kvLoopEndIdx = 0;
runParam.s2LoopEndIdx = 0;
return true;
}
uint32_t s2BaseSize = constInfo.s2BaseSize;
if (constInfo.sparseMode == sparseModeZero) {
runParam.s2LineStartIdx = 0;
runParam.s2LineEndIdx = Min(runParam.actualS2Size, constInfo.sparseBlockCount);
} else if (constInfo.sparseMode == sparseModeThree) {
runParam.s2LineStartIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(
runParam.cubeSOuterOffset - runParam.preTokensPerBatch, 0, runParam.actualS2Size);
runParam.s2LineEndIdx = ClipSInnerTokenCube<TEMPLATE_INTF_ARGS>(
runParam.cubeSOuterOffset + runParam.nextTokensPerBatch +
runParam.s1RealSize, 0, runParam.actualS2Size);
runParam.s2LineEndIdx = Min(runParam.s2LineEndIdx, constInfo.sparseBlockCount); // 当前LI输出的block size只可能是1
}
runParam.kvLoopEndIdx = (runParam.s2LineEndIdx + s2BaseSize - 1) / s2BaseSize;
runParam.s2LoopEndIdx = runParam.kvLoopEndIdx;
return false;
}
TEMPLATE_INTF
__aicore__ inline void InitTaskParamByRun(const RunParamStr& runParam, RunInfo &runInfo)
{
runInfo.boIdx = runParam.boIdx;
runInfo.preTokensPerBatch = runParam.preTokensPerBatch;
runInfo.nextTokensPerBatch = runParam.nextTokensPerBatch;
runInfo.actualS1Size = runParam.actualS1Size;
runInfo.actualS2Size = runParam.actualS2Size;
runInfo.softmaxLseOffset = runParam.softmaxLseOffset;
runInfo.qSNumInOneBlock = runParam.qSNumInOneBlock;
runInfo.kvLoopEndIdx = runParam.kvLoopEndIdx;
}
#endif // SPARSE_FLASH_ATTENTION_KVCACHE_H

View File

@@ -0,0 +1,387 @@
/**
 * 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_flash_attention_service_cube_mla.h
* \brief
*/
#ifndef SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
#define SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_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_flash_attention_common_arch35.h"
#include "util_regbase.h"
#if __has_include("../../common/op_kernel/offset_calculator.h")
#include "../../common/op_kernel/offset_calculator.h"
#else
#include "../common/offset_calculator.h"
#endif
#if __has_include("../../common/op_kernel/matmul.h")
#include "../../common/op_kernel/matmul.h"
#else
#include "../common/matmul.h"
#endif
#if __has_include("../../common/op_kernel/FixpipeOut.h")
#include "../../common/op_kernel/FixpipeOut.h"
#else
#include "../common/FixpipeOut.h"
#endif
#if __has_include("../../common/op_kernel/CopyInL1.h")
#include "../../common/op_kernel/CopyInL1.h"
#else
#include "../common/CopyInL1.h"
#endif
using namespace AscendC;
using namespace AscendC::Impl::Detail;
using namespace regbaseutil;
using namespace fa_base_matmul;
namespace BaseApi {
template <SFA_LAYOUT LAYOUT>
__aicore__ inline constexpr GmFormat GetQueryGmFormat()
{
if constexpr (LAYOUT == SFA_LAYOUT::BSND) {
return GmFormat::BSNGD;
} else {
return GmFormat::TNGD;
}
}
TEMPLATES_DEF
class SFAMatmulService {
public:
/* =================编译期常量的基本块信息================= */
static constexpr uint32_t s1BaseSize = 64;
static constexpr uint32_t s2BaseSize = 128;
static constexpr uint32_t dBaseSize = 576;
static constexpr uint32_t dBaseMatmulSize = 128;
__aicore__ inline SFAMatmulService() {};
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> &l1BuffMgr,
__gm__ uint8_t *query, __gm__ uint8_t *queryRope);
__aicore__ inline void InitCubeInput(__gm__ uint8_t *key, __gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, const ConstInfo& constInfo);
__aicore__ inline void IterateBmm1(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &output,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
ConstInfo &constInfo);
private:
__aicore__ inline void InitLocalBuffer(BufferManager<BufferType::L1> &l1BuffMgr);
__aicore__ inline void InitGmTensor(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
__aicore__ inline void IterateBmm1SFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
RunInfo &runInfo, ConstInfo &constInfo);
// --------------------Bmm2--------------------------
__aicore__ inline void IterateBmm2SFA(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
ConstInfo &constInfo);
TPipe *tPipe;
/* =====================GM变量==================== */
static constexpr GmFormat Q_FORMAT = GetQueryGmFormat<LAYOUT_T>();
FaGmTensor<Q_T, Q_FORMAT, int32_t> queryGm;
FaGmTensor<Q_T, Q_FORMAT, int32_t> queryRopeGm;
FaGmTensor<KV_T, GmFormat::PA_BnBsND> keyGm;
GlobalTensor<int32_t> blockTableGm;
FaGmTensor<KV_T, GmFormat::PA_BnBsND> curKvGm;
GlobalTensor<int32_t> cuSeqlensQGm;
/* =====================运行时变量==================== */
uint32_t kvCacheBlockSize = 0;
uint32_t maxBlockNumPerBatch = 0;
TEventID mte1ToMte2Id[3];
TEventID mte2ToMte1Id[3];
/* =====================LocalBuffer变量==================== */
BufferManager<BufferType::L0A> l0aBufferManager;
BufferManager<BufferType::L0B> l0bBufferManager;
BufferManager<BufferType::L0C> l0cBufferManager;
// D小于等于256 mm1左矩阵Q,GS1循环内左矩阵复用, GS1循环间开pingpong;D大于256使用单块Buffer,S1循环间驻留;fp32场景单块不驻留
BuffersPolicySingleBuffer<BufferType::L1> l1QBuffers;
// L0A
BuffersPolicyDB<BufferType::L0A> mmL0ABuffers;
// L0B
BuffersPolicyDB<BufferType::L0B> mmL0BBuffers;
// L0C
BuffersPolicyDB<BufferType::L0C> mmL0CBuffers;
};
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::InitCubeBlock(
TPipe *pipe, BufferManager<BufferType::L1> &l1BuffMgr, __gm__ uint8_t *query, __gm__ uint8_t *queryRope)
{
if ASCEND_IS_AIC {
tPipe = pipe;
this->queryGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)query);
this->queryRopeGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)queryRope);
InitLocalBuffer(l1BuffMgr);
}
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::InitCubeInput(__gm__ uint8_t *key, __gm__ uint8_t *keyRope,
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ,
const ConstInfo& constInfo)
{
if ASCEND_IS_AIC {
mte1ToMte2Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
mte1ToMte2Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
mte1ToMte2Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE1>();
mte2ToMte1Id[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
mte2ToMte1Id[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
mte2ToMte1Id[2] = GetTPipePtr()->AllocEventID<HardEvent::MTE1_MTE2>();
InitGmTensor(actualSeqLengthsQ, constInfo);
}
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void
SFAMatmulService<TEMPLATE_ARGS>::InitLocalBuffer(BufferManager<BufferType::L1> &l1BuffMgr)
{
constexpr uint32_t mm1LeftSize = s1BaseSize * dBaseSize * sizeof(Q_T);
l1QBuffers.Init(l1BuffMgr, mm1LeftSize);
// L0A B C 当前写死,能否通过基础api获取
l0aBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
l0bBufferManager.Init(tPipe, L0AB_SHARED_SIZE_64K);
l0cBufferManager.Init(tPipe, L0C_SHARED_SIZE_256K);
mmL0ABuffers.Init(l0aBufferManager, BUFFER_SIZE_16K); // db类型,填入数值是总大小的一半
mmL0BBuffers.Init(l0bBufferManager, BUFFER_SIZE_32K);
mmL0CBuffers.Init(l0cBufferManager, BUFFER_SIZE_128K);
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void
SFAMatmulService<TEMPLATE_ARGS>::InitGmTensor(__gm__ uint8_t *actualSeqLengthsQ, const ConstInfo& constInfo)
{
if constexpr (LAYOUT_T == SFA_LAYOUT::BSND) {
this->queryGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize,
constInfo.s1Size, constInfo.dSize);
this->queryRopeGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize,
constInfo.s1Size, constInfo.dSizeRope);
} else { // SFA_LAYOUT::TND
GlobalTensor<int32_t> actualSeqQLen;
actualSeqQLen.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ);
this->queryGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSize,
actualSeqQLen, constInfo.actualSeqLenSize);
this->queryRopeGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSizeRope,
actualSeqQLen, constInfo.actualSeqLenSize);
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::IterateBmm1(
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, RunInfo &runInfo,
ConstInfo &constInfo)
{
IterateBmm1SFA(outputBuf, inputRightBuf, v0ResGm, runInfo, constInfo);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::IterateBmm2(
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
ConstInfo &constInfo)
{
IterateBmm2SFA(outputBuf, inputLeftBuffers, inputRightBuf, runInfo, constInfo);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::IterateBmm1SFA(
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, RunInfo &runInfo,
ConstInfo &constInfo)
{
Buffer<BufferType::L1> inputLeftBuf;
// 左矩阵复用,S2的第一次循环加载左矩阵
// 加载左矩阵到L1, 全载
if (unlikely(runInfo.s2LoopCount == 0)) { // sOuter循环第一个基本块:搬运Q
inputLeftBuf = l1QBuffers.Get();
inputLeftBuf.Wait<HardEvent::MTE1_MTE2>(); // 占用L1A
LocalTensor<Q_T> inputLeftTensor = inputLeftBuf.GetTensor<Q_T>();
uint32_t s1Coord = runInfo.s1oIdx * runInfo.qSNumInOneBlock;
uint64_t queryGmOffset = this->queryGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx,
runInfo.goIdx, s1Coord, 0);
uint64_t queryRopeGmOffset = this->queryRopeGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx,
runInfo.goIdx, s1Coord, 0);
CopyToL1Nd2Nz<Q_T>(inputLeftTensor, this->queryGm.gmTensor[queryGmOffset],
runInfo.mRealSize, 512, 512); // 64 constInfo.dSize constInfo.mm1Ka
CopyToL1Nd2Nz<Q_T>(inputLeftTensor[Align16Func(runInfo.mRealSize) * 512],
this->queryRopeGm.gmTensor[queryRopeGmOffset], runInfo.mRealSize,
64, 64); // constInfo.dSize constInfo.mm1Ka
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
} else { // 非S2的第一次循环直接复用Q
inputLeftBuf = l1QBuffers.GetPre();
// 左矩阵复用时,sinner循环内不需要MTE2同步等待
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
}
inputRightBuf.WaitCrossCore();
SetFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
WaitFlag<HardEvent::MTE1_MTE2>(mte2ToMte1Id[runInfo.taskIdMod3]);
LocalTensor<Q_T> dst = inputRightBuf.GetTensor<Q_T>();
v0ResGm.WaitCrossCore();
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
CopyToL1Nd2Nz<Q_T>(dst, v0ResGmTensor, runInfo.s2RealSize, 576, 576);
SetFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
WaitFlag<HardEvent::MTE2_MTE1>(mte1ToMte2Id[runInfo.taskIdMod3]);
inputLeftBuf.Wait<HardEvent::MTE2_MTE1>(); // 等待L1A
Buffer<BufferType::L0C> mm1ResL0C = mmL0CBuffers.Get();
mm1ResL0C.Wait<HardEvent::FIX_M>(); // 占用
MMParam param = {static_cast<uint32_t>(runInfo.mRealSize), // singleM
static_cast<uint32_t>(runInfo.s2RealSize), // singleN
static_cast<uint32_t>(constInfo.dSizeNope + constInfo.dSizeRope), // singleK
0, // isLeftTranspose
1 // isRightTranspose
};
MatmulK<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>(
inputLeftBuf.GetTensor<Q_T>(), inputRightBuf.GetTensor<Q_T>(), // mm1B直接用tensor的数据
mmL0ABuffers, mmL0BBuffers,
mm1ResL0C.GetTensor<T>(),
param);
if (unlikely(runInfo.s2LoopCount == runInfo.s2LoopLimit)) {
inputLeftBuf.Set<HardEvent::MTE1_MTE2>(); // 释放L1A
}
mm1ResL0C.Set<HardEvent::M_FIX>(); // 通知
mm1ResL0C.Wait<HardEvent::M_FIX>(); // 等待L0C
outputBuf.WaitCrossCore();
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB
// L0C上的bmm1结果矩阵N方向的size大小; 同mmadParams.n; 为什么要8个元素对齐(32B对齐) // 128
fixpipeParams.nSize = Align8Func(runInfo.s2RealSize);
// 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小(必须为偶数) // 128
fixpipeParams.mSize = Align2Func(runInfo.mRealSize);
// L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔), 单位为16*sizeof(T)
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize);
// mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到, ND方案(S1*S2)时脏数据用mask剔除
fixpipeParams.dstStride = s2BaseSize;
fixpipeParams.dualDstCtl = 1; // 双目标模式,按M维度拆分,M / 2 * N写入每个UB, M必须为2的倍数
fixpipeParams.params.ndNum = 1;
fixpipeParams.params.srcNdStride = 0;
fixpipeParams.params.dstNdStride = 0;
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(), \
mm1ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
mm1ResL0C.Set<HardEvent::FIX_M>(); // 释放L0C
outputBuf.SetCrossCore();
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void SFAMatmulService<TEMPLATE_ARGS>::IterateBmm2SFA(
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
BuffersPolicy3buff<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
ConstInfo &constInfo)
{
inputRightBuf.WaitCrossCore();
Buffer<BufferType::L0C> mm2ResL0C = mmL0CBuffers.Get();
mm2ResL0C.Wait<HardEvent::FIX_M>(); // 占用
MMParam param = {static_cast<uint32_t>(runInfo.mRealSize), // singleM
static_cast<uint32_t>(constInfo.dSizeNope), // singleN
static_cast<uint32_t>(runInfo.s2RealSize), // singleK
0, // isLeftTranspose
0 // isRightTranspose
};
MatmulN<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>(
inputRightBuf.GetTensor<Q_T>(s2BaseSize * constInfo.dSizeNope), // 左矩阵P 来自rope位置
inputRightBuf.GetTensor<Q_T>(), // 右矩阵V nope
mmL0ABuffers,
mmL0BBuffers,
mm2ResL0C.GetTensor<T>(),
param);
inputRightBuf.SetCrossCore(); // bmm2才释放KV,在这里释放
mm2ResL0C.Set<HardEvent::M_FIX>(); // 通知
mm2ResL0C.Wait<HardEvent::M_FIX>(); // 等待
outputBuf.WaitCrossCore();
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C→UB;FixpipeParamsM300:L0C→UB
fixpipeParams.nSize = Align8Func(constInfo.dSizeNope); // L0C上的bmm1结果矩阵N方向的size大小, 分档计算且vector2中通过mask筛选出实际有效值
fixpipeParams.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可; L0C上的bmm1结果矩阵M方向的size大小;
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔)
fixpipeParams.dstStride = Align16Func(constInfo.dSizeNope);
fixpipeParams.dualDstCtl = 1;
fixpipeParams.params.ndNum = 1;
fixpipeParams.params.srcNdStride = 0;
fixpipeParams.params.dstNdStride = 0;
Fixpipe<T, T, PFA_CFG_ROW_MAJOR_UB>(outputBuf.template GetTensor<T>(),
mm2ResL0C.GetTensor<T>(), fixpipeParams); // 将matmul结果从L0C搬运到UB
mm2ResL0C.Set<HardEvent::FIX_M>(); // 释放
outputBuf.SetCrossCore();
}
TEMPLATES_DEF
class SFAMatmulServiceDummy {
public:
__aicore__ inline SFAMatmulServiceDummy() {};
__aicore__ inline void InitCubeBlock(TPipe *pipe,
BufferManager<BufferType::L1> &l1BuffMgr, __gm__ uint8_t *query, __gm__ uint8_t *queryRope) {}
__aicore__ inline void InitCubeInput(__gm__ uint8_t *key, __gm__ uint8_t *keyRope,
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *blockTable,
__gm__ uint8_t *actualSeqLengthsQ, const ConstInfo& constInfo) {}
__aicore__ inline void IterateBmm1(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf,
RunInfo &runInfo, ConstInfo &constInfo) {}
__aicore__ inline void IterateBmm2(Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &outputBuf,
BuffersPolicyDB<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputLeftBuffers,
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &inputRightBuf, RunInfo &runInfo,
ConstInfo &constInfo) {}
};
template <typename T>
struct CubeBlockTraits; // 声明
/* 生成CubeBlockTraits */
#define GEN_TRAIT_TYPE(name, ...) using name##_TRAITS = name;
#define GEN_TRAIT_CONST(name, type, ...) static constexpr type name##Traits = name;
#define DEFINE_CUBE_BLOCK_TRAITS(CUBE_BLOCK_CLASS) \
TEMPLATES_DEF_NO_DEFAULT \
struct CubeBlockTraits<CUBE_BLOCK_CLASS<TEMPLATE_ARGS>> { \
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TRAIT_TYPE) \
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TRAIT_CONST) \
}
DEFINE_CUBE_BLOCK_TRAITS(SFAMatmulService);
DEFINE_CUBE_BLOCK_TRAITS(SFAMatmulServiceDummy);
// /* 生成Arg Traits, kernel中只需要调用ARGS_TRAITS就可以获取所有CubeBlock中的模板参数 */
#define GEN_ARGS_TYPE(name, ...) using name = typename CubeBlockTraits<CubeBlockType>::name##_TRAITS;
#define GEN_ARGS_CONST(name, type, ...) static constexpr type name = CubeBlockTraits<CubeBlockType>::name##Traits;
#define ARGS_TRAITS \
CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARGS_TYPE) \
CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARGS_CONST)
}
#endif // SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H

View File

@@ -0,0 +1,879 @@
/**
 * 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_flash_attention_service_vector_mla.h
* \brief
*/
#ifndef SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
#define SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
#include "util_regbase.h"
#include "sparse_flash_attention_common_arch35.h"
#include "kernel_operator_list_tensor_intf.h"
#include "lib/matmul_intf.h"
#include "lib/matrix/matmul/tiling.h"
#if __has_include("../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h")
#include "../../common/op_kernel/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h"
#else
#include "../../common/arch35/vf/vf_mul_sel_softmaxflashv2_cast_nz_sfa.h"
#endif
#if __has_include("../../common/op_kernel/arch35/vf/vf_flashupdate_new.h")
#include "../../common/op_kernel/arch35/vf/vf_flashupdate_new.h"
#else
#include "../../common/arch35/vf/vf_flashupdate_new.h"
#endif
#if __has_include("../../common/op_kernel/buffers_policy.h")
#include "../../common/op_kernel/buffers_policy.h"
#else
#include "../../common/buffers_policy.h"
#endif
#if __has_include("../../common/op_kernel/buffer_manager.h")
#include "../../common/op_kernel/buffer_manager.h"
#else
#include "../../common/buffer_manager.h"
#endif
#if __has_include("../../common/op_kernel/buffer.h")
#include "../../common/op_kernel/buffer.h"
#else
#include "../../common/buffer.h"
#endif
using namespace AscendC;
using namespace FaVectorApi;
using namespace AscendC::Impl::Detail;
using namespace regbaseutil;
using namespace matmul;
namespace BaseApi {
TEMPLATES_DEF
class SFAVectorService {
public:
// BUFFER的字节数
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
/* =================编译期常量的基本块信息================= */
static constexpr uint32_t s1BaseSize = 64;
static constexpr uint32_t s2BaseSize = 128;
static constexpr uint32_t vec1Srcstride = (s1BaseSize >> 1) + 1;
static constexpr uint32_t dVTemplateType = 512;
static constexpr uint32_t dTemplateAlign64 = Align64Func(dVTemplateType);
static constexpr uint32_t dVTemplateTypeInput = 576;
static constexpr float R0 = 1.0f;
static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6;
// ==================== Functions ======================
__aicore__ inline SFAVectorService() {};
__aicore__ inline void InitVecBlock(TPipe *pipe, const SparseFlashAttentionTilingDataMla *__restrict tiling,
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx,
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths)
{
if ASCEND_IS_AIV {
tPipe = pipe;
tilingData = tiling;
if (actualSeqLengthsQ != nullptr) {
cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ);
}
if (actualSeqLengths != nullptr) {
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengths);
}
this->InitCubeVecSharedParams(sharedParams, aicIdx, subBlockIdx);
this->GetExtremeValue(this->negativeFloatScalar);
}
}
// 初始化LocalTensor
__aicore__ inline void InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo);
// 初始化attentionOutGM
__aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum, ConstInfo &constInfo);
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum);
__aicore__ inline void InitOutputSingleCore(ConstInfo &constInfo);
__aicore__ inline void ProcessVec0(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
const RunInfo &runInfo, ConstInfo &constInfo, int32_t startPos);
__aicore__ inline void ProcessVec1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf, RunInfo &runInfo,
ConstInfo &constInfo);
using mm2ResPos = Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH>;
__aicore__ inline void ProcessVec2(mm2ResPos &bmm2ResBuf, RunInfo &runInfo,
ConstInfo &constInfo);
private:
__aicore__ inline void ProcessSparseKv(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
const RunInfo &runInfo, ConstInfo &constInfo, int32_t startPos);
__aicore__ inline int64_t GetkeyOffset(int64_t s2Idx, const RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void GetRealCmpS2Idx(int64_t &token0Idx, int64_t &token1Idx, int64_t s2IdxInBase,
const RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void CopyInKvNotSparse(LocalTensor<KV_T> kvMergUb, int64_t v0Loop, int64_t dealRow,
int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline uint32_t CopyInKvSparse(LocalTensor<KV_T> kvInUb, int64_t startRow, int64_t token0Idx,
int64_t token1Idx, const RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void CopyOutKvUb2Gm(Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
LocalTensor<Q_T> kvOutUb, int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo,
ConstInfo &constInfo);
__aicore__ inline void CopyInSingleKv(LocalTensor<KV_T> kvInUb,
int64_t startRow, int64_t keyOffset, ConstInfo &constInfo);
/* VEC2_RES_T 表示bmm2ResUb当前的类型,VEC2_RES_T = Q_T那么不需要做Cast。另外,无效行场景当前默认需要做Cast */
using VEC2_RES_T = T;
template <typename VEC2_RES_T>
__aicore__ inline void Bmm2DataCopyOut(RunInfo &runInfo, ConstInfo &constInfo,
LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t vec2CalcSize = 0);
template <typename VEC2_RES_T>
__aicore__ inline void CopyOutAttentionOut(RunInfo &runInfo,
ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb,
int64_t vec2S1Idx, int64_t vec2CalcSize);
__aicore__ inline void SoftmaxInitBuffer();
__aicore__ inline void CopyFALseToGm(RunInfo &runInfo, ConstInfo &constInfo);
__aicore__ inline void InitCubeVecSharedParams(CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx);
__aicore__ inline void GetExtremeValue(T &negativeScalar);
__aicore__ inline void InitSinksBuffer(ConstInfo &constInfo);
TPipe *tPipe;
const SparseFlashAttentionTilingDataMla *__restrict tilingData;
GlobalTensor<OUTPUT_T> attentionOutGm;
GlobalTensor<KV_T> keyGm;
GlobalTensor<KV_T> keyRopeGm;
GlobalTensor<int32_t> sparseIndicesGm;
GlobalTensor<int32_t> blockTableGm;
GlobalTensor<int32_t> cuSeqlensQGm;
GlobalTensor<int32_t> actualSeqLengthsKVGm;
GlobalTensor<float> softmaxMaxGm;
GlobalTensor<float> softmaxSumGm;
LocalTensor<float> lseUb;
TBuf<> commonTBuf; // common的复用空间
TBuf<> sinksBuf;
TQue<QuePosition::VECOUT, 1> stage1OutQue[2]; // 2份表示可能存在pingpong
TBuf<> stage0OutBuf[2];
TBuf<> stage2OutBuf;
TEventID mte3ToVAttnOutId; // 存放MTE3_V的eventId, 用于V2 attentionOut拷出阶段的同步
TEventID vToMte3AttnOutId; // 存放V_MTE3的eventId, 用于V2 attentionOut拷出阶段的同步
TEventID mte3ToVLseOutId; // 存放MTE3_V的eventId, 用于V1 LSE拷出阶段的同步
TEventID vToMte3LseOutId; // 存放V_MTE3的eventId, 用于V1 LSE拷出阶段的同步
TBuf<> softmaxMaxBuf[2];
TBuf<> softmaxSumBuf[2];
TBuf<> softmaxExpBuf[2];
TBuf<> dequantScaleBuff;
TBuf<> lseBuf;
TEventID mte2ToV;
TEventID mte2ToMte3[2];
TEventID mte3ToMte2[2];
bool isSinks = false;
T negativeFloatScalar;
uint32_t maxBlockNumPerBatch;
uint32_t blockSize;
int64_t sparseCalSize;
int64_t sparseS2Start;
int64_t sparseS2End;
};
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void
SFAVectorService<TEMPLATE_ARGS>::GetRealCmpS2Idx(int64_t &token0Idx, int64_t &token1Idx,
int64_t s2IdxInBase, const RunInfo &runInfo, ConstInfo &constInfo)
{
int64_t topkBS1Idx = 0;
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
uint64_t actualSeqQPrefixSum = runInfo.boIdx == 0 ? 0 : cuSeqlensQGm.GetValue(runInfo.boIdx - 1);
topkBS1Idx += (actualSeqQPrefixSum + runInfo.s1oIdx) * constInfo.sparseBlockCount; // T, N2(1), K
} else {
topkBS1Idx += runInfo.boIdx * constInfo.s1Size * constInfo.sparseBlockCount +
runInfo.s1oIdx * constInfo.sparseBlockCount; // B, S1, N2(1), K
}
int64_t cmpS2LoopCnt = runInfo.s2LoopCount;
int64_t topkKIdx = s2IdxInBase + cmpS2LoopCnt * constInfo.s2BaseSize;
if (unlikely(topkKIdx >= constInfo.sparseBlockCount)) {
token0Idx = -1;
} else {
token0Idx = sparseIndicesGm.GetValue(topkBS1Idx + topkKIdx) + runInfo.s2StartIdx;
}
topkKIdx += 1;
if (unlikely((topkKIdx >= constInfo.sparseBlockCount) || (s2IdxInBase + 1 >= sparseS2End))) {
token1Idx = -1;
} else {
token1Idx = sparseIndicesGm.GetValue(topkBS1Idx + topkKIdx) + runInfo.s2StartIdx;
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
int64_t SFAVectorService<TEMPLATE_ARGS>::GetkeyOffset(
int64_t s2Idx, const RunInfo &runInfo, ConstInfo &constInfo)
{
if (s2Idx < 0) {
return -1;
}
int64_t realkeyOffset = 0;
if constexpr (isPa) {
int64_t blkTableIdx = s2Idx / blockSize;
int64_t blkTableOffset = s2Idx % blockSize;
realkeyOffset = blockTableGm.GetValue(runInfo.boIdx * maxBlockNumPerBatch + blkTableIdx) *
static_cast<int64_t>(blockSize) +
blkTableOffset; // BlockNum, BlockSize, N(1), D
} else {
if constexpr (LAYOUT_T == SFA_LAYOUT::BSND) {
realkeyOffset = (runInfo.boIdx * constInfo.s2Size + s2Idx); // BSN(1)D
} else if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
int64_t batchKvStart = (runInfo.boIdx == 0) ? 0 : actualSeqLengthsKVGm.GetValue(runInfo.boIdx - 1);
realkeyOffset = (batchKvStart + s2Idx);
}
}
return realkeyOffset;
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void
SFAVectorService<TEMPLATE_ARGS>::CopyInSingleKv(LocalTensor<KV_T> kvInUb, int64_t startRow,
int64_t keyOffset, ConstInfo &constInfo)
{
if (keyOffset < 0) {
return;
}
DataCopyExtParams intriParams;
intriParams.blockCount = 1;
intriParams.dstStride = 0;
intriParams.srcStride = 0;
DataCopyPadExtParams<KV_T> padParams;
// 当前仅支持COMBINE模式
uint32_t combineBytes = 512 * sizeof(KV_T);
intriParams.blockLen = combineBytes;
uint32_t combineDim = combineBytes / sizeof(KV_T);
uint32_t combineDimAlign = CeilAlign(combineBytes, BUFFER_SIZE_BYTE_32B) / sizeof(KV_T);
padParams.isPad = true;
padParams.leftPadding = 0;
padParams.rightPadding = combineDimAlign - combineDim;
padParams.paddingValue = 0;
DataCopyPad(kvInUb[startRow * 576], keyGm[keyOffset * 512], intriParams, padParams);
intriParams.blockLen = constInfo.sparseBlockSize * constInfo.dSizeRope *sizeof(KV_T);
intriParams.dstStride = 512 / BUFFER_SIZE_BYTE_32B;
DataCopyPad(kvInUb[startRow * 576 + 512], keyRopeGm[keyOffset * 64], intriParams, padParams); // combineDimAlign
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
uint32_t SFAVectorService<TEMPLATE_ARGS>::CopyInKvSparse(
LocalTensor<KV_T> kvInUb, int64_t startRow, int64_t token0Idx,
int64_t token1Idx, const RunInfo &runInfo, ConstInfo &constInfo)
{
int64_t keyOffset0 = GetkeyOffset(token0Idx, runInfo, constInfo);
int64_t keyOffset1 = GetkeyOffset(token1Idx, runInfo, constInfo);
if (unlikely(keyOffset0 < 0 && keyOffset1 < 0)) {
return 0;
}
int64_t blkTableSrcStride =
((keyOffset0 > keyOffset1 ? (keyOffset0 - keyOffset1) :
(keyOffset1 - keyOffset0)) - constInfo.sparseBlockSize);
int64_t keySrcStride = blkTableSrcStride * constInfo.dSizeNope * sizeof(KV_T);
int64_t keyRopeSrcStride = blkTableSrcStride * constInfo.dSizeRope * sizeof(KV_T);
if (unlikely(keyOffset1 < 0)) {
CopyInSingleKv(kvInUb, startRow, keyOffset0, constInfo);
} else if (keySrcStride >= INT32_MAX || keySrcStride < 0 || constInfo.sparseBlockSize > 1) {
// stride溢出、stride为负数、s2超长等异常场景,还原成2条搬运指令
CopyInSingleKv(kvInUb, startRow, keyOffset0, constInfo);
CopyInSingleKv(kvInUb, startRow + 1, keyOffset1, constInfo);
} else {
DataCopyExtParams intriParams;
intriParams.blockCount = (keyOffset0 >= 0) + (keyOffset1 >= 0);
intriParams.blockLen = constInfo.sparseBlockSize * constInfo.dSizeNope *sizeof(KV_T);
intriParams.dstStride = constInfo.dSizeRope * sizeof(KV_T) / BUFFER_SIZE_BYTE_32B;
intriParams.srcStride = keySrcStride;
DataCopyPadExtParams<KV_T> padParams;
int64_t keyOffset = keyOffset0 > -1 ? keyOffset0 : keyOffset1;
if (keyOffset1 > -1 && keyOffset1 < keyOffset0) {
keyOffset = keyOffset1;
}
DataCopyPad(kvInUb[startRow * 576], keyGm[keyOffset * constInfo.dSizeNope],
intriParams, padParams); // combineDimAlign
intriParams.blockLen = constInfo.sparseBlockSize * constInfo.dSizeRope *sizeof(KV_T);
intriParams.dstStride = constInfo.dSizeNope * sizeof(KV_T) / BUFFER_SIZE_BYTE_32B;
intriParams.srcStride = keyRopeSrcStride;
DataCopyPad(kvInUb[startRow * 576 + 512], keyRopeGm[keyOffset * constInfo.dSizeRope],
intriParams, padParams); // combineDimAlign
}
return (keyOffset0 > -1) + (keyOffset1 > -1);
}
// fp8->fp32
static constexpr MicroAPI::CastTrait castTraitFp8_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
// fp8->fp32
static constexpr MicroAPI::CastTrait castTraitFp8_2 = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::UNKNOWN,
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
// fp32->fp16
static constexpr MicroAPI::CastTrait castTraitFp8_3 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT};
// fp32->fp16
static constexpr MicroAPI::CastTrait castTraitFp8_4 = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT,
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_RINT};
template <typename Q_T, typename KV_T>
__simd_vf__ void CastScaleImpl(__ubuf__ float* ubDstAddr, __ubuf__ int8_t* ubSrcAddr, uint32_t dealRowCount)
{
MicroAPI::RegTensor<fp8_e8m0_t> vScale0;
MicroAPI::RegTensor<fp8_e8m0_t> vScale1;
MicroAPI::RegTensor<bfloat16_t> vScalebf16Res0;
MicroAPI::RegTensor<bfloat16_t> vScalebf16Res1;
MicroAPI::RegTensor<float> vScalefp32Res0;
MicroAPI::RegTensor<float> vScalefp32Res1;
__ubuf__ int8_t* ubScaleSrcAddrTemp = ubSrcAddr;
__ubuf__ float* ubDstAddrTmp = ubDstAddr;
MicroAPI::MaskReg bf16TypeMaskAll = MicroAPI::CreateMask<bfloat16_t, MicroAPI::MaskPattern::ALL>();
MicroAPI::MaskReg fp32MaskAll = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::ALL>();
for (uint16_t i = 0; i < static_cast<uint16_t>(dealRowCount); i++) {
// load scale
MicroAPI::LoadAlign<int8_t, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_UNPACK4_B8>(
(MicroAPI::RegTensor<int8_t>&)vScale0, ubScaleSrcAddrTemp, 640);
MicroAPI::Cast<bfloat16_t, fp8_e8m0_t, castTraitFp8_1>(vScalebf16Res0, vScale0, bf16TypeMaskAll);
MicroAPI::Cast<float, bfloat16_t, castTraitFp8_1>(vScalefp32Res0, vScalebf16Res0, fp32MaskAll);
MicroAPI::StoreAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
ubDstAddrTmp, vScalefp32Res0, 64, bf16TypeMaskAll);
}
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::CopyOutKvUb2Gm(
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, LocalTensor<Q_T> kvOutUb,
int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo)
{
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
DataCopy(v0ResGmTensor[s2StartIdx * 576], kvOutUb, dealRow * 576);
}
TEMPLATES_DEF_NO_DEFAULT
__aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo)
{
if constexpr (IS_SPLIT_G) {
uint32_t aicIdx = constInfo.aivIdx >> 1U;
uint32_t v0S2SizeFirstCore = CeilDiv(runInfo.s2RealSize, 2);
uint32_t v0S2SizeSecondCore = runInfo.s2RealSize - v0S2SizeFirstCore;
if (aicIdx % 2U == 0) {
if (GetSubBlockIdx() == 0) {
sparseCalSize = CeilDiv(v0S2SizeFirstCore, 2);
sparseS2Start = 0;
} else {
sparseCalSize = v0S2SizeFirstCore - CeilDiv(v0S2SizeFirstCore, 2);
sparseS2Start = CeilDiv(v0S2SizeFirstCore, 2);
}
} else {
if (GetSubBlockIdx() == 0) {
sparseCalSize = CeilDiv(v0S2SizeSecondCore, 2);
sparseS2Start = v0S2SizeFirstCore;
} else {
sparseCalSize = v0S2SizeSecondCore - CeilDiv(v0S2SizeSecondCore, 2);
sparseS2Start = v0S2SizeFirstCore + CeilDiv(v0S2SizeSecondCore, 2);
}
}
sparseS2End = sparseS2Start + sparseCalSize;
} else {
int64_t s2PerVecLoop = 2LL;
int64_t vecNum = 2LL;
int64_t s2Loops = CeilDiv(CeilDiv(runInfo.s2RealSize, vecNum), s2PerVecLoop);
sparseS2Start = GetSubBlockIdx() == 0 ? 0 : Min(s2Loops * s2PerVecLoop, runInfo.s2RealSize);
sparseS2End = GetSubBlockIdx() == 0 ? Min(s2Loops * s2PerVecLoop, runInfo.s2RealSize) : runInfo.s2RealSize;
sparseCalSize = sparseS2End - sparseS2Start;
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::ProcessVec0(
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
const RunInfo &runInfo, ConstInfo &constInfo, int32_t startPos)
{
blockSize = constInfo.oriBlockSize;
maxBlockNumPerBatch = constInfo.oriMaxBlockNumPerBatch;
CalSparseCalSize(runInfo, constInfo);
ProcessSparseKv(outputL1, v0ResGm, runInfo, constInfo, startPos);
if constexpr (IS_SPLIT_G) {
CrossCoreSetFlag<0, PIPE_MTE3>(15);
CrossCoreWaitFlag<0, PIPE_MTE3>(15);
}
outputL1.SetCrossCore();
v0ResGm.SetCrossCore();
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::ProcessSparseKv(
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
const RunInfo &runInfo, ConstInfo &constInfo, int32_t startPos)
{
if (sparseCalSize == 0) {
return;
}
bool meetEnd = false;
int64_t s2Start = sparseS2Start;
int64_t s2 = sparseS2Start;
int64_t token0Idx;
int64_t token1Idx; // 拷贝进入的两个token的index
// 处理一个s2的base块
uint32_t pingPong = 0;
while ((s2 < sparseS2End) && !meetEnd) { // 拷贝到s2End或者遇到-1
int64_t dealRow = 0;
// 1、copy kv in, gm ->ub
LocalTensor<Q_T> stage0OutUb = this->stage0OutBuf[pingPong].template Get<Q_T>();
WaitFlag<HardEvent::MTE3_MTE2>(mte3ToMte2[pingPong]);
while (dealRow < Min(16, sparseCalSize) && s2 < sparseS2End) { // 拷贝满16行或者遇到-1
GetRealCmpS2Idx(token0Idx, token1Idx, s2, runInfo, constInfo);
s2 += 2; // 每次搬运2行
if (token0Idx == -1 && token1Idx == -1) {
meetEnd = true;
break;
}
dealRow += CopyInKvSparse(stage0OutUb, dealRow, token0Idx, token1Idx, runInfo, constInfo);
if (token1Idx == -1) {
meetEnd = true;
break;
}
}
if (dealRow == 0) {
SetFlag<HardEvent::MTE3_MTE2>(mte3ToMte2[pingPong]);
pingPong ^= 1;
return;
}
SetFlag<HardEvent::MTE2_MTE3>(mte2ToMte3[pingPong]);
WaitFlag<HardEvent::MTE2_MTE3>(mte2ToMte3[pingPong]);
// 2、copy kv out, ub -> l1
CopyOutKvUb2Gm(v0ResGm, stage0OutUb, dealRow, s2Start, runInfo, constInfo);
SetFlag<HardEvent::MTE3_MTE2>(mte3ToMte2[pingPong]);
s2Start += dealRow;
pingPong ^= 1;
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::ProcessVec1(
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf, RunInfo &runInfo,
ConstInfo &constInfo)
{
bmm1ResBuf.WaitCrossCore();
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
LocalTensor<float> maxUb = this->softmaxMaxBuf[runInfo.multiCoreIdxMod2].template Get<float>();
LocalTensor<float> expUb = this->softmaxExpBuf[runInfo.taskIdMod2].template Get<T>();
int64_t stage1Offset = runInfo.taskIdMod2;
auto stage1CastTensor = this->stage1OutQue[stage1Offset].template AllocTensor<Q_T>();
LocalTensor<T> apiTmpBuffer = this->commonTBuf.template Get<T>();
LocalTensor<T> mmRes = bmm1ResBuf.template GetTensor<T>();
// loopCount = 0 但传入sinks时走update分支,maxUb通过sinks初始化,sumUb初始化为1.0
if (runInfo.s2LoopCount == 0) { // sink 丢失首token信息,sink会增加首token信息,维度是n1
if (likely(runInfo.s2RealSize == 128)) { // s2RealSize等于128分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::EQ_128_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
} else if (runInfo.s2RealSize <= 64) { // s2RealSize小于等于64分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_0_AND_LTE_64_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize,
runInfo.s2RealSize, // 实际的计算有效元素,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
} else if (runInfo.s2RealSize < 128) { // s2RealSize小于128分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, false, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_64_AND_LTE_128_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
}
} else {
if (likely(runInfo.s2RealSize == 128)) { // s2RealSize等于128分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::EQ_128_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
} else if (runInfo.s2RealSize <= 64) { // s2RealSize小于等于64分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_0_AND_LTE_64_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
} else if (runInfo.s2RealSize < 128) { // s2RealSize小于128分档, VF内常量化减少if判断
ProcessVec1Vf<T, Q_T, true, s1BaseSize, s2BaseSize, FaVectorApi::OriginNRange::GT_64_AND_LTE_128_SFA>(
stage1CastTensor, mmRes, sumUb, maxUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize, runInfo.s2RealSize,
static_cast<T>(constInfo.softmaxScale), negativeFloatScalar);
}
}
bmm1ResBuf.SetCrossCore();
// ===================DataCopy to L1 ====================
this->stage1OutQue[stage1Offset].template EnQue(stage1CastTensor);
this->stage1OutQue[stage1Offset].template DeQue<Q_T>();
LocalTensor<Q_T> mm2AL1Tensor = outputBuf.GetTensor<Q_T>(s2BaseSize * constInfo.dSizeV);
if (likely(runInfo.halfMRealSize != 0)) {
DataCopy(mm2AL1Tensor[constInfo.subBlockIdx * \
(BLOCK_BYTE / sizeof(Q_T)) * (runInfo.mRealSize - runInfo.halfMRealSize)],
stage1CastTensor, {s2BaseSize / 16, (uint16_t)runInfo.halfMRealSize,
(uint16_t)(vec1Srcstride - runInfo.halfMRealSize),
(uint16_t)(Align16Func(runInfo.mRealSize) - runInfo.halfMRealSize)});
}
this->stage1OutQue[stage1Offset].template FreeTensor(stage1CastTensor);
outputBuf.SetCrossCore();
if (runInfo.s2LoopCount != 0) {
SFAUpdateExpSumAndExpMax<T>(sumUb, maxUb, expUb, sumUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize);
}
if (constInfo.returnSoftmaxLse && runInfo.s2LoopCount == runInfo.s2LoopLimit) {
CopyFALseToGm(runInfo, constInfo);
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::ProcessVec2(
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm2ResBuf, RunInfo &runInfo,
ConstInfo &constInfo)
{
bmm2ResBuf.WaitCrossCore();
if (unlikely(runInfo.vec2MBaseSize == 0)) {
bmm2ResBuf.SetCrossCore();
return;
}
runInfo.vec2S1RealSize = runInfo.vec2S1BaseSize;
runInfo.vec2MRealSize = runInfo.vec2MBaseSize;
int64_t vec2CalcSize = runInfo.vec2MRealSize * dTemplateAlign64;
LocalTensor<T> vec2ResUb = this->stage2OutBuf.template Get<T>();
LocalTensor<T> mmRes = bmm2ResBuf.template GetTensor<T>();
WaitFlag<HardEvent::MTE3_V>(mte3ToVAttnOutId);
if (unlikely(runInfo.s2LoopCount == 0)) {
DataCopy(vec2ResUb, mmRes, vec2CalcSize);
} else {
LocalTensor<T> expUb = softmaxExpBuf[runInfo.taskIdMod2].template Get<T>();
if (runInfo.s2LoopCount < runInfo.s2LoopLimit) {
FlashUpdateNew<T, Q_T, OUTPUT_T, dTemplateAlign64, false, false>(
vec2ResUb, mmRes, vec2ResUb, expUb, expUb, runInfo.vec2MRealSize, dTemplateAlign64, 1.0, 1.0);
} else {
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
FlashUpdateLastNew<T, Q_T, OUTPUT_T, dTemplateAlign64, false, false>(
vec2ResUb, mmRes, vec2ResUb, expUb, expUb, sumUb, runInfo.vec2MRealSize, dTemplateAlign64, 1.0, 1.0);
}
}
bmm2ResBuf.SetCrossCore();
if (runInfo.s2LoopCount == runInfo.s2LoopLimit) {
if (unlikely(runInfo.s2LoopCount == 0)) {
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
LastDivNew<T, Q_T, OUTPUT_T, dTemplateAlign64, false>(
vec2ResUb, vec2ResUb, sumUb, runInfo.vec2MRealSize, dTemplateAlign64, 1.0);
}
this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, vec2CalcSize);
}
SetFlag<HardEvent::MTE3_V>(mte3ToVAttnOutId);
}
TEMPLATES_DEF_NO_DEFAULT
template <typename VEC2_RES_T>
__aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::Bmm2DataCopyOut (RunInfo &runInfo, ConstInfo &constInfo,
LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t vec2CalcSize)
{
LocalTensor<OUTPUT_T> attenOut;
int64_t dSizeAligned64 = (int64_t)dTemplateAlign64;
attenOut.SetAddr(vec2ResUb.address_);
Cast(attenOut, vec2ResUb, RoundMode::CAST_ROUND, vec2CalcSize);
SetFlag<HardEvent::V_MTE3>(vToMte3AttnOutId);
WaitFlag<HardEvent::V_MTE3>(vToMte3AttnOutId);
DataCopyExtParams dataCopyParams;
dataCopyParams.blockLen = constInfo.dSizeV * sizeof(OUTPUT_T);
dataCopyParams.srcStride = (dSizeAligned64 - constInfo.dSizeV) >> 4; // 以32B为单位偏移,bf16类型即偏移16个数,右移4
dataCopyParams.dstStride = constInfo.attentionOutStride;
dataCopyParams.blockCount = runInfo.vec2MRealSize;
DataCopyPad(this->attentionOutGm[runInfo.attentionOutOffset], attenOut, dataCopyParams);
}
TEMPLATES_DEF_NO_DEFAULT
template <typename VEC2_RES_T>
__aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::CopyOutAttentionOut(
RunInfo &runInfo, ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t vec2CalcSize)
{
this->Bmm2DataCopyOut(runInfo, constInfo, vec2ResUb, vec2S1Idx, vec2CalcSize);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
void SFAVectorService<TEMPLATE_ARGS>::InitOutputSingleCore(ConstInfo &constInfo)
{
uint32_t coreNum = GetBlockNum();
uint64_t totalOutputSize = 0;
// n2 = 1, n1 = gn2 = gSize
if constexpr (LAYOUT_T == SFA_LAYOUT::BSND) {
totalOutputSize = constInfo.bSize * constInfo.gSize * constInfo.s1Size * constInfo.dSizeV;
} else if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
totalOutputSize = constInfo.s1Size * constInfo.gSize * constInfo.dSizeV;
}
if (coreNum != 0) {
uint64_t singleCoreSize = (totalOutputSize + (CV_RATIO * coreNum) - 1) / (CV_RATIO * coreNum);
uint64_t tailSize = totalOutputSize - constInfo.aivIdx * singleCoreSize;
uint64_t singleInitOutputSize = tailSize < singleCoreSize ? tailSize : singleCoreSize;
if (constInfo.aivIdx * singleCoreSize < totalOutputSize && singleInitOutputSize > 0) {
matmul::InitOutput<OUTPUT_T>(
this->attentionOutGm[constInfo.aivIdx * singleCoreSize], singleInitOutputSize, 0);
}
}
if (constInfo.returnSoftmaxLse) {
uint64_t totalReturnSoftmaxSize = 0;
if constexpr (LAYOUT_T == SFA_LAYOUT::BSND) {
totalReturnSoftmaxSize = constInfo.bSize * constInfo.n2Size * constInfo.s1Size * constInfo.gSize;
} else if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
totalReturnSoftmaxSize = constInfo.n2Size * constInfo.s1Size * constInfo.gSize; // (N2,T1,G)
}
if (coreNum != 0 && totalReturnSoftmaxSize > 0) {
uint64_t singleCoreSoftmaxSize = (totalReturnSoftmaxSize + (CV_RATIO * coreNum) - 1) / (CV_RATIO * coreNum);
uint64_t tailSoftmaxSize = totalReturnSoftmaxSize - constInfo.aivIdx * singleCoreSoftmaxSize;
uint64_t singleInitSoftmaxSize = tailSoftmaxSize < singleCoreSoftmaxSize ?
tailSoftmaxSize : singleCoreSoftmaxSize;
if (constInfo.aivIdx * singleCoreSoftmaxSize < totalReturnSoftmaxSize && singleInitSoftmaxSize > 0) {
matmul::InitOutput<float>(this->softmaxSumGm[constInfo.aivIdx * singleCoreSoftmaxSize], singleInitSoftmaxSize, 0);
matmul::InitOutput<float>(this->softmaxMaxGm[constInfo.aivIdx * singleCoreSoftmaxSize], singleInitSoftmaxSize, 0);
}
}
}
SyncAll();
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
void SFAVectorService<TEMPLATE_ARGS>::CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum, ConstInfo &constInfo)
{
if ASCEND_IS_AIV {
this->attentionOutGm.SetGlobalBuffer((__gm__ OUTPUT_T *)attentionOut);
this->softmaxSumGm.SetGlobalBuffer((__gm__ float *)(softmaxSum));
this->softmaxMaxGm.SetGlobalBuffer((__gm__ float *)(softmaxMax));
if (constInfo.needInit == 1) {
InitOutputSingleCore(constInfo);
}
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
void SFAVectorService<TEMPLATE_ARGS>::InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum)
{
keyGm.SetGlobalBuffer((__gm__ KV_T *)(key));
if constexpr (isPa) {
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);;
}
sparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)sparseIndices);
keyRopeGm.SetGlobalBuffer((__gm__ KV_T *)(keyRope));
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::SoftmaxInitBuffer()
{
constexpr uint32_t softmaxBufSize = 256; // VF单次操作256Byte
tPipe->InitBuffer(softmaxSumBuf[0], softmaxBufSize);
tPipe->InitBuffer(softmaxSumBuf[1], softmaxBufSize);
tPipe->InitBuffer(softmaxMaxBuf[0], softmaxBufSize);
tPipe->InitBuffer(softmaxMaxBuf[1], softmaxBufSize);
tPipe->InitBuffer(softmaxExpBuf[0], softmaxBufSize);
tPipe->InitBuffer(softmaxExpBuf[1], softmaxBufSize);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
void SFAVectorService<TEMPLATE_ARGS>::CopyFALseToGm(RunInfo &runInfo, ConstInfo &constInfo)
{
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
LocalTensor<float> maxUb = this->softmaxMaxBuf[runInfo.multiCoreIdxMod2].template Get<float>();
size_t alignedSize = (sizeof(float) * runInfo.halfMRealSize + 31) / 32 * 32 / sizeof(float);
int64_t lseOffset = runInfo.softmaxLseOffset;
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = 1;
dataCopyParams.blockLen = sizeof(float) * runInfo.halfMRealSize;
dataCopyParams.srcStride = 0;
dataCopyParams.dstStride = 0;
// 拷贝 softmaxMaxUb -> GM
WaitFlag<HardEvent::MTE3_V>(mte3ToVLseOutId);
DataCopy(lseUb, maxUb, alignedSize);
SetFlag<HardEvent::V_MTE3>(vToMte3LseOutId);
WaitFlag<HardEvent::V_MTE3>(vToMte3LseOutId);
DataCopyPad(this->softmaxMaxGm[lseOffset], lseUb, dataCopyParams);
SetFlag<HardEvent::MTE3_V>(mte3ToVLseOutId);
// 拷贝 softmaxSumUb -> GM
WaitFlag<HardEvent::MTE3_V>(mte3ToVLseOutId);
DataCopy(lseUb, sumUb, alignedSize);
SetFlag<HardEvent::V_MTE3>(vToMte3LseOutId);
WaitFlag<HardEvent::V_MTE3>(vToMte3LseOutId);
DataCopyPad(this->softmaxSumGm[lseOffset], lseUb, dataCopyParams);
SetFlag<HardEvent::MTE3_V>(mte3ToVLseOutId);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::InitSinksBuffer(ConstInfo &constInfo)
{
LocalTensor<T> sinksUb = this->sinksBuf.template Get<T>();
const uint32_t maxN = constInfo.gSize; // N最大支持128, sink shape是[N]
DataCopyExtParams dataCopyParams;
dataCopyParams.blockCount = 1U;
dataCopyParams.blockLen = maxN * sizeof(T);
dataCopyParams.srcStride = 0U;
dataCopyParams.dstStride = 0U;
DataCopyPadExtParams<T> padParams;
DataCopyPad(sinksUb, this->sinksGm, dataCopyParams, padParams);
SetFlag<AscendC::HardEvent::MTE2_V>(SYNC_SINKS_BUF_FLAG);
WaitFlag<AscendC::HardEvent::MTE2_V>(SYNC_SINKS_BUF_FLAG);
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline
void SFAVectorService<TEMPLATE_ARGS>::InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo)
{
SoftmaxInitBuffer();
tPipe->InitBuffer(commonTBuf, 512); // commonTBuf内存申请512B
tPipe->InitBuffer(sinksBuf, 512); // sinksBuf内存申请512B
tPipe->InitBuffer(lseBuf, 512); // lseBuf内存申请512B
lseUb = this->lseBuf.template Get<float>();
tPipe->InitBuffer(stage0OutBuf[0], 576 * 16 * sizeof(KV_T));
tPipe->InitBuffer(stage0OutBuf[1], 576 * 16 * sizeof(KV_T));
tPipe->InitBuffer(stage1OutQue[0], 1, vec1Srcstride * s2BaseSize * sizeof(Q_T));
tPipe->InitBuffer(stage1OutQue[1], 1, vec1Srcstride * s2BaseSize * sizeof(Q_T));
tPipe->InitBuffer(stage2OutBuf, (s1BaseSize / CV_RATIO) * dTemplateAlign64 * sizeof(T));
mte3ToVAttnOutId = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
mte3ToVLseOutId = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
SetFlag<HardEvent::MTE3_V>(mte3ToVAttnOutId);
SetFlag<HardEvent::MTE3_V>(mte3ToVLseOutId);
vToMte3AttnOutId = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
vToMte3LseOutId = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
mte2ToV = GetTPipePtr()->AllocEventID<HardEvent::MTE2_V>();
mte3ToMte2[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
mte3ToMte2[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_MTE2>();
SetFlag<HardEvent::MTE3_MTE2>(mte3ToMte2[0]);
SetFlag<HardEvent::MTE3_MTE2>(mte3ToMte2[1]);
mte2ToMte3[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE3>();
mte2ToMte3[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE2_MTE3>();
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::InitCubeVecSharedParams(
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx)
{
// TODO参数整改
auto &sparseAttnSharedkvBaseParams = this->tilingData->baseParams;
sharedParams.bSize = sparseAttnSharedkvBaseParams.batchSize;
sharedParams.n2Size = 1;
sharedParams.gSize = sparseAttnSharedkvBaseParams.nNumOfQInOneGroup;
sharedParams.s1Size = sparseAttnSharedkvBaseParams.qSeqSize;
sharedParams.s2Size = sparseAttnSharedkvBaseParams.seqSize;
sharedParams.sparseBlockCount = sparseAttnSharedkvBaseParams.sparseBlockCount;
sharedParams.cmpRatio = 1; // 走sparse, 但不压缩
sharedParams.oriMaskMode = sparseAttnSharedkvBaseParams.sparseMode;
sharedParams.oriWinLeft = -1;
sharedParams.oriWinRight = 0;
sharedParams.layoutType = sparseAttnSharedkvBaseParams.outputLayout;
sharedParams.dSizeRope = 64;
sharedParams.softmaxScale = sparseAttnSharedkvBaseParams.scaleValue;
sharedParams.dSize = 512;
sharedParams.dSizeVInput = 512;
sharedParams.usedCoreNum = this->tilingData->singleCoreParams.usedCoreNum;
// pageAttention, rope在C侧搬运时使用
if constexpr (isPa) {
sharedParams.oriBlockSize = sparseAttnSharedkvBaseParams.blockSize;
sharedParams.oriMaxBlockNumPerBatch = sparseAttnSharedkvBaseParams.maxBlockNumPerBatch;
}
// actQ->TND, actKV pa场景任意layout均有
sharedParams.isActualSeqLengthsNull = sparseAttnSharedkvBaseParams.isActualLenDimsNull;
sharedParams.isActualSeqLengthsKVNull = sparseAttnSharedkvBaseParams.isActualLenDimsKVNull;
sharedParams.returnSoftmaxLse = sparseAttnSharedkvBaseParams.returnSoftmaxLse;
sharedParams.needInit = 0;
for (uint32_t bIdx = 0; bIdx < sharedParams.bSize; bIdx++) {
int64_t s2Size;
if constexpr (KV_LAYOUT_T == SFA_LAYOUT::TND) {
s2Size = bIdx == 0 ? actualSeqLengthsKVGm.GetValue(bIdx) : \
actualSeqLengthsKVGm.GetValue(bIdx) - actualSeqLengthsKVGm.GetValue(bIdx - 1);
} else {
if (sharedParams.isActualSeqLengthsKVNull) {
s2Size = sharedParams.s2Size;
} else {
s2Size = actualSeqLengthsKVGm.GetValue(bIdx);
}
}
int64_t s1Size;
if constexpr (LAYOUT_T == SFA_LAYOUT::TND) {
s1Size = bIdx == 0 ? cuSeqlensQGm.GetValue(bIdx) : \
cuSeqlensQGm.GetValue(bIdx) - cuSeqlensQGm.GetValue(bIdx - 1);
} else {
if (sharedParams.isActualSeqLengthsNull) {
s1Size = sharedParams.s1Size;
} else {
s1Size = cuSeqlensQGm.GetValue(bIdx);
}
}
if (s1Size > s2Size || (LAYOUT_T == SFA_LAYOUT::BSND && s1Size < sharedParams.s1Size)) {
sharedParams.needInit = 1;
break;
}
}
if ASCEND_IS_AIV {
if (subBlockIdx == 0) {
auto tempTilingSSbuf = reinterpret_cast<__ssbuf__ uint32_t*>(0); // 从ssbuf的0地址开始拷贝
auto tempTiling = reinterpret_cast<uint32_t *>(&sharedParams);
#pragma unroll
for (int i = 0; i < sizeof(CVSharedParams) / sizeof(uint32_t); ++i, ++tempTilingSSbuf, ++tempTiling) {
*tempTilingSSbuf = *tempTiling;
}
CrossCoreSetFlag<SYNC_MODE, PIPE_S>(15);
}
}
}
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void SFAVectorService<TEMPLATE_ARGS>::GetExtremeValue(
T &negativeScalar)
{
uint32_t tmp1 = NEGATIVE_MIN_VALUE_FP32;
negativeScalar = *((float *)&tmp1);
}
TEMPLATES_DEF class SFAVectorServiceDummy {
public:
__aicore__ inline SFAVectorServiceDummy() {};
__aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum, ConstInfo &constInfo) {}
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value,
__gm__ uint8_t *keyRope, __gm__ uint8_t *sparseIndices,
__gm__ uint8_t *blockTable, __gm__ uint8_t *softmaxMax,
__gm__ uint8_t *softmaxSum) {}
__aicore__ inline void InitVecBlock(TPipe *pipe, const SparseFlashAttentionTilingDataMla *__restrict tiling,
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx, __gm__ uint8_t *actualSeqLengthsQ,
__gm__ uint8_t *actualSeqLengths) {};
__aicore__ inline void InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo) {}
__aicore__ inline void ProcessVec1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputBuf,
Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> &bmm1ResBuf, RunInfo &runInfo,
ConstInfo &constInfo) {}
using mm2ResPos = Buffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH>;
__aicore__ inline void ProcessVec2(mm2ResPos &bmm2ResBuf, RunInfo &runInfo,
ConstInfo &constInfo) {}
};
}
#endif // SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H

View File

@@ -0,0 +1,265 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file util_regbase.h
* \brief
*/
#ifndef UTIL_REGBASE_H
#define UTIL_REGBASE_H
#include "util.h"
using AscendC::TQue;
using AscendC::QuePosition;
namespace regbaseutil {
constexpr int64_t MAX_PRE_NEXT_TOKENS = 0x7FFFFFFF;
#define COMMON_RUN_PARAM \
int64_t boIdx; \
int64_t s1oIdx; \
int64_t n2oIdx; \
int64_t goIdx; \
int64_t s2LoopEndIdx; /* S2方向的循环控制信息 souter层确定 */ \
int64_t s2LineStartIdx = 0; /* S2方向按行的起始位置 */ \
int64_t s2LineEndIdx; /* S2方向按行的结束位置 */ \
int64_t s2CmpLineEndIdx; \
/* cube视角的sOuter,在SAMEAB场景中cubeSOuterSize为两倍的 halfS1RealSize souter层确定 */ \
uint32_t s1RealSize; \
uint32_t halfS1RealSize; \
uint32_t firstHalfS1RealSize; \
uint32_t mRealSize; \
uint32_t halfMRealSize; \
uint32_t firstHalfMRealSize; \
int64_t attentionOutOffset; /* attentionOut的offset souter层确定 */ \
int32_t actualS1Size; /* Q的actualSeqLength */ \
int32_t actualS2Size; /* KV的actualSeqLength */ \
int64_t tensorQOffset; \
int64_t tensorQRopeOffset; \
int64_t qBOffset; \
int64_t qRopeBOffset;
struct RunParamStr { // 分核与切块需要使用到参数
COMMON_RUN_PARAM;
/* 推理新增 */
int64_t gs1LoopStartIdx;
int64_t gs1LoopEndIdx;
// BN循环生产的数据
int64_t preTokensPerBatch = MAX_PRE_NEXT_TOKENS; // 左上顶点的pretoken
int64_t nextTokensPerBatch = MAX_PRE_NEXT_TOKENS; // 左上顶点的nexttoken
// NBS1循环生产的数据
int64_t sOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定
int64_t cubeSOuterOffset; // 单个S内 souter的 souterIdx * halfS1RealSize souter层确定
int64_t mOuterOffset;
int64_t cubeMOuterOffset;
// lse 输出offset
int64_t softmaxLseOffset; // souter层确定
int64_t qSNumInOneBlock;
int64_t kvLoopEndIdx;
};
#define COMMON_RUN_INFO \
int64_t s2StartIdx; /* s2的起始位置,sparse场景下可能不是0 */ \
int64_t s2EndIdx; \
int64_t s2LoopCount; /* s2循环当前的循环index */ \
int64_t s2LoopLimit; \
int64_t s1oIdx = 0; /* s1轴的index */ \
int64_t loop = 0; /* for v0 perload loop */ \
int64_t boIdx = 0; /* b轴的index */ \
int64_t n2oIdx = 0; /* n2轴的index */ \
int64_t goIdx = 0; /* g轴的index */ \
int32_t s1RealSize; \
int32_t halfS1RealSize; /* vector侧实际的s1基本块大小,如果Cube基本块=128,那么halfS1RealSize=64 */ \
int32_t firstHalfS1RealSize; /* 当s1RealSize不是2的整数倍时,v0比v1少计算一行,计算subblock偏移的时候需要使用v0的s1 size */ \
int32_t mRealSize; \
int32_t halfMRealSize; \
int32_t firstHalfMRealSize; \
int32_t s2RealSize; /* s2方向基本块的真实长度 */ \
int64_t s2AlignedSize; /* s2方向基本块对齐到16之后的长度 */ \
int32_t vec2S1BaseSize; /* vector2侧开循环之后,经过切分的S1大小,例如把64切分成两份32 */ \
int32_t vec2S1RealSize; /* vector2侧开循环之后,经过切分的S1的尾块大小,例如把63切分成两份32和31,第二份的实际大小是31 */ \
int32_t vec2MBaseSize; \
int32_t vec2MRealSize; \
int64_t taskId; \
int64_t multiCoreInnerIdx = 0; \
int64_t attentionOutOffset; \
int32_t actualS1Size; /* 非TND场景=总s1Size, Tnd场景下当前batch对应的s1 */ \
int32_t actualS2Size; /* 非TND场景=总s2Size, Tnd场景下当前batch对应的s2 */ \
int64_t preTokensPerBatch; /* vector2 左上顶点的pretoken */ \
int64_t nextTokensPerBatch; /* vector2 左上顶点的nexttoken */ \
uint8_t taskIdMod2; \
uint8_t taskIdMod3; \
uint8_t multiCoreIdxMod2 = 0; \
uint8_t multiCoreIdxMod3 = 0; \
int64_t sOuterOffset; \
int64_t mOuterOffset; \
int64_t queryOffset; \
int64_t queryRopeOffset
struct RunInfo {
COMMON_RUN_INFO;
// 推理新增
// lse 输出offset
int64_t softmaxLseOffset;
int64_t qSNumInOneBlock;
int64_t kvLoopEndIdx;
};
#define COMMON_CONST_INFO \
/* 全局的基本块信息 */ \
uint32_t bSize; \
uint32_t needInit; \
uint32_t s1BaseSize; \
uint32_t s2BaseSize; \
int64_t dSize; /* query d 512 */ \
int64_t dSizeV; /* key d 512 */ \
int64_t dSizeVInput; /* key inpue d 656 = rope + nope + scale + pad */ \
int64_t dSizeNope; /* key nope d 448 */ \
int64_t dSizeRope; /* key rope d 64 */ \
int64_t tileSize; /* 64 */ \
int64_t sparseMode; \
int64_t gSize; /* g轴的大小 */ \
int64_t n2Size; \
int64_t s1Size; /* s1总大小 */ \
int64_t s2Size; /* s2总大小 */ \
/* 轴的乘积 */ \
int64_t s1D; \
int64_t gS1D; \
int64_t n2GS1D; \
int64_t s2D; \
int64_t n2S2D; \
int64_t s1Dv; \
int64_t gS1Dv; \
int64_t n2GS1Dv; \
int64_t s2Dv; \
int64_t n2S2Dv; \
int64_t s1S2; \
int64_t gS1; \
int64_t gD; \
int64_t n2D; \
int64_t bN2D; \
int64_t gDv; \
int64_t n2Dv; \
int64_t bN2Dv; \
int64_t n2G; \
int64_t n2GD; \
int64_t bN2GD; \
int64_t n2GDv; \
int64_t bN2GDv; \
int64_t gS2; \
int64_t s1Dr; \
int64_t gS1Dr; \
int64_t n2GS1Dr; \
int64_t s2Dr; \
int64_t n2S2Dr; \
int64_t gDr; \
int64_t n2Dr; \
int64_t bN2Dr; \
int64_t n2GDr; \
int64_t bN2GDr; \
int32_t s2BaseN2D; \
int32_t s1BaseN2GD; \
int64_t s2BaseBN2D; \
int64_t s1BaseBN2GD; \
int32_t s1BaseD; \
int32_t s2BaseD; \
int64_t s2BaseN2Dv; \
int64_t s2BaseBN2Dv; \
int64_t s1BaseN2GDv; \
int64_t s1BaseBN2GDv; \
int32_t s1BaseDv; \
int32_t s2BaseDv; \
bool returnSoftmaxLse; \
/* matmul跳读参数 */ \
int64_t mm1Ka; \
/* dq 或者attentionOut的Stride */ \
int64_t attentionOutStride; \
uint32_t aivIdx; \
uint8_t layoutType; \
uint8_t subBlockIdx;\
/* 分核相关 */ \
uint32_t s2Start; \
uint32_t s2End; \
uint32_t bN2Start; \
uint32_t bN2End; \
uint32_t gS1Start; \
uint32_t gS1End
#define INFER_CONST_INFO \
/* 推理 */ \
bool isActualLenDimsNull; /* 判断是否有actualseq */ \
bool isActualLenDimsKVNull; /* 判断是否有actualseq_kv */ \
bool isSoftmaxLseEnable; \
bool rsvd1; \
uint32_t sparseBlockCount; \
uint32_t actualSeqLenSize; /* 用户输入的actualseq的长度 */ \
uint32_t actualSeqLenKVSize; /* 用户输入的actualseq_kv的长度 */ \
/* service mm1 mm2 pageAttention */ \
uint32_t oriBlockSize; \
uint32_t cmpBlockSize; \
uint32_t paLayoutType; \
uint32_t oriMaxBlockNumPerBatch; \
uint32_t cmpMaxBlockNumPerBatch; \
int32_t oriWinLeft; \
int32_t oriWinRight; \
uint32_t sparseBlockSize; \
uint32_t cmpRatio; \
float softmaxScale
#define CV_SHARED_PARAMS \
/* base params */ \
uint32_t s1BaseSize; \
uint32_t s2BaseSize; \
uint32_t bSize; \
uint32_t n2Size; \
uint32_t gSize; \
uint32_t s1Size; \
uint32_t s2Size; \
uint32_t dSize : 10; \
int64_t dSizeVInput : 12; \
uint32_t needInit : 4; \
uint32_t layoutType : 4; \
uint32_t isActualSeqLengthsNull : 1; \
uint32_t isActualSeqLengthsKVNull : 1; \
uint32_t sparseBlockCount; \
float softmaxScale; \
uint32_t cmpRatio : 9; \
uint32_t dSizeRope : 11; \
uint32_t oriMaskMode : 6; \
uint32_t cmpMaskMode : 6; \
int32_t oriWinLeft; \
int32_t oriWinRight; \
uint32_t tileSize : 8; \
/* pa params */ \
uint32_t oriBlockSize : 12; \
uint32_t cmpBlockSize : 12; \
uint32_t oriMaxBlockNumPerBatch; \
uint32_t cmpMaxBlockNumPerBatch; \
uint32_t usedCoreNum; \
bool returnSoftmaxLse
struct ConstInfo {
COMMON_CONST_INFO;
INFER_CONST_INFO;
};
/* only support b32 or b64 */
struct CVSharedParams {
CV_SHARED_PARAMS;
};
}
#endif // UTIL_REGBASE_H