@@ -0,0 +1,130 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_flash_attention_common_arch35.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
#include <type_traits>
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
|
||||
#if __has_include("../../sparse_flash_attention/arch35/common/util_regbase.h")
|
||||
#include "../../sparse_flash_attention/arch35/common/util_regbase.h"
|
||||
#else
|
||||
#include "../../../sparse_flash_attention/op_kernel/arch35/common/util_regbase.h"
|
||||
#endif
|
||||
|
||||
#if __has_include("../../common/op_kernel/buffer.h")
|
||||
#include "../../common/op_kernel/buffer.h"
|
||||
#else
|
||||
#include "../../common/buffer.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/buffers_policy.h")
|
||||
#include "../../common/op_kernel/buffers_policy.h"
|
||||
#else
|
||||
#include "../../common/buffers_policy.h"
|
||||
#endif
|
||||
|
||||
constexpr uint64_t BLOCK_BYTE = 32;
|
||||
constexpr uint32_t NEGATIVE_MIN_VALUE_FP32 = 0xFF7FFFFF;
|
||||
|
||||
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 L0AB_SHARED_SIZE_64K = 65536; // 65536表示64*1024
|
||||
constexpr uint32_t L0C_SHARED_SIZE_256K = 262144; // 262144表示256 * 1024
|
||||
|
||||
constexpr uint32_t CV_RATIO = 2;
|
||||
constexpr uint64_t SYNC_MODE = 4;
|
||||
|
||||
static constexpr uint32_t QSFA_SYNC_MODE0 = 0;
|
||||
|
||||
enum class QSFA_LAYOUT {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2,
|
||||
};
|
||||
|
||||
enum class QSFATemplateMode {
|
||||
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, QSFA_LAYOUT LAYOUT_T, \
|
||||
QSFA_LAYOUT KV_LAYOUT_T, QSFATemplateMode 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 QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(X) \
|
||||
X(Q_T) \
|
||||
X(KV_T) \
|
||||
X(T) \
|
||||
X(OUTPUT_T) \
|
||||
|
||||
#define QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(X) \
|
||||
X(isFd, bool, false) \
|
||||
X(isPa, bool, true) \
|
||||
X(LAYOUT_T, QSFA_LAYOUT, QSFA_LAYOUT::BSND) \
|
||||
X(KV_LAYOUT_T, QSFA_LAYOUT, QSFA_LAYOUT::PA_BSND) \
|
||||
X(TEMPLATE_MODE, QSFATemplateMode, QSFATemplateMode::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 <QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TYPE_PARAM) \
|
||||
QSFA_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 <QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TEMPLATE_TYPE_NODEF) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TEMPLATE_CONST_NODEF) bool end>
|
||||
|
||||
/* 3. 生成有默认值的Args */
|
||||
#define GEN_ARG_NAME(name, ...) name,
|
||||
#define TEMPLATE_ARGS \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARG_NAME) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARG_NAME) end
|
||||
|
||||
#endif //KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_ARCH35_H
|
||||
@@ -0,0 +1,707 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_flash_attention_kernel_mla.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
|
||||
#define KV_QUANT_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 "kv_quant_sparse_flash_attention_service_cube_mla.h"
|
||||
#include "kv_quant_sparse_flash_attention_service_vector_mla.h"
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
#include "kv_quant_sparse_flash_attention_kvcache.h"
|
||||
#if __has_include("../../common/op_kernel/CopyInL1.h")
|
||||
#include "../../common/op_kernel/CopyInL1.h"
|
||||
#else
|
||||
#include "../common/CopyInL1.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
|
||||
|
||||
using matmul::MatmulType;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace BaseApi {
|
||||
template <typename CubeBlockType, typename VecBlockType> class KvQuantSparseFlashAttentionMla {
|
||||
public:
|
||||
ARGS_TRAITS;
|
||||
|
||||
__aicore__ inline KvQuantSparseFlashAttentionMla(){};
|
||||
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale,
|
||||
__gm__ uint8_t* valueScale, __gm__ uint8_t *blockTable,
|
||||
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
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 *sparseIndices, __gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *workspace, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void InitLocalBuffer();
|
||||
__aicore__ inline void ComputeConstexpr();
|
||||
__aicore__ inline void InitMMResBuf(__gm__ uint8_t *workspace);
|
||||
__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 InitUniqueRunInfo(const RunParamStr &runParam, RunInfo &runInfo);
|
||||
__aicore__ inline void ComputeAxisIdxByBnAndGs1(int64_t bnIndex, int64_t gS1Index, RunParamStr &runParam);
|
||||
__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 KvQuantSparseFlashAttentionTilingDataMla *__restrict tilingData;
|
||||
static constexpr uint64_t SYNC_MODE = 4;
|
||||
static constexpr uint32_t PRELOAD_NUM = 2;
|
||||
/* 核间通道 */
|
||||
BufferManager<BufferType::UB> ubBufferManager;
|
||||
BuffersPolicyDB<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm1Buffers;
|
||||
BuffersPolicySingleBuffer<BufferType::UB, SyncType::CROSS_CORE_SYNC_BOTH> bmm2Buffers;
|
||||
BufferManager<BufferType::GM> gmBufferManager;
|
||||
|
||||
// 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;
|
||||
|
||||
/* 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 KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::Init(
|
||||
__gm__ uint8_t *query,
|
||||
__gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale,
|
||||
__gm__ uint8_t* valueScale, __gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengths, __gm__ uint8_t *attentionOut, __gm__ uint8_t *workspace,
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
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.needInit = this->sharedParams.needInit;
|
||||
constInfo.dSizeV = 512;
|
||||
}
|
||||
vecBlock.CleanOutput(attentionOut, constInfo);
|
||||
/* cube侧不依赖sharedParams的scalar前置 */
|
||||
InitMMResBuf(workspace);
|
||||
if ASCEND_IS_AIC {
|
||||
cubeBlock.InitCubeBlock(pipe, &l1BufferManager, query);
|
||||
/* 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, sparseIndices, blockTable, actualSeqLengthsQ, actualSeqLengths,
|
||||
workspace, tiling, tPipe); // gm设置
|
||||
this->InitCalcParamsEach();
|
||||
this->InitLocalBuffer();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitCalcParamsEach()
|
||||
{
|
||||
// 计算总的基本块
|
||||
maxS2LoopCnt = 0; // 所有核中最大累计s2Loop
|
||||
uint32_t qsfaTotalBaseNum = 0;
|
||||
uint32_t actBatchS2 = 1;
|
||||
uint32_t coreNum = GetBlockNum(); // G128时相邻两个cube核处理一个s1,coreNum减半
|
||||
uint32_t currCoreIdx = aicIdx;
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
currCoreIdx = currCoreIdx >> 1;
|
||||
coreNum = coreNum >> 1;
|
||||
}
|
||||
|
||||
uint32_t actBatchS1 = 1;
|
||||
for (uint32_t bIdx = 0; bIdx < constInfo.bSize; bIdx++) {
|
||||
uint32_t actBatchS1 = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx); //不切S2,只关注S1
|
||||
qsfaTotalBaseNum += actBatchS1 * actBatchS2;
|
||||
}
|
||||
|
||||
uint32_t avgBaseNum = 1;
|
||||
if (qsfaTotalBaseNum > coreNum) {
|
||||
avgBaseNum = (qsfaTotalBaseNum + coreNum - 1) / coreNum;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
usedCoreNum = ((qsfaTotalBaseNum + avgBaseNum - 1) / avgBaseNum) << 1;
|
||||
}
|
||||
} else {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
usedCoreNum = qsfaTotalBaseNum << 1;
|
||||
} else {
|
||||
usedCoreNum = qsfaTotalBaseNum;
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt = avgBaseNum * (Min(constInfo.sparseBlockCount, constInfo.s2Size) +
|
||||
constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
}
|
||||
|
||||
if (aicIdx >= usedCoreNum) {
|
||||
return;
|
||||
}
|
||||
// 计算当前核的基本块
|
||||
uint32_t qsfaAccumBaseNum = 0; // qsfa当前累积的基本块数
|
||||
uint32_t targetBaseNum = 0;
|
||||
uint32_t qsfaLastValidBIdx = 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++) {
|
||||
qsfaAccumBaseNum += 1;
|
||||
if (!setStart && qsfaAccumBaseNum >= targetStartBaseNum) {
|
||||
constInfo.bN2Start = bN2Idx;
|
||||
constInfo.gS1Start = s1GIdx;
|
||||
setStart = true;
|
||||
}
|
||||
if (qsfaAccumBaseNum >= targetBaseNum) {
|
||||
// 更新当前核的End分核信息
|
||||
constInfo.s2End = 0;
|
||||
constInfo.bN2End = bN2Idx;
|
||||
constInfo.gS1End = s1GIdx;
|
||||
|
||||
if (currCoreIdx != 0) {
|
||||
GetAxisStartIdx(constInfo.bN2Start, constInfo.gS1Start, 0);
|
||||
}
|
||||
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if ((actBatchS1 > 0) && (actBatchS2 > 0)) {
|
||||
qsfaLastValidBIdx = bIdx;
|
||||
lastValidactBatchS1 = actBatchS1;
|
||||
}
|
||||
}
|
||||
if (!setStart) {
|
||||
constInfo.bN2Start = qsfaLastValidBIdx;
|
||||
constInfo.gS1Start = lastValidactBatchS1 - 1;
|
||||
}
|
||||
if (qsfaAccumBaseNum < targetBaseNum) {
|
||||
// 更新最后一个核的End分核信息
|
||||
constInfo.bN2End = qsfaLastValidBIdx;
|
||||
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 KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::\
|
||||
GetBalanceActualSeqLengths(GlobalTensor<int32_t> &actualSeqLengths, uint32_t bIdx)
|
||||
{
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
if (bIdx == 0) {
|
||||
return actualSeqQlenAddr[0];
|
||||
} else if (bIdx > 0) {
|
||||
return actualSeqQlenAddr[bIdx] - actualSeqQlenAddr[bIdx - 1];
|
||||
} else {
|
||||
return 0;
|
||||
}
|
||||
} else {
|
||||
if (constInfo.isActualLenDimsNull == 0) {
|
||||
return actualSeqQlenAddr[bIdx];
|
||||
} else {
|
||||
return constInfo.s1Size;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::GetAxisStartIdx(uint32_t bN2EndPrev,
|
||||
uint32_t s1GEndPrev,
|
||||
uint32_t s2EndPrev)
|
||||
{
|
||||
uint32_t qsfaBEndPrev = bN2EndPrev / constInfo.n2Size;
|
||||
uint32_t actualSeqQPrev = GetBalanceActualSeqLengths(actualSeqLengthsQGm, qsfaBEndPrev);
|
||||
uint32_t s1GPrevBaseNum = actualSeqQPrev;
|
||||
constInfo.bN2Start = bN2EndPrev;
|
||||
constInfo.gS1Start = s1GEndPrev;
|
||||
constInfo.s2Start = 0;
|
||||
if (s1GEndPrev >= s1GPrevBaseNum - 1) { // 上个核把S1G处理完了
|
||||
constInfo.bN2Start++;
|
||||
constInfo.gS1Start = 0;
|
||||
} else {
|
||||
constInfo.gS1Start++;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitGlobalBuffer(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths,
|
||||
__gm__ uint8_t *workspace, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
if (actualSeqLengthsQ != nullptr) {
|
||||
actualSeqQlenAddr = (__gm__ int32_t *)actualSeqLengthsQ;
|
||||
}
|
||||
|
||||
if (actualSeqLengths != nullptr) {
|
||||
actualSeqKvlenAddr = (__gm__ int32_t *)actualSeqLengths;
|
||||
}
|
||||
|
||||
vecBlock.InitGlobalBuffer(key, value, sparseIndices, blockTable);
|
||||
cubeBlock.InitCubeInput(actualSeqLengthsQ, constInfo);
|
||||
}
|
||||
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitMMResBuf(
|
||||
__gm__ uint8_t *workspace)
|
||||
{
|
||||
uint32_t mm1RightSize = constInfo.s2BaseSize * 576 * sizeof(Q_T);
|
||||
l1BufferManager.Init(pipe, 524288); // 512 * 1024
|
||||
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();
|
||||
}
|
||||
uint32_t mm1ResultSize = constInfo.s1BaseSize / CV_RATIO * constInfo.s2BaseSize * sizeof(T);
|
||||
uint32_t mm2ResultSize = constInfo.s1BaseSize / CV_RATIO * 512 * sizeof(T);
|
||||
ubBufferManager.Init(pipe, mm1ResultSize * 2 + mm2ResultSize);
|
||||
|
||||
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();
|
||||
}
|
||||
|
||||
bmm2Buffers.Init(ubBufferManager, mm2ResultSize);
|
||||
bmm2Buffers.Get().SetCrossCoreID(crossCoreSyncBufId, crossCoreSyncBufId);
|
||||
crossCoreSyncBufId++;
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
bmm2Buffers.Get().SetCrossCore();
|
||||
}
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
uint32_t v0ResSize = constInfo.s2BaseSize * 576U * sizeof(Q_T);
|
||||
int64_t totalOffset = v0ResSize * 3 * (aicIdx >> 1U);
|
||||
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 KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitLocalBuffer()
|
||||
{
|
||||
vecBlock.InitLocalBuffer(pipe, constInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<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.needInit = this->sharedParams.needInit;
|
||||
constInfo.dSizeV = 512;
|
||||
}
|
||||
constInfo.n2Size = sharedParams.n2Size;
|
||||
constInfo.s2Size = sharedParams.s2Size;
|
||||
constInfo.dSize = sharedParams.dSize;
|
||||
constInfo.dSizeVInput = sharedParams.dSizeVInput;
|
||||
constInfo.dSizeRope = sharedParams.dSizeRope;
|
||||
constInfo.dSizeNope = constInfo.dSize - constInfo.dSizeRope;
|
||||
constInfo.tileSize = sharedParams.tileSize;
|
||||
constInfo.sparseBlockCount = sharedParams.sparseBlockCount;
|
||||
constInfo.sparseBlockSize = 1;
|
||||
|
||||
constInfo.sparseMode = sharedParams.maskMode;
|
||||
constInfo.n2G = constInfo.n2Size * constInfo.gSize;
|
||||
|
||||
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.n2S2Dv = constInfo.n2Size * constInfo.s2Dv;
|
||||
constInfo.n2GDv = constInfo.n2Size * constInfo.gDv;
|
||||
constInfo.s2BaseN2Dv = constInfo.s2BaseSize * constInfo.n2Dv;
|
||||
constInfo.layoutType = sharedParams.layoutType;
|
||||
|
||||
constInfo.isActualLenDimsNull = sharedParams.isActualSeqLengthsNull;
|
||||
constInfo.isActualLenDimsKVNull = sharedParams.isActualSeqLengthsKVNull;
|
||||
|
||||
if constexpr (LAYOUT_T == QSFA_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 == QSFA_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.blockSize = sharedParams.blockSize;
|
||||
constInfo.softmaxScale = sharedParams.softmaxScale;
|
||||
constInfo.maxBlockNumPerBatch = sharedParams.maxBlockNumPerBatch;
|
||||
}
|
||||
|
||||
InitUniqueConstInfo();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<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 KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::Process()
|
||||
{
|
||||
// SyncAll Cube和Vector都需要调用
|
||||
if (this->sharedParams.needInit) {
|
||||
SyncAll<false>();
|
||||
}
|
||||
|
||||
ProcessMainLoop();
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<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<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
// 适配分核左闭右开
|
||||
uint32_t bIdx = constInfo.bN2End / constInfo.n2Size;
|
||||
uint32_t qsfaActS1Size = GetBalanceActualSeqLengths(actualSeqLengthsQGm, bIdx);
|
||||
uint32_t gS1max = qsfaActS1Size;
|
||||
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 qsfaBN2StartIdx = 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++;
|
||||
}
|
||||
|
||||
RunInfo runInfo[3];
|
||||
RunParamStr runParam;
|
||||
int64_t taskId = 0;
|
||||
bool notLast = true;
|
||||
int64_t multiCoreInnerIdx = 1;
|
||||
for (int64_t qsfaBnIdx = qsfaBN2StartIdx; qsfaBnIdx < bN2EndIdx; qsfaBnIdx++) {
|
||||
bool lastBN = (qsfaBnIdx == bN2EndIdx - 1);
|
||||
runParam.boIdx = qsfaBnIdx;
|
||||
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 qsfaExtraGS1 = gS1Index - runParam.gs1LoopEndIdx;
|
||||
switch (qsfaExtraGS1) {
|
||||
case 0:
|
||||
notLastTwoLoop = false;
|
||||
break;
|
||||
case 1:
|
||||
notLastTwoLoop = false;
|
||||
notLast = false;
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (notLastTwoLoop) {
|
||||
this->ComputeAxisIdxByBnAndGs1(qsfaBnIdx, gS1Index, runParam);
|
||||
bool s1NoNeedCalc = ComputeParamS1<TEMPLATE_INTF_ARGS>(
|
||||
runParam, this->constInfo, gS1Index, this->actualSeqQlenAddr);
|
||||
bool s2NoNeedCalc =
|
||||
ComputeS2LoopInfo<TEMPLATE_INTF_ARGS>(runParam, this->constInfo);
|
||||
// s1和s2有任意一个不需要算, 则continue, 如果是当前核最后一次循环,则补充计算taskIdx+2的部分
|
||||
if (s1NoNeedCalc || s2NoNeedCalc) {
|
||||
continue;
|
||||
}
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
maxS2LoopCnt -= runParam.s2LoopEndIdx;
|
||||
}
|
||||
s2LoopLimit = runParam.s2LoopEndIdx - 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(),
|
||||
this->v0ResGmBuffers.Get(), runInfo1, this->constInfo);
|
||||
} else {
|
||||
this->vecBlock.ProcessVec0(this->l1RightBuffers.Get(), this->v0ResGmBuffers.Get(),
|
||||
runInfo1, this->constInfo);
|
||||
}
|
||||
} else {
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
if (maxS2LoopCnt > 0) {
|
||||
maxS2LoopCnt--;
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, 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 &qsfaRunInfo3 = runInfo[(taskId + 1) % 3];
|
||||
this->vecBlock.ProcessVec2(this->bmm2Buffers.Get(), qsfaRunInfo3, this->constInfo);
|
||||
}
|
||||
}
|
||||
++taskId;
|
||||
}
|
||||
++multiCoreInnerIdx;
|
||||
}
|
||||
gS1StartIdx = 0;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
for (int64_t qsfaLoopCnt = 0; qsfaLoopCnt < maxS2LoopCnt; qsfaLoopCnt++) {
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<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 KvQuantSparseFlashAttentionMla<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.boIdx = runParam.boIdx;
|
||||
runInfo.s1oIdx = runParam.s1oIdx;
|
||||
runInfo.n2oIdx = runParam.n2oIdx;
|
||||
runInfo.goIdx = runParam.goIdx;
|
||||
|
||||
runInfo.multiCoreInnerIdx = multiCoreInnerIdx;
|
||||
runInfo.multiCoreIdxMod2 = multiCoreInnerIdx & 1;
|
||||
runInfo.multiCoreIdxMod3 = multiCoreInnerIdx % 3;
|
||||
}
|
||||
|
||||
runInfo.s2LoopLimit = s2LoopLimit;
|
||||
runInfo.taskId = taskId;
|
||||
runInfo.taskIdMod2 = taskId & 1;
|
||||
runInfo.taskIdMod3 = taskId % 3;
|
||||
|
||||
runInfo.sOuterOffset = runParam.sOuterOffset;
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.attentionOutOffset = runParam.attentionOutOffset;
|
||||
this->ComputeBmm1Tail(runInfo, runParam);
|
||||
InitUniqueRunInfo(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::InitUniqueRunInfo(
|
||||
const RunParamStr &runParam, RunInfo &runInfo)
|
||||
{
|
||||
InitTaskParamByRun<TEMPLATE_INTF_ARGS>(runParam, runInfo);
|
||||
}
|
||||
|
||||
template <typename CubeBlockType, typename VecBlockType>
|
||||
__aicore__ inline void KvQuantSparseFlashAttentionMla<CubeBlockType, VecBlockType>::ComputeBmm1Tail(
|
||||
RunInfo &runInfo, RunParamStr &runParam)
|
||||
{
|
||||
// ------------------------S1 Base Related---------------------------
|
||||
runInfo.s1RealSize = runParam.s1RealSize;
|
||||
runInfo.halfS1RealSize = runParam.halfS1RealSize;
|
||||
runInfo.firstHalfS1RealSize = runParam.firstHalfS1RealSize;
|
||||
|
||||
runInfo.halfMRealSize = runParam.halfMRealSize;
|
||||
runInfo.firstHalfMRealSize = runParam.firstHalfMRealSize;
|
||||
runInfo.mRealSize = runParam.mRealSize;
|
||||
|
||||
runInfo.vec2S1BaseSize = runInfo.halfS1RealSize;
|
||||
runInfo.vec2MBaseSize = runInfo.halfMRealSize;
|
||||
|
||||
// ------------------------S2 Base Related----------------------------
|
||||
runInfo.s2RealSize = constInfo.s2BaseSize;
|
||||
runInfo.s2AlignedSize = runInfo.s2RealSize;
|
||||
|
||||
if (runInfo.s2StartIdx + (runInfo.s2LoopCount + 1) * runInfo.s2RealSize > runInfo.s2EndIdx) {
|
||||
runInfo.s2RealSize = runInfo.s2EndIdx - runInfo.s2LoopCount * runInfo.s2RealSize - runInfo.s2StartIdx;
|
||||
runInfo.s2AlignedSize = Align(runInfo.s2RealSize);
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_KERNEL_MLA_H
|
||||
@@ -0,0 +1,256 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_flash_attention_kvcache.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
|
||||
using namespace matmul;
|
||||
using namespace regbaseutil;
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
static constexpr uint32_t sparseModeThree = 3;
|
||||
static constexpr uint32_t sparseModeZero = 0;
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void GetSingleCoreParam(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
__gm__ int32_t *actualSeqQlenAddr, __gm__ int32_t * actualSeqKvlenAddr)
|
||||
{
|
||||
int32_t qsfaActualS1Size = 0;
|
||||
int32_t qsfaActualS2Size = 0;
|
||||
int32_t actualSeqMin = 1;
|
||||
int32_t actualSeqKVMin = 1;
|
||||
int32_t sIdx = runParam.boIdx;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
// actual seq length first
|
||||
if (actualSeqQlenAddr != nullptr) {
|
||||
qsfaActualS1Size = (sIdx == 0) ? actualSeqQlenAddr[0] :
|
||||
actualSeqQlenAddr[sIdx] - actualSeqQlenAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaActualS1Size = constInfo.s1Size;
|
||||
}
|
||||
} else {
|
||||
qsfaActualS1Size = (actualSeqQlenAddr == nullptr) ? constInfo.s1Size :
|
||||
actualSeqQlenAddr[sIdx];
|
||||
}
|
||||
|
||||
if (constInfo.isActualLenDimsKVNull) {
|
||||
qsfaActualS2Size = constInfo.s2Size;
|
||||
} else {
|
||||
if constexpr (isPa) {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaActualS2Size = actualSeqKvlenAddr[sIdx];
|
||||
} else {
|
||||
qsfaActualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
|
||||
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
|
||||
}
|
||||
} else {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaActualS2Size = (sIdx == 0) ? actualSeqKvlenAddr[0] :
|
||||
actualSeqKvlenAddr[sIdx] - actualSeqKvlenAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaActualS2Size = (constInfo.actualSeqLenKVSize == actualSeqKVMin) ?
|
||||
actualSeqKvlenAddr[0] : actualSeqKvlenAddr[sIdx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
runParam.actualS1Size = qsfaActualS1Size;
|
||||
runParam.actualS2Size = qsfaActualS2Size;
|
||||
runParam.preTokensPerBatch = runParam.actualS1Size;
|
||||
if (constInfo.sparseMode == sparseModeZero) {
|
||||
runParam.nextTokensPerBatch = MAX_PRE_NEXT_TOKENS;
|
||||
} else {
|
||||
runParam.nextTokensPerBatch = runParam.actualS2Size - 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.gs1LoopStartIdx = gS1StartIdx;
|
||||
runParam.qSNumInOneBlock = 1; // qsfa 不切G轴, 计算每个基本块可以拷贝多少行s
|
||||
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
uint64_t invalidTokenCount = static_cast<uint64_t>(-(runParam.nextTokensPerBatch + 1)) + 1ULL;
|
||||
int64_t gs1LoopStartIdx =
|
||||
invalidTokenCount / runParam.qSNumInOneBlock * runParam.qSNumInOneBlock;
|
||||
if (gs1LoopStartIdx > gS1StartIdx) {
|
||||
runParam.gs1LoopStartIdx = gs1LoopStartIdx;
|
||||
}
|
||||
}
|
||||
|
||||
int32_t qsfaGs1LoopEndIdx = runParam.actualS1Size; // qsfa 不切G轴, 每次拷贝一行的topk,只算一行的qs
|
||||
|
||||
// 不是最后一个bn, 赋值souterBlockNum
|
||||
if (!lastBN) {
|
||||
runParam.gs1LoopEndIdx = qsfaGs1LoopEndIdx;
|
||||
} else { // 最后一个bn, 从数组下一个元素取值
|
||||
runParam.gs1LoopEndIdx = nextGs1Idx == 0 ? qsfaGs1LoopEndIdx : 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 qsfaCubeSOuterOffset = sOuterLoopIdx * runParam.qSNumInOneBlock;
|
||||
if (runParam.actualS1Size == 0) {
|
||||
runParam.s1RealSize = 0;
|
||||
runParam.mRealSize = 0;
|
||||
} else {
|
||||
runParam.s1RealSize = Min(runParam.qSNumInOneBlock, runParam.actualS1Size - qsfaCubeSOuterOffset);
|
||||
runParam.mRealSize = runParam.s1RealSize * constInfo.gSize;
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
runParam.mRealSize = runParam.mRealSize >> 1;
|
||||
}
|
||||
}
|
||||
|
||||
runParam.cubeMOuterOffset = qsfaCubeSOuterOffset * constInfo.gSize;
|
||||
runParam.halfMRealSize = (runParam.mRealSize + 1) >> 1;
|
||||
runParam.firstHalfMRealSize = runParam.halfMRealSize;
|
||||
if (constInfo.subBlockIdx == 0) {
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset;
|
||||
} else {
|
||||
runParam.halfMRealSize = runParam.mRealSize - runParam.halfMRealSize;
|
||||
runParam.mOuterOffset = runParam.cubeMOuterOffset + runParam.firstHalfMRealSize;
|
||||
}
|
||||
runParam.halfS1RealSize = (runParam.s1RealSize + 1) >> 1;
|
||||
runParam.firstHalfS1RealSize = runParam.halfS1RealSize;
|
||||
|
||||
if (constInfo.subBlockIdx == 1) {
|
||||
runParam.halfS1RealSize = runParam.s1RealSize - runParam.halfS1RealSize;
|
||||
runParam.sOuterOffset = qsfaCubeSOuterOffset + runParam.halfMRealSize / constInfo.gSize;
|
||||
} else {
|
||||
runParam.sOuterOffset = qsfaCubeSOuterOffset;
|
||||
}
|
||||
runParam.cubeSOuterOffset = qsfaCubeSOuterOffset;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void LoopSOuterOffsetInit(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
int32_t sIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
int64_t qsfaSeqOffset = 0;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
qsfaSeqOffset = sIdx == 0 ? 0 : cuSeqlensQAddr[sIdx - 1];
|
||||
} else {
|
||||
qsfaSeqOffset = sIdx * constInfo.s1Size;
|
||||
}
|
||||
|
||||
int64_t attentionOutSeqOffset = qsfaSeqOffset * constInfo.n2GDv;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline bool ComputeParamS1(RunParamStr& runParam, const ConstInfo &constInfo,
|
||||
uint32_t sOuterLoopIdx, __gm__ int32_t *cuSeqlensQAddr)
|
||||
{
|
||||
if (runParam.nextTokensPerBatch < 0) {
|
||||
uint64_t invalidTokenCount = static_cast<uint64_t>(-(runParam.nextTokensPerBatch + 1)) + 1ULL;
|
||||
if (runParam.s1oIdx <
|
||||
invalidTokenCount / 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 == QSFA_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 qsfaSInnerToken, int64_t minValue, int64_t maxValue)
|
||||
{
|
||||
qsfaSInnerToken = qsfaSInnerToken > minValue ? qsfaSInnerToken : minValue;
|
||||
qsfaSInnerToken = qsfaSInnerToken < maxValue ? qsfaSInnerToken : maxValue;
|
||||
return qsfaSInnerToken;
|
||||
}
|
||||
|
||||
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 qsfaS2BaseSize = 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 + qsfaS2BaseSize - 1) / qsfaS2BaseSize;
|
||||
runParam.s2LoopEndIdx = runParam.kvLoopEndIdx;
|
||||
return false;
|
||||
}
|
||||
|
||||
TEMPLATE_INTF
|
||||
__aicore__ inline void InitTaskParamByRun(const RunParamStr& runParam, RunInfo &runInfo)
|
||||
{
|
||||
runInfo.boIdx = runParam.boIdx;
|
||||
runInfo.actualS1Size = runParam.actualS1Size;
|
||||
runInfo.actualS2Size = runParam.actualS2Size;
|
||||
runInfo.preTokensPerBatch = runParam.preTokensPerBatch;
|
||||
runInfo.nextTokensPerBatch = runParam.nextTokensPerBatch;
|
||||
runInfo.softmaxLseOffset = runParam.softmaxLseOffset;
|
||||
runInfo.qSNumInOneBlock = runParam.qSNumInOneBlock;
|
||||
runInfo.kvLoopEndIdx = runParam.kvLoopEndIdx;
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_KVCACHE_H
|
||||
@@ -0,0 +1,388 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_service_cube_mla.h
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
#define KV_QUANT_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 "kv_quant_sparse_flash_attention_common_arch35.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/CopyInL1.h")
|
||||
#include "../../common/op_kernel/CopyInL1.h"
|
||||
#else
|
||||
#include "../common/CopyInL1.h"
|
||||
#endif
|
||||
#if __has_include("../../common/op_kernel/FixpipeOut.h")
|
||||
#include "../../common/op_kernel/FixpipeOut.h"
|
||||
#else
|
||||
#include "../common/FixpipeOut.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
|
||||
using namespace fa_base_matmul;
|
||||
using namespace regbaseutil;
|
||||
|
||||
namespace BaseApi {
|
||||
struct CubeCoordInfo {
|
||||
uint32_t curBIdx;
|
||||
uint32_t s1Coord;
|
||||
uint32_t s2Coord;
|
||||
};
|
||||
|
||||
template <QSFA_LAYOUT LAYOUT>
|
||||
__aicore__ inline constexpr GmFormat GetQueryGmFormat()
|
||||
{
|
||||
if constexpr (LAYOUT == QSFA_LAYOUT::BSND) {
|
||||
return GmFormat::BSNGD;
|
||||
} else {
|
||||
return GmFormat::TNGD;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF
|
||||
class QSFAMatmulService {
|
||||
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 QSFAMatmulService() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr,
|
||||
__gm__ uint8_t *query);
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, 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();
|
||||
__aicore__ inline void InitGmTensor(__gm__ uint8_t *cuSeqlensQ, const ConstInfo& constInfo);
|
||||
__aicore__ inline void CalcS1Coord(RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
__aicore__ inline void IterateBmm1QSFA(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);
|
||||
__aicore__ inline void PrepareLeftMatrixBmm1QSFA(Buffer<BufferType::L1> &inputLeftBuf,
|
||||
RunInfo &runInfo, ConstInfo &constInfo);
|
||||
|
||||
// --------------------Bmm2--------------------------
|
||||
__aicore__ inline void IterateBmm2QSFA(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;
|
||||
|
||||
/* =====================运行时变量==================== */
|
||||
CubeCoordInfo coordInfo[3];
|
||||
TEventID mte1ToMte2Id[3];
|
||||
TEventID mte2ToMte1Id[3];
|
||||
|
||||
/* =====================LocalBuffer变量==================== */
|
||||
// D小于等于256 mm1左矩阵Q,GS1循环内左矩阵复用, GS1循环间开pingpong;D大于256使用单块Buffer,S1循环间驻留;fp32场景单块不驻留
|
||||
BuffersPolicySingleBuffer<BufferType::L1> l1QBuffers;
|
||||
// L0空间buffer manager
|
||||
BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr;
|
||||
BufferManager<BufferType::L0A> l0aBufferManager;
|
||||
BufferManager<BufferType::L0B> l0bBufferManager;
|
||||
BufferManager<BufferType::L0C> l0cBufferManager;
|
||||
// L0A
|
||||
BuffersPolicyDB<BufferType::L0A> mmL0ABuffers;
|
||||
// L0B
|
||||
BuffersPolicyDB<BufferType::L0B> mmL0BBuffers;
|
||||
// L0C
|
||||
BuffersPolicyDB<BufferType::L0C> mmL0CBuffers;
|
||||
};
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::InitCubeBlock(
|
||||
TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BuffMgr, __gm__ uint8_t *query)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
tPipe = pipe;
|
||||
qsfaL1BufferManagerPtr = qsfaL1BuffMgr;
|
||||
this->queryGm.gmTensor.SetGlobalBuffer((__gm__ Q_T *)query);
|
||||
InitLocalBuffer();
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitCubeInput(__gm__ uint8_t *qsfaActualSeqLengthsQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if ASCEND_IS_AIC {
|
||||
InitGmTensor(qsfaActualSeqLengthsQ, constInfo);
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
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>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitLocalBuffer()
|
||||
{
|
||||
constexpr uint32_t mm1LeftSize = s1BaseSize * dBaseSize * sizeof(Q_T);
|
||||
l1QBuffers.Init((*qsfaL1BufferManagerPtr), 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
|
||||
QSFAMatmulService<TEMPLATE_ARGS>::InitGmTensor(__gm__ uint8_t *qsfaActualSeqLengthsQ, const ConstInfo& constInfo)
|
||||
{
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
this->queryGm.offsetCalculator.Init(constInfo.bSize, constInfo.n2Size, constInfo.gSize,
|
||||
constInfo.s1Size, constInfo.dSize);
|
||||
} else { // QSFA_LAYOUT::TND
|
||||
GlobalTensor<int32_t> actualSeqQLen;
|
||||
actualSeqQLen.SetGlobalBuffer((__gm__ int32_t *)qsfaActualSeqLengthsQ);
|
||||
this->queryGm.offsetCalculator.Init(constInfo.n2Size, constInfo.gSize, constInfo.dSize,
|
||||
actualSeqQLen, constInfo.actualSeqLenSize);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::CalcS1Coord(RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
// 计算s1方向偏移
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord = runInfo.s1oIdx * runInfo.qSNumInOneBlock;
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<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)
|
||||
{
|
||||
CalcS1Coord(runInfo, constInfo);
|
||||
|
||||
IterateBmm1QSFA(outputBuf, inputRightBuf, v0ResGm, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<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)
|
||||
{
|
||||
IterateBmm2QSFA(outputBuf, inputLeftBuffers, inputRightBuf, runInfo, constInfo);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm1QSFA(
|
||||
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;
|
||||
PrepareLeftMatrixBmm1QSFA(inputLeftBuf, runInfo, constInfo);
|
||||
|
||||
// 加载当前轮的右矩阵到L1
|
||||
inputRightBuf.WaitCrossCore(); // 核间同步,这里需要根据V0操作处理同步,确保取tensor时,数据已经准备好
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
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>();
|
||||
DataCopy(dst, v0ResGmTensor, Align16Func(runInfo.s2RealSize) * constInfo.dSize);
|
||||
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.dSize), // singleK
|
||||
0, // isLeftTranspose
|
||||
1 // isRightTranspose
|
||||
};
|
||||
|
||||
MatmulK<Q_T, Q_T, T, s1BaseSize, s2BaseSize, dBaseMatmulSize, ABLayout::MK, ABLayout::KN>( // m,n不切,k切128
|
||||
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
|
||||
fixpipeParams.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可;
|
||||
fixpipeParams.nSize = Align8Func(runInfo.s2RealSize); // L0C上的bmm1结果矩阵N方向的size大小; 同mmadParams.n; 为什么要8个元素对齐(32B对齐) // 128
|
||||
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔), 单位为16*sizeof(T) // 源Nz矩阵中相邻大Z排布的起始地址偏移
|
||||
fixpipeParams.dstStride = s2BaseSize; // mmResUb上两行之间的间隔,单位:element。 // 128:根据比对dump文件得到, ND方案(S1*S2)时脏数据用mask剔除
|
||||
fixpipeParams.dualDstCtl = 1; // 双目标模式,按M维度拆分,M / 2 * N写入每个UB, M必须为2的倍数
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
|
||||
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 QSFAMatmulService<TEMPLATE_ARGS>::PrepareLeftMatrixBmm1QSFA(
|
||||
Buffer<BufferType::L1> &inputLeftBuf,
|
||||
RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
// 左矩阵复用,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>();
|
||||
uint64_t gmOffset = this->queryGm.offsetCalculator.GetOffset(runInfo.boIdx, runInfo.n2oIdx, runInfo.goIdx,
|
||||
coordInfo[runInfo.taskIdMod3].s1Coord, 0);
|
||||
CopyToL1Nd2Nz<Q_T>(inputLeftTensor, this->queryGm.gmTensor[gmOffset], runInfo.mRealSize, constInfo.dSize,
|
||||
constInfo.mm1Ka);
|
||||
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
} else { // 非S2的第一次循环直接复用Q
|
||||
inputLeftBuf = l1QBuffers.GetPre();
|
||||
// 左矩阵复用时,sinner循环内不需要MTE2同步等待
|
||||
inputLeftBuf.Set<HardEvent::MTE2_MTE1>(); // 通知
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAMatmulService<TEMPLATE_ARGS>::IterateBmm2QSFA(
|
||||
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 qsfaParam = {static_cast<uint32_t>(runInfo.mRealSize), // singleM 64
|
||||
static_cast<uint32_t>(constInfo.dSizeNope), // singleN 576->512
|
||||
static_cast<uint32_t>(runInfo.s2RealSize), // singleK 128
|
||||
0, 0};
|
||||
|
||||
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>(), qsfaParam);
|
||||
|
||||
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.mSize = Align2Func(runInfo.mRealSize); // 有效数据不足16行,只需要输出部分行即可;
|
||||
fixpipeParams.nSize = Align8Func(constInfo.dSizeNope); // L0C上的bmm1结果矩阵N方向的size大小, 分档计算且vector2中通过mask筛选出实际有效值
|
||||
fixpipeParams.srcStride = Align16Func(fixpipeParams.mSize); // L0C上bmm1结果相邻连续数据片段间隔(前面一个数据块的头与后面数据块的头的间隔)
|
||||
fixpipeParams.dstStride = Align16Func(constInfo.dSizeNope);
|
||||
fixpipeParams.dualDstCtl = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
|
||||
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 QSFAMatmulServiceDummy {
|
||||
public:
|
||||
__aicore__ inline QSFAMatmulServiceDummy() {};
|
||||
__aicore__ inline void InitCubeBlock(TPipe *pipe, BufferManager<BufferType::L1> *qsfaL1BufferManagerPtr,
|
||||
__gm__ uint8_t *query) {}
|
||||
__aicore__ inline void InitCubeInput(__gm__ uint8_t *cuSeqlensQ, 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_QSFA_CUBE_BLOCK_TRAITS(CUBE_BLOCK_CLASS) \
|
||||
TEMPLATES_DEF_NO_DEFAULT \
|
||||
struct CubeBlockTraits<CUBE_BLOCK_CLASS<TEMPLATE_ARGS>> { \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_TRAIT_TYPE) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_TRAIT_CONST) \
|
||||
}
|
||||
|
||||
DEFINE_QSFA_CUBE_BLOCK_TRAITS(QSFAMatmulService);
|
||||
DEFINE_QSFA_CUBE_BLOCK_TRAITS(QSFAMatmulServiceDummy);
|
||||
|
||||
// /* 生成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 \
|
||||
QSFA_CUBE_BLOCK_TRAITS_TYPE_FIELDS(GEN_ARGS_TYPE) \
|
||||
QSFA_CUBE_BLOCK_TRAITS_CONST_FIELDS(GEN_ARGS_CONST)
|
||||
}
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
@@ -0,0 +1,894 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_flash_attention_service_vector_mla.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
|
||||
#include "kv_quant_sparse_flash_attention_common_arch35.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "lib/matmul_intf.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
|
||||
|
||||
using namespace AscendC;
|
||||
using namespace FaVectorApi;
|
||||
using namespace AscendC::Impl::Detail;
|
||||
using namespace regbaseutil;
|
||||
using namespace matmul;
|
||||
|
||||
namespace BaseApi {
|
||||
|
||||
TEMPLATES_DEF
|
||||
class QSFAVectorService {
|
||||
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 qsfaDTemplateAlign64 = Align64Func(dVTemplateType);
|
||||
static constexpr uint32_t dVTemplateTypeInput = 672;
|
||||
static constexpr float R0 = 1.0f;
|
||||
static constexpr uint64_t SYNC_SINKS_BUF_FLAG = 6;
|
||||
|
||||
// ==================== Functions ======================
|
||||
__aicore__ inline QSFAVectorService() {};
|
||||
__aicore__ inline void InitVecBlock(TPipe *pipe, const KvQuantSparseFlashAttentionTilingDataMla *__restrict tiling,
|
||||
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx, __gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengths)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
tilingData = tiling;
|
||||
tPipe = pipe;
|
||||
if (actualSeqLengths != nullptr) {
|
||||
actualSeqLengthsKVGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengths);
|
||||
}
|
||||
if (actualSeqLengthsQ != nullptr) {
|
||||
cuSeqlensQGm.SetGlobalBuffer((__gm__ int32_t *)actualSeqLengthsQ);
|
||||
}
|
||||
|
||||
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, ConstInfo &constInfo);
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable);
|
||||
__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);
|
||||
__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 ProcessVec1SoftmaxDispatchQSFA(LocalTensor<Q_T> &stage1CastTensor,
|
||||
LocalTensor<T> &mmRes, LocalTensor<float> &sumUb, LocalTensor<float> &maxUb,
|
||||
LocalTensor<T> &apiTmpBuffer, RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__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);
|
||||
__aicore__ inline void CalSparseCalSize(const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__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 DequantKv(LocalTensor<Q_T> antiKvTensorAsB16, LocalTensor<KV_T> srcTensor, int64_t dealRow,
|
||||
ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutKvUb2L1(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutKvUb2Gm(Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo,
|
||||
ConstInfo &constInfo);
|
||||
__aicore__ inline void CopyOutMrgeResult(Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
int64_t mte2Size, int64_t mte3Size, int64_t s2keyOffset, int64_t mergeMte3Idx, const RunInfo &runInfo);
|
||||
__aicore__ inline void CopyInSingleKv(LocalTensor<KV_T> kvInUb, int64_t startRow,
|
||||
int64_t keyOffset, uint32_t combineBytes);
|
||||
/* 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 qsfaVec2CalcSize = 0);
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void CopyOutAttentionOut(
|
||||
RunInfo &runInfo, ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx,
|
||||
int64_t qsfaVec2CalcSize);
|
||||
__aicore__ inline void SoftmaxInitBuffer();
|
||||
__aicore__ inline void InitCubeVecSharedParams(CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx);
|
||||
__aicore__ inline void ComputeNeedInitQSFA(CVSharedParams &sharedParams) const;
|
||||
__aicore__ inline void GetExtremeValue(T &negativeScalar);
|
||||
|
||||
TPipe *tPipe;
|
||||
const KvQuantSparseFlashAttentionTilingDataMla *__restrict tilingData;
|
||||
|
||||
GlobalTensor<OUTPUT_T> attentionOutGm;
|
||||
GlobalTensor<KV_T> keyGm;
|
||||
GlobalTensor<int32_t> SparseIndicesGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
GlobalTensor<int32_t> cuSeqlensQGm;
|
||||
GlobalTensor<int32_t> actualSeqLengthsKVGm;
|
||||
|
||||
TBuf<> commonTBuf; // common的复用空间
|
||||
TQue<QuePosition::VECOUT, 1> stage1OutQue[2]; // 2份表示可能存在pingpong
|
||||
TQue<QuePosition::VECIN, 2> stage0InQue; // for v0 input, 2份表示可能存在pingpong
|
||||
TQue<QuePosition::VECOUT, 2> stage0OutQue; // for v0 output, 2份表示可能存在pingpong
|
||||
TBuf<> stage2OutBuf;
|
||||
TEventID mte3ToVId[2]; // 存放MTE3_V的eventId, 2份表示可能存在pingpong
|
||||
TEventID vToMte3Id[2]; // 存放V_MTE3的eventId, 2份表示可能存在pingpong
|
||||
TBuf<> softmaxMaxBuf[2];
|
||||
TBuf<> softmaxSumBuf[2];
|
||||
TBuf<> softmaxExpBuf[2];
|
||||
|
||||
T negativeFloatScalar;
|
||||
uint32_t maxBlockNumPerBatch;
|
||||
uint32_t blockSize;
|
||||
int64_t qsfaSparseCalSize;
|
||||
int64_t sparseS2Start;
|
||||
int64_t sparseS2End;
|
||||
};
|
||||
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<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 == QSFA_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 qsfaCmpS2LoopCnt = runInfo.s2LoopCount;
|
||||
int64_t qsfaTopkIdx = s2IdxInBase + qsfaCmpS2LoopCnt * constInfo.s2BaseSize;
|
||||
|
||||
if (unlikely(qsfaTopkIdx >= constInfo.sparseBlockCount)) {
|
||||
token0Idx = -1;
|
||||
} else {
|
||||
token0Idx = SparseIndicesGm.GetValue(topkBS1Idx + qsfaTopkIdx) + runInfo.s2StartIdx;
|
||||
}
|
||||
qsfaTopkIdx += 1;
|
||||
if (unlikely((qsfaTopkIdx >= constInfo.sparseBlockCount) || (s2IdxInBase + 1 >= sparseS2End))) {
|
||||
token1Idx = -1;
|
||||
} else {
|
||||
token1Idx = SparseIndicesGm.GetValue(topkBS1Idx + qsfaTopkIdx) + runInfo.s2StartIdx;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline int64_t QSFAVectorService<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) * constInfo.dSizeVInput +
|
||||
blkTableOffset * constInfo.dSizeVInput; // BlockNum, BlockSize, N(1), D
|
||||
} else {
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
realkeyOffset = (runInfo.boIdx * constInfo.s2Size + s2Idx) * constInfo.dSizeVInput; // BSN(1)D
|
||||
} else if constexpr (LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
int64_t batchKvStart = (runInfo.boIdx == 0) ? 0 : actualSeqLengthsKVGm.GetValue(runInfo.boIdx - 1);
|
||||
realkeyOffset = (batchKvStart + s2Idx) * constInfo.dSizeVInput;
|
||||
}
|
||||
}
|
||||
return realkeyOffset;
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void
|
||||
QSFAVectorService<TEMPLATE_ARGS>::CopyInSingleKv(LocalTensor<KV_T> kvInUb, int64_t startRow,
|
||||
int64_t keyOffset, uint32_t combineBytes)
|
||||
{
|
||||
if (keyOffset < 0) {
|
||||
return;
|
||||
}
|
||||
DataCopyExtParams intriParams;
|
||||
|
||||
intriParams.blockCount = 1;
|
||||
intriParams.dstStride = 0;
|
||||
intriParams.srcStride = 0;
|
||||
DataCopyPadExtParams<KV_T> padParams;
|
||||
// 当前仅支持COMBINE模式
|
||||
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 * combineDimAlign], keyGm[keyOffset], intriParams, padParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline uint32_t QSFAVectorService<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;
|
||||
}
|
||||
uint32_t combineBytes = constInfo.dSizeVInput * sizeof(KV_T);
|
||||
int64_t keySrcStride = (keyOffset0 > keyOffset1 ? (keyOffset0 - keyOffset1) * sizeof(KV_T):
|
||||
(keyOffset1 - keyOffset0)) * sizeof(KV_T) - combineBytes;
|
||||
if (keySrcStride >= INT32_MAX || keySrcStride < 0 || constInfo.sparseBlockSize > 1) {
|
||||
// stride溢出、stride为负数、s2超长等异常场景,还原成2条搬运指令
|
||||
CopyInSingleKv(kvInUb, startRow, keyOffset0, combineBytes);
|
||||
CopyInSingleKv(kvInUb, startRow + 1, keyOffset1, combineBytes);
|
||||
} else {
|
||||
DataCopyExtParams intriParams;
|
||||
intriParams.blockCount = (keyOffset0 >= 0) + (keyOffset1 >= 0);
|
||||
intriParams.blockLen = combineBytes;
|
||||
intriParams.dstStride = 0;
|
||||
intriParams.srcStride = keySrcStride;
|
||||
DataCopyPadExtParams<KV_T> padParams;
|
||||
|
||||
int64_t keyOffset = keyOffset0 > -1 ? keyOffset0 : keyOffset1;
|
||||
if (keyOffset1 > -1 && keyOffset1 < keyOffset0) {
|
||||
keyOffset = keyOffset1;
|
||||
}
|
||||
|
||||
// 当前仅支持COMBINE模式
|
||||
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 * combineDimAlign], keyGm[keyOffset], intriParams, padParams);
|
||||
}
|
||||
return (keyOffset0 > -1) + (keyOffset1 > -1);
|
||||
}
|
||||
|
||||
// fp8->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraitFp8_1 = {MicroAPI::RegLayout::ZERO, 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};
|
||||
|
||||
// int8->half
|
||||
static constexpr MicroAPI::CastTrait castTraitint8_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
// half->fp32
|
||||
static constexpr MicroAPI::CastTrait castTraithalf_1 = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::UNKNOWN,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
template <typename Q_T, typename KV_T>
|
||||
__simd_vf__ void AntiquantVFImplFp8D448(__ubuf__ int8_t* ubSrcAddr, __ubuf__ Q_T* ubDstAddr, // output first
|
||||
__ubuf__ float* ubScaleSrcAddr, uint32_t dealRowCount)
|
||||
{
|
||||
uint32_t combineDim = 672; // 128对齐 640->672
|
||||
MicroAPI::RegTensor<KV_T> vKvData0;
|
||||
MicroAPI::RegTensor<KV_T> vKvData1;
|
||||
MicroAPI::RegTensor<half> vKvDataHalf0;
|
||||
MicroAPI::RegTensor<half> vKvDataHalf1;
|
||||
MicroAPI::RegTensor<half> vCastHalfRes0;
|
||||
MicroAPI::RegTensor<half> vCastHalfRes1;
|
||||
MicroAPI::RegTensor<float> vCastFp32Res0;
|
||||
MicroAPI::RegTensor<float> vCastFp32Res1;
|
||||
MicroAPI::RegTensor<float> vMulRes0;
|
||||
MicroAPI::RegTensor<float> vMulRes1;
|
||||
MicroAPI::RegTensor<float> vScale0;
|
||||
MicroAPI::RegTensor<float> vScale1;
|
||||
MicroAPI::RegTensor<Q_T> vCastRes0;
|
||||
MicroAPI::RegTensor<Q_T> vCastRes1;
|
||||
MicroAPI::RegTensor<Q_T> vCastResPack0;
|
||||
MicroAPI::RegTensor<Q_T> vCastResPack1;
|
||||
|
||||
MicroAPI::MaskReg kvTypeMaskAll = MicroAPI::CreateMask<KV_T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg kvRopeTypeMaskAll = MicroAPI::CreateMask<Q_T, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg int8MaskAll = MicroAPI::CreateMask<half, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg fp32MaskAll = MicroAPI::CreateMask<float, MicroAPI::MaskPattern::ALL>();
|
||||
uint32_t blockStride = 17; // +1 to solve bank conflict
|
||||
uint32_t repeatStride = 1;
|
||||
const uint32_t nopeDim = 512; // 448->512 64
|
||||
const uint32_t kvNumPerLoop = 128;
|
||||
const uint32_t scaleNumPerLoop = 1;
|
||||
const uint32_t tileSize = 128;
|
||||
static constexpr bool isKvInt8 = (IsSameType<KV_T, int8_t>::value);
|
||||
// tilesize is 128, deal 128 b8 kv, deal 1 fp32 scale
|
||||
for (uint16_t j = 0; j < (nopeDim / kvNumPerLoop); j++) {
|
||||
__ubuf__ int8_t* ubSrcTemp = ubSrcAddr + j * kvNumPerLoop;
|
||||
__ubuf__ float* ubScaleSrcAddrTemp = ubScaleSrcAddr + j * scaleNumPerLoop;
|
||||
__ubuf__ Q_T* ubDstAddrTmp = ubDstAddr + j * kvNumPerLoop * blockStride;
|
||||
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>&)vKvData0, ubSrcTemp, tileSize / 2);
|
||||
MicroAPI::LoadAlign<int8_t, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_UNPACK4_B8>(
|
||||
(MicroAPI::RegTensor<int8_t>&)vKvData1, ubSrcTemp, combineDim - tileSize / 2);
|
||||
|
||||
MicroAPI::LoadAlign<float, MicroAPI::PostLiteral::POST_MODE_UPDATE, MicroAPI::LoadDist::DIST_BRC_B32>(
|
||||
(MicroAPI::RegTensor<float>&)vScale0, ubScaleSrcAddrTemp, combineDim / 4);
|
||||
|
||||
if constexpr (isKvInt8) {
|
||||
// int8 -> half
|
||||
MicroAPI::Cast<half, KV_T, castTraitint8_1>(vCastHalfRes0, vKvData0, int8MaskAll);
|
||||
MicroAPI::Cast<half, KV_T, castTraitint8_1>(vCastHalfRes1, vKvData1, int8MaskAll);
|
||||
// half -> float
|
||||
MicroAPI::Cast<float, half, castTraithalf_1>(vCastFp32Res0, vCastHalfRes0, fp32MaskAll);
|
||||
MicroAPI::Cast<float, half, castTraithalf_1>(vCastFp32Res1, vCastHalfRes1, fp32MaskAll);
|
||||
} else {
|
||||
MicroAPI::Cast<float, KV_T, castTraitFp8_1>(vCastFp32Res0, vKvData0, fp32MaskAll);
|
||||
MicroAPI::Cast<float, KV_T, castTraitFp8_1>(vCastFp32Res1, vKvData1, fp32MaskAll);
|
||||
}
|
||||
|
||||
MicroAPI::Mul<float, MicroAPI::MaskMergeMode::ZEROING>(vMulRes0, vCastFp32Res0, vScale0, fp32MaskAll);
|
||||
MicroAPI::Mul<float, MicroAPI::MaskMergeMode::ZEROING>(vMulRes1, vCastFp32Res1, vScale0, fp32MaskAll);
|
||||
|
||||
MicroAPI::Cast<Q_T, float, castTraitFp8_3>(vCastRes0, vMulRes0, fp32MaskAll);
|
||||
MicroAPI::Cast<Q_T, float, castTraitFp8_3>(vCastRes1, vMulRes1, fp32MaskAll);
|
||||
|
||||
MicroAPI::DeInterleave(vCastResPack0, vCastResPack1, vCastRes0, vCastRes1);
|
||||
|
||||
MicroAPI::StoreAlign<Q_T, MicroAPI::DataCopyMode::DATA_BLOCK_COPY, MicroAPI::PostLiteral::POST_MODE_UPDATE>(
|
||||
ubDstAddrTmp, vCastResPack0, blockStride, repeatStride, kvRopeTypeMaskAll);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Q_T, typename KV_T>
|
||||
__aicore__ inline void AntiquantVFFp8D448(LocalTensor<Q_T>& outputUb, LocalTensor<KV_T>& inputUb, uint32_t dealRowCount)
|
||||
{
|
||||
__ubuf__ int8_t* ubSrcAddr = (__ubuf__ int8_t*)(inputUb.GetPhyAddr()); // nope改成在左,所以起始位置是0
|
||||
__ubuf__ Q_T* ubDstAddr = (__ubuf__ Q_T*)(outputUb.GetPhyAddr());
|
||||
__ubuf__ float* ubScaleAddr = (__ubuf__ float*)(inputUb[512 + 64 * 2].GetPhyAddr());
|
||||
|
||||
AntiquantVFImplFp8D448<Q_T, KV_T>(ubSrcAddr, ubDstAddr, ubScaleAddr, dealRowCount);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::DequantKv(LocalTensor<Q_T> antiKvTensorAsB16,
|
||||
LocalTensor<KV_T> srcTensor, int64_t dealRow, ConstInfo &constInfo)
|
||||
{
|
||||
// srcTensor是nope(512) + nope(64) + scale + pad, dstTensor是nope(512) + rope(64)
|
||||
AntiquantVFFp8D448<Q_T, KV_T>(antiKvTensorAsB16, srcTensor, dealRow);
|
||||
|
||||
LocalTensor<Q_T> kRopeUb = srcTensor[constInfo.dSizeNope].template ReinterpretCast<Q_T>();
|
||||
LocalTensor<Q_T> kRopeUbNz = antiKvTensorAsB16[constInfo.dSizeNope * (16 + 1)]; // V0单次处理16行数据
|
||||
Copy(kRopeUbNz, kRopeUb,
|
||||
constInfo.dSizeRope, // mask 处理多少列数据
|
||||
static_cast<uint8_t>(dealRow), // repeatTime, 每次处理多少个block
|
||||
{
|
||||
17, // dst stride
|
||||
1, // src stride
|
||||
1, // dst repeat stride
|
||||
21 // src repeat stride, 640 / 32 // 640 -> 672 : 20 -> 21
|
||||
});
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CopyOutKvUb2L1(
|
||||
Buffer<BufferType::L1, SyncType::CROSS_CORE_SYNC_FORWARD> &outputL1,
|
||||
LocalTensor<Q_T> antiKvTensorAsB16, int64_t dealRow, int64_t s2StartIdx,
|
||||
const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
uint64_t blockElementNum = 16;
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = (constInfo.dSizeNope + constInfo.dSizeRope) / blockElementNum;
|
||||
dataCopyParams.blockLen = dealRow;
|
||||
dataCopyParams.srcGap = blockElementNum + 1 - dealRow;
|
||||
dataCopyParams.dstGap = Align16Func(runInfo.s2RealSize) - dealRow;
|
||||
|
||||
LocalTensor<Q_T> dst = outputL1.GetTensor<Q_T>();
|
||||
DataCopy(dst[s2StartIdx * blockElementNum], antiKvTensorAsB16, dataCopyParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CopyOutKvUb2Gm(
|
||||
Buffer<BufferType::GM, SyncType::CROSS_CORE_SYNC_BACKWARD> &v0ResGm, LocalTensor<Q_T> antiKvTensorAsB16,
|
||||
int64_t dealRow, int64_t s2StartIdx, const RunInfo &runInfo, ConstInfo &constInfo)
|
||||
{
|
||||
GlobalTensor<Q_T> v0ResGmTensor = v0ResGm.template GetTensor<Q_T>();
|
||||
uint64_t blockElementNum = 16;
|
||||
DataCopyParams dataCopyParams;
|
||||
dataCopyParams.blockCount = (constInfo.dSizeNope + constInfo.dSizeRope) / blockElementNum;
|
||||
dataCopyParams.blockLen = dealRow;
|
||||
dataCopyParams.srcGap = blockElementNum + 1 - dealRow;
|
||||
dataCopyParams.dstGap = Align16Func(runInfo.s2RealSize) - dealRow;
|
||||
DataCopy(v0ResGmTensor[s2StartIdx * blockElementNum], antiKvTensorAsB16, dataCopyParams);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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) {
|
||||
qsfaSparseCalSize = CeilDiv(v0S2SizeFirstCore, 2); // 2: Vector split size for first core (first half)
|
||||
sparseS2Start = 0;
|
||||
} else {
|
||||
// 2: Vector split size for first core (second half)
|
||||
qsfaSparseCalSize = v0S2SizeFirstCore - CeilDiv(v0S2SizeFirstCore, 2);
|
||||
sparseS2Start = CeilDiv(v0S2SizeFirstCore, 2); // 2: Start offset for second half of first core
|
||||
}
|
||||
} else {
|
||||
if (GetSubBlockIdx() == 0) {
|
||||
qsfaSparseCalSize = CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
sparseS2Start = v0S2SizeFirstCore;
|
||||
} else {
|
||||
qsfaSparseCalSize = v0S2SizeSecondCore - CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
sparseS2Start = v0S2SizeFirstCore + CeilDiv(v0S2SizeSecondCore, 2); // 2: Same as above
|
||||
}
|
||||
}
|
||||
sparseS2End = sparseS2Start + qsfaSparseCalSize;
|
||||
} else {
|
||||
int64_t s2PerVecLoop = 2LL;
|
||||
int64_t vecNum = 2LL;
|
||||
int64_t s2Loops = CeilDiv(CeilDiv(runInfo.s2RealSize, vecNum), s2PerVecLoop);
|
||||
sparseS2Start = GetSubBlockIdx() == 0 ? 0 : s2Loops * s2PerVecLoop;
|
||||
sparseS2End = GetSubBlockIdx() == 0 ? s2Loops * s2PerVecLoop : runInfo.s2RealSize;
|
||||
qsfaSparseCalSize = sparseS2End - sparseS2Start;
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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)
|
||||
{
|
||||
outputL1.WaitCrossCore(); // 核间同步
|
||||
blockSize = constInfo.blockSize;
|
||||
maxBlockNumPerBatch = constInfo.maxBlockNumPerBatch;
|
||||
|
||||
CalSparseCalSize(runInfo, constInfo);
|
||||
ProcessSparseKv(outputL1, v0ResGm, runInfo, constInfo);
|
||||
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
CrossCoreSetFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15); // 15: 跨核同步标志位值
|
||||
CrossCoreWaitFlag<QSFA_SYNC_MODE0, PIPE_MTE3>(15); // 15: 跨核同步标志位值
|
||||
}
|
||||
|
||||
outputL1.SetCrossCore(); // 核间同步
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
v0ResGm.SetCrossCore();
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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)
|
||||
{
|
||||
if (qsfaSparseCalSize == 0) {
|
||||
return;
|
||||
}
|
||||
// Left-closed, right-open interval
|
||||
// 4x = 2x + 2x
|
||||
// 4x + 1 = (2x + 2) + (2x - 1)
|
||||
// 4x + 2 = (2x + 2) + (2x)
|
||||
// 4x + 3 = (2x + 2) + (2x + 1)
|
||||
int64_t s2Start = sparseS2Start;
|
||||
int64_t s2 = sparseS2Start;
|
||||
bool meetEnd = false;
|
||||
int64_t token0Idx, token1Idx; // 拷贝进入的两个token的index
|
||||
// 处理一个s2的base块
|
||||
while ((s2 < sparseS2End) && !meetEnd) { // 拷贝到s2End或者遇到-1
|
||||
int64_t dealRow = 0;
|
||||
// 1、copy kv in, gm ->ub
|
||||
LocalTensor<KV_T> kvInUb = stage0InQue.AllocTensor<KV_T>();
|
||||
while (dealRow < Min(16, qsfaSparseCalSize) && s2<sparseS2End) { // 拷贝满16行或者遇到-1
|
||||
GetRealCmpS2Idx(token0Idx, token1Idx, s2, runInfo, constInfo);
|
||||
s2 += 2; // 每次搬运2行
|
||||
if (token0Idx== -1 && token1Idx == -1) {
|
||||
meetEnd = true;
|
||||
break;
|
||||
}
|
||||
dealRow += CopyInKvSparse(kvInUb, dealRow, token0Idx, token1Idx, runInfo, constInfo);
|
||||
if (token1Idx == -1) {
|
||||
meetEnd = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (dealRow == 0) {
|
||||
stage0InQue.FreeTensor(kvInUb);
|
||||
return;
|
||||
}
|
||||
stage0InQue.EnQue(kvInUb);
|
||||
kvInUb = stage0InQue.DeQue<KV_T>();
|
||||
|
||||
// 2、dequant by vf
|
||||
LocalTensor<Q_T> kvDequantOutUb = stage0OutQue.AllocTensor<Q_T>();
|
||||
DequantKv(kvDequantOutUb, kvInUb, dealRow, constInfo);
|
||||
stage0InQue.FreeTensor(kvInUb);
|
||||
stage0OutQue.EnQue(kvDequantOutUb);
|
||||
kvDequantOutUb = stage0OutQue.DeQue<Q_T>();
|
||||
|
||||
// 3、copy kv out, ub -> l1
|
||||
if constexpr (IS_SPLIT_G) {
|
||||
CopyOutKvUb2Gm(v0ResGm, kvDequantOutUb, dealRow, s2Start, runInfo, constInfo);
|
||||
} else {
|
||||
CopyOutKvUb2L1(outputL1, kvDequantOutUb, dealRow, s2Start, runInfo, constInfo);
|
||||
}
|
||||
s2Start += dealRow;
|
||||
stage0OutQue.FreeTensor(kvDequantOutUb);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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> qsfaExpUb = 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>();
|
||||
|
||||
ProcessVec1SoftmaxDispatchQSFA(stage1CastTensor, mmRes, sumUb, maxUb, apiTmpBuffer, runInfo, constInfo);
|
||||
|
||||
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, qsfaExpUb, sumUb, maxUb, apiTmpBuffer, runInfo.halfMRealSize);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ProcessVec1SoftmaxDispatchQSFA(
|
||||
LocalTensor<Q_T> &stage1CastTensor, LocalTensor<T> &mmRes, LocalTensor<float> &sumUb,
|
||||
LocalTensor<float> &maxUb, LocalTensor<T> &apiTmpBuffer, RunInfo &runInfo,
|
||||
ConstInfo &constInfo)
|
||||
{
|
||||
if (runInfo.s2LoopCount == 0) {
|
||||
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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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.vec2MRealSize = runInfo.vec2MBaseSize;
|
||||
runInfo.vec2S1RealSize = runInfo.vec2S1BaseSize;
|
||||
int64_t qsfaVec2CalcSize = runInfo.vec2MRealSize * qsfaDTemplateAlign64;
|
||||
|
||||
LocalTensor<T> vec2ResUb = this->stage2OutBuf.template Get<T>();
|
||||
LocalTensor<T> mmRes = bmm2ResBuf.template GetTensor<T>();
|
||||
|
||||
WaitFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
if (unlikely(runInfo.s2LoopCount == 0)) {
|
||||
DataCopy(vec2ResUb, mmRes, qsfaVec2CalcSize);
|
||||
} else {
|
||||
LocalTensor<T> qsfaExpUb = softmaxExpBuf[runInfo.taskIdMod2].template Get<T>();
|
||||
if (runInfo.s2LoopCount < runInfo.s2LoopLimit) {
|
||||
FlashUpdateNew<T, Q_T, OUTPUT_T, qsfaDTemplateAlign64, false, false>(
|
||||
vec2ResUb, mmRes, vec2ResUb, qsfaExpUb, qsfaExpUb, runInfo.vec2MRealSize,
|
||||
qsfaDTemplateAlign64, 1.0, 1.0);
|
||||
} else {
|
||||
LocalTensor<float> sumUb = this->softmaxSumBuf[runInfo.multiCoreIdxMod2].template Get<float>();
|
||||
FlashUpdateLastNew<T, Q_T, OUTPUT_T, qsfaDTemplateAlign64, false, false>(
|
||||
vec2ResUb, mmRes, vec2ResUb, qsfaExpUb, qsfaExpUb, sumUb, runInfo.vec2MRealSize,
|
||||
qsfaDTemplateAlign64, 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, qsfaDTemplateAlign64, false>(
|
||||
vec2ResUb, vec2ResUb, sumUb, runInfo.vec2MRealSize, qsfaDTemplateAlign64, 1.0);
|
||||
}
|
||||
|
||||
this->CopyOutAttentionOut(runInfo, constInfo, vec2ResUb, 0, qsfaVec2CalcSize);
|
||||
}
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
template <typename VEC2_RES_T>
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::Bmm2DataCopyOut (RunInfo &runInfo, ConstInfo &constInfo,
|
||||
LocalTensor<VEC2_RES_T> &vec2ResUb, int64_t vec2S1Idx, int64_t qsfaVec2CalcSize)
|
||||
{
|
||||
LocalTensor<OUTPUT_T> attenOut;
|
||||
int64_t dSizeAligned64 = (int64_t)qsfaDTemplateAlign64;
|
||||
|
||||
attenOut.SetAddr(vec2ResUb.address_);
|
||||
Cast(attenOut, vec2ResUb, RoundMode::CAST_ROUND, qsfaVec2CalcSize);
|
||||
SetFlag<HardEvent::V_MTE3>(vToMte3Id[0]);
|
||||
WaitFlag<HardEvent::V_MTE3>(vToMte3Id[0]);
|
||||
|
||||
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 QSFAVectorService<TEMPLATE_ARGS>::CopyOutAttentionOut(
|
||||
RunInfo &runInfo, ConstInfo &constInfo, LocalTensor<VEC2_RES_T> &vec2ResUb,
|
||||
int64_t vec2S1Idx, int64_t qsfaVec2CalcSize)
|
||||
{
|
||||
this->Bmm2DataCopyOut(runInfo, constInfo, vec2ResUb, vec2S1Idx, qsfaVec2CalcSize);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitOutputSingleCore(ConstInfo &constInfo)
|
||||
{
|
||||
uint32_t coreNum = GetBlockNum();
|
||||
uint64_t totalOutputSize = 0;
|
||||
|
||||
// n2 = 1, n1 = gn2 = gSize
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND) {
|
||||
totalOutputSize = constInfo.bSize * constInfo.gSize * constInfo.s1Size * constInfo.dSizeV;
|
||||
} else if constexpr (LAYOUT_T == QSFA_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 (singleInitOutputSize > 0) {
|
||||
matmul::InitOutput<OUTPUT_T>(this->attentionOutGm[constInfo.aivIdx * singleCoreSize], singleInitOutputSize, 0);
|
||||
}
|
||||
}
|
||||
SyncAll();
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::CleanOutput(__gm__ uint8_t *attentionOut, ConstInfo &constInfo)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
this->attentionOutGm.SetGlobalBuffer((__gm__ OUTPUT_T *)attentionOut);
|
||||
if (constInfo.needInit == 1) {
|
||||
InitOutputSingleCore(constInfo);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitGlobalBuffer(__gm__ uint8_t *key,
|
||||
__gm__ uint8_t *value, __gm__ uint8_t *sparseIndices, __gm__ uint8_t *blockTable)
|
||||
{
|
||||
keyGm.SetGlobalBuffer((__gm__ KV_T *)(key));
|
||||
SparseIndicesGm.SetGlobalBuffer((__gm__ int32_t *)sparseIndices);
|
||||
if constexpr (isPa) {
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT
|
||||
__aicore__ inline void QSFAVectorService<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 QSFAVectorService<TEMPLATE_ARGS>::InitLocalBuffer(TPipe *pipe, ConstInfo &constInfo)
|
||||
{
|
||||
// ub buffer
|
||||
SoftmaxInitBuffer();
|
||||
|
||||
tPipe->InitBuffer(commonTBuf, 512); // commonTBuf内存申请512B
|
||||
tPipe->InitBuffer(stage0InQue, 2, dVTemplateTypeInput * 16 * sizeof(KV_T)); // V0阶段每次处理16个seq, 开2 buffer
|
||||
// 576: 模型特征维度(dSize)
|
||||
tPipe->InitBuffer(stage0OutQue, 2, 576 * (16 + 1) * sizeof(Q_T)); // kv输入D轴640, V0阶段每次处理16个seq, 开2 buffer
|
||||
|
||||
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) * qsfaDTemplateAlign64 * sizeof(T));
|
||||
|
||||
mte3ToVId[0] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
|
||||
mte3ToVId[1] = GetTPipePtr()->AllocEventID<HardEvent::MTE3_V>();
|
||||
|
||||
vToMte3Id[0] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
|
||||
vToMte3Id[1] = GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>();
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[0]);
|
||||
SetFlag<HardEvent::MTE3_V>(mte3ToVId[1]);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::InitCubeVecSharedParams(
|
||||
CVSharedParams &sharedParams, int32_t aicIdx, uint8_t subBlockIdx)
|
||||
{
|
||||
auto &sparseAttnSharedkvBaseParams = this->tilingData->baseParams;
|
||||
sharedParams.bSize = sparseAttnSharedkvBaseParams.batchSize;
|
||||
sharedParams.n2Size = 1;
|
||||
sharedParams.s1Size = sparseAttnSharedkvBaseParams.qSeqSize;
|
||||
sharedParams.s2Size = sparseAttnSharedkvBaseParams.seqSize;
|
||||
sharedParams.gSize = sparseAttnSharedkvBaseParams.nNumOfQInOneGroup;
|
||||
|
||||
sharedParams.sparseBlockCount = sparseAttnSharedkvBaseParams.sparseBlockCount;
|
||||
sharedParams.maskMode = sparseAttnSharedkvBaseParams.sparseMode;
|
||||
sharedParams.layoutType = sparseAttnSharedkvBaseParams.outputLayout;
|
||||
sharedParams.dSizeRope = 64; // 64: 编码维度
|
||||
sharedParams.softmaxScale = sparseAttnSharedkvBaseParams.scaleValue;
|
||||
sharedParams.dSize = 576; // 576: 模型特征维度(dSize)
|
||||
sharedParams.dSizeVInput = sparseAttnSharedkvBaseParams.dSizeVInput;
|
||||
sharedParams.usedCoreNum = this->tilingData->singleCoreParams.usedCoreNum;
|
||||
if constexpr (isPa) {
|
||||
sharedParams.blockSize = sparseAttnSharedkvBaseParams.blockSize;
|
||||
sharedParams.maxBlockNumPerBatch = sparseAttnSharedkvBaseParams.maxBlockNumPerBatch;
|
||||
}
|
||||
|
||||
sharedParams.isActualSeqLengthsNull = sparseAttnSharedkvBaseParams.isActualLenDimsNull;
|
||||
sharedParams.isActualSeqLengthsKVNull = sparseAttnSharedkvBaseParams.isActualLenDimsKVNull;
|
||||
|
||||
ComputeNeedInitQSFA(sharedParams);
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
if (subBlockIdx == 0) {
|
||||
auto qsfaTempTilingSSbuf = 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, ++qsfaTempTilingSSbuf, ++tempTiling) {
|
||||
*qsfaTempTilingSSbuf = *tempTiling;
|
||||
}
|
||||
|
||||
CrossCoreSetFlag<SYNC_MODE, PIPE_S>(15);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::ComputeNeedInitQSFA(
|
||||
CVSharedParams &sharedParams) const
|
||||
{
|
||||
sharedParams.needInit = 0;
|
||||
for (uint32_t bIdx = 0; bIdx < sharedParams.bSize; bIdx++) {
|
||||
int64_t s2Size;
|
||||
if constexpr (KV_LAYOUT_T == QSFA_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 == QSFA_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 == QSFA_LAYOUT::BSND && s1Size < sharedParams.s1Size)) {
|
||||
sharedParams.needInit = 1;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
TEMPLATES_DEF_NO_DEFAULT __aicore__ inline void QSFAVectorService<TEMPLATE_ARGS>::GetExtremeValue(
|
||||
T &negativeScalar)
|
||||
{
|
||||
uint32_t tmp1 = NEGATIVE_MIN_VALUE_FP32;
|
||||
negativeScalar = *((float *)&tmp1);
|
||||
}
|
||||
|
||||
TEMPLATES_DEF class QSFAVectorServiceDummy {
|
||||
public:
|
||||
__aicore__ inline QSFAVectorServiceDummy() {};
|
||||
__aicore__ inline void CleanOutput(__gm__ uint8_t *attentionOut, ConstInfo &constInfo) {}
|
||||
__aicore__ inline void InitGlobalBuffer(__gm__ uint8_t *key, __gm__ uint8_t *value, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *blockTable) {}
|
||||
__aicore__ inline void InitVecBlock(TPipe *pipe, const KvQuantSparseFlashAttentionTilingDataMla *__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 // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_VECTOR_MLA_H
|
||||
@@ -0,0 +1,147 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kv_quant_sparse_flash_attention_template_tiling_key.h"
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
#include "arch35/kv_quant_sparse_flash_attention_kernel_mla.h"
|
||||
#else
|
||||
#include "kv_quant_sparse_flash_attention_kernel_mla.h"
|
||||
#endif
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
#if defined(__DAV_C310_CUBE__)
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAMatmulService<__VA_ARGS__>, BaseApi::QSFAMatmulServiceDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAVectorServiceDummy<__VA_ARGS__>, BaseApi::QSFAVectorService<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, user, nullptr, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#else
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
using CubeBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAMatmulService<__VA_ARGS__>, BaseApi::QSFAMatmulServiceDummy<__VA_ARGS__>>::type; \
|
||||
using VecBlockType = typename std::conditional<g_coreType == AscendC::AIC, \
|
||||
BaseApi::QSFAVectorServiceDummy<__VA_ARGS__>, BaseApi::QSFAVectorService<__VA_ARGS__>>::type; \
|
||||
templateClass<CubeBlockType, VecBlockType> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tilingDataIn, tiling); \
|
||||
const tilingdataClass *__restrict tilingData = &tilingDataIn; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, user, tilingData, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#endif
|
||||
#else
|
||||
#define QSFA_OP_IMPL(templateClass, tilingdataClass, ...) \
|
||||
do { \
|
||||
templateClass<QSFAType<__VA_ARGS__>> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(tilingdataClass, tiling_data_in, tiling); \
|
||||
const tilingdataClass *__restrict tiling_data = &tiling_data_in; \
|
||||
op.Init(query, key, value, sparseIndices, keyScale, valueScale, blocktable, \
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, \
|
||||
attentionOut, softmaxMax, softmaxSum, user, tiling_data, tiling, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
template<int FLASH_DECODE, int PAGE_ATTENTION, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE, int IS_SPLIT_G>
|
||||
__aicore__ inline void DispatchKernelDtype310(
|
||||
__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *keyScale, __gm__ uint8_t *valueScale,
|
||||
__gm__ uint8_t *blocktable, __gm__ uint8_t *actualSeqLengthsQuery,
|
||||
__gm__ uint8_t *actualSeqLengthsKV, __gm__ uint8_t *attentionOut,
|
||||
__gm__ uint8_t *user, __gm__ uint8_t *tiling, TPipe &tPipe)
|
||||
{
|
||||
if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_FLOAT8_E4M3FN &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, fp8_e4m3fn_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_HIFLOAT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, hifloat8_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_BF16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_BF16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
bfloat16_t, int8_t, float, bfloat16_t, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_FLOAT8_E4M3FN &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, fp8_e4m3fn_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_HIFLOAT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, hifloat8_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
} else if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(BaseApi::KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla,
|
||||
half, int8_t, float, half, FLASH_DECODE, PAGE_ATTENTION,
|
||||
static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
static_cast<QSFATemplateMode>(TEMPLATE_MODE), IS_SPLIT_G);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
template<int FLASH_DECODE, int PAGE_ATTENTION, int LAYOUT_T, int KV_LAYOUT_T, int TEMPLATE_MODE, int IS_SPLIT_G>
|
||||
__global__ __aicore__ void
|
||||
kv_quant_sparse_flash_attention(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *value,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t* keyScale, __gm__ uint8_t* valueScale,
|
||||
__gm__ uint8_t *blocktable, __gm__ uint8_t *actualSeqLengthsQuery,
|
||||
__gm__ uint8_t *actualSeqLengthsKV, __gm__ uint8_t *attentionOut,
|
||||
__gm__ uint8_t *softmaxMax, __gm__ uint8_t *softmaxSum,
|
||||
__gm__ uint8_t *workspace, __gm__ uint8_t *tiling)
|
||||
{
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
|
||||
TPipe tPipe;
|
||||
__gm__ uint8_t *user = GetUserWorkspace(workspace);
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
DispatchKernelDtype310<FLASH_DECODE, PAGE_ATTENTION, LAYOUT_T, KV_LAYOUT_T, TEMPLATE_MODE, IS_SPLIT_G>(
|
||||
query, key, value, sparseIndices, keyScale, valueScale, blocktable,
|
||||
actualSeqLengthsQuery, actualSeqLengthsKV, attentionOut, user, tiling, tPipe);
|
||||
#else
|
||||
if constexpr (ORIG_DTYPE_QUERY == DT_FLOAT16 && ORIG_DTYPE_KEY == DT_INT8 &&
|
||||
ORIG_DTYPE_ATTENTION_OUT == DT_FLOAT16) {
|
||||
QSFA_OP_IMPL(KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla, half, int8_t,
|
||||
half, FLASH_DECODE, static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
TEMPLATE_MODE);
|
||||
} else { // bf16
|
||||
QSFA_OP_IMPL(KvQuantSparseFlashAttentionMla, KvQuantSparseFlashAttentionTilingDataMla, bfloat16_t, int8_t,
|
||||
bfloat16_t, FLASH_DECODE, static_cast<QSFA_LAYOUT>(LAYOUT_T), static_cast<QSFA_LAYOUT>(KV_LAYOUT_T),
|
||||
TEMPLATE_MODE);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_common.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
|
||||
using namespace AscendC;
|
||||
// 将isCheckTiling设置为false, 输入输出的max&sum&exp的shape为(m, 1)
|
||||
constexpr SoftmaxConfig QSFA_SOFTMAX_FLASHV2_CFG_WITHOUT_BRC = {false, 0, 0, SoftmaxMode::SOFTMAX_OUTPUT_WITHOUT_BRC};
|
||||
|
||||
enum class QSFA_LAYOUT {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2,
|
||||
};
|
||||
|
||||
enum class QUANT_MODE {
|
||||
PER_CHANNEL = 0, // GQA支持
|
||||
PER_TOKEN_HEAD = 1, // GQA支持
|
||||
PER_TILE = 2, // MLA支持
|
||||
};
|
||||
|
||||
enum class ATTENTION_MODE {
|
||||
GQA_MHA = 0, // QKV headDim相等
|
||||
MLA_NATIVE = 1, // Dn=128, Dr=64
|
||||
MLA_ABSORB = 2, // Dn=512, Dr=64
|
||||
};
|
||||
|
||||
enum class QUANT_SCALE_REPO_MODE {
|
||||
SEPARATE = 0, // 分开存储
|
||||
COMBINE = 1, // 合并存储,量化模式是PER_TOKEN_HEAD/PER_TILE时支持COMBINE模式,参数顺序为:Nope+Rope+DequantScale
|
||||
};
|
||||
|
||||
template <typename Q_T, typename KV_T, typename OUT_T, const bool FLASH_DECODE = false,
|
||||
QSFA_LAYOUT LAYOUT_T = QSFA_LAYOUT::BSND, QSFA_LAYOUT KV_LAYOUT_T = QSFA_LAYOUT::BSND,
|
||||
const int TEMPLATE_MODE = C_TEMPLATE, typename... Args>
|
||||
struct QSFAType {
|
||||
using queryType = Q_T;
|
||||
using kvType = KV_T;
|
||||
using kRopeType = Q_T;
|
||||
using outputType = OUT_T;
|
||||
static constexpr bool flashDecode = FLASH_DECODE;
|
||||
static constexpr QSFA_LAYOUT layout = LAYOUT_T;
|
||||
static constexpr QSFA_LAYOUT kvLayout = KV_LAYOUT_T;
|
||||
static constexpr int templateMode = TEMPLATE_MODE;
|
||||
static constexpr bool pageAttention = (KV_LAYOUT_T == QSFA_LAYOUT::PA_BSND);
|
||||
};
|
||||
|
||||
// ================================Util functions==================================
|
||||
template <typename T> __aicore__ inline T QSFAAlign(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
template <typename T> __aicore__ inline size_t BlockAlign(size_t s)
|
||||
{
|
||||
if constexpr (IsSameType<T, int4b_t>::value) {
|
||||
return (s + 63) / 64 * 64;
|
||||
}
|
||||
size_t n = (32 / sizeof(T));
|
||||
return (s + n - 1) / n * n;
|
||||
}
|
||||
|
||||
template <typename T1, typename T2> __aicore__ inline T1 Min(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (b) : (a);
|
||||
}
|
||||
|
||||
struct RunInfo {
|
||||
uint32_t loop;
|
||||
uint32_t bIdx;
|
||||
uint32_t gIdx;
|
||||
uint32_t s1Idx;
|
||||
uint32_t s2Idx;
|
||||
uint32_t bn2IdxInCurCore;
|
||||
uint32_t curSInnerLoopTimes;
|
||||
uint32_t s2BatchOffset;
|
||||
|
||||
uint64_t tndBIdxOffsetForQ;
|
||||
uint64_t tndBIdxOffsetForKV;
|
||||
uint64_t tensorAOffset;
|
||||
uint64_t tensorBOffset;
|
||||
uint64_t tensorARopeOffset;
|
||||
uint64_t tensorBRopeOffset;
|
||||
uint64_t attenOutOffset;
|
||||
uint64_t topKBaseOffset;
|
||||
uint64_t attenMaskOffset;
|
||||
|
||||
uint32_t actualSingleProcessSInnerSize;
|
||||
uint32_t actualSingleProcessSInnerSizeAlign;
|
||||
uint32_t gSize;
|
||||
uint32_t s1Size;
|
||||
uint32_t s2Size;
|
||||
uint32_t mSize;
|
||||
uint32_t mSizeV;
|
||||
uint32_t mSizeVStart;
|
||||
uint32_t tndIsS2SplitCore;
|
||||
uint32_t tndCoreStartKVSplitPos;
|
||||
bool isBmm2Output;
|
||||
bool isValid = false;
|
||||
bool isFirstSInnerLoop;
|
||||
bool isChangeBatch;
|
||||
static constexpr uint32_t n2Idx = 0;
|
||||
|
||||
uint64_t actS1Size = 1;
|
||||
uint64_t curActualSeqLenOri = 0ULL;
|
||||
uint64_t actS2Size = 1;
|
||||
|
||||
uint32_t gS1Idx;
|
||||
uint32_t actMBaseSize;
|
||||
int32_t nextTokensPerBatch = 0;
|
||||
bool isLastS2Loop;
|
||||
uint8_t resv[3];
|
||||
int64_t threshold;
|
||||
};
|
||||
|
||||
struct ConstInfo {
|
||||
// CUBE与VEC核间同步的模式
|
||||
static constexpr uint32_t QSFA_SYNC_MODE2 = 2;
|
||||
// BUFFER的字节数
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512;
|
||||
// FP32的0值和极大值
|
||||
static constexpr float FLOAT_ZERO = 0;
|
||||
static constexpr float FLOAT_MAX = 3.402823466e+38F;
|
||||
|
||||
// preLoad的总次数
|
||||
uint32_t preLoadNum = 0U;
|
||||
uint32_t nBufferMBaseSize = 0U;
|
||||
// CUBE和VEC的核间同步EventID
|
||||
uint32_t syncV0C1 = 0U;
|
||||
uint32_t syncC1V1 = 0U;
|
||||
uint32_t syncV1C2 = 0U;
|
||||
uint32_t syncC2V2 = 0U;
|
||||
uint32_t syncC2V1 = 0U;
|
||||
uint32_t syncV1NupdateC2 = 0U;
|
||||
|
||||
uint32_t mmResUbSize = 0U; // Matmul1输出结果GM上的大小
|
||||
uint32_t vec1ResUbSize = 0U; // Vector1输出结果GM上的大小
|
||||
uint32_t bmm2ResUbSize = 0U; // Matmul2输出结果GM上的大小
|
||||
uint64_t gSize = 0ULL;
|
||||
uint64_t batchSize = 0ULL;
|
||||
uint64_t qHeadNum = 0ULL;
|
||||
uint64_t kvHeadNum;
|
||||
uint64_t headDim;
|
||||
uint64_t headDimRope;
|
||||
uint64_t combineHeadDim; // quantScaleRepoMode为Combine模式时=headDim+headDimRope, 否则=headDim
|
||||
uint64_t kvSeqSize = 0ULL; // kv最大S长度
|
||||
uint64_t qSeqSize = 1ULL; // q最大S长度
|
||||
int64_t kvCacheBlockSize = 0; // PA场景的block size
|
||||
uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number
|
||||
uint32_t splitKVNum = 0U; // S2核间切分的切分份数
|
||||
QSFA_LAYOUT outputLayout; // 输出的Transpose格式
|
||||
uint32_t sparseMode = 0;
|
||||
bool returnSoftmaxLse = false;
|
||||
bool needInit = false;
|
||||
|
||||
// FlashDecoding
|
||||
uint64_t combineLseOffset = 0ULL;
|
||||
uint64_t combineAccumOutOffset = 0ULL;
|
||||
uint32_t actualCombineLoopSize = 0U; // FlashDecoding场景, S2在核间切分的最大份数
|
||||
|
||||
uint32_t actualLenDimsQ = 0U; // query的actualSeqLength 的维度
|
||||
uint32_t actualLenDimsKV = 0U; // KV 的actualSeqLength 的维度
|
||||
|
||||
// TND
|
||||
uint32_t s2Start = 0U; // TND场景下,S2的起始位置
|
||||
uint32_t s2End = 0U; // 单核TND场景下S2循环index上限
|
||||
|
||||
uint32_t bN2Start = 0U;
|
||||
uint32_t bN2End = 0U;
|
||||
uint32_t gS1Start = 0U;
|
||||
uint32_t gS1End = 0U;
|
||||
|
||||
uint32_t mBaseSize = 1ULL;
|
||||
uint32_t s2BaseSize = 1ULL;
|
||||
|
||||
uint32_t tndFDCoreArrLen = 0U; // TNDFlashDecoding相关分核信息array的长度
|
||||
uint32_t coreStartKVSplitPos = 0U; // TNDFlashDecoding kv起始位置
|
||||
|
||||
// sparse attr
|
||||
uint32_t sparseBlockCount = 0;
|
||||
int64_t sparseBlockSize = 0;
|
||||
|
||||
// attention模式与量化模式
|
||||
ATTENTION_MODE attentionMode = ATTENTION_MODE::MLA_ABSORB;
|
||||
QUANT_MODE keyQuantMode = QUANT_MODE::PER_TILE;
|
||||
QUANT_MODE valueQuantMode = QUANT_MODE::PER_TILE;
|
||||
QUANT_SCALE_REPO_MODE quantScaleRepoMode = QUANT_SCALE_REPO_MODE::COMBINE;
|
||||
uint64_t tileSize = 128ULL;
|
||||
};
|
||||
|
||||
struct MSplitInfo {
|
||||
uint32_t nBufferIdx = 0U;
|
||||
uint32_t nBufferStartM = 0U;
|
||||
uint32_t nBufferDealM = 0U;
|
||||
uint32_t vecStartM = 0U;
|
||||
uint32_t vecDealM = 0U;
|
||||
};
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_COMMON_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,943 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_flash_attention_service_cube_mla.h
|
||||
* \brief use 7 buffer for matmul l1, better pipeline
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
#define KV_QUANT_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 "kv_quant_sparse_flash_attention_common.h"
|
||||
|
||||
struct Position {
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx;
|
||||
uint32_t s2Idx;
|
||||
uint32_t dIdx;
|
||||
};
|
||||
|
||||
struct PAShape {
|
||||
uint32_t blockSize;
|
||||
uint32_t headNum; // 一般为kv的head num,对应n2
|
||||
uint32_t headDim; // mla下rope为64,nope为512, 对应d
|
||||
uint32_t maxblockNumPerBatch; // block table 每一行的最大个数
|
||||
uint32_t actHeadDim; // 实际拷贝col大小,考虑到N切块 s*d, 对应d
|
||||
uint32_t copyRowNum; // 总共要拷贝的行数
|
||||
uint32_t copyRowNumAlign;
|
||||
};
|
||||
|
||||
// 场景:query、queryRope、key、value GM to L1
|
||||
// GM按ND格式存储
|
||||
// L1按NZ格式存储
|
||||
// GM的行、列、列的stride
|
||||
template <typename T>
|
||||
__aicore__ inline void DataCopyGmNDToL1(LocalTensor<T> &l1Tensor, GlobalTensor<T> &gmTensor,
|
||||
uint32_t rowAct, uint32_t rowAlign,
|
||||
uint32_t col, // D
|
||||
uint32_t colStride) // D or N*D
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = rowAct; // nd矩阵的行数
|
||||
// T为int4场景下,dValue = col / 2,srcDValue = colStride / 2
|
||||
nd2nzPara.srcDValue = colStride; // 同一nd矩阵相邻行起始地址间的偏移
|
||||
nd2nzPara.dValue = col; // nd矩阵的列数
|
||||
nd2nzPara.dstNzC0Stride = rowAlign;
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
DataCopy(l1Tensor, gmTensor, nd2nzPara);
|
||||
}
|
||||
|
||||
/*
|
||||
适用PA数据从GM拷贝到L1,支持ND、NZ数据;
|
||||
PA的layout分 BNBD(blockNum,N,blockSize,D) BBH(blockNum,blockSize,N*D
|
||||
BSH\BSND\TND 为BBH
|
||||
shape.copyRowNumAlign 需要16字节对齐,如拷贝k矩阵,一次拷贝128*512,遇到尾块 10*512 需对齐到16*512
|
||||
*/
|
||||
template <typename T, QSFA_LAYOUT SRC_LAYOUT>
|
||||
__aicore__ inline void DataCopyPA(LocalTensor<T> &dstTensor, // l1
|
||||
GlobalTensor<T> &srcTensor, // gm
|
||||
GlobalTensor<int32_t> &blockTableGm,
|
||||
const PAShape &shape, // blockSize, headNum, headDim
|
||||
const Position &startPos) // bacthIdx nIdx curSeqIdx
|
||||
{
|
||||
uint32_t copyFinishRowCnt = 0;
|
||||
uint64_t blockTableBaseOffset = startPos.bIdx * shape.maxblockNumPerBatch;
|
||||
uint32_t curS2Idx = startPos.s2Idx;
|
||||
uint32_t blockElementCnt = 32 / sizeof(T);
|
||||
while (copyFinishRowCnt < shape.copyRowNum) {
|
||||
uint64_t blockIdOffset = curS2Idx / shape.blockSize; // 获取block table上的索引
|
||||
uint64_t reaminRowCnt = curS2Idx % shape.blockSize; // 获取在单个块上超出的行数
|
||||
// 从block table上的获取编号
|
||||
uint64_t idInBlockTable = blockTableGm.GetValue(blockTableBaseOffset + blockIdOffset);
|
||||
// 计算可以拷贝行数
|
||||
uint32_t copyRowCnt = shape.blockSize - reaminRowCnt; // 一次只能处理一个Block
|
||||
if (copyFinishRowCnt + copyRowCnt > shape.copyRowNum) {
|
||||
copyRowCnt = shape.copyRowNum - copyFinishRowCnt; // 一个block未拷满
|
||||
}
|
||||
uint64_t offset = idInBlockTable * shape.blockSize * shape.headNum * shape.headDim ; // PA的偏移
|
||||
|
||||
uint64_t dStride = shape.headDim;
|
||||
if constexpr (SRC_LAYOUT == QSFA_LAYOUT::BSND || SRC_LAYOUT == QSFA_LAYOUT::TND) {
|
||||
offset += (uint64_t)(startPos.n2Idx * shape.headDim) +
|
||||
reaminRowCnt * shape.headDim * shape.headNum + startPos.dIdx;
|
||||
dStride = shape.headDim * shape.headNum;
|
||||
} else {
|
||||
offset += (uint64_t)(startPos.n2Idx * shape.headDim * shape.blockSize) +
|
||||
reaminRowCnt * shape.headDim + startPos.dIdx;
|
||||
}
|
||||
|
||||
uint32_t srcDValue = dStride;
|
||||
uint32_t dValue = shape.actHeadDim;
|
||||
LocalTensor<T> tmpDstTensor = dstTensor[copyFinishRowCnt * blockElementCnt];
|
||||
GlobalTensor<T> tmpSrcTensor = srcTensor[offset];
|
||||
|
||||
DataCopyGmNDToL1<T>(tmpDstTensor, tmpSrcTensor, copyRowCnt, shape.copyRowNumAlign, dValue, srcDValue);
|
||||
copyFinishRowCnt += copyRowCnt;
|
||||
curS2Idx += copyRowCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QSFAT> class QSFAMatmulService {
|
||||
public:
|
||||
// 中间计算数据类型为float, 高精度模式
|
||||
using T = float;
|
||||
using Q_T = typename QSFAT::queryType;
|
||||
using KV_T = typename QSFAT::kvType;
|
||||
using K_ROPE_T = typename QSFAT::kRopeType;
|
||||
using OUT_T = typename QSFAT::outputType;
|
||||
using MM_OUT_T = T;
|
||||
|
||||
__aicore__ inline QSFAMatmulService(){};
|
||||
__aicore__ inline void InitParams(const ConstInfo &constInfo);
|
||||
__aicore__ inline void InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<Q_T> qRopeGm,
|
||||
GlobalTensor<KV_T> keyGm, GlobalTensor<K_ROPE_T> kRopeGm,
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm);
|
||||
__aicore__ inline void InitMm2GlobalTensor(GlobalTensor<K_ROPE_T> vec1ResGm, GlobalTensor<KV_T> valueGm,
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm);
|
||||
__aicore__ inline void InitPageAttentionInfo(const GlobalTensor<K_ROPE_T>& kvMergeGm,
|
||||
GlobalTensor<int32_t> blockTableGm, GlobalTensor<int32_t> topKGm,
|
||||
uint32_t blockSize, uint32_t maxBlockNumPerBatch);
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void UpdateKey(GlobalTensor<KV_T> keyGm);
|
||||
__aicore__ inline void UpdateValue(GlobalTensor<KV_T> valueGm);
|
||||
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
__aicore__ inline void CalcTopKBlockInfo(const RunInfo &info, uint32_t &curTopKIdx,
|
||||
uint64_t &curOffsetInSparseBlock, uint32_t curSeqIdx,
|
||||
uint32_t ©RowCnt, uint64_t &idInTopK);
|
||||
__aicore__ inline void ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo);
|
||||
__aicore__ inline void ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo);
|
||||
|
||||
private:
|
||||
static constexpr bool PAGE_ATTENTION = QSFAT::pageAttention;
|
||||
static constexpr int TEMPLATE_MODE = QSFAT::templateMode;
|
||||
static constexpr bool FLASH_DECODE = QSFAT::flashDecode;
|
||||
static constexpr QSFA_LAYOUT LAYOUT_T = QSFAT::layout;
|
||||
static constexpr QSFA_LAYOUT KV_LAYOUT_T = QSFAT::kvLayout;
|
||||
|
||||
static constexpr uint32_t M_SPLIT_SIZE = 128; // m方向切分
|
||||
static constexpr uint32_t N_SPLIT_SIZE = 128; // n方向切分
|
||||
static constexpr uint32_t N_WORKSPACE_SIZE = 512; // n方向切分
|
||||
static constexpr uint32_t K_SPLIT_SIZE = 288; // K方向切分
|
||||
|
||||
static constexpr uint32_t L1_BLOCK_SIZE = (64 * (512 + 64) * sizeof(Q_T));
|
||||
static constexpr uint32_t L1_BLOCK_OFFSET = 64 * (512 + 64); // 72K的元素个数
|
||||
|
||||
static constexpr uint32_t L0A_PP_SIZE = (32 * 1024);
|
||||
static constexpr uint32_t L0B_PP_SIZE = (32 * 1024);
|
||||
static constexpr uint32_t L0C_PP_SIZE = (64 * 1024);
|
||||
|
||||
// m <> mte1 EventID
|
||||
static constexpr uint32_t L0AB_EVENT0 = EVENT_ID3;
|
||||
static constexpr uint32_t L0AB_EVENT1 = EVENT_ID4;
|
||||
|
||||
// mte2 <> mte1 EventID
|
||||
// L1 3buf, 使用3个eventId
|
||||
static constexpr uint32_t L1_EVENT0 = EVENT_ID2;
|
||||
static constexpr uint32_t L1_EVENT1 = EVENT_ID3;
|
||||
static constexpr uint32_t L1_EVENT2 = EVENT_ID4;
|
||||
static constexpr uint32_t L1_EVENT3 = EVENT_ID5;
|
||||
static constexpr uint32_t L1_EVENT4 = EVENT_ID6;
|
||||
static constexpr uint32_t L1_EVENT5 = EVENT_ID7;
|
||||
static constexpr uint32_t L1_EVENT6 = EVENT_ID1;
|
||||
|
||||
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
|
||||
static constexpr uint32_t mte21QPIds[4] = {L1_EVENT0, L1_EVENT1, L1_EVENT2, L1_EVENT3}; // mte12复用
|
||||
static constexpr uint32_t mte21KVIds[3] = {L1_EVENT4, L1_EVENT5, L1_EVENT6};
|
||||
|
||||
static constexpr uint32_t BLOCK_ELEMENT_NUM = ConstInfo::BUFFER_SIZE_BYTE_32B / sizeof(K_ROPE_T);
|
||||
|
||||
uint32_t kvCacheBlockSize = 0;
|
||||
uint32_t maxBlockNumPerBatch = 0;
|
||||
ConstInfo constInfo{};
|
||||
|
||||
// L1分成3块buf, 用于记录
|
||||
uint32_t qpL1BufIter = 0;
|
||||
uint32_t kvL1BufIter = -1;
|
||||
uint32_t abL0BufIter = 0;
|
||||
uint32_t cL0BufIter = 0;
|
||||
|
||||
// mm1
|
||||
GlobalTensor<Q_T> queryGm;
|
||||
GlobalTensor<Q_T> qRopeGm;
|
||||
GlobalTensor<KV_T> keyGm;
|
||||
GlobalTensor<K_ROPE_T> kRopeGm;
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm;
|
||||
GlobalTensor<K_ROPE_T> kvMergeGm_;
|
||||
|
||||
// mm2
|
||||
GlobalTensor<K_ROPE_T> vec1ResGm;
|
||||
GlobalTensor<KV_T> valueGm;
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm;
|
||||
GlobalTensor<OUT_T> attentionOutGm;
|
||||
|
||||
// block_table
|
||||
GlobalTensor<int32_t> topKGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
|
||||
TBuf<TPosition::A1> bufQPL1;
|
||||
TBuf<TPosition::A1> bufKVL1;
|
||||
TBuf<TPosition::A2> tmpBufL0A;
|
||||
TBuf<TPosition::B2> tmpBufL0B;
|
||||
TBuf<TPosition::CO1> tmpBufL0C;
|
||||
|
||||
LocalTensor<K_ROPE_T> aL0TensorPingPong;
|
||||
LocalTensor<K_ROPE_T> bL0TensorPingPong;
|
||||
LocalTensor<MM_OUT_T> cL0TensorPingPong;
|
||||
LocalTensor<Q_T> l1QPTensor;
|
||||
LocalTensor<Q_T> l1KVTensor;
|
||||
|
||||
// L0AB m <> mte1 EventID
|
||||
__aicore__ inline uint32_t Mte1MmABEventId(uint32_t idx)
|
||||
{
|
||||
return (L0AB_EVENT0 + idx);
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t GetQPL1RealIdx(uint32_t mIdx, uint32_t k1Idx)
|
||||
{
|
||||
uint32_t idxMap[] = {0, 2}; // 确保0块和1块连在一起, 2和3块连在一起, 来保证同一m块的地址相连
|
||||
return idxMap[mIdx % 2] + k1Idx;
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyGmToL1(LocalTensor<K_ROPE_T> &l1Tensor, GlobalTensor<K_ROPE_T> &gmSrcTensor,
|
||||
uint32_t srcN, uint32_t srcD, uint32_t srcDstride);
|
||||
__aicore__ inline void CopyInMm1AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct, uint32_t headSize, uint32_t headOffset);
|
||||
__aicore__ inline void CopyInMm1ARopeToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct);
|
||||
__aicore__ inline void CopyInMm1BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize);
|
||||
__aicore__ inline void CopyInMm1BRopeToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize);
|
||||
__aicore__ inline void CopyInMm2AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t subMSizeAct, uint32_t nSize, uint32_t nOffset);
|
||||
__aicore__ inline void CopyInMm2BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t valueGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t copyStartColumnCount,
|
||||
uint32_t copyColumnCount);
|
||||
__aicore__ inline void LoadDataMm1A(LocalTensor<K_ROPE_T> &aL0Tensor, LocalTensor<K_ROPE_T> &aL1Tensor,
|
||||
uint32_t idx, uint32_t kSplitSize, uint32_t mSize, uint32_t kSize);
|
||||
__aicore__ inline void LoadDataMm1B(LocalTensor<K_ROPE_T> &bL0Tensor, LocalTensor<K_ROPE_T> &bL1Tensor,
|
||||
uint32_t idx, uint32_t kSplitSize, uint32_t kSize, uint32_t nSize);
|
||||
};
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::InitParams(const ConstInfo &constInfo)
|
||||
{
|
||||
this->constInfo = constInfo;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitMm1GlobalTensor(GlobalTensor<Q_T> queryGm, GlobalTensor<Q_T> qRopeGm,
|
||||
GlobalTensor<KV_T> keyGm, GlobalTensor<K_ROPE_T> kRopeGm,
|
||||
GlobalTensor<MM_OUT_T> mm1ResGm)
|
||||
{
|
||||
// mm1
|
||||
this->queryGm = queryGm;
|
||||
this->qRopeGm = qRopeGm;
|
||||
this->keyGm = keyGm;
|
||||
this->kRopeGm = kRopeGm;
|
||||
this->mm1ResGm = mm1ResGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitMm2GlobalTensor(GlobalTensor<K_ROPE_T> vec1ResGm, GlobalTensor<KV_T> valueGm,
|
||||
GlobalTensor<MM_OUT_T> mm2ResGm, GlobalTensor<OUT_T> attentionOutGm)
|
||||
{
|
||||
// mm2
|
||||
this->vec1ResGm = vec1ResGm;
|
||||
this->valueGm = valueGm;
|
||||
this->mm2ResGm = mm2ResGm;
|
||||
this->attentionOutGm = attentionOutGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::InitPageAttentionInfo(const GlobalTensor<K_ROPE_T>& kvMergeGm,
|
||||
GlobalTensor<int32_t> blockTableGm, GlobalTensor<int32_t> topKGm,
|
||||
uint32_t blockSize, uint32_t maxBlockNumPerBatch)
|
||||
{
|
||||
this->blockTableGm = blockTableGm;
|
||||
this->topKGm = topKGm;
|
||||
this->kvCacheBlockSize = blockSize;
|
||||
this->maxBlockNumPerBatch = maxBlockNumPerBatch;
|
||||
this->kvMergeGm_ = kvMergeGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(bufQPL1, L1_BLOCK_SIZE * 4); // (64K + 8K) * 4
|
||||
l1QPTensor = bufQPL1.Get<Q_T>();
|
||||
pipe->InitBuffer(bufKVL1, L1_BLOCK_SIZE * 3); // (64K + 8K) * 3
|
||||
l1KVTensor = bufKVL1.Get<K_ROPE_T>();
|
||||
|
||||
// L0A
|
||||
pipe->InitBuffer(tmpBufL0A, L0A_PP_SIZE * 2); // 64K
|
||||
aL0TensorPingPong = tmpBufL0A.Get<K_ROPE_T>();
|
||||
// L0B
|
||||
pipe->InitBuffer(tmpBufL0B, L0B_PP_SIZE * 2); // 64K
|
||||
bL0TensorPingPong = tmpBufL0B.Get<K_ROPE_T>();
|
||||
// L0C
|
||||
pipe->InitBuffer(tmpBufL0C, L0C_PP_SIZE * 2); // 128K
|
||||
cL0TensorPingPong = tmpBufL0C.Get<MM_OUT_T>();
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::UpdateKey(GlobalTensor<KV_T> keyGm)
|
||||
{
|
||||
this->keyGm = keyGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::UpdateValue(GlobalTensor<KV_T> valueGm)
|
||||
{
|
||||
this->valueGm = valueGm;
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::AllocEventID()
|
||||
{
|
||||
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
|
||||
SetFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
|
||||
}
|
||||
|
||||
template <typename QSFAT> __aicore__ inline void QSFAMatmulService<QSFAT>::FreeEventID()
|
||||
{
|
||||
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT0);
|
||||
WaitFlag<HardEvent::M_MTE1>(L0AB_EVENT1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT2);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT3);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT4);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT5);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(L1_EVENT6);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyGmToL1(LocalTensor<K_ROPE_T> &l1Tensor,
|
||||
GlobalTensor<K_ROPE_T> &gmSrcTensor, uint32_t srcN,
|
||||
uint32_t srcD, uint32_t srcDstride)
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.dValue = srcD;
|
||||
nd2nzPara.nValue = srcN; // 行数
|
||||
nd2nzPara.srcDValue = srcDstride;
|
||||
nd2nzPara.dstNzC0Stride = (srcN + 15) / 16 * 16; // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
DataCopy(l1Tensor, gmSrcTensor, nd2nzPara);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm1AToL1(LocalTensor<K_ROPE_T> &l1Tensor, const RunInfo &info,
|
||||
uint32_t mSeqIdx, uint32_t mSizeAct,
|
||||
uint32_t headSize, uint32_t headOffset)
|
||||
{
|
||||
auto srcGm = queryGm[info.tensorAOffset + mSeqIdx * constInfo.combineHeadDim + headOffset];
|
||||
CopyGmToL1(l1Tensor, srcGm, mSizeAct, headSize, headSize);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm1ARopeToL1(LocalTensor<K_ROPE_T> &l1Tensor,
|
||||
const RunInfo &info, uint32_t mSeqIdx,
|
||||
uint32_t mSizeAct)
|
||||
{
|
||||
auto srcGm = qRopeGm[info.tensorARopeOffset + mSeqIdx * constInfo.headDimRope];
|
||||
CopyGmToL1(l1Tensor, srcGm, mSizeAct, constInfo.headDimRope, constInfo.headDimRope);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::CopyInMm1BToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t keyGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize)
|
||||
{
|
||||
uint64_t dStride = constInfo.headDim;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
dStride = constInfo.headDim * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = headSize;
|
||||
mm1Nd2NzParamsForB.srcDValue = dStride;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], keyGm[keyGmBaseOffset], mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void
|
||||
QSFAMatmulService<QSFAT>::CopyInMm1BRopeToL1(LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t kRopeGmBaseOffset,
|
||||
uint32_t copyTotalRowCntAlign, uint32_t copyStartRowCnt,
|
||||
uint32_t nActCopyRowCount, uint32_t headSize)
|
||||
{
|
||||
uint64_t dStride = constInfo.headDimRope;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
dStride = constInfo.headDimRope * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = headSize;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.srcDValue = dStride;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], kRopeGm[kRopeGmBaseOffset], mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::LoadDataMm1A(LocalTensor<K_ROPE_T> &aL0Tensor,
|
||||
LocalTensor<K_ROPE_T> &aL1Tensor, uint32_t idx,
|
||||
uint32_t kSplitSize, uint32_t mSize, uint32_t kSize)
|
||||
{
|
||||
LocalTensor<K_ROPE_T> srcTensor = aL1Tensor[mSize * kSplitSize * idx];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParams;
|
||||
// SetFmatrixParams
|
||||
loadData3DParams.l1H = mSize / 16; // Hin=M1=8
|
||||
loadData3DParams.l1W = 16; // Win=M0
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
// SetLoadToA0Params
|
||||
loadData3DParams.mExtension = mSize; // M
|
||||
loadData3DParams.kExtension = kSize; // K
|
||||
loadData3DParams.mStartPt = 0;
|
||||
loadData3DParams.kStartPt = 0;
|
||||
loadData3DParams.strideH = 1;
|
||||
loadData3DParams.strideW = 1;
|
||||
loadData3DParams.filterW = 1;
|
||||
loadData3DParams.filterSizeW = (1 >> 8) & 255;
|
||||
loadData3DParams.filterH = 1;
|
||||
loadData3DParams.filterSizeH = (1 >> 8) & 255;
|
||||
loadData3DParams.dilationFilterH = 1;
|
||||
loadData3DParams.dilationFilterW = 1;
|
||||
loadData3DParams.fMatrixCtrl = 0;
|
||||
loadData3DParams.channelSize = kSize; // Cin=K
|
||||
loadData3DParams.enTranspose = 0;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(aL0Tensor, srcTensor, loadData3DParams);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::LoadDataMm1B(LocalTensor<K_ROPE_T> &l0Tensor,
|
||||
LocalTensor<K_ROPE_T> &l1Tensor, uint32_t idx,
|
||||
uint32_t kSplitSize, uint32_t kSize, uint32_t nSize)
|
||||
{
|
||||
// N 方向全载
|
||||
LocalTensor<K_ROPE_T> srcTensor = l1Tensor[nSize * kSplitSize * idx];
|
||||
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0;
|
||||
loadData2DParams.repeatTimes = (nSize + 15) / 16 * kSize / (32 / sizeof(K_ROPE_T));
|
||||
loadData2DParams.srcStride = 1;
|
||||
loadData2DParams.dstGap = 0;
|
||||
loadData2DParams.ifTranspose = false;
|
||||
LoadData(l0Tensor, srcTensor, loadData2DParams);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm2AToL1(LocalTensor<K_ROPE_T> &aL1Tensor, const RunInfo &info,
|
||||
uint32_t mSeqIdx, uint32_t subMSizeAct,
|
||||
uint32_t nSize, uint32_t nOffset)
|
||||
{
|
||||
auto srcGm = vec1ResGm[(info.loop % constInfo.preLoadNum) * constInfo.mmResUbSize +
|
||||
mSeqIdx * info.actualSingleProcessSInnerSizeAlign + nOffset];
|
||||
CopyGmToL1(aL1Tensor, srcGm, subMSizeAct, nSize, info.actualSingleProcessSInnerSizeAlign);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CopyInMm2BToL1(
|
||||
LocalTensor<K_ROPE_T> &bL1Tensor, const uint64_t valueGmBaseOffset, uint32_t copyTotalRowCntAlign,
|
||||
uint32_t copyStartRowCnt, uint32_t nActCopyRowCount, uint32_t copyStartColumnCount, uint32_t copyColumnCount)
|
||||
{
|
||||
uint64_t step = constInfo.headDim;
|
||||
if constexpr (LAYOUT_T == QSFA_LAYOUT::BSND || LAYOUT_T == QSFA_LAYOUT::TND) {
|
||||
step = constInfo.headDim * constInfo.kvHeadNum;
|
||||
}
|
||||
|
||||
uint32_t blockElementCnt = 32 / sizeof(K_ROPE_T);
|
||||
|
||||
Nd2NzParams mm1Nd2NzParamsForB;
|
||||
mm1Nd2NzParamsForB.ndNum = 1;
|
||||
mm1Nd2NzParamsForB.nValue = nActCopyRowCount;
|
||||
mm1Nd2NzParamsForB.dValue = copyColumnCount;
|
||||
mm1Nd2NzParamsForB.srcDValue = step;
|
||||
mm1Nd2NzParamsForB.dstNzNStride = 1;
|
||||
mm1Nd2NzParamsForB.dstNzC0Stride = copyTotalRowCntAlign;
|
||||
mm1Nd2NzParamsForB.srcNdMatrixStride = 0;
|
||||
mm1Nd2NzParamsForB.dstNzMatrixStride = 0;
|
||||
DataCopy(bL1Tensor[copyStartRowCnt * blockElementCnt], valueGm[valueGmBaseOffset + copyStartColumnCount],
|
||||
mm1Nd2NzParamsForB);
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::CalcTopKBlockInfo(
|
||||
const RunInfo &info, uint32_t &curTopKIdx, uint64_t &curOffsetInSparseBlock,
|
||||
uint32_t curSeqIdx, uint32_t ©RowCnt, uint64_t &idInTopK)
|
||||
{
|
||||
if (curTopKIdx == 0 && curOffsetInSparseBlock == 0 && copyRowCnt == 0) {
|
||||
uint64_t sparseLen = 0;
|
||||
for (uint64_t qsfaTopkidx = 0; qsfaTopkidx < constInfo.sparseBlockCount; qsfaTopkidx++) {
|
||||
int32_t qsfaSparseIndices = topKGm.GetValue(info.topKBaseOffset + qsfaTopkidx);
|
||||
if (qsfaSparseIndices == -1) {
|
||||
break;
|
||||
}
|
||||
uint64_t qsfaBlockBegin = qsfaSparseIndices * constInfo.sparseBlockSize;
|
||||
if (qsfaBlockBegin >= info.threshold) {
|
||||
continue;
|
||||
}
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.curActualSeqLenOri) ?
|
||||
info.curActualSeqLenOri : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = (qsfaBlockEnd <= info.threshold) ? \
|
||||
qsfaBlockEnd - qsfaBlockBegin : info.threshold - qsfaBlockBegin;
|
||||
sparseLen += qsfaBlockLen;
|
||||
if (sparseLen >= curSeqIdx + 1) {
|
||||
curTopKIdx = qsfaTopkidx;
|
||||
idInTopK = qsfaSparseIndices;
|
||||
curOffsetInSparseBlock = qsfaBlockLen - (sparseLen - curSeqIdx);
|
||||
copyRowCnt = sparseLen - curSeqIdx;
|
||||
break;
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
uint64_t qsfaBlockBegin = idInTopK * constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.threshold) ?
|
||||
info.threshold : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = qsfaBlockEnd - qsfaBlockBegin;
|
||||
if (curOffsetInSparseBlock + copyRowCnt < qsfaBlockLen) {
|
||||
curOffsetInSparseBlock += copyRowCnt;
|
||||
copyRowCnt = qsfaBlockLen - curOffsetInSparseBlock;
|
||||
} else {
|
||||
for (uint64_t qsfaTopkidx = curTopKIdx + 1; qsfaTopkidx < constInfo.sparseBlockCount; qsfaTopkidx++) {
|
||||
int64_t qsfaSparseIndices = topKGm.GetValue(info.topKBaseOffset + qsfaTopkidx);
|
||||
if (qsfaSparseIndices == -1) {
|
||||
break;
|
||||
}
|
||||
|
||||
uint64_t qsfaBlockBegin = qsfaSparseIndices * constInfo.sparseBlockSize;
|
||||
if (qsfaBlockBegin >= info.threshold) {
|
||||
continue;
|
||||
}
|
||||
uint64_t qsfaBlockEnd = (qsfaBlockBegin + constInfo.sparseBlockSize > info.threshold) ?
|
||||
info.threshold : qsfaBlockBegin + constInfo.sparseBlockSize;
|
||||
uint64_t qsfaBlockLen = qsfaBlockEnd - qsfaBlockBegin;
|
||||
curTopKIdx = qsfaTopkidx;
|
||||
idInTopK = qsfaSparseIndices;
|
||||
curOffsetInSparseBlock = 0;
|
||||
copyRowCnt = qsfaBlockLen;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::ComputeMm1(const RunInfo &info, const MSplitInfo mSplitInfo)
|
||||
{
|
||||
// 最外层还需要一层m的循环
|
||||
uint32_t mSize = mSplitInfo.nBufferDealM;
|
||||
uint32_t mL1Size = M_SPLIT_SIZE;
|
||||
uint32_t mL1SizeAlign = QSFAAlign(M_SPLIT_SIZE, 16U);
|
||||
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
|
||||
|
||||
uint32_t nSize = info.actualSingleProcessSInnerSize;
|
||||
uint32_t nL1Size = N_SPLIT_SIZE;
|
||||
uint32_t nL1SizeAlign = QSFAAlign(N_SPLIT_SIZE, 16U);
|
||||
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
|
||||
|
||||
uint32_t kSize = 576;
|
||||
uint32_t kL1Size = 288;
|
||||
uint32_t kL1Loops = 2; // 2 : 576/288, mla专用 这里不考虑d泛化
|
||||
|
||||
uint32_t kL0Size = 96;
|
||||
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size; // 288 / 96 = 3 kloops
|
||||
|
||||
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
|
||||
uint32_t ka = 0, kb = 0;
|
||||
for (uint32_t mL1 = 0; mL1 < mL1Loops; mL1++) {
|
||||
mL1Size = M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(M_SPLIT_SIZE, 16U);
|
||||
if (mL1 == (mL1Loops - 1)) {
|
||||
// 尾块重新计算size
|
||||
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(mL1Size, 16U);
|
||||
}
|
||||
|
||||
// 左矩阵L1选择12块还是34块的index, 由m l1 index决定
|
||||
// 左矩阵L1选择12块或34块的前一块还是后一块, 由k l1 index决定
|
||||
uint32_t mIdx = qpL1BufIter + mL1;
|
||||
ka = GetQPL1RealIdx(mIdx, 0);
|
||||
LocalTensor<Q_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
|
||||
CopyInMm1AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE, mL1Size, 576, 0);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // L1切n, 512/128=4
|
||||
if (nL1 == (nL1Loops - 1)) {
|
||||
// 尾块重新计算size
|
||||
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
|
||||
nL1SizeAlign = QSFAAlign(nL1Size, 16U);
|
||||
}
|
||||
|
||||
// 使用unitflag同步
|
||||
// 需要保证cL0BufIter和m步调一致
|
||||
LocalTensor cL0Tensor = cL0TensorPingPong[(cL0BufIter % 2) * (L0C_PP_SIZE / sizeof(MM_OUT_T))];
|
||||
for (uint32_t kL1 = 0; kL1 < kL1Loops; kL1++) { // L1切k, 576/288, 这里不考虑d泛化
|
||||
kvL1BufIter++;
|
||||
uint32_t kb = kvL1BufIter % 3;
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]);
|
||||
// 从k当中取当前的块
|
||||
LocalTensor<K_ROPE_T> bL1Tensor = l1KVTensor[kb * L1_BLOCK_OFFSET];
|
||||
if constexpr (TEMPLATE_MODE == V_TEMPLATE) {
|
||||
if (kL1 == 0) {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockCount = 288 / BLOCK_ELEMENT_NUM;
|
||||
copyParams.blockLen = nL1Size;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - nL1Size;
|
||||
copyParams.dstStride = nL1SizeAlign - nL1Size;
|
||||
DataCopy(bL1Tensor, kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize +
|
||||
nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM], copyParams);
|
||||
} else {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockCount = 224 / BLOCK_ELEMENT_NUM;
|
||||
copyParams.blockLen = nL1Size;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - nL1Size;
|
||||
copyParams.dstStride = nL1SizeAlign - nL1Size;
|
||||
DataCopy(bL1Tensor, kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize +
|
||||
288 * constInfo.s2BaseSize + nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM], copyParams);
|
||||
copyParams.blockCount = constInfo.headDimRope / BLOCK_ELEMENT_NUM;
|
||||
DataCopy(
|
||||
bL1Tensor[224 * nL1SizeAlign],
|
||||
kvMergeGm_[info.loop % 4 * N_WORKSPACE_SIZE * kSize + N_WORKSPACE_SIZE * constInfo.headDim +
|
||||
nL1 * N_SPLIT_SIZE * BLOCK_ELEMENT_NUM],
|
||||
copyParams);
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[kb]);
|
||||
|
||||
aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET + kL1 * mL1SizeAlign * K_SPLIT_SIZE];
|
||||
for (uint32_t kL0 = 0; kL0 < kL0Loops; kL0++) {
|
||||
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
LocalTensor<K_ROPE_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadDataMm1A(aL0Tensor, aL1Tensor, kL0, kL0Size, mL1SizeAlign, kL0Size);
|
||||
LocalTensor<K_ROPE_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadDataMm1B(bL0Tensor, bL1Tensor, kL0, kL0Size, kL0Size, nL1SizeAlign);
|
||||
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
|
||||
// m == 1的时候需要特殊处理
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = mL1SizeAlign;
|
||||
mmadParams.n = nL1SizeAlign;
|
||||
mmadParams.k = kL0Size;
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.cmatrixInitVal = (kL1 == 0 && kL0 == 0);
|
||||
mmadParams.unitFlag =
|
||||
(kL1 == 1 && kL0 == (kL0Loops - 1)) ? 0b11 : 0b10; // 累加最后一次翻转flag, 表示可以搬出
|
||||
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
|
||||
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
abL0BufIter++;
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[kb]); // 反向同步, 表示L1已经被mte1消费完
|
||||
}
|
||||
FixpipeParamsV220 fixParams;
|
||||
fixParams.mSize = mL1SizeAlign;
|
||||
fixParams.nSize = nL1SizeAlign;
|
||||
fixParams.srcStride = mL1SizeAlign;
|
||||
fixParams.ndNum = 1; // 输出ND
|
||||
// 改成nSizeAlign
|
||||
fixParams.dstStride = info.actualSingleProcessSInnerSizeAlign; // mm1ResGm两行之间的间隔
|
||||
fixParams.unitFlag = 0b11;
|
||||
|
||||
// 输出偏移info.loop % (constInfo.preLoadNum)) * mmResUbSize是否在matmul里计算
|
||||
Fixpipe(mm1ResGm[(info.loop % (constInfo.preLoadNum)) * constInfo.mmResUbSize + nL1 * N_SPLIT_SIZE +
|
||||
(mSplitInfo.nBufferStartM + mL1 * M_SPLIT_SIZE) *
|
||||
info.actualSingleProcessSInnerSizeAlign],
|
||||
cL0Tensor, fixParams);
|
||||
cL0BufIter++;
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
|
||||
}
|
||||
qpL1BufIter += mL1Loops;
|
||||
}
|
||||
|
||||
template <typename QSFAT>
|
||||
__aicore__ inline void QSFAMatmulService<QSFAT>::ComputeMm2(const RunInfo &info, const MSplitInfo mSplitInfo)
|
||||
{
|
||||
uint32_t mSize = mSplitInfo.nBufferDealM;
|
||||
uint32_t mSizeAlign = (mSize + 16 - 1) / 16;
|
||||
uint32_t mL1Loops = (mSize + M_SPLIT_SIZE - 1) / M_SPLIT_SIZE;
|
||||
uint32_t mL1SizeAlign = M_SPLIT_SIZE; // 16对齐
|
||||
uint32_t mL1Size = M_SPLIT_SIZE; // m的实际大小
|
||||
|
||||
uint32_t nSize = BlockAlign<K_ROPE_T>(constInfo.headDim);
|
||||
uint32_t nL1Loops = (nSize + N_SPLIT_SIZE - 1) / N_SPLIT_SIZE;
|
||||
uint32_t nL1SizeAlign = N_SPLIT_SIZE; // 16对齐
|
||||
uint32_t nL1Size = N_SPLIT_SIZE; // n的实际大小
|
||||
|
||||
uint32_t kSize = info.actualSingleProcessSInnerSize;
|
||||
uint32_t kL1Size = 256;
|
||||
uint32_t kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
uint32_t kL1Loops = (kSize + kL1Size - 1) / kL1Size;
|
||||
uint32_t kL0Size = 128;
|
||||
uint32_t kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
|
||||
uint32_t kL0SizeAlign = kL0Size;
|
||||
LocalTensor<K_ROPE_T> bL1Tensor;
|
||||
LocalTensor<K_ROPE_T> subvTensor;
|
||||
|
||||
// ka表示左矩阵4buf选择哪一块buf, kb表示右矩阵3buf选择哪一块buf
|
||||
uint32_t ka = 0, qsfaKb = 0;
|
||||
uint32_t mBaseIdx = qpL1BufIter;
|
||||
for (uint32_t nL1 = 0; nL1 < nL1Loops; nL1++) { // n切L1
|
||||
if (nL1 == (nL1Loops - 1)) {
|
||||
// 尾块
|
||||
nL1Size = nSize - (nL1Loops - 1) * N_SPLIT_SIZE;
|
||||
nL1SizeAlign = QSFAAlign(nL1Size, 16U);
|
||||
}
|
||||
|
||||
// k l1写成一个循环, 和mm1保持一致
|
||||
kL1Size = 256;
|
||||
kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
for (uint32_t k1 = 0; k1 < kL1Loops; k1++) { // k切L1, 这里套了一层l0来操作
|
||||
if (k1 == (kL1Loops - 1)) {
|
||||
// 尾块
|
||||
kL1Size = kSize - (kL1Loops - 1) * 256;
|
||||
kL1SizeAlign = QSFAAlign(kL1Size, 16U);
|
||||
}
|
||||
kvL1BufIter++;
|
||||
uint32_t qsfaKb = kvL1BufIter % 3;
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21KVIds[qsfaKb]);
|
||||
bL1Tensor = l1KVTensor[qsfaKb * L1_BLOCK_OFFSET];
|
||||
uint32_t qsfaKOffset = k1 * kL0Loops;
|
||||
kL0Size = 128;
|
||||
// 此处必须先初始化kL0Size, 再求kL0Loops, 否则由于循环会改变kL0Size大小, 导致kL0Loops错误
|
||||
kL0Loops = (kL1Size + kL0Size - 1) / kL0Size;
|
||||
kL0SizeAlign = kL0Size;
|
||||
for (uint32_t qsfaKL1 = qsfaKOffset; qsfaKL1 < kL0Loops + qsfaKOffset; qsfaKL1++) { // 128 循环搬pa
|
||||
if (qsfaKL1 == qsfaKOffset + kL0Loops - 1) {
|
||||
// 尾块
|
||||
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
|
||||
kL0SizeAlign = QSFAAlign(kL0Size, 16U);
|
||||
}
|
||||
if constexpr (TEMPLATE_MODE == V_TEMPLATE) {
|
||||
DataCopyParams copyParams;
|
||||
copyParams.blockLen = kL0Size;
|
||||
copyParams.blockCount = nL1Size / BLOCK_ELEMENT_NUM;
|
||||
copyParams.srcStride = constInfo.s2BaseSize - kL0Size;
|
||||
copyParams.dstStride = kL0SizeAlign - kL0Size;
|
||||
DataCopy(bL1Tensor[(qsfaKL1 - qsfaKOffset) * 128 * N_SPLIT_SIZE], kvMergeGm_[info.loop % 4 *
|
||||
N_WORKSPACE_SIZE * 576 + qsfaKL1 * 128 * BLOCK_ELEMENT_NUM + nL1 * N_SPLIT_SIZE *
|
||||
constInfo.s2BaseSize], copyParams);
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21KVIds[qsfaKb]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21KVIds[qsfaKb]);
|
||||
mL1SizeAlign = M_SPLIT_SIZE;
|
||||
mL1Size = M_SPLIT_SIZE; // m的实际大小
|
||||
for (uint32_t qsfaML1 = 0; qsfaML1 < mL1Loops; qsfaML1++) {
|
||||
if (qsfaML1 == (mL1Loops - 1)) {
|
||||
// 尾块
|
||||
mL1Size = mSize - (mL1Loops - 1) * M_SPLIT_SIZE;
|
||||
mL1SizeAlign = QSFAAlign(mL1Size, 16U);
|
||||
}
|
||||
|
||||
uint32_t mIdx = mBaseIdx + qsfaML1;
|
||||
ka = GetQPL1RealIdx(mIdx, k1);
|
||||
LocalTensor<K_ROPE_T> aL1Tensor = l1QPTensor[ka * L1_BLOCK_OFFSET];
|
||||
if (nL1 == 0) {
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]);
|
||||
CopyInMm2AToL1(aL1Tensor, info, mSplitInfo.nBufferStartM + qsfaML1 * M_SPLIT_SIZE, mL1Size, kL1Size,
|
||||
256 * k1);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(mte21QPIds[ka]);
|
||||
}
|
||||
|
||||
LocalTensor cL0Tensor =
|
||||
cL0TensorPingPong[(cL0BufIter % 2) *
|
||||
(L0C_PP_SIZE / sizeof(MM_OUT_T))]; // 需要保证cL0BufIter和m步调一致
|
||||
uint32_t qsfaBaseK = 128;
|
||||
uint32_t qsfaBaseN = 128;
|
||||
kL0Size = 128;
|
||||
kL0SizeAlign = kL0Size;
|
||||
for (uint32_t qsfaKL0 = 0; qsfaKL0 < kL0Loops; qsfaKL0++) {
|
||||
if (qsfaKL0 + 1 == kL0Loops) {
|
||||
kL0Size = kL1Size - (kL0Loops - 1) * kL0Size;
|
||||
kL0SizeAlign = QSFAAlign(kL0Size, 16U);
|
||||
}
|
||||
WaitFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
LocalTensor<K_ROPE_T> bL0Tensor = bL0TensorPingPong[(abL0BufIter % 2) * (L0B_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParamsForB;
|
||||
loadData3DParamsForB.l1H = kL0SizeAlign / 16; // 源操作数height
|
||||
loadData3DParamsForB.l1W = 16; // 源操作数weight=16,目的height=l1H*L1W
|
||||
loadData3DParamsForB.padList[0] = 0;
|
||||
loadData3DParamsForB.padList[1] = 0;
|
||||
loadData3DParamsForB.padList[2] = 0;
|
||||
loadData3DParamsForB.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParamsForB.mExtension = kL0SizeAlign; // 在目的操作数height维度的传输长度
|
||||
loadData3DParamsForB.kExtension = nL1SizeAlign; // 在目的操作数width维度的传输长度
|
||||
loadData3DParamsForB.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParamsForB.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParamsForB.strideH = 1;
|
||||
loadData3DParamsForB.strideW = 1;
|
||||
loadData3DParamsForB.filterW = 1;
|
||||
loadData3DParamsForB.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParamsForB.filterH = 1;
|
||||
loadData3DParamsForB.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParamsForB.dilationFilterH = 1; // 卷积核height膨胀系数
|
||||
loadData3DParamsForB.dilationFilterW = 1; // 卷积核width膨胀系数
|
||||
loadData3DParamsForB.enTranspose = 1; // 是否启用转置功能
|
||||
// 使用FMATRIX_LEFT还是使用FMATRIX_RIGHT,=0使用FMATRIX_LEFT,=1使用FMATRIX_RIGHT 1
|
||||
loadData3DParamsForB.fMatrixCtrl = 0;
|
||||
// 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
loadData3DParamsForB.channelSize = nL1SizeAlign;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(bL0Tensor, bL1Tensor[qsfaKL0 * qsfaBaseK * qsfaBaseN],
|
||||
loadData3DParamsForB);
|
||||
LocalTensor<K_ROPE_T> aL0Tensor = aL0TensorPingPong[(abL0BufIter % 2) * (L0A_PP_SIZE /
|
||||
sizeof(K_ROPE_T))];
|
||||
LoadData3DParamsV2<K_ROPE_T> loadData3DParamsForA;
|
||||
loadData3DParamsForA.l1H = mL1SizeAlign / 16; // 源操作数height
|
||||
loadData3DParamsForA.l1W = 16; // 源操作数weight
|
||||
loadData3DParamsForA.padList[0] = 0;
|
||||
loadData3DParamsForA.padList[1] = 0;
|
||||
loadData3DParamsForA.padList[2] = 0;
|
||||
loadData3DParamsForA.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
loadData3DParamsForA.mExtension = mL1SizeAlign; // 在目的操作数height维度的传输长度
|
||||
loadData3DParamsForA.kExtension = kL0SizeAlign; // 在目的操作数width维度的传输长度
|
||||
loadData3DParamsForA.mStartPt = 0; // 卷积核在目的操作数width维度的起点
|
||||
loadData3DParamsForA.kStartPt = 0; // 卷积核在目的操作数height维度的起点
|
||||
loadData3DParamsForA.strideW = 1; // 卷积核在源操作数width维度滑动的步长
|
||||
loadData3DParamsForA.strideH = 1; // 卷积核在源操作数height维度滑动的步长
|
||||
loadData3DParamsForA.filterW = 1; // 卷积核width
|
||||
loadData3DParamsForA.filterSizeW = false; // 是否在filterW的基础上将卷积核width增加256个元素
|
||||
loadData3DParamsForA.filterH = 1; // 卷积核height
|
||||
loadData3DParamsForA.filterSizeH = false; // 是否在filterH的基础上将卷积核height增加256个元素
|
||||
loadData3DParamsForA.dilationFilterW = 1; // 卷积核width膨胀系数
|
||||
loadData3DParamsForA.dilationFilterH = 1; // 卷积核height膨胀系数
|
||||
loadData3DParamsForA.enTranspose = 0; // 是否启用转置功能,对整个目标矩阵进行转置
|
||||
loadData3DParamsForA.fMatrixCtrl = 0;
|
||||
// 源操作数的通道数。膨胀系数为1时,目的weight为filterW*filterH*channelSize
|
||||
loadData3DParamsForA.channelSize = kL0SizeAlign;
|
||||
LoadData<K_ROPE_T, LOAD3DV2_CONFIG>(aL0Tensor, aL1Tensor[qsfaKL0 * qsfaBaseK * mL1SizeAlign],
|
||||
loadData3DParamsForA);
|
||||
SetFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
WaitFlag<HardEvent::MTE1_M>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = mL1SizeAlign;
|
||||
mmadParams.n = nL1SizeAlign;
|
||||
mmadParams.k = kL0Size;
|
||||
mmadParams.cmatrixInitVal = (qsfaKL0 == 0 && k1 == 0);
|
||||
mmadParams.cmatrixSource = false;
|
||||
mmadParams.unitFlag = ((k1 == (kL1Loops - 1)) && (qsfaKL0 == (kL0Loops - 1))) ? 0b11 : 0b10;
|
||||
|
||||
Mmad(cL0Tensor, aL0Tensor, bL0Tensor, mmadParams);
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
SetFlag<HardEvent::M_MTE1>(Mte1MmABEventId(abL0BufIter % 2));
|
||||
abL0BufIter++;
|
||||
}
|
||||
|
||||
if (nL1 == (nL1Loops - 1)) { // nL1最后一轮, 需要将B驻留在L1中, 用于下一轮的计算?
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21QPIds[ka]); // 反向同步, 表示L1中的A已经被mte1消费完
|
||||
}
|
||||
|
||||
if (k1 == (kL1Loops - 1)) {
|
||||
// ND
|
||||
FixpipeParamsV220 fixParams;
|
||||
fixParams.nSize = nL1SizeAlign;
|
||||
fixParams.mSize = mL1SizeAlign;
|
||||
fixParams.srcStride = mL1SizeAlign;
|
||||
fixParams.dstStride = nSize; // mm2ResGm两行之间的间隔
|
||||
fixParams.ndNum = 1; // 输出ND
|
||||
fixParams.unitFlag = 0b11;
|
||||
|
||||
uint64_t qsfaMm2Offset = (mSplitInfo.nBufferStartM + qsfaML1 * M_SPLIT_SIZE) * nSize +
|
||||
nL1 * N_SPLIT_SIZE;
|
||||
Fixpipe(mm2ResGm[(info.loop % (constInfo.preLoadNum)) *
|
||||
constInfo.bmm2ResUbSize + qsfaMm2Offset], cL0Tensor, fixParams);
|
||||
}
|
||||
|
||||
if (mL1Loops == 2) {
|
||||
cL0BufIter++;
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(mte21KVIds[qsfaKb]); // 反向同步, 表示L1已经被mte1消费完
|
||||
}
|
||||
// cL0BufIter已经不在使用
|
||||
if (mL1Loops == 1) {
|
||||
cL0BufIter++;
|
||||
}
|
||||
}
|
||||
qpL1BufIter += mL1Loops;
|
||||
}
|
||||
|
||||
#endif // KV_QUANT_SPARSE_FLASH_ATTENTION_SERVICE_CUBE_MLA_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,82 @@
|
||||
/**
|
||||
* 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 kv_quant_sparse_flash_attention_template_tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_FLASH_ATTENTION_TEMPLATE_TILING_KEY_H
|
||||
#define KV_QUANT_SPARSE_FLASH_ATTENTION_TEMPLATE_TILING_KEY_H
|
||||
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define QSFA_LAYOUT_BSND 0
|
||||
#define QSFA_LAYOUT_TND 1
|
||||
#define QSFA_LAYOUT_PA_BSND 2
|
||||
|
||||
#define ASCENDC_TPL_4_BW 4
|
||||
|
||||
#define C_TEMPLATE 0
|
||||
#define V_TEMPLATE 1
|
||||
|
||||
// 模板参数支持的范围定义
|
||||
ASCENDC_TPL_ARGS_DECL(KvQuantSparseFlashAttention, // 算子OpType
|
||||
ASCENDC_TPL_BOOL_DECL(FLASH_DECODE, 0, 1),
|
||||
ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 0, 1),
|
||||
ASCENDC_TPL_UINT_DECL(LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QSFA_LAYOUT_BSND, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_DECL(KV_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QSFA_LAYOUT_BSND, QSFA_LAYOUT_TND, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_DECL(TEMPLATE_MODE, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, C_TEMPLATE, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_DECL(IS_SPLIT_G, 0, 1),
|
||||
);
|
||||
|
||||
// 支持的模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
|
||||
ASCENDC_TPL_ARGS_SEL(
|
||||
ASCENDC_TPL_BOOL_SEL(FLASH_DECODE, 0),
|
||||
ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(KV_LAYOUT_T, ASCENDC_TPL_UI_LIST, QSFA_LAYOUT_PA_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(TEMPLATE_MODE, ASCENDC_TPL_UI_LIST, V_TEMPLATE),
|
||||
ASCENDC_TPL_BOOL_SEL(IS_SPLIT_G, 0, 1),
|
||||
),
|
||||
);
|
||||
|
||||
#endif // TEMPLATE_TILING_KEY
|
||||
Reference in New Issue
Block a user