@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user