@@ -0,0 +1,148 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_common.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_COMMON_H
|
||||
#define QUANT_LIGHTNING_INDEXER_COMMON_H
|
||||
|
||||
namespace QLICommon {
|
||||
|
||||
// 与tiling的layout保持一致
|
||||
enum class LI_LAYOUT : uint32_t {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2
|
||||
};
|
||||
|
||||
template <typename Q_T, typename K_T, typename OUT_T, const bool PAGE_ATTENTION = false,
|
||||
LI_LAYOUT Q_LAYOUT_T = LI_LAYOUT::BSND, LI_LAYOUT K_LAYOUT_T = LI_LAYOUT::PA_BSND, typename... Args>
|
||||
struct QLIType {
|
||||
using queryType = Q_T;
|
||||
using keyType = K_T;
|
||||
using outputType = OUT_T;
|
||||
static constexpr bool pageAttention = PAGE_ATTENTION;
|
||||
static constexpr LI_LAYOUT layout = Q_LAYOUT_T;
|
||||
static constexpr LI_LAYOUT keyLayout = K_LAYOUT_T;
|
||||
};
|
||||
|
||||
struct RunInfo {
|
||||
uint32_t loop;
|
||||
uint32_t bN2Idx;
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx = 0;
|
||||
uint32_t gS1Idx;
|
||||
uint32_t s2Idx;
|
||||
|
||||
uint32_t actS1Size = 1;
|
||||
uint32_t actS2Size = 1;
|
||||
uint32_t actS2SizeOrig = 1;
|
||||
uint32_t actMBaseSize;
|
||||
uint32_t actualSingleProcessSInnerSize;
|
||||
uint32_t actualSingleProcessSInnerSizeAlign;
|
||||
|
||||
uint64_t tensorQueryOffset;
|
||||
uint64_t tensorKeyOffset;
|
||||
uint64_t tensorKeyScaleOffset;
|
||||
uint64_t tensorWeightsOffset;
|
||||
uint64_t indiceOutOffset;
|
||||
|
||||
bool isFirstS2InnerLoop;
|
||||
bool isLastS2InnerLoop;
|
||||
bool isValid = false;
|
||||
};
|
||||
|
||||
struct ConstInfo {
|
||||
// CUBE与VEC核间同步的模式
|
||||
static constexpr uint32_t FIA_SYNC_MODE2 = 2;
|
||||
// BUFFER的字节数
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768;
|
||||
// 无效索引
|
||||
static constexpr int INVALID_IDX = -1;
|
||||
|
||||
// CUBE和VEC的核间同步EventID
|
||||
uint32_t syncC1V1 = 0U;
|
||||
uint32_t syncC1V0 = 2U;
|
||||
uint32_t syncV1C1 = 0U;
|
||||
uint32_t syncV0C1 = 1U;
|
||||
|
||||
// 基本块大小
|
||||
uint32_t mBaseSize = 1ULL;
|
||||
uint32_t s1BaseSize = 1ULL;
|
||||
uint32_t s2BaseSize = 1ULL;
|
||||
|
||||
uint64_t batchSize = 0ULL;
|
||||
uint64_t gSize = 0ULL;
|
||||
uint64_t qHeadNum = 0ULL;
|
||||
uint64_t kHeadNum;
|
||||
uint64_t headDim;
|
||||
uint64_t sparseCount; // topK选取大小
|
||||
uint64_t kSeqSize = 0ULL; // kv最大S长度
|
||||
uint64_t qSeqSize = 1ULL; // q最大S长度
|
||||
uint32_t kCacheBlockSize = 0; // PA场景的block size
|
||||
uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number
|
||||
LI_LAYOUT outputLayout; // 输出的格式
|
||||
bool attenMaskFlag = false;
|
||||
uint32_t cmpRatio = 1; // 压缩率
|
||||
bool batchSupperFlag = false; // Qactual_se长度是否为B+1
|
||||
int64_t stride = 1;
|
||||
int64_t scaleStride = 1;
|
||||
|
||||
uint32_t actualLenQDims = 0U; // query的actualSeqLength 的维度
|
||||
uint32_t actualLenDims = 0U; // KV 的actualSeqLength 的维度
|
||||
bool isAccumSeqS1 = false; // 是否累加模式
|
||||
bool isAccumSeqS2 = false; // 是否累加模式
|
||||
|
||||
uint32_t s2Start = 0U;
|
||||
uint32_t s2End = 0U;
|
||||
uint32_t bN2Start = 0U;
|
||||
uint32_t bN2End = 0U;
|
||||
uint32_t gS1Start = 0U;
|
||||
uint32_t gS1End = 0U;
|
||||
uint32_t coreEnable = 0U;
|
||||
};
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Align(T1 num, T2 rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd) - 1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Min(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (b) : (a);
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Max(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (a) : (b);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T CeilDiv(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd)));
|
||||
}
|
||||
} // namespace QLICommon
|
||||
|
||||
#endif // QUANT_LIGHTNING_INDEXER_COMMON_H
|
||||
@@ -0,0 +1,667 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_KERNEL_H
|
||||
#define QUANT_LIGHTNING_INDEXER_KERNEL_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
#include "quant_lightning_indexer_service_vector.h"
|
||||
#include "quant_lightning_indexer_service_cube.h"
|
||||
#include "../vllm_quant_lightning_indexer_metadata.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
using namespace QLIServiceVec;
|
||||
using namespace matmul;
|
||||
using namespace optiling::detail;
|
||||
using namespace optiling;
|
||||
using AscendC::CacheMode;
|
||||
using AscendC::CrossCoreSetFlag;
|
||||
using AscendC::CrossCoreWaitFlag;
|
||||
|
||||
// 由于S2循环前,RunInfo还没有赋值,使用TempLoopInfo临时存放B、N、S1轴相关的信息;同时减少重复计算
|
||||
struct TempLoopInfo {
|
||||
uint32_t bN2Idx = 0;
|
||||
uint32_t bIdx = 0U;
|
||||
uint32_t n2Idx = 0U;
|
||||
uint32_t gS1Idx = 0U;
|
||||
uint32_t gS1LoopEnd = 0U; // gS1方向循环的结束Idx
|
||||
uint32_t s2LoopEnd = 0U; // S2方向循环的结束Idx
|
||||
uint32_t actS1Size = 1ULL; // 当前Batch循环处理的S1轴的实际大小
|
||||
uint32_t actS2Size = 0ULL;
|
||||
uint32_t actS2SizeOrig = 0ULL;
|
||||
bool curActSeqLenIsZero = false;
|
||||
bool needDealActS1LessThanS1 = false; // S1的实际长度小于shape的S1长度时,是否需要清理输出
|
||||
uint32_t actMBaseSize = 0U; // m轴(gS1)方向实际大小
|
||||
uint32_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小
|
||||
uint32_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小
|
||||
uint32_t validS2Len = 0U;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
class QLIPreload {
|
||||
public:
|
||||
__aicore__ inline QLIPreload(){};
|
||||
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights,
|
||||
__gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, __gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengthsK, __gm__ uint8_t *blockTable, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *workspace,
|
||||
const QLITilingData *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void Process();
|
||||
|
||||
// =================================类型定义区=================================
|
||||
using Q_T = typename QLIT::queryType;
|
||||
using K_T = typename QLIT::keyType;
|
||||
using OUT_T = typename QLIT::outputType;
|
||||
static constexpr bool PAGE_ATTENTION = QLIT::pageAttention;
|
||||
static constexpr LI_LAYOUT Q_LAYOUT_T = QLIT::layout;
|
||||
static constexpr LI_LAYOUT K_LAYOUT_T = QLIT::keyLayout;
|
||||
|
||||
using MM1_OUT_T = float;
|
||||
|
||||
QLIMatmul<QLIT> matmulService;
|
||||
QLIVector<QLIT> vectorService;
|
||||
|
||||
// =================================常量区=================================
|
||||
static constexpr uint32_t SYNC_C1_V1_FLAG = 4;
|
||||
static constexpr uint32_t SYNC_V1_C1_FLAG = 5;
|
||||
|
||||
static constexpr uint32_t M_BASE_SIZE = 256;
|
||||
static constexpr uint32_t S2_BASE_SIZE = 2048;
|
||||
static constexpr uint32_t HEAD_DIM = 128;
|
||||
static constexpr uint32_t K_HEAD_NUM = 1;
|
||||
static constexpr uint32_t GM_ALIGN_BYTES = 512;
|
||||
static constexpr uint32_t LI_QUANT_PRELOAD_TASK_CACHE_SIZE = 2;
|
||||
|
||||
// for workspace double
|
||||
static constexpr uint32_t WS_DOBULE = 2;
|
||||
static constexpr uint32_t ELE_NUM_PER_BLOCK = 16;
|
||||
|
||||
protected:
|
||||
TPipe *pipe = nullptr;
|
||||
|
||||
// offset
|
||||
uint64_t queryCoreOffset = 0ULL;
|
||||
uint64_t keyCoreOffset = 0ULL;
|
||||
uint64_t keyScaleCoreOffset = 0ULL;
|
||||
uint64_t weightsCoreOffset = 0ULL;
|
||||
uint64_t indiceOutCoreOffset = 0ULL;
|
||||
uint32_t coreZeroEnable = 1U;
|
||||
|
||||
// ================================Global Buffer区=================================
|
||||
GlobalTensor<Q_T> queryGm;
|
||||
GlobalTensor<K_T> keyGm;
|
||||
GlobalTensor<half> weightsGm;
|
||||
GlobalTensor<uint32_t> metadataGm;
|
||||
GlobalTensor<int32_t> indiceOutGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
|
||||
GlobalTensor<uint32_t> actualSeqLengthsGmQ;
|
||||
GlobalTensor<uint32_t> actualSeqLengthsGm;
|
||||
|
||||
// ================================类成员变量====================================
|
||||
// aic、aiv核信息
|
||||
uint32_t tmpBlockIdx = 0U;
|
||||
uint32_t aiCoreIdx = 0U;
|
||||
|
||||
QLICommon::ConstInfo constInfo{};
|
||||
TempLoopInfo tempLoopInfo{};
|
||||
|
||||
// ================================Init functions==================================
|
||||
__aicore__ inline void InitTilingData(const QLITilingData *__restrict tilingData);
|
||||
__aicore__ inline void InitBuffers();
|
||||
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsK);
|
||||
// ================================Split Core================================
|
||||
__aicore__ inline void SplitCore();
|
||||
__aicore__ inline uint32_t GetS2BaseBlockNumOnMask(uint32_t s1gIdx, uint32_t actS1Size, uint32_t actS2SizeOrig,
|
||||
uint32_t &validS2Len);
|
||||
__aicore__ inline uint32_t GetTotalBaseBlockNum();
|
||||
// ================================Process functions================================
|
||||
__aicore__ inline void ProcessMain();
|
||||
__aicore__ inline void ProcessBaseBlock(uint32_t loop, uint64_t s2LoopIdx,
|
||||
QLICommon::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE]);
|
||||
__aicore__ inline void ProcessInvalid();
|
||||
// ================================Params Calc=====================================
|
||||
__aicore__ inline void CalcGS1LoopParams(uint32_t bN2Idx);
|
||||
__aicore__ inline void GetBN2Idx(uint32_t bN2Idx);
|
||||
__aicore__ inline uint32_t GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm, uint32_t defaultSeqLen);
|
||||
__aicore__ inline uint32_t GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm, uint32_t defaultSeqLen, uint32_t cmpRatio);
|
||||
__aicore__ inline void GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, uint32_t &actS2Size, uint32_t &actS2SizeOrig);
|
||||
__aicore__ inline void CalcS2LoopParams(uint32_t bN2LoopIdx, uint32_t gS1LoopIdx);
|
||||
__aicore__ inline void CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void DealActSeqLenIsZero(uint32_t bIdx, uint32_t n2Idx, uint32_t s1Start);
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitTilingData(const QLITilingData *__restrict tilingData)
|
||||
{
|
||||
constInfo.batchSize = tilingData->bSize;
|
||||
constInfo.qHeadNum = constInfo.gSize = tilingData->gSize;
|
||||
constInfo.kSeqSize = tilingData->s2Size;
|
||||
constInfo.qSeqSize = tilingData->s1Size;
|
||||
constInfo.attenMaskFlag = (tilingData->sparseMode == 3);
|
||||
constInfo.kCacheBlockSize = tilingData->blockSize;
|
||||
constInfo.maxBlockNumPerBatch = tilingData->maxBlockNumPerBatch;
|
||||
constInfo.sparseCount = tilingData->sparseCount;
|
||||
constInfo.cmpRatio = tilingData->cmpRatio;
|
||||
constInfo.batchSupperFlag = tilingData->batchSupperFlag;
|
||||
constInfo.stride = tilingData->stride;
|
||||
constInfo.scaleStride = tilingData->scaleStride;
|
||||
|
||||
constInfo.outputLayout = Q_LAYOUT_T; // 输出和输入形状一致
|
||||
if (Q_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
constInfo.isAccumSeqS1 = true;
|
||||
}
|
||||
if (K_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
constInfo.isAccumSeqS2 = true;
|
||||
}
|
||||
|
||||
constInfo.kHeadNum = K_HEAD_NUM;
|
||||
constInfo.headDim = HEAD_DIM;
|
||||
|
||||
constInfo.mBaseSize = M_BASE_SIZE;
|
||||
constInfo.s2BaseSize = S2_BASE_SIZE;
|
||||
constInfo.s1BaseSize = (constInfo.mBaseSize + constInfo.gSize - 1) / constInfo.gSize;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitBuffers()
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.InitBuffers(pipe);
|
||||
} else {
|
||||
matmulService.InitBuffers(pipe);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengthsK)
|
||||
{
|
||||
if (actualSeqLengthsQ == nullptr) {
|
||||
constInfo.actualLenQDims = 0;
|
||||
} else {
|
||||
constInfo.actualLenQDims = (constInfo.batchSupperFlag) ? constInfo.batchSize + 1 : constInfo.batchSize;
|
||||
actualSeqLengthsGmQ.SetGlobalBuffer((__gm__ uint32_t *)actualSeqLengthsQ, constInfo.actualLenQDims);
|
||||
}
|
||||
if (actualSeqLengthsK == nullptr) {
|
||||
constInfo.actualLenDims = 0;
|
||||
} else {
|
||||
constInfo.actualLenDims = constInfo.batchSize;
|
||||
actualSeqLengthsGm.SetGlobalBuffer((__gm__ uint32_t *)actualSeqLengthsK, constInfo.actualLenDims);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm,
|
||||
uint32_t defaultSeqLen)
|
||||
{
|
||||
if (actualLenDims == 0) {
|
||||
return defaultSeqLen;
|
||||
} else if (constInfo.batchSupperFlag) {
|
||||
return actualSeqLengthsGm.GetValue(bIdx + 1) - actualSeqLengthsGm.GetValue(bIdx);
|
||||
} else if (isAccumSeq && bIdx > 0) {
|
||||
return actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1);
|
||||
} else {
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm,
|
||||
uint32_t defaultSeqLen, uint32_t cmpRatio)
|
||||
{
|
||||
if (actualLenDims == 0) {
|
||||
return defaultSeqLen * cmpRatio;
|
||||
} else if (isAccumSeq && bIdx > 0) {
|
||||
return actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1);
|
||||
} else {
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, uint32_t &actS2Size, uint32_t &actS2SizeOrig)
|
||||
{
|
||||
actS1Size = GetActualSeqLen(bIdx, constInfo.actualLenQDims, constInfo.isAccumSeqS1, actualSeqLengthsGmQ,
|
||||
constInfo.qSeqSize);
|
||||
actS2SizeOrig =
|
||||
GetActualSeqLenKey(bIdx, constInfo.actualLenDims, constInfo.isAccumSeqS2, actualSeqLengthsGm, constInfo.kSeqSize, constInfo.cmpRatio); // 压缩前的actS2Size
|
||||
actS2Size = actS2SizeOrig / constInfo.cmpRatio; // 真实使用的压缩后S2长度
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetS2BaseBlockNumOnMask(uint32_t s1gIdx, uint32_t actS1Size,
|
||||
uint32_t actS2SizeOrig, uint32_t &validS2Len)
|
||||
{
|
||||
if (actS2SizeOrig / constInfo.cmpRatio == 0) {
|
||||
validS2Len = 0;
|
||||
return 0;
|
||||
}
|
||||
uint32_t s1Offset = constInfo.s1BaseSize * s1gIdx;
|
||||
int32_t validS2LenBase = static_cast<int32_t>(actS2SizeOrig) - static_cast<int32_t>(actS1Size); // 压缩前的validS2LenBase
|
||||
validS2Len = (static_cast<int32_t>(s1Offset) + validS2LenBase + static_cast<int32_t>(constInfo.s1BaseSize)) / static_cast<int32_t>(constInfo.cmpRatio);
|
||||
validS2Len = Min(validS2Len, static_cast<int32_t>(actS2SizeOrig) / constInfo.cmpRatio);
|
||||
validS2Len = Max(validS2Len, 1);
|
||||
return (validS2Len + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetTotalBaseBlockNum()
|
||||
{
|
||||
uint32_t totalBlockNum = 0;
|
||||
uint32_t actS1Size, actS2Size, actS2SizeOrig;
|
||||
uint32_t s1GBaseNum, s2BaseNum;
|
||||
uint32_t validS2Len = 0;
|
||||
for (uint32_t bIdx = 0; bIdx < constInfo.batchSize; bIdx++) {
|
||||
GetS1S2ActualSeqLen(bIdx, actS1Size, actS2Size, actS2SizeOrig);
|
||||
s1GBaseNum = CeilDiv(actS1Size, constInfo.s1BaseSize);
|
||||
if (!constInfo.attenMaskFlag) {
|
||||
s2BaseNum = CeilDiv(actS2Size, constInfo.s2BaseSize);
|
||||
totalBlockNum += s1GBaseNum * s2BaseNum * constInfo.kHeadNum;
|
||||
continue;
|
||||
}
|
||||
for (uint32_t s1gIdx = 0; s1gIdx < s1GBaseNum; s1gIdx++) {
|
||||
s2BaseNum = GetS2BaseBlockNumOnMask(s1gIdx, actS1Size, actS2SizeOrig, validS2Len);
|
||||
totalBlockNum += s2BaseNum * constInfo.kHeadNum;
|
||||
}
|
||||
}
|
||||
return totalBlockNum;
|
||||
}
|
||||
|
||||
// 多核版本,双闭区间。基本原则:计算每个核最少处理的块数, 剩余的部分前面的核每个核多处理一块
|
||||
template <typename QLIT>
|
||||
__aicore__ void inline QLIPreload<QLIT>::SplitCore()
|
||||
{
|
||||
constInfo.coreEnable = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_CORE_ENABLE_INDEX, false));
|
||||
if (aiCoreIdx != 0) {
|
||||
constInfo.bN2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_BN2_START_INDEX, false));
|
||||
constInfo.gS1Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_M_START_INDEX, false));
|
||||
constInfo.s2Start = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_S2_START_INDEX, false));
|
||||
}
|
||||
constInfo.bN2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_BN2_END_INDEX, false));
|
||||
constInfo.gS1End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_M_END_INDEX, false));
|
||||
constInfo.s2End = metadataGm.GetValue(GetAttrAbsIndex(aiCoreIdx, LI_S2_END_INDEX, false));
|
||||
|
||||
// 如果0核都没有启动,说明所有核都没启动
|
||||
coreZeroEnable = metadataGm.GetValue(GetAttrAbsIndex(0, LI_CORE_ENABLE_INDEX, false));
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::DealActSeqLenIsZero(uint32_t bIdx, uint32_t n2Idx, uint32_t s1Start)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
if (constInfo.outputLayout == LI_LAYOUT::TND) {
|
||||
uint32_t tSizeIdx = (constInfo.batchSupperFlag) ? constInfo.batchSize : constInfo.batchSize - 1;
|
||||
uint32_t tBaseIdx = (constInfo.batchSupperFlag) ? bIdx : bIdx - 1;
|
||||
uint32_t tSize = actualSeqLengthsGmQ.GetValue(constInfo.batchSize - 1);
|
||||
uint32_t tBase = bIdx == 0 ? 0 : actualSeqLengthsGmQ.GetValue(tBaseIdx);
|
||||
uint32_t s1Count = tempLoopInfo.actS1Size;
|
||||
|
||||
for (uint32_t s1Idx = s1Start; s1Idx < s1Count; s1Idx++) {
|
||||
uint64_t indiceOutOffset =
|
||||
(tBase + s1Idx) * constInfo.kHeadNum * constInfo.sparseCount + // T轴、s1轴偏移
|
||||
n2Idx * constInfo.sparseCount; // N2轴偏移
|
||||
vectorService.CleanInvalidOutput(indiceOutOffset);
|
||||
}
|
||||
} else if (constInfo.outputLayout == LI_LAYOUT::BSND) {
|
||||
for (uint32_t s1Idx = s1Start; s1Idx < constInfo.qSeqSize; s1Idx++) {
|
||||
// B,S1,N2,K
|
||||
uint64_t indiceOutOffset = bIdx * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount +
|
||||
s1Idx * constInfo.kHeadNum * constInfo.sparseCount + // B轴、S1轴偏移
|
||||
n2Idx * constInfo.sparseCount; // N2轴偏移
|
||||
vectorService.CleanInvalidOutput(indiceOutOffset);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights,
|
||||
__gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale,
|
||||
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsK,
|
||||
__gm__ uint8_t *blockTable, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *workspace,
|
||||
const QLITilingData *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
tmpBlockIdx = GetBlockIdx(); // vec:0-47
|
||||
aiCoreIdx = tmpBlockIdx / 2;
|
||||
} else {
|
||||
tmpBlockIdx = GetBlockIdx(); // cube:0-23
|
||||
aiCoreIdx = tmpBlockIdx;
|
||||
}
|
||||
|
||||
InitTilingData(tiling);
|
||||
InitActualSeqLen(actualSeqLengthsQ, actualSeqLengthsK);
|
||||
|
||||
if (metadata != nullptr) {
|
||||
metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata);
|
||||
// 计算分核
|
||||
SplitCore();
|
||||
}
|
||||
|
||||
pipe = tPipe;
|
||||
// workspace 内存排布
|
||||
// |mm1ResGm(存S)
|
||||
uint64_t offset = 0;
|
||||
|
||||
// mm1开DoubleBuffer
|
||||
GlobalTensor<MM1_OUT_T> mm1ResGm; // 存放S
|
||||
uint64_t singleCoreMm1ResSize = WS_DOBULE * constInfo.s1BaseSize * constInfo.s2BaseSize * sizeof(MM1_OUT_T);
|
||||
mm1ResGm.SetGlobalBuffer((__gm__ MM1_OUT_T *)(workspace + aiCoreIdx * singleCoreMm1ResSize));
|
||||
offset += GetBlockNum() * singleCoreMm1ResSize;
|
||||
|
||||
GlobalTensor<half> weightWorkspaceGm; // v1阶段处理w*scale后的结果
|
||||
uint64_t weightMemSize = BLOCK_CUBE * constInfo.mBaseSize * WS_DOBULE * sizeof(half);
|
||||
weightWorkspaceGm.SetGlobalBuffer((__gm__ half *)(workspace + offset + aiCoreIdx * weightMemSize));
|
||||
offset += GetBlockNum() * weightMemSize;
|
||||
|
||||
GlobalTensor<half> qScaleGm;
|
||||
GlobalTensor<half> kScaleGm;
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.InitParams(constInfo, tiling);
|
||||
indiceOutGm.SetGlobalBuffer((__gm__ int32_t *)sparseIndices);
|
||||
weightsGm.SetGlobalBuffer((__gm__ half *)weights);
|
||||
qScaleGm.SetGlobalBuffer((__gm__ half *)queryScale);
|
||||
kScaleGm.SetGlobalBuffer((__gm__ half *)keyScale);
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
vectorService.InitVecInputTensor(weightsGm, qScaleGm, kScaleGm, indiceOutGm, blockTableGm);
|
||||
vectorService.InitVecWorkspaceTensor(weightWorkspaceGm, mm1ResGm);
|
||||
} else {
|
||||
matmulService.InitParams(constInfo);
|
||||
queryGm.SetGlobalBuffer((__gm__ Q_T *)query);
|
||||
if constexpr (PAGE_ATTENTION) {
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
}
|
||||
keyGm.SetGlobalBuffer((__gm__ K_T *)key);
|
||||
matmulService.InitMm1GlobalTensor(blockTableGm, keyGm, queryGm, mm1ResGm, weightWorkspaceGm);
|
||||
}
|
||||
InitBuffers();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::GetBN2Idx(uint32_t bN2Idx)
|
||||
{
|
||||
tempLoopInfo.bN2Idx = bN2Idx;
|
||||
tempLoopInfo.bIdx = bN2Idx / constInfo.kHeadNum;
|
||||
tempLoopInfo.n2Idx = bN2Idx % constInfo.kHeadNum;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcS2LoopParams(uint32_t bN2LoopIdx, uint32_t gS1LoopIdx)
|
||||
{
|
||||
tempLoopInfo.gS1Idx = gS1LoopIdx;
|
||||
tempLoopInfo.actMBaseSize = constInfo.mBaseSize;
|
||||
uint32_t remainedGS1Size = tempLoopInfo.actS1Size * constInfo.gSize - tempLoopInfo.gS1Idx * constInfo.mBaseSize;
|
||||
if (remainedGS1Size <= constInfo.mBaseSize && remainedGS1Size > 0) {
|
||||
tempLoopInfo.actMBaseSize = tempLoopInfo.mBasicSizeTail;
|
||||
}
|
||||
|
||||
bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == tempLoopInfo.gS1LoopEnd);
|
||||
uint32_t s2BlockNum;
|
||||
uint32_t validS2Len = 0;
|
||||
if (constInfo.attenMaskFlag) {
|
||||
s2BlockNum = GetS2BaseBlockNumOnMask(gS1LoopIdx, tempLoopInfo.actS1Size, tempLoopInfo.actS2SizeOrig,
|
||||
tempLoopInfo.validS2Len);
|
||||
} else {
|
||||
s2BlockNum = (tempLoopInfo.actS2Size + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
tempLoopInfo.validS2Len = tempLoopInfo.actS2Size;
|
||||
}
|
||||
tempLoopInfo.s2LoopEnd = (isEnd && constInfo.s2End != 0) ? constInfo.s2End : s2BlockNum;
|
||||
tempLoopInfo.s2BasicSizeTail = tempLoopInfo.validS2Len % constInfo.s2BaseSize;
|
||||
tempLoopInfo.s2BasicSizeTail = (tempLoopInfo.s2BasicSizeTail == 0) ?
|
||||
constInfo.s2BaseSize : tempLoopInfo.s2BasicSizeTail;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcGS1LoopParams(uint32_t bN2LoopIdx)
|
||||
{
|
||||
GetBN2Idx(bN2LoopIdx);
|
||||
GetS1S2ActualSeqLen(tempLoopInfo.bIdx, tempLoopInfo.actS1Size, tempLoopInfo.actS2Size, tempLoopInfo.actS2SizeOrig);
|
||||
if ((tempLoopInfo.actS2Size == 0) || (tempLoopInfo.actS1Size == 0)) {
|
||||
tempLoopInfo.curActSeqLenIsZero = true;
|
||||
return;
|
||||
}
|
||||
tempLoopInfo.curActSeqLenIsZero = false;
|
||||
tempLoopInfo.mBasicSizeTail = (tempLoopInfo.actS1Size * constInfo.gSize) % constInfo.mBaseSize;
|
||||
tempLoopInfo.mBasicSizeTail =
|
||||
(tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail;
|
||||
|
||||
uint32_t gS1SplitNum = (tempLoopInfo.actS1Size * constInfo.gSize + constInfo.mBaseSize - 1) / constInfo.mBaseSize;
|
||||
tempLoopInfo.gS1LoopEnd = (bN2LoopIdx + 1 == constInfo.bN2End && constInfo.gS1End != 0) ? constInfo.gS1End : gS1SplitNum;
|
||||
if constexpr (Q_LAYOUT_T == LI_LAYOUT::BSND) {
|
||||
if (tempLoopInfo.gS1LoopEnd == gS1SplitNum && constInfo.qSeqSize > tempLoopInfo.actS1Size) {
|
||||
tempLoopInfo.needDealActS1LessThanS1 = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
runInfo.loop = loop;
|
||||
runInfo.bIdx = tempLoopInfo.bIdx;
|
||||
runInfo.gS1Idx = tempLoopInfo.gS1Idx;
|
||||
runInfo.s2Idx = s2LoopIdx;
|
||||
runInfo.bN2Idx = tempLoopInfo.bN2Idx;
|
||||
runInfo.isValid = s2LoopIdx < tempLoopInfo.s2LoopEnd;
|
||||
|
||||
if (!runInfo.isValid) {
|
||||
return; // 需要验证, v1 时候需要runInfo
|
||||
}
|
||||
|
||||
runInfo.actS1Size = tempLoopInfo.actS1Size;
|
||||
runInfo.actS2Size = tempLoopInfo.actS2Size;
|
||||
runInfo.actS2SizeOrig = tempLoopInfo.actS2SizeOrig;
|
||||
// 计算实际基本块size
|
||||
runInfo.actMBaseSize = tempLoopInfo.actMBaseSize;
|
||||
runInfo.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
|
||||
uint32_t s2SplitNum = (tempLoopInfo.validS2Len + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
if (runInfo.s2Idx == s2SplitNum - 1) {
|
||||
runInfo.actualSingleProcessSInnerSize = tempLoopInfo.s2BasicSizeTail;
|
||||
}
|
||||
runInfo.actualSingleProcessSInnerSizeAlign =
|
||||
QLICommon::Align((uint32_t)runInfo.actualSingleProcessSInnerSize, QLICommon::ConstInfo::BUFFER_SIZE_BYTE_32B);
|
||||
|
||||
runInfo.isFirstS2InnerLoop = s2LoopIdx == constInfo.s2Start;
|
||||
runInfo.isLastS2InnerLoop = (s2LoopIdx + 1 == tempLoopInfo.s2LoopEnd);
|
||||
|
||||
if (runInfo.isFirstS2InnerLoop) {
|
||||
uint64_t actualSeqQPrefixSum;
|
||||
if constexpr (Q_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
uint32_t actualSeqLengthsGmQIdx = (constInfo.batchSupperFlag) ? runInfo.bIdx : runInfo.bIdx - 1;
|
||||
actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : actualSeqLengthsGmQ.GetValue(actualSeqLengthsGmQIdx);
|
||||
} else { // BSND
|
||||
actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : runInfo.bIdx * constInfo.qSeqSize;
|
||||
}
|
||||
uint64_t tndBIdxOffset = actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim;
|
||||
// B,S1,N1(N2,G),D
|
||||
queryCoreOffset = tndBIdxOffset + runInfo.gS1Idx * constInfo.mBaseSize * constInfo.headDim;
|
||||
// B,S1,N1(N2,G)/T,N1(N2,G)
|
||||
weightsCoreOffset = actualSeqQPrefixSum * constInfo.qHeadNum + runInfo.n2Idx * constInfo.gSize;
|
||||
// B,S1,N2,k/T,N2,k
|
||||
indiceOutCoreOffset =
|
||||
actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.sparseCount + runInfo.n2Idx * constInfo.sparseCount;
|
||||
}
|
||||
uint64_t actualSeqKPrefixSum;
|
||||
if constexpr (K_LAYOUT_T == LI_LAYOUT::TND) { // T N2 D
|
||||
actualSeqKPrefixSum = (runInfo.bIdx <= 0) ? 0 : actualSeqLengthsGm.GetValue(runInfo.bIdx - 1);
|
||||
} else {
|
||||
actualSeqKPrefixSum = (runInfo.bIdx <= 0) ? 0 : runInfo.bIdx * constInfo.kSeqSize;
|
||||
}
|
||||
uint64_t tndBIdxOffsetForK = actualSeqKPrefixSum * constInfo.kHeadNum * constInfo.headDim;
|
||||
keyCoreOffset = tndBIdxOffsetForK + runInfo.s2Idx * constInfo.s2BaseSize * constInfo.kHeadNum * constInfo.headDim;
|
||||
keyScaleCoreOffset = (actualSeqKPrefixSum + runInfo.s2Idx * constInfo.s2BaseSize) * constInfo.kHeadNum;
|
||||
runInfo.tensorQueryOffset = queryCoreOffset;
|
||||
runInfo.tensorKeyOffset = keyCoreOffset;
|
||||
runInfo.tensorKeyScaleOffset = keyScaleCoreOffset;
|
||||
runInfo.tensorWeightsOffset = weightsCoreOffset;
|
||||
runInfo.indiceOutOffset = indiceOutCoreOffset;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::Process()
|
||||
{
|
||||
// 没有计算任务,直接清理输出
|
||||
if (coreZeroEnable == 0) {
|
||||
ProcessInvalid();
|
||||
return;
|
||||
}
|
||||
ProcessMain();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessInvalid()
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
uint32_t aivCoreNum = GetBlockNum() * 2; // 2 means c:v = 1:2
|
||||
uint64_t totalOutputSize =
|
||||
constInfo.batchSize * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount;
|
||||
uint64_t singleCoreSize =
|
||||
QLICommon::Align((totalOutputSize + aivCoreNum - 1) / aivCoreNum, GM_ALIGN_BYTES / sizeof(OUT_T));
|
||||
uint64_t baseSize = tmpBlockIdx * singleCoreSize;
|
||||
if (baseSize < totalOutputSize) {
|
||||
uint64_t dealSize =
|
||||
(baseSize + singleCoreSize <= totalOutputSize) ? singleCoreSize : totalOutputSize - baseSize;
|
||||
GlobalTensor<OUT_T> output = indiceOutGm[baseSize];
|
||||
AscendC::InitGlobalMemory(output, dealSize, constInfo.INVALID_IDX);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessMain()
|
||||
{
|
||||
// 无任务核直接返回
|
||||
if (constInfo.coreEnable == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.AllocEventID();
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_MTE2>(constInfo.syncV1C1);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_MTE2>(constInfo.syncV1C1);
|
||||
} else {
|
||||
matmulService.AllocEventID();
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V0);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V0);
|
||||
}
|
||||
|
||||
QLICommon::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE];
|
||||
|
||||
// 适配左闭右开
|
||||
if (constInfo.bN2Start == constInfo.bN2End) {
|
||||
if (constInfo.gS1Start != constInfo.gS1End || constInfo.s2Start != constInfo.s2End) {
|
||||
constInfo.bN2End += 1;
|
||||
}
|
||||
} else if ((constInfo.gS1End != 0) || (constInfo.s2End != 0)){
|
||||
constInfo.bN2End += 1;
|
||||
}
|
||||
|
||||
uint32_t gloop = 0;
|
||||
for (uint32_t bN2LoopIdx = constInfo.bN2Start; bN2LoopIdx < constInfo.bN2End; bN2LoopIdx++) {
|
||||
CalcGS1LoopParams(bN2LoopIdx);
|
||||
if (tempLoopInfo.curActSeqLenIsZero) {
|
||||
DealActSeqLenIsZero(tempLoopInfo.bIdx, tempLoopInfo.n2Idx, 0U);
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
if (bN2LoopIdx + 1 == constInfo.bN2End && gloop > 0) {
|
||||
CrossCoreWaitFlag(constInfo.syncC1V1);
|
||||
vectorService.ProcessVec1(runInfo[1 - gloop % LI_QUANT_PRELOAD_TASK_CACHE_SIZE]);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_MTE3>(
|
||||
constInfo.syncV1C1); // 反向同步 1
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
for (uint32_t gS1LoopIdx = constInfo.gS1Start; gS1LoopIdx < tempLoopInfo.gS1LoopEnd; gS1LoopIdx++) {
|
||||
CalcS2LoopParams(bN2LoopIdx, gS1LoopIdx);
|
||||
bool isEnd = (bN2LoopIdx + 1 == constInfo.bN2End) && (gS1LoopIdx + 1 == tempLoopInfo.gS1LoopEnd);
|
||||
uint32_t extraLoop = isEnd ? LI_QUANT_PRELOAD_TASK_CACHE_SIZE - 1 : 0; // 只preload一轮
|
||||
|
||||
for (uint32_t s2LoopIdx = constInfo.s2Start; s2LoopIdx < (tempLoopInfo.s2LoopEnd + extraLoop); s2LoopIdx++) {
|
||||
ProcessBaseBlock(gloop, s2LoopIdx, runInfo);
|
||||
++gloop;
|
||||
}
|
||||
constInfo.s2Start = 0;
|
||||
}
|
||||
if (tempLoopInfo.needDealActS1LessThanS1) {
|
||||
DealActSeqLenIsZero(tempLoopInfo.bIdx, tempLoopInfo.n2Idx, tempLoopInfo.actS1Size);
|
||||
}
|
||||
constInfo.gS1Start = 0;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.FreeEventID();
|
||||
CrossCoreWaitFlag(constInfo.syncC1V0);
|
||||
CrossCoreWaitFlag(constInfo.syncC1V0);
|
||||
} else {
|
||||
matmulService.FreeEventID();
|
||||
CrossCoreWaitFlag(constInfo.syncV1C1);
|
||||
CrossCoreWaitFlag(constInfo.syncV1C1);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessBaseBlock(uint32_t loop, uint64_t s2LoopIdx,
|
||||
QLICommon::RunInfo runInfo[LI_QUANT_PRELOAD_TASK_CACHE_SIZE])
|
||||
{
|
||||
int32_t curTaskId = loop % LI_QUANT_PRELOAD_TASK_CACHE_SIZE;
|
||||
QLICommon::RunInfo &curRunInfo = runInfo[curTaskId];
|
||||
QLICommon::RunInfo &lastRunInfo = runInfo[1 - curTaskId];
|
||||
|
||||
CalcRunInfo(loop, s2LoopIdx, curRunInfo);
|
||||
|
||||
if (curRunInfo.isValid) {
|
||||
if ASCEND_IS_AIC {
|
||||
if (curRunInfo.isFirstS2InnerLoop) {
|
||||
CrossCoreWaitFlag(constInfo.syncV0C1);
|
||||
}
|
||||
CrossCoreWaitFlag(constInfo.syncV1C1); // 反向同步 1
|
||||
matmulService.ComputeMm1(curRunInfo);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V1);
|
||||
if (curRunInfo.isLastS2InnerLoop) {
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_FIX>(constInfo.syncC1V0); // 反向同步 0
|
||||
}
|
||||
} else {
|
||||
if (curRunInfo.isFirstS2InnerLoop) {
|
||||
CrossCoreWaitFlag(constInfo.syncC1V0); // 反向同步 0
|
||||
vectorService.ProcessVec0(curRunInfo);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_MTE3>(constInfo.syncV0C1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (lastRunInfo.isValid) {
|
||||
if ASCEND_IS_AIV {
|
||||
CrossCoreWaitFlag(constInfo.syncC1V1);
|
||||
vectorService.ProcessVec1(lastRunInfo);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::FIA_SYNC_MODE2, PIPE_MTE3>(constInfo.syncV1C1); // 反向同步 1
|
||||
}
|
||||
lastRunInfo.isValid = false;
|
||||
}
|
||||
}
|
||||
} // namespace QLIKernel
|
||||
#endif // QUANT_LIGHTNING_INDEXER_KERNEL_H
|
||||
@@ -0,0 +1,613 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_service_cube.h
|
||||
* \brief use 5 buffer for matmul l1, better pipeline
|
||||
*/
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_SERVICE_CUBE_H
|
||||
#define QUANT_LIGHTNING_INDEXER_SERVICE_CUBE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
struct MmInfo {
|
||||
int64_t s2L0LoopId;
|
||||
int64_t s1gL0LoopId;
|
||||
int64_t s2L0RealSize;
|
||||
int64_t s2GmOffset;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
class QLIMatmul {
|
||||
public:
|
||||
using Q_T = typename QLIT::queryType;
|
||||
using K_T = typename QLIT::keyType;
|
||||
|
||||
__aicore__ inline QLIMatmul(){};
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void InitMm1GlobalTensor(const GlobalTensor<int32_t> &blkTableGm, const GlobalTensor<K_T> &keyGm,
|
||||
const GlobalTensor<Q_T> &queryGm, const GlobalTensor<float> &mm1ResGm,
|
||||
const GlobalTensor<half> &weightWorkspaceGm);
|
||||
__aicore__ inline void InitParams(const ConstInfo &constInfo);
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
__aicore__ inline void ComputeMm1(const QLICommon::RunInfo &runInfo);
|
||||
|
||||
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
|
||||
static constexpr uint64_t DOUBLE_BUF_NUM = 2;
|
||||
static constexpr uint64_t L0AB_BUF_NUM = 4;
|
||||
|
||||
static constexpr uint32_t KEY_MTE1_MTE2_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t QW_MTE1_MTE2_EVENT = EVENT_ID5; // KEY_MTE1_MTE2_EVENT + DOUBLE_BUF_NUM;
|
||||
static constexpr uint32_t M_MTE1_EVENT = EVENT_ID3;
|
||||
static constexpr uint32_t M_FIX_EVENT = EVENT_ID0;
|
||||
static constexpr uint32_t FIX_M_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t FIX_MTE1_EVENT = EVENT_ID4;
|
||||
|
||||
static constexpr uint64_t S8_BLOCK_CUBE = 32;
|
||||
|
||||
static constexpr uint32_t MTE2_MTE1_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t MTE1_M_EVENT = EVENT_ID2;
|
||||
|
||||
static constexpr uint64_t D_BASIC_BLOCK = 128;
|
||||
static constexpr uint64_t S1G_BASIC_BLOCK_L1 = 256;
|
||||
|
||||
static constexpr uint64_t S1G_BASIC_BLOCK_L0 = 128;
|
||||
static constexpr uint64_t S2_BASIC_BLOCK_L0 = 128;
|
||||
|
||||
static constexpr uint64_t QUERY_BUFFER_OFFSET = S1G_BASIC_BLOCK_L1 * D_BASIC_BLOCK;
|
||||
static constexpr uint64_t SL1_BUFFER_OFFSET = S1G_BASIC_BLOCK_L0 * S2_BASIC_BLOCK_L0;
|
||||
static constexpr uint64_t KEY_BUFFER_OFFSET = S2_BASIC_BLOCK_L0 * D_BASIC_BLOCK;
|
||||
static constexpr uint64_t WEIGHT_BUFFER_OFFSET = S1G_BASIC_BLOCK_L1 * BLOCK_CUBE;
|
||||
static constexpr uint64_t L0AB_BUFFER_OFFSET_S8_16K = 16 * 1024;
|
||||
static constexpr uint64_t L0AB_BUFFER_OFFSET_FP16_16K = 16 * 512;
|
||||
static constexpr uint64_t L0C_BUFFER_OFFSET = 64 * 256;
|
||||
|
||||
private:
|
||||
__aicore__ inline void WeightDmaCopy(uint64_t s1gL1RealSize, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void LoadKeyToL0b(uint64_t s2L0RealSize);
|
||||
__aicore__ inline void LoadQueryToL0a(uint64_t s1gL1Offset, uint64_t s1gL1RealSize, uint64_t s1gL0RealSize);
|
||||
__aicore__ inline void QueryNd2Nz(uint64_t s1gL1RealSize, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void KeyNd2Nz(uint64_t s2L1RealSize, const MmInfo &mmInfo, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void FixpSToL1(uint64_t s1gL0RealSize, uint64_t s2L0RealSize);
|
||||
__aicore__ inline void LoadSToL0b(uint64_t s1gL1RealSize, uint64_t s2L0RealSize, uint64_t sL1BufIdx,
|
||||
int64_t mStartPt);
|
||||
__aicore__ inline void LoadWeightToL0a(uint64_t s1gL1Offset);
|
||||
__aicore__ inline void ComputeWs(uint64_t s1gL0RealSize, uint64_t s2L0RealSize, int64_t s1gOffset);
|
||||
__aicore__ inline void FixpResToGm(uint64_t s1L0RealCount, uint64_t s2L0RealSize, uint64_t s1GmOffset,
|
||||
uint64_t s2GmOffset, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void ComputeQk(uint64_t s1gL0RealSize, uint64_t s2L0RealSize);
|
||||
__aicore__ inline void ProcessWs(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, uint64_t sL1BufIdx,
|
||||
const MmInfo &mmInfo, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void ProcessQk(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, uint64_t s1L0LoopCnt,
|
||||
const MmInfo &mmInfo, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void CalcMmInfo(MmInfo &mmInfo, uint64_t loopIdx, uint64_t s1L0LoopCnt, const MmInfo &lastMmInfo,
|
||||
const QLICommon::RunInfo &runInfo);
|
||||
static constexpr LI_LAYOUT Q_LAYOUT_T = QLIT::layout;
|
||||
static constexpr LI_LAYOUT K_LAYOUT_T = QLIT::keyLayout;
|
||||
GlobalTensor<int32_t> blkTableGm_;
|
||||
GlobalTensor<K_T> keyGm_;
|
||||
GlobalTensor<Q_T> queryGm_;
|
||||
GlobalTensor<half> weightGm_;
|
||||
GlobalTensor<float> mm1ResGm_;
|
||||
|
||||
TBuf<TPosition::A1> bufQL1_;
|
||||
LocalTensor<Q_T> queryL1_;
|
||||
TBuf<TPosition::B1> bufKeyL1_;
|
||||
LocalTensor<K_T> keyL1_;
|
||||
TBuf<TPosition::A1> bufWeightL1_;
|
||||
LocalTensor<half> weightL1_;
|
||||
TBuf<TPosition::B1> bufSL1_;
|
||||
LocalTensor<half> sL1_;
|
||||
|
||||
TBuf<TPosition::A2> bufL0A_;
|
||||
LocalTensor<Q_T> l0a_;
|
||||
TBuf<TPosition::B2> bufL0B_;
|
||||
LocalTensor<K_T> l0b_;
|
||||
|
||||
TBuf<TPosition::CO1> bufL0C_;
|
||||
LocalTensor<int32_t> cL0_;
|
||||
|
||||
uint64_t keyL1BufIdx_ = 0;
|
||||
uint64_t qwL1Mte2BufIdx_ = 0;
|
||||
uint64_t sL1BufIdx_ = 0;
|
||||
uint64_t l0BufIdx_ = 0;
|
||||
uint64_t l0cBufIdx_ = 0;
|
||||
|
||||
ConstInfo constInfo_;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::InitParams(const ConstInfo &constInfo)
|
||||
{
|
||||
constInfo_ = constInfo;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(bufQL1_, DOUBLE_BUF_NUM * S1G_BASIC_BLOCK_L1 * D_BASIC_BLOCK * sizeof(Q_T));
|
||||
queryL1_ = bufQL1_.Get<Q_T>();
|
||||
pipe->InitBuffer(bufKeyL1_, DOUBLE_BUF_NUM * S2_BASIC_BLOCK_L0 * D_BASIC_BLOCK * sizeof(K_T));
|
||||
keyL1_ = bufKeyL1_.Get<K_T>();
|
||||
|
||||
pipe->InitBuffer(bufWeightL1_, DOUBLE_BUF_NUM * S1G_BASIC_BLOCK_L1 * BLOCK_CUBE * sizeof(half));
|
||||
weightL1_ = bufWeightL1_.Get<half>();
|
||||
pipe->InitBuffer(bufSL1_, DOUBLE_BUF_NUM * S2_BASIC_BLOCK_L0 * S1G_BASIC_BLOCK_L0 * sizeof(half));
|
||||
sL1_ = bufSL1_.Get<half>();
|
||||
|
||||
pipe->InitBuffer(bufL0A_, 64 * 1024);
|
||||
l0a_ = bufL0A_.Get<Q_T>();
|
||||
pipe->InitBuffer(bufL0B_, 64 * 1024);
|
||||
l0b_ = bufL0B_.Get<K_T>();
|
||||
|
||||
pipe->InitBuffer(bufL0C_, 128 * 1024);
|
||||
cL0_ = bufL0C_.Get<int32_t>();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::InitMm1GlobalTensor(const GlobalTensor<int32_t> &blkTableGm,
|
||||
const GlobalTensor<K_T> &keyGm,
|
||||
const GlobalTensor<Q_T> &queryGm,
|
||||
const GlobalTensor<float> &mm1ResGm,
|
||||
const GlobalTensor<half> &weightWorkspaceGm)
|
||||
{
|
||||
blkTableGm_ = blkTableGm;
|
||||
keyGm_ = keyGm;
|
||||
queryGm_ = queryGm;
|
||||
mm1ResGm_ = mm1ResGm;
|
||||
weightGm_ = weightWorkspaceGm;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ProcessWs(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, uint64_t sL1BufIdx,
|
||||
const MmInfo &mmInfo, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0cBufIdx_ % DOUBLE_BUF_NUM);
|
||||
for (int64_t s1gOffset = 0; s1gOffset < s1gL0RealSize; s1gOffset += constInfo_.gSize) {
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0AB_BUF_NUM);
|
||||
LoadSToL0b(s1gL0RealSize, mmInfo.s2L0RealSize, sL1BufIdx, s1gOffset);
|
||||
LoadWeightToL0a(s1gOffset + s1gL1Offset);
|
||||
|
||||
ComputeWs(s1gL0RealSize, mmInfo.s2L0RealSize, s1gOffset);
|
||||
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0AB_BUF_NUM);
|
||||
l0BufIdx_++;
|
||||
}
|
||||
|
||||
FixpResToGm(s1gL0RealSize / constInfo_.gSize, mmInfo.s2L0RealSize, s1gL1Offset / constInfo_.gSize,
|
||||
mmInfo.s2L0LoopId * S2_BASIC_BLOCK_L0, runInfo);
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0cBufIdx_ % DOUBLE_BUF_NUM);
|
||||
l0cBufIdx_++;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ProcessQk(uint64_t s1gL0RealSize, uint64_t s1gL1Offset, uint64_t s1L0LoopCnt,
|
||||
const MmInfo &mmInfo, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
if (mmInfo.s1gL0LoopId == 0) {
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + keyL1BufIdx_ % DOUBLE_BUF_NUM);
|
||||
if constexpr (K_LAYOUT_T == LI_LAYOUT::PA_BSND) {
|
||||
KeyNd2NzForPA(mmInfo.s2L0RealSize, runInfo.s2Idx * constInfo_.s2BaseSize + mmInfo.s2GmOffset, runInfo);
|
||||
} else {
|
||||
KeyNd2Nz(mmInfo.s2L0RealSize, mmInfo, runInfo);
|
||||
}
|
||||
|
||||
SetFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
}
|
||||
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0AB_BUF_NUM);
|
||||
LoadQueryToL0a(s1gL1Offset, runInfo.actMBaseSize, s1gL0RealSize);
|
||||
LoadKeyToL0b(mmInfo.s2L0RealSize);
|
||||
|
||||
if (mmInfo.s1gL0LoopId + 1 >= s1L0LoopCnt) {
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + keyL1BufIdx_ % DOUBLE_BUF_NUM);
|
||||
keyL1BufIdx_++;
|
||||
}
|
||||
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0cBufIdx_ % DOUBLE_BUF_NUM);
|
||||
ComputeQk(s1gL0RealSize, mmInfo.s2L0RealSize);
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0AB_BUF_NUM);
|
||||
|
||||
FixpSToL1(s1gL0RealSize, mmInfo.s2L0RealSize);
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0cBufIdx_ % DOUBLE_BUF_NUM);
|
||||
l0BufIdx_++;
|
||||
l0cBufIdx_++;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::CalcMmInfo(MmInfo &mmInfo, uint64_t loopIdx, uint64_t s1L0LoopCnt,
|
||||
const MmInfo &lastMmInfo, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
mmInfo.s2L0LoopId = loopIdx / s1L0LoopCnt;
|
||||
mmInfo.s1gL0LoopId = loopIdx % s1L0LoopCnt;
|
||||
|
||||
if (mmInfo.s1gL0LoopId == 0) {
|
||||
mmInfo.s2GmOffset = mmInfo.s2L0LoopId * S2_BASIC_BLOCK_L0;
|
||||
mmInfo.s2L0RealSize = mmInfo.s2GmOffset + S2_BASIC_BLOCK_L0 > runInfo.actualSingleProcessSInnerSize
|
||||
? runInfo.actualSingleProcessSInnerSize - mmInfo.s2GmOffset
|
||||
: S2_BASIC_BLOCK_L0;
|
||||
} else {
|
||||
mmInfo.s2L0RealSize = lastMmInfo.s2L0RealSize;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ComputeMm1(const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
if (runInfo.isFirstS2InnerLoop) {
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM);
|
||||
QueryNd2Nz(runInfo.actMBaseSize, runInfo); // 256 * 128 // L1BasicBlock
|
||||
WeightDmaCopy(runInfo.actMBaseSize, runInfo);
|
||||
}
|
||||
int64_t loopIdx = 0;
|
||||
int64_t s2L0LoopCnt = CeilDiv(runInfo.actualSingleProcessSInnerSize, S2_BASIC_BLOCK_L0); // 2048取128
|
||||
int64_t s1L0LoopCnt = CeilDiv(runInfo.actMBaseSize, S1G_BASIC_BLOCK_L0); // 256取128
|
||||
int64_t s1gL1Offset[2] = {0, static_cast<int64_t>(S1G_BASIC_BLOCK_L0)};
|
||||
int64_t s1gL0RealSize[2] = {s1L0LoopCnt > 1 ? static_cast<int64_t>(S1G_BASIC_BLOCK_L0) : runInfo.actMBaseSize,
|
||||
runInfo.actMBaseSize - s1gL1Offset[1]};
|
||||
MmInfo mmInfo[2];
|
||||
CalcMmInfo(mmInfo[loopIdx & 1], loopIdx, s1L0LoopCnt, mmInfo[(loopIdx + 1) & 1], runInfo);
|
||||
|
||||
ProcessQk(s1gL0RealSize[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt],
|
||||
s1gL1Offset[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt], s1L0LoopCnt, mmInfo[loopIdx & 1],
|
||||
runInfo);
|
||||
|
||||
SetFlag<HardEvent::FIX_MTE1>(FIX_MTE1_EVENT + sL1BufIdx_ % DOUBLE_BUF_NUM);
|
||||
sL1BufIdx_++;
|
||||
loopIdx++;
|
||||
|
||||
while (loopIdx < s2L0LoopCnt * s1L0LoopCnt) {
|
||||
CalcMmInfo(mmInfo[loopIdx & 1], loopIdx, s1L0LoopCnt, mmInfo[(loopIdx + 1) & 1], runInfo);
|
||||
|
||||
ProcessQk(s1gL0RealSize[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt],
|
||||
s1gL1Offset[mmInfo[loopIdx & 1].s1gL0LoopId % s1L0LoopCnt], s1L0LoopCnt, mmInfo[loopIdx & 1],
|
||||
runInfo);
|
||||
|
||||
SetFlag<HardEvent::FIX_MTE1>(FIX_MTE1_EVENT + sL1BufIdx_ % DOUBLE_BUF_NUM);
|
||||
sL1BufIdx_++;
|
||||
|
||||
WaitFlag<HardEvent::FIX_MTE1>(FIX_MTE1_EVENT + sL1BufIdx_ % DOUBLE_BUF_NUM);
|
||||
|
||||
ProcessWs(s1gL0RealSize[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt],
|
||||
s1gL1Offset[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt], sL1BufIdx_,
|
||||
mmInfo[(loopIdx + 1) & 1], runInfo);
|
||||
loopIdx++;
|
||||
}
|
||||
|
||||
WaitFlag<HardEvent::FIX_MTE1>(FIX_MTE1_EVENT + (sL1BufIdx_ + 1) % DOUBLE_BUF_NUM);
|
||||
|
||||
ProcessWs(s1gL0RealSize[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt],
|
||||
s1gL1Offset[mmInfo[(loopIdx + 1) & 1].s1gL0LoopId % s1L0LoopCnt], sL1BufIdx_ - 1,
|
||||
mmInfo[(loopIdx + 1) & 1], runInfo);
|
||||
|
||||
if (runInfo.isLastS2InnerLoop) {
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM);
|
||||
qwL1Mte2BufIdx_++;
|
||||
}
|
||||
}
|
||||
|
||||
// blkNum, blkSize, N2, D
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
uint64_t s2L1Offset = 0;
|
||||
while (s2L1Offset < s2L1RealSize) {
|
||||
uint64_t s2BlkId = (s2L1Offset + s2GmOffset) / constInfo_.kCacheBlockSize;
|
||||
uint64_t s2BlkOffset = (s2L1Offset + s2GmOffset) % constInfo_.kCacheBlockSize;
|
||||
uint64_t keyGmOffset = blkTableGm_.GetValue(runInfo.bIdx * constInfo_.maxBlockNumPerBatch + s2BlkId) *
|
||||
constInfo_.stride +
|
||||
s2BlkOffset * constInfo_.headDim;
|
||||
uint64_t s2Mte2Size = s2L1RealSize - s2L1Offset;
|
||||
s2Mte2Size = s2BlkOffset + s2Mte2Size >= constInfo_.kCacheBlockSize ? constInfo_.kCacheBlockSize - s2BlkOffset
|
||||
: s2Mte2Size;
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s2Mte2Size; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = constInfo_.headDim;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
DataCopy(keyL1_[(keyL1BufIdx_ % DOUBLE_BUF_NUM) * KEY_BUFFER_OFFSET + s2L1Offset * S8_BLOCK_CUBE],
|
||||
keyGm_[keyGmOffset], nd2nzPara);
|
||||
|
||||
s2L1Offset += s2Mte2Size;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::KeyNd2Nz(uint64_t s2L1RealSize, const MmInfo &mmInfo,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
uint64_t dStride = constInfo_.headDim;
|
||||
if constexpr (K_LAYOUT_T == LI_LAYOUT::BSND || K_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
dStride = constInfo_.headDim * constInfo_.kHeadNum; // constInfo_.kHeadNum
|
||||
}
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s2L1RealSize; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = dStride;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
// 默认一块buf最多放两份
|
||||
DataCopy(keyL1_[(keyL1BufIdx_ % DOUBLE_BUF_NUM) * KEY_BUFFER_OFFSET],
|
||||
keyGm_[runInfo.tensorKeyOffset + mmInfo.s2GmOffset * constInfo_.headDim], nd2nzPara);
|
||||
}
|
||||
|
||||
// batch, s1, g, 1
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::WeightDmaCopy(uint64_t s1gL1RealSize, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
DataCopyParams copyInParams;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.blockLen = s1gL1RealSize;
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
DataCopy(weightL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * WEIGHT_BUFFER_OFFSET],
|
||||
weightGm_[runInfo.loop % DOUBLE_BUF_NUM * BLOCK_CUBE * constInfo_.mBaseSize], copyInParams);
|
||||
}
|
||||
|
||||
// batch, s1, n2, g, d
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::QueryNd2Nz(uint64_t s1gL1RealSize, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s1gL1RealSize; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = constInfo_.headDim;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s1gL1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
// 默认一块buf最多放两份
|
||||
DataCopy(queryL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * QUERY_BUFFER_OFFSET], queryGm_[runInfo.tensorQueryOffset],
|
||||
nd2nzPara);
|
||||
}
|
||||
|
||||
// s1g, d
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadQueryToL0a(uint64_t s1gL1Offset, uint64_t s1gL1RealSize,
|
||||
uint64_t s1gL0RealSize)
|
||||
{
|
||||
LoadData3DParamsV2<Q_T> loadData3DParams;
|
||||
// SetFmatrixParams
|
||||
loadData3DParams.l1H = CeilDiv(s1gL1RealSize, BLOCK_CUBE); // Hin=M1=8
|
||||
loadData3DParams.l1W = BLOCK_CUBE; // Win=M0
|
||||
loadData3DParams.channelSize = constInfo_.headDim; // Cin=K
|
||||
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
// SetLoadToA0Params
|
||||
loadData3DParams.mExtension = CeilAlign(s1gL0RealSize, BLOCK_CUBE); // M height维度目的
|
||||
loadData3DParams.kExtension = constInfo_.headDim; // K width维度目的
|
||||
loadData3DParams.mStartPt = s1gL1Offset;
|
||||
loadData3DParams.kStartPt = 0;
|
||||
loadData3DParams.strideW = 1;
|
||||
loadData3DParams.strideH = 1;
|
||||
loadData3DParams.filterW = 1;
|
||||
loadData3DParams.filterSizeW = (1 >> 8) & 255;
|
||||
loadData3DParams.filterH = 1;
|
||||
loadData3DParams.filterSizeH = (1 >> 8) & 255;
|
||||
loadData3DParams.dilationFilterW = 1;
|
||||
loadData3DParams.dilationFilterH = 1;
|
||||
loadData3DParams.enTranspose = 0;
|
||||
loadData3DParams.fMatrixCtrl = 0;
|
||||
|
||||
LoadData<Q_T, LOAD3DV2_CONFIG>(l0a_[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_S8_16K],
|
||||
queryL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * QUERY_BUFFER_OFFSET],
|
||||
loadData3DParams);
|
||||
}
|
||||
|
||||
// s1, g, s2 --> 2 * 64* 128
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadSToL0b(uint64_t s1gL1RealSize, uint64_t s2L0RealSize, uint64_t sL1BufIdx,
|
||||
int64_t mStartPt)
|
||||
{
|
||||
LoadData3DParamsV2<half> loadData3DParams;
|
||||
// SetFmatrixParams
|
||||
loadData3DParams.l1H = S1G_BASIC_BLOCK_L0 / BLOCK_CUBE; // Hin=M1=8
|
||||
loadData3DParams.l1W = BLOCK_CUBE; // Win=M0
|
||||
loadData3DParams.channelSize = CeilAlign(s2L0RealSize, BLOCK_CUBE); // Cin=K
|
||||
|
||||
loadData3DParams.padList[0] = 0;
|
||||
loadData3DParams.padList[1] = 0;
|
||||
loadData3DParams.padList[2] = 0;
|
||||
loadData3DParams.padList[3] = 255; // 尾部数据不影响滑窗的结果
|
||||
|
||||
// SetLoadToA0Params
|
||||
loadData3DParams.mExtension = constInfo_.gSize; // M height维度目的
|
||||
loadData3DParams.kExtension = CeilAlign(s2L0RealSize, BLOCK_CUBE); // K width维度目的
|
||||
loadData3DParams.kStartPt = 0;
|
||||
loadData3DParams.strideW = 1;
|
||||
loadData3DParams.strideH = 1;
|
||||
loadData3DParams.filterW = 1;
|
||||
loadData3DParams.filterSizeW = (1 >> 8) & 255;
|
||||
loadData3DParams.filterH = 1;
|
||||
loadData3DParams.filterSizeH = (1 >> 8) & 255;
|
||||
loadData3DParams.dilationFilterW = 1;
|
||||
loadData3DParams.dilationFilterH = 1;
|
||||
loadData3DParams.enTranspose = 1;
|
||||
loadData3DParams.fMatrixCtrl = 0;
|
||||
|
||||
loadData3DParams.mStartPt = mStartPt;
|
||||
LoadData<half, LOAD3DV2_CONFIG>(
|
||||
l0b_.template ReinterpretCast<half>()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K],
|
||||
sL1_[(sL1BufIdx % DOUBLE_BUF_NUM) * SL1_BUFFER_OFFSET], loadData3DParams);
|
||||
}
|
||||
|
||||
// s1,g,1(16), 2,64,16
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadWeightToL0a(uint64_t s1gL1Offset)
|
||||
{
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0;
|
||||
loadData2DParams.repeatTimes = CeilDiv(constInfo_.gSize, BLOCK_CUBE);
|
||||
loadData2DParams.srcStride = 1;
|
||||
loadData2DParams.dstGap = 0;
|
||||
loadData2DParams.ifTranspose = true;
|
||||
LoadData(l0a_.template ReinterpretCast<half>()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K],
|
||||
weightL1_[(qwL1Mte2BufIdx_ % DOUBLE_BUF_NUM) * WEIGHT_BUFFER_OFFSET + s1gL1Offset* BLOCK_CUBE],
|
||||
loadData2DParams);
|
||||
}
|
||||
|
||||
// s2, d -> 128,128
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadKeyToL0b(uint64_t s2L0RealSize)
|
||||
{
|
||||
LoadData2DParams loadData2DParams;
|
||||
loadData2DParams.startIndex = 0;
|
||||
loadData2DParams.repeatTimes = CeilDiv(s2L0RealSize, BLOCK_CUBE) * CeilDiv(constInfo_.headDim, S8_BLOCK_CUBE);
|
||||
loadData2DParams.srcStride = 1;
|
||||
loadData2DParams.dstGap = 0;
|
||||
loadData2DParams.ifTranspose = false;
|
||||
LoadData(l0b_[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_S8_16K],
|
||||
keyL1_[(keyL1BufIdx_ % DOUBLE_BUF_NUM) * KEY_BUFFER_OFFSET], loadData2DParams);
|
||||
}
|
||||
|
||||
// A: s1,g,1(16) B: s1,g,s2 C: s1, 1(16), s2
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ComputeWs(uint64_t s1gL0RealSize, uint64_t s2L0RealSize, int64_t s1gOffset)
|
||||
{
|
||||
SetFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
WaitFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = BLOCK_CUBE;
|
||||
mmadParams.n = s2L0RealSize;
|
||||
mmadParams.k = constInfo_.gSize;
|
||||
mmadParams.cmatrixInitVal = true;
|
||||
mmadParams.cmatrixSource = false;
|
||||
Mmad(cL0_.template ReinterpretCast<float>()[(l0cBufIdx_ % DOUBLE_BUF_NUM) * L0C_BUFFER_OFFSET +
|
||||
s1gOffset * S2_BASIC_BLOCK_L0],
|
||||
l0a_.template ReinterpretCast<half>()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K],
|
||||
l0b_.template ReinterpretCast<half>()[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_FP16_16K],
|
||||
mmadParams);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ComputeQk(uint64_t s1gL0RealSize, uint64_t s2L0RealSize)
|
||||
{
|
||||
SetFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
WaitFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = CeilAlign(s1gL0RealSize, BLOCK_CUBE);
|
||||
mmadParams.n = s2L0RealSize;
|
||||
mmadParams.k = constInfo_.headDim;
|
||||
mmadParams.cmatrixInitVal = true;
|
||||
mmadParams.cmatrixSource = false;
|
||||
Mmad(cL0_[(l0cBufIdx_ % DOUBLE_BUF_NUM) * L0C_BUFFER_OFFSET],
|
||||
l0a_[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_S8_16K],
|
||||
l0b_[(l0BufIdx_ % L0AB_BUF_NUM) * L0AB_BUFFER_OFFSET_S8_16K], mmadParams);
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::FixpSToL1(uint64_t s1gL0RealSize, uint64_t s2L0RealSize)
|
||||
{
|
||||
SetFlag<HardEvent::M_FIX>(M_FIX_EVENT);
|
||||
WaitFlag<HardEvent::M_FIX>(M_FIX_EVENT);
|
||||
DataCopyCO12DstParams params;
|
||||
params.mSize = CeilAlign(s1gL0RealSize, BLOCK_CUBE);
|
||||
params.nSize = CeilAlign(s2L0RealSize, BLOCK_CUBE);
|
||||
params.dstStride = S1G_BASIC_BLOCK_L0;
|
||||
params.srcStride = params.mSize;
|
||||
params.quantPre = QuantMode_t::DEQF16;
|
||||
params.reluPre = 1;
|
||||
params.channelSplit = 0;
|
||||
params.nz2ndEn = 0;
|
||||
SetFixpipePreQuantFlag(0x3a800000);
|
||||
DataCopy(sL1_[(sL1BufIdx_ % DOUBLE_BUF_NUM) * SL1_BUFFER_OFFSET],
|
||||
cL0_[(l0cBufIdx_ % DOUBLE_BUF_NUM) * L0C_BUFFER_OFFSET], params);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::FixpResToGm(uint64_t s1L0RealCount, uint64_t s2L0RealSize, uint64_t s1GmOffset,
|
||||
uint64_t s2GmOffset, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
SetFlag<HardEvent::M_FIX>(M_FIX_EVENT);
|
||||
WaitFlag<HardEvent::M_FIX>(M_FIX_EVENT);
|
||||
|
||||
AscendC::DataCopyCO12DstParams intriParams;
|
||||
intriParams.mSize = 1;
|
||||
intriParams.nSize = s2L0RealSize;
|
||||
intriParams.dstStride = constInfo_.s2BaseSize;
|
||||
intriParams.srcStride = 16;
|
||||
// set mode according to dtype
|
||||
intriParams.quantPre = QuantMode_t::NoQuant;
|
||||
intriParams.nz2ndEn = true;
|
||||
intriParams.reluPre = 0;
|
||||
AscendC::SetFixpipeNz2ndFlag(s1L0RealCount, CeilDiv(constInfo_.gSize, BLOCK_CUBE) * S2_BASIC_BLOCK_L0 / BLOCK_CUBE,
|
||||
2048);
|
||||
AscendC::DataCopy(mm1ResGm_[(runInfo.loop % 2) * constInfo_.mBaseSize / constInfo_.gSize * constInfo_.s2BaseSize +
|
||||
s1GmOffset * intriParams.dstStride + s2GmOffset],
|
||||
cL0_.template ReinterpretCast<float>()[(l0cBufIdx_ % DOUBLE_BUF_NUM) * L0C_BUFFER_OFFSET],
|
||||
intriParams);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::AllocEventID()
|
||||
{
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 2);
|
||||
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + 0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + 1);
|
||||
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 0);
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 1);
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 2);
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 3);
|
||||
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + 0);
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + 1);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::FreeEventID()
|
||||
{
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 2);
|
||||
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + 0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QW_MTE1_MTE2_EVENT + 1);
|
||||
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 0);
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 1);
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 2);
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 3);
|
||||
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + 0);
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + 1);
|
||||
}
|
||||
} // namespace QLIKernel
|
||||
#endif // QUANT_LIGHTNING_INDEXER_SERVICE_CUBE_H
|
||||
@@ -0,0 +1,437 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_service_vector.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_SERVICE_VECTOR_H
|
||||
#define QUANT_LIGHTNING_INDEXER_SERVICE_VECTOR_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
#include "quant_lightning_indexer_vector.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
using namespace QLIServiceVec;
|
||||
constexpr uint32_t BASE_TOPK = 2048;
|
||||
constexpr uint32_t BASE_TOPK_VALUE_IDX_SIZE = 4096;
|
||||
constexpr uint32_t ELE_NUM_32 = 32;
|
||||
constexpr uint32_t ELE_NUM_128 = 128;
|
||||
constexpr uint32_t ELE_NUM_512 = 512;
|
||||
|
||||
template <typename QLIT>
|
||||
class QLIVector {
|
||||
public:
|
||||
// =================================类型定义区=================================
|
||||
static constexpr LI_LAYOUT Q_LAYOUT_T = QLIT::layout;
|
||||
static constexpr LI_LAYOUT K_LAYOUT_T = QLIT::keyLayout;
|
||||
static constexpr bool PAGE_ATTENTION = QLIT::pageAttention;
|
||||
// MM输出数据类型, 当前只支持float
|
||||
using MM1_OUT_T = float;
|
||||
|
||||
__aicore__ inline QLIVector(){};
|
||||
__aicore__ inline void ProcessVec0(const QLICommon::RunInfo &info);
|
||||
__aicore__ inline void ProcessVec1(const QLICommon::RunInfo &info);
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void InitParams(const struct QLICommon::ConstInfo &constInfo,
|
||||
const QLITilingData *__restrict tilingData);
|
||||
__aicore__ inline void InitVecWorkspaceTensor(GlobalTensor<half> vec0OutGm, GlobalTensor<MM1_OUT_T> mm1ResGm);
|
||||
__aicore__ inline void InitVecInputTensor(GlobalTensor<half> weightsGm, GlobalTensor<half> qScaleGm,
|
||||
GlobalTensor<half> kScaleGm, GlobalTensor<int32_t> indiceOutGm,
|
||||
GlobalTensor<int32_t> blockTableGm);
|
||||
__aicore__ inline void CleanInvalidOutput(int64_t invalidS1offset);
|
||||
__aicore__ inline int32_t AlignS2(int32_t cuS2Len);
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
|
||||
protected:
|
||||
GlobalTensor<MM1_OUT_T> mm1ResGm;
|
||||
GlobalTensor<half> weightsGm;
|
||||
GlobalTensor<half> qScaleGm;
|
||||
GlobalTensor<half> kScaleGm;
|
||||
GlobalTensor<half> vec0OutGm;
|
||||
GlobalTensor<int32_t> indiceOutGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
// =================================常量区=================================
|
||||
|
||||
private:
|
||||
__aicore__ inline void GetKeyScale(const QLICommon::RunInfo &runInfo, const LocalTensor<half> &resUb,
|
||||
int64_t batchId, int64_t startS2, int64_t getLen);
|
||||
// ================================Local Buffer区====================================
|
||||
// queue
|
||||
TQue<QuePosition::VECIN, 1> inQueue_;
|
||||
TQue<QuePosition::VECOUT, 1> outQueue_;
|
||||
|
||||
// tmp buff for vector
|
||||
TBuf<TPosition::VECCALC> sortOutBuf_;
|
||||
TBuf<TPosition::VECCALC> indexBuf_;
|
||||
TBuf<TPosition::VECCALC> tmpBuf_;
|
||||
|
||||
LocalTensor<int32_t> globalTopkIndice_;
|
||||
LocalTensor<float> globalTopkUb_;
|
||||
|
||||
int32_t blockId_ = -1;
|
||||
// para for vector
|
||||
int32_t groupInner_ = 0;
|
||||
int32_t globalTopkNum_ = 0;
|
||||
int64_t blockS2StartIdx_ = 0;
|
||||
int32_t gSize_ = 0;
|
||||
int32_t kSeqSize_ = 0;
|
||||
int32_t kHeadNum_ = 0;
|
||||
int32_t qHeadNum_ = 0;
|
||||
int32_t s1BaseSize_ = 0;
|
||||
int32_t s2BaseSize_ = 0;
|
||||
int32_t kCacheBlockSize_ = 0;
|
||||
int32_t maxBlockNumPerBatch_ = 0;
|
||||
|
||||
struct QLICommon::ConstInfo constInfo_;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::GetKeyScale(const QLICommon::RunInfo &runInfo, const LocalTensor<half> &resUb,
|
||||
int64_t batchId, int64_t startS2, int64_t getLen)
|
||||
{
|
||||
// startS2一定能整除kCacheBlockSize_
|
||||
AscendC::DataCopyPadExtParams<half> padParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams copyInParams;
|
||||
if constexpr (PAGE_ATTENTION) {
|
||||
int32_t startBlockTableIdx = startS2 / kCacheBlockSize_;
|
||||
int32_t startBlockTableOffset = startS2 % kCacheBlockSize_;
|
||||
int32_t blockTableBatchOffset = batchId * maxBlockNumPerBatch_;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
int32_t resUbBaseOffset = 0;
|
||||
if (startBlockTableOffset > 0) {
|
||||
int32_t firstPartLen =
|
||||
kCacheBlockSize_ - startBlockTableOffset > getLen ? getLen : kCacheBlockSize_ - startBlockTableOffset;
|
||||
copyInParams.blockLen = firstPartLen * sizeof(half);
|
||||
int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx);
|
||||
SetWaitFlag<HardEvent::S_MTE2>(HardEvent::S_MTE2);
|
||||
AscendC::DataCopyPad(resUb, kScaleGm[blockId * constInfo_.scaleStride + startBlockTableOffset],
|
||||
copyInParams, padParams);
|
||||
startBlockTableIdx++;
|
||||
getLen = getLen - firstPartLen;
|
||||
resUbBaseOffset = firstPartLen;
|
||||
}
|
||||
int32_t getLoopNum = CeilDiv(getLen, kCacheBlockSize_);
|
||||
copyInParams.blockLen = kCacheBlockSize_ * sizeof(half);
|
||||
for (int32_t i = 0; i < getLoopNum; i++) {
|
||||
if (i == getLoopNum - 1) {
|
||||
copyInParams.blockLen = (getLen - i * kCacheBlockSize_) * sizeof(half);
|
||||
}
|
||||
int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx + i);
|
||||
SetWaitFlag<HardEvent::S_MTE2>(HardEvent::S_MTE2);
|
||||
AscendC::DataCopyPad(resUb[resUbBaseOffset + i * kCacheBlockSize_], kScaleGm[blockId * constInfo_.scaleStride],
|
||||
copyInParams, padParams);
|
||||
}
|
||||
} else {
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.blockLen = getLen * sizeof(half);
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
AscendC::DataCopyPad(resUb, kScaleGm[runInfo.tensorKeyScaleOffset], copyInParams, padParams);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(inQueue_, 2, s2BaseSize_ * sizeof(float) * 2); // 32KB
|
||||
pipe->InitBuffer(outQueue_, 1, BASE_TOPK * sizeof(float)); // 8 KB
|
||||
pipe->InitBuffer(indexBuf_, s2BaseSize_ * sizeof(int32_t)); // 8 KB
|
||||
pipe->InitBuffer(tmpBuf_, 64 * 1024); // 64KB
|
||||
pipe->InitBuffer(sortOutBuf_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE * sizeof(float)); // 32KB
|
||||
|
||||
globalTopkIndice_ = indexBuf_.Get<int32_t>();
|
||||
globalTopkUb_ = sortOutBuf_.Get<float>();
|
||||
globalTopkNum_ = 0;
|
||||
|
||||
// 基本块执行前初始化UB和GM
|
||||
// step1. 初始化一个有序索引 0 - s2BaseSize_
|
||||
ArithProgression<int32_t>(globalTopkIndice_, 0, 1, s2BaseSize_);
|
||||
// step2. globalTopkUb_ [CeilDiv(s1BaseSize_, 2), BASE_TOPK, 2] -inf,-1
|
||||
InitSortOutBuf(globalTopkUb_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitParams(const struct QLICommon::ConstInfo &constInfo,
|
||||
const QLITilingData *__restrict tilingData)
|
||||
{
|
||||
this->constInfo_ = constInfo;
|
||||
blockS2StartIdx_ = 0;
|
||||
gSize_ = constInfo.gSize;
|
||||
kSeqSize_ = constInfo.kSeqSize;
|
||||
// define N2 para
|
||||
kHeadNum_ = constInfo.kHeadNum;
|
||||
qHeadNum_ = constInfo.qHeadNum;
|
||||
// define MMBase para
|
||||
s1BaseSize_ = constInfo.s1BaseSize; // 4
|
||||
s2BaseSize_ = constInfo.s2BaseSize; // 2048
|
||||
kCacheBlockSize_ = constInfo.kCacheBlockSize;
|
||||
maxBlockNumPerBatch_ = constInfo.maxBlockNumPerBatch;
|
||||
blockId_ = GetBlockIdx();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitVecInputTensor(GlobalTensor<half> weightsGm, GlobalTensor<half> qScaleGm,
|
||||
GlobalTensor<half> kScaleGm,
|
||||
GlobalTensor<int32_t> indiceOutGm,
|
||||
GlobalTensor<int32_t> blockTableGm)
|
||||
{
|
||||
this->weightsGm = weightsGm;
|
||||
this->qScaleGm = qScaleGm;
|
||||
this->kScaleGm = kScaleGm;
|
||||
this->indiceOutGm = indiceOutGm;
|
||||
this->blockTableGm = blockTableGm;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitVecWorkspaceTensor(GlobalTensor<half> vec0OutGm,
|
||||
GlobalTensor<MM1_OUT_T> mm1ResGm)
|
||||
{
|
||||
this->mm1ResGm = mm1ResGm;
|
||||
this->vec0OutGm = vec0OutGm;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::AllocEventID()
|
||||
{
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::FreeEventID()
|
||||
{
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::CleanInvalidOutput(int64_t invalidS1offset)
|
||||
{
|
||||
// init -1 and copy to output
|
||||
LocalTensor<float> valueULocal = outQueue_.AllocTensor<float>();
|
||||
LocalTensor<int32_t> idxULocal1 = valueULocal.template ReinterpretCast<int32_t>();
|
||||
Duplicate(idxULocal1, constInfo_.INVALID_IDX, constInfo_.sparseCount);
|
||||
outQueue_.EnQue<float>(valueULocal);
|
||||
valueULocal = outQueue_.DeQue<float>();
|
||||
QLIServiceVec::CopyOut(indiceOutGm[invalidS1offset], idxULocal1, constInfo_.sparseCount);
|
||||
outQueue_.FreeTensor(valueULocal);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::ProcessVec0(const QLICommon::RunInfo &info)
|
||||
{
|
||||
// 只需要一个v核做
|
||||
if (blockId_ % 2 != 0) {
|
||||
return;
|
||||
}
|
||||
int32_t cuBaseS1Idx = info.gS1Idx * s1BaseSize_;
|
||||
// 计算输出w基地址偏移 偶数循环 -> 0 + aic_offset 奇数循环 -> 4*64 + aic_offset
|
||||
int64_t vec0OutGmOffset = (info.loop % 2) * ((s1BaseSize_ * gSize_ * BLOCK_CUBE));
|
||||
// 计算输入weight的地址偏移,qScale的地址偏移与weight相同
|
||||
int64_t weightGmOffset = info.tensorWeightsOffset + cuBaseS1Idx * qHeadNum_;
|
||||
// 当前需要计算的S1行数,处理尾块场景
|
||||
int32_t cuS1ProcNum = cuBaseS1Idx + s1BaseSize_ > info.actS1Size ? info.actS1Size % s1BaseSize_ : s1BaseSize_;
|
||||
int32_t cuProcEleNum = cuS1ProcNum * gSize_;
|
||||
|
||||
LocalTensor<half> inWeightsUb = inQueue_.AllocTensor<half>();
|
||||
LocalTensor<half> inQScaleUb = inWeightsUb[cuProcEleNum];
|
||||
AscendC::DataCopyPadExtParams<half> padParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams copyInParams;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.blockLen = cuProcEleNum * sizeof(half);
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
AscendC::DataCopyPad(inWeightsUb, weightsGm[weightGmOffset], copyInParams, padParams);
|
||||
AscendC::DataCopyPad(inQScaleUb, qScaleGm[weightGmOffset], copyInParams, padParams);
|
||||
|
||||
inQueue_.EnQue<half>(inWeightsUb);
|
||||
inWeightsUb = inQueue_.DeQue<half>();
|
||||
AscendC::Mul(inWeightsUb, inWeightsUb, inQScaleUb, cuProcEleNum);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<half> resUb = outQueue_.AllocTensor<half>();
|
||||
AscendC::Brcb(resUb, inWeightsUb, static_cast<uint8_t>(cuProcEleNum / 8), {1, 8});
|
||||
inQueue_.FreeTensor(inWeightsUb);
|
||||
|
||||
outQueue_.EnQue<half>(resUb);
|
||||
resUb = outQueue_.DeQue<half>();
|
||||
AscendC::DataCopyParams copyOutParams;
|
||||
copyOutParams.blockCount = 1;
|
||||
copyOutParams.blockLen = cuProcEleNum * BLOCK_CUBE * sizeof(half);
|
||||
copyOutParams.srcStride = 0;
|
||||
copyOutParams.dstStride = 0;
|
||||
AscendC::DataCopyPad(vec0OutGm[vec0OutGmOffset], resUb, copyOutParams);
|
||||
outQueue_.FreeTensor(resUb);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline int32_t QLIVector<QLIT>::AlignS2(int32_t cuS2Len)
|
||||
{
|
||||
// 限制:当前cuS2Len最大为2048,暂不考虑更长
|
||||
// 该函数目的是将cuS2Len对齐到形如 32*(4^n)*m 的形式 (m ∈ [1, 3]),方便后续sort/merge
|
||||
if (cuS2Len <= ELE_NUM_128) {
|
||||
return Align(cuS2Len, ELE_NUM_32);
|
||||
} else if (cuS2Len <= ELE_NUM_512) {
|
||||
return Align(cuS2Len, ELE_NUM_128);
|
||||
} else {
|
||||
return Align(cuS2Len, ELE_NUM_512);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::ProcessVec1(const QLICommon::RunInfo &info)
|
||||
{
|
||||
int32_t cuBaseS1Idx = info.gS1Idx * s1BaseSize_;
|
||||
int32_t cuBaseS2Idx = info.s2Idx * s2BaseSize_;
|
||||
|
||||
// 计算基本块基地址偏移 偶数循环 -> 0 + aic_offset 奇数循环 -> 4*2048 + aic_offset
|
||||
int64_t mmGmOffset = (info.loop % 2) * (s1BaseSize_ * s2BaseSize_);
|
||||
|
||||
// cuS1BeginIdxPerAiv: 每个AIV的S1起始偏移
|
||||
int32_t cuS1BeginIdxPerAiv = cuBaseS1Idx;
|
||||
int32_t cuS1ProcNum =
|
||||
cuS1BeginIdxPerAiv + s1BaseSize_ > info.actS1Size ? info.actS1Size % s1BaseSize_ : s1BaseSize_;
|
||||
// cuS1ProcNumPerAiv: 每个AIv的S1计算量
|
||||
int32_t cuS1ProcNumPerAiv = blockId_ % 2 == 0 ? CeilDiv(cuS1ProcNum, 2) : (cuS1ProcNum / 2);
|
||||
cuS1BeginIdxPerAiv += (blockId_ % 2) * CeilDiv(cuS1ProcNum, 2);
|
||||
// 基本块基地址偏移奇数核加一个S1地址偏移
|
||||
mmGmOffset += (blockId_ % 2) * CeilDiv(cuS1ProcNum, 2) * s2BaseSize_;
|
||||
// 非首个基本块, M(S1)轴发生切换需要初始化
|
||||
if (info.loop != 0 && info.s2Idx == 0) {
|
||||
// globalTopkUb_ value,index=-inf,-1
|
||||
InitSortOutBuf(globalTopkUb_, CeilDiv(s1BaseSize_, 2) * BASE_TOPK_VALUE_IDX_SIZE);
|
||||
blockS2StartIdx_ = 0;
|
||||
} else if (info.loop == 0) {
|
||||
blockS2StartIdx_ = info.s2Idx;
|
||||
}
|
||||
// cuRealAcSeq: 当前基本块S1对应的AcSeq
|
||||
int32_t cuRealAcSeq = info.actS2Size;
|
||||
int32_t cuRealAcSeqCount = 0;
|
||||
if (constInfo_.attenMaskFlag) {
|
||||
// attenMask true场景
|
||||
cuRealAcSeq = info.actS2SizeOrig - info.actS1Size + cuS1BeginIdxPerAiv;
|
||||
}
|
||||
int32_t cuRealAcSeqIni = cuRealAcSeq;
|
||||
|
||||
|
||||
// LD输出S1方向偏移,保证2个Vector输出的内容连续
|
||||
uint32_t ldS1Offset = (blockId_ % 2 == 0) ? s1BaseSize_ / 2 - cuS1ProcNumPerAiv : 0;
|
||||
for (int innerS1Idx = 0; innerS1Idx < cuS1ProcNumPerAiv; innerS1Idx++) {
|
||||
if (constInfo_.attenMaskFlag) {
|
||||
cuRealAcSeqCount += 1;
|
||||
cuRealAcSeq = (cuRealAcSeqCount + cuRealAcSeqIni) / static_cast<int32_t>(constInfo_.cmpRatio);
|
||||
}
|
||||
int32_t cuS2Len = cuBaseS2Idx + s2BaseSize_ >= cuRealAcSeq ? cuRealAcSeq - cuBaseS2Idx : s2BaseSize_;
|
||||
int32_t cuS1Idx = cuS1BeginIdxPerAiv + innerS1Idx;
|
||||
if (cuRealAcSeq > 0 && cuS2Len > 0) {
|
||||
int32_t cuS2LenVecAlign = AlignS2(cuS2Len);
|
||||
LocalTensor<float> mmInUb = inQueue_.AllocTensor<float>();
|
||||
LocalTensor<float> kScaleUb = mmInUb[cuS2LenVecAlign];
|
||||
LocalTensor<half> kScaleTUb = kScaleUb.template ReinterpretCast<half>()[cuS2LenVecAlign];
|
||||
AscendC::DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyPadExtParams<half> padTParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams copyInParams;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.blockLen = cuS2Len * sizeof(float);
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
AscendC::DataCopyPad(mmInUb, mm1ResGm[mmGmOffset + innerS1Idx * s2BaseSize_], copyInParams, padParams);
|
||||
GetKeyScale(info, kScaleTUb, info.bIdx, cuBaseS2Idx, cuS2Len);
|
||||
inQueue_.EnQue<float>(mmInUb);
|
||||
mmInUb = inQueue_.DeQue<float>();
|
||||
AscendC::Cast(kScaleUb, kScaleTUb, RoundMode::CAST_NONE, cuS2Len);
|
||||
PipeBarrier<PIPE_V>();
|
||||
AscendC::Mul(mmInUb, mmInUb, kScaleUb, cuS2Len);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<float> sortBuff = tmpBuf_.Get<float>();
|
||||
LocalTensor<float> sortScoreUb = sortBuff;
|
||||
LocalTensor<float> sortIndiceUb = sortBuff[cuS2LenVecAlign];
|
||||
PipeBarrier<PIPE_V>();
|
||||
Duplicate(sortScoreUb.template ReinterpretCast<int32_t>(), QLIServiceVec::NEG_INF, cuS2LenVecAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Adds(sortScoreUb, mmInUb, 0.0f, cuS2Len);
|
||||
PipeBarrier<PIPE_V>();
|
||||
inQueue_.FreeTensor(mmInUb);
|
||||
LocalTensor<int32_t> sortIndiceUbInt = sortIndiceUb.template ReinterpretCast<int32_t>();
|
||||
// 无效数据索引填充为-1
|
||||
if (cuS2LenVecAlign != cuS2Len) {
|
||||
Duplicate(sortIndiceUbInt, -1, cuS2LenVecAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
Adds(sortIndiceUbInt, globalTopkIndice_, static_cast<int32_t>(cuBaseS2Idx), cuS2Len);
|
||||
PipeBarrier<PIPE_V>();
|
||||
LocalTensor<float> tmpSortBuf = sortBuff[2 * cuS2LenVecAlign];
|
||||
QLIServiceVec::SortAll(sortBuff, tmpSortBuf, cuS2LenVecAlign);
|
||||
PipeBarrier<PIPE_V>();
|
||||
QLIServiceVec::MergeSort(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK, sortBuff,
|
||||
cuS2LenVecAlign, tmpSortBuf);
|
||||
PipeBarrier<PIPE_V>();
|
||||
bool isS2End = cuBaseS2Idx + s2BaseSize_ >= cuRealAcSeq;
|
||||
bool needCopyOutGm = blockS2StartIdx_ == 0 && isS2End;
|
||||
if (needCopyOutGm) {
|
||||
LocalTensor<uint32_t> idxULocal = outQueue_.AllocTensor<uint32_t>();
|
||||
ExtractIndex(idxULocal,
|
||||
globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE].template ReinterpretCast<uint32_t>(),
|
||||
BASE_TOPK);
|
||||
PipeBarrier<PIPE_V>();
|
||||
InitSortOutBuf(globalTopkUb_[innerS1Idx * BASE_TOPK_VALUE_IDX_SIZE], BASE_TOPK_VALUE_IDX_SIZE);
|
||||
outQueue_.EnQue<uint32_t>(idxULocal);
|
||||
idxULocal = outQueue_.DeQue<uint32_t>();
|
||||
QLIServiceVec::CopyOut(indiceOutGm[info.indiceOutOffset + cuS1Idx * constInfo_.sparseCount],
|
||||
idxULocal.template ReinterpretCast<int32_t>(), constInfo_.sparseCount);
|
||||
outQueue_.FreeTensor(idxULocal);
|
||||
}
|
||||
} else if (cuRealAcSeq <= 0) {
|
||||
CleanInvalidOutput(info.indiceOutOffset + cuS1Idx * constInfo_.sparseCount);
|
||||
}
|
||||
}
|
||||
|
||||
// BNSD场景无效S1 输出-1
|
||||
if (Q_LAYOUT_T == LI_LAYOUT::BSND) {
|
||||
// 最后一个S1的基本块, 需要 >= info.actS1Size
|
||||
bool isS1LoopEnd = (cuBaseS1Idx + s1BaseSize_) >= info.actS1Size;
|
||||
int32_t invalidS1Num = constInfo_.qSeqSize - info.actS1Size;
|
||||
// blockS2StartIdx_ == 0 控制S2从开始的核去做冗余清理
|
||||
if (invalidS1Num > 0 && isS1LoopEnd && blockS2StartIdx_ == 0) {
|
||||
int32_t s1NumPerAiv = blockId_ % 2 == 0 ? CeilDiv(invalidS1Num, 2) : (invalidS1Num / 2);
|
||||
int32_t s1OffsetPerAiv = info.actS1Size + (blockId_ % 2) * CeilDiv(invalidS1Num, 2);
|
||||
for (int innerS1Idx = 0; innerS1Idx < s1NumPerAiv; innerS1Idx++) {
|
||||
CleanInvalidOutput(info.indiceOutOffset + (s1OffsetPerAiv + innerS1Idx) * constInfo_.sparseCount);
|
||||
}
|
||||
}
|
||||
|
||||
int32_t invalidS1Num2 = info.actS1Size - info.actS2SizeOrig;
|
||||
if (invalidS1Num2 > 0 && isS1LoopEnd && blockS2StartIdx_ == 0 && constInfo_.attenMaskFlag) {
|
||||
int32_t s1NumPerAiv = blockId_ % 2 == 0 ? CeilDiv(invalidS1Num2, 2) : (invalidS1Num2 / 2);
|
||||
int32_t s1OffsetPerAiv = (blockId_ % 2) * CeilDiv(invalidS1Num2, 2);
|
||||
for (int innerS1Idx = 0; innerS1Idx < s1NumPerAiv; innerS1Idx++) {
|
||||
CleanInvalidOutput((info.bN2Idx * constInfo_.qSeqSize + s1OffsetPerAiv + innerS1Idx) *
|
||||
constInfo_.sparseCount);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (info.isLastS2InnerLoop) {
|
||||
// S2最后一个Loop后, 下一个基本块初始从0开始
|
||||
blockS2StartIdx_ = 0;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace QLIKernel
|
||||
#endif // QUANT_LIGHTNING_INDEXER_SERVICE_VECTOR_H
|
||||
@@ -0,0 +1,193 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_vector.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_VECTOR_H
|
||||
#define QUANT_LIGHTNING_INDEXER_VECTOR_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "quant_lightning_indexer_vector.h"
|
||||
|
||||
namespace QLIServiceVec {
|
||||
using namespace AscendC;
|
||||
|
||||
constexpr int32_t NEG_INF = 0xFF800000;
|
||||
constexpr int32_t INVALID_INDEX = -1;
|
||||
constexpr uint8_t VEC_REPEAT_MAX = 255;
|
||||
constexpr uint8_t B32_VEC_ELM_NUM = 64;
|
||||
constexpr uint8_t B32_BLOCK_ALIGN_NUM = 8;
|
||||
constexpr uint8_t B32_VEC_REPEAT_STRIDE = 8;
|
||||
constexpr uint64_t VEC_REPEAT_BYTES = 256;
|
||||
constexpr int32_t CONST_TWO = 2;
|
||||
constexpr int64_t VALUE_AND_INDEX_NUM = 2;
|
||||
constexpr int64_t BLOCK_BYTES = 32;
|
||||
constexpr int64_t MRG_QUE_0 = 0;
|
||||
constexpr int64_t MRG_QUE_1 = 1;
|
||||
constexpr int64_t MRG_QUE_2 = 2;
|
||||
constexpr int64_t MRG_QUE_3 = 3;
|
||||
constexpr int64_t MRG_BLOCK_2 = 2;
|
||||
constexpr int64_t MRG_BLOCK_3 = 3;
|
||||
constexpr int64_t MRG_BLOCK_4 = 4;
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyOut(const GlobalTensor<T> &dstGm, const LocalTensor<T> &srcUb, int64_t copyCount)
|
||||
{
|
||||
AscendC::DataCopyParams dataCopyOutyParams;
|
||||
dataCopyOutyParams.blockCount = 1;
|
||||
dataCopyOutyParams.blockLen = copyCount * sizeof(T);
|
||||
dataCopyOutyParams.srcStride = 0;
|
||||
dataCopyOutyParams.dstStride = 0;
|
||||
AscendC::DataCopyPad(dstGm, srcUb, dataCopyOutyParams);
|
||||
}
|
||||
|
||||
/**
|
||||
src: 传入的初始化空间
|
||||
eleNum: 需要初始化的元素个数需为64整数倍,元素将被初始化为交错排布的-inf,-1
|
||||
*/
|
||||
__aicore__ inline void InitSortOutBuf(const LocalTensor<float> &src, int64_t eleNum)
|
||||
{
|
||||
uint64_t mask1[2] = {0x5555555555555555, 0};
|
||||
uint64_t mask0[2] = {0xaaaaaaaaaaaaaaaa, 0};
|
||||
int64_t repeatNum = eleNum / B32_VEC_ELM_NUM;
|
||||
int64_t forLoop = repeatNum / VEC_REPEAT_MAX;
|
||||
int64_t forRemain = repeatNum % VEC_REPEAT_MAX;
|
||||
for (int i = 0; i < forLoop; i++) {
|
||||
AscendC::Duplicate(src.template ReinterpretCast<int32_t>(), NEG_INF, mask1, VEC_REPEAT_MAX, 1,
|
||||
B32_VEC_REPEAT_STRIDE);
|
||||
AscendC::Duplicate(src.template ReinterpretCast<int32_t>(), INVALID_INDEX, mask0, VEC_REPEAT_MAX, 1,
|
||||
B32_VEC_REPEAT_STRIDE);
|
||||
}
|
||||
if (forRemain > 0) {
|
||||
AscendC::Duplicate(src.template ReinterpretCast<int32_t>()[forLoop * VEC_REPEAT_MAX * B32_VEC_ELM_NUM], NEG_INF,
|
||||
mask1, forRemain, 1, B32_VEC_REPEAT_STRIDE);
|
||||
AscendC::Duplicate(src.template ReinterpretCast<int32_t>()[forLoop * VEC_REPEAT_MAX * B32_VEC_ELM_NUM],
|
||||
INVALID_INDEX, mask0, forRemain, 1, B32_VEC_REPEAT_STRIDE);
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
/**
|
||||
src: logits和索引,前logitsNum为logits,后logitsNum为索引
|
||||
tmp: 计算使用到的临时空间,大小与src一致
|
||||
logitsNum: 排序的元素个数, 暂只支持[128,256,384,512,1024,2048]
|
||||
*/
|
||||
__aicore__ inline void SortAll(LocalTensor<float> &src, LocalTensor<float> &tmp, int64_t logitsNum)
|
||||
{
|
||||
int64_t sort32Repeats = logitsNum / BLOCK_BYTES;
|
||||
AscendC::Sort32(tmp, src, src[logitsNum].ReinterpretCast<uint32_t>(), sort32Repeats);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
int64_t mrgGroups = sort32Repeats;
|
||||
int64_t mrgElements = BLOCK_BYTES;
|
||||
int64_t i = 0;
|
||||
AscendC::LocalTensor<float> srcTensor;
|
||||
AscendC::LocalTensor<float> dstTensor;
|
||||
while (true) {
|
||||
if (i % CONST_TWO == 0) {
|
||||
srcTensor = tmp;
|
||||
dstTensor = src;
|
||||
} else {
|
||||
srcTensor = src;
|
||||
dstTensor = tmp;
|
||||
}
|
||||
AscendC::MrgSort4Info params;
|
||||
params.elementLengths[0] = mrgElements;
|
||||
params.elementLengths[MRG_QUE_1] = mrgElements;
|
||||
params.elementLengths[MRG_QUE_2] = mrgElements;
|
||||
params.elementLengths[MRG_QUE_3] = mrgElements;
|
||||
params.ifExhaustedSuspension = false;
|
||||
params.validBit = 0b1111;
|
||||
|
||||
AscendC::MrgSortSrcList<float> srcList;
|
||||
srcList.src1 = srcTensor[0];
|
||||
srcList.src2 = srcTensor[MRG_QUE_1 * VALUE_AND_INDEX_NUM * mrgElements];
|
||||
srcList.src3 = srcTensor[MRG_QUE_2 * VALUE_AND_INDEX_NUM * mrgElements];
|
||||
srcList.src4 = srcTensor[MRG_QUE_3 * VALUE_AND_INDEX_NUM * mrgElements];
|
||||
if (mrgGroups <= MRG_BLOCK_4) {
|
||||
params.repeatTimes = 1;
|
||||
if (mrgGroups == 1) {
|
||||
break;
|
||||
} else if (mrgGroups == MRG_BLOCK_2) {
|
||||
params.validBit = 0b0011;
|
||||
} else if (mrgGroups == MRG_BLOCK_3) {
|
||||
params.validBit = 0b0111;
|
||||
} else if (mrgGroups == MRG_BLOCK_4) {
|
||||
params.validBit = 0b1111;
|
||||
}
|
||||
AscendC::MrgSort<float>(dstTensor, srcList, params);
|
||||
i += 1;
|
||||
break;
|
||||
} else {
|
||||
params.repeatTimes = mrgGroups / MRG_BLOCK_4;
|
||||
AscendC::MrgSort<float>(dstTensor, srcList, params);
|
||||
i += 1;
|
||||
mrgElements = mrgElements * MRG_BLOCK_4;
|
||||
mrgGroups = mrgGroups / MRG_BLOCK_4;
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
if (i % CONST_TWO == 0) {
|
||||
AscendC::DataCopy(src, tmp, logitsNum * VALUE_AND_INDEX_NUM);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
mrgDst: 合并进的Tensor
|
||||
mrgSrc: 待合并的Tensor
|
||||
tmpTensor:空间为mrgDst+mrgSrc
|
||||
*/
|
||||
__aicore__ inline void MergeSort(const LocalTensor<float> &mrgDst, int32_t mrgDstNum, LocalTensor<float> &mrgSrc,
|
||||
int32_t mrgSrcNum, LocalTensor<float> &tmpTensor)
|
||||
{
|
||||
AscendC::MrgSort4Info params;
|
||||
params.elementLengths[0] = mrgSrcNum;
|
||||
params.elementLengths[1] = mrgDstNum;
|
||||
params.ifExhaustedSuspension = false;
|
||||
params.validBit = 0b0011;
|
||||
params.repeatTimes = 1;
|
||||
|
||||
AscendC::MrgSortSrcList<float> srcList;
|
||||
srcList.src1 = mrgSrc;
|
||||
srcList.src2 = mrgDst;
|
||||
|
||||
AscendC::MrgSort<float>(tmpTensor, srcList, params);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::DataCopy(mrgDst, tmpTensor, mrgDstNum * VALUE_AND_INDEX_NUM);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
__aicore__ inline void ExtractIndex(const LocalTensor<uint32_t> &idxULocal, const LocalTensor<uint32_t> &sortLocal,
|
||||
int64_t extractNum)
|
||||
{
|
||||
AscendC::GatherMaskParams gatherMaskParams;
|
||||
gatherMaskParams.repeatTimes = Ceil(extractNum * sizeof(float) * VALUE_AND_INDEX_NUM, VEC_REPEAT_BYTES);
|
||||
gatherMaskParams.src0BlockStride = 1;
|
||||
gatherMaskParams.src0RepeatStride = B32_VEC_REPEAT_STRIDE;
|
||||
gatherMaskParams.src1RepeatStride = 0;
|
||||
uint64_t rsvdCnt = 0; // 用于保存筛选后保留下来的元素个数
|
||||
uint8_t src1Pattern = 2; // 固定模式2,表示筛选出奇数索引的数
|
||||
AscendC::GatherMask(idxULocal, sortLocal, src1Pattern, false, static_cast<uint32_t>(0), gatherMaskParams, rsvdCnt);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
template <HardEvent event>
|
||||
__aicore__ inline void SetWaitFlag(HardEvent evt)
|
||||
{
|
||||
event_t eventId = static_cast<event_t>(GetTPipePtr()->FetchEventID(evt));
|
||||
AscendC::SetFlag<event>(eventId);
|
||||
AscendC::WaitFlag<event>(eventId);
|
||||
}
|
||||
|
||||
} // namespace QLIServiceVec
|
||||
#endif // QUANT_LIGHTNING_INDEXER_VECTOR_H
|
||||
@@ -0,0 +1,179 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_common.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef quant_lightning_indexer_COMMON_H
|
||||
#define quant_lightning_indexer_COMMON_H
|
||||
using namespace AscendC;
|
||||
namespace QLICommon {
|
||||
|
||||
// 与tiling的layout保持一致
|
||||
enum class LI_LAYOUT : uint32_t {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_BSND = 2
|
||||
};
|
||||
|
||||
template <typename Q_T, typename K_T, typename QK_T, typename SCORE_T, typename OUT_T, const bool PAGE_ATTENTION = false,
|
||||
LI_LAYOUT Q_LAYOUT_T = LI_LAYOUT::BSND, LI_LAYOUT K_LAYOUT_T = LI_LAYOUT::PA_BSND, typename... Args>
|
||||
struct QLIType {
|
||||
static_assert(
|
||||
(std::is_same_v<QK_T, float> &&
|
||||
(std::is_same_v<SCORE_T, uint32_t> || std::is_same_v<SCORE_T, uint16_t>)) ||
|
||||
(std::is_same_v<QK_T, bfloat16_t> &&
|
||||
std::is_same_v<SCORE_T, uint16_t>),
|
||||
"Invalid combination of QK_T and SCORE_T"
|
||||
);
|
||||
using queryType = Q_T;
|
||||
using keyType = K_T;
|
||||
using queryKeyType = QK_T;
|
||||
using scoreType = SCORE_T;
|
||||
using outputType = OUT_T;
|
||||
|
||||
static constexpr bool pageAttention = PAGE_ATTENTION;
|
||||
static constexpr LI_LAYOUT layout = Q_LAYOUT_T;
|
||||
static constexpr LI_LAYOUT keyLayout = K_LAYOUT_T;
|
||||
};
|
||||
|
||||
struct RunInfo {
|
||||
uint32_t loop;
|
||||
uint32_t bN2Idx;
|
||||
uint32_t bIdx;
|
||||
uint32_t n2Idx = 0;
|
||||
uint32_t gS1Idx;
|
||||
uint32_t s2Idx;
|
||||
|
||||
uint32_t actS1Size = 1;
|
||||
uint32_t actS2Size = 1;
|
||||
uint32_t actS2SizeOrig = 1;
|
||||
uint32_t actMBaseSize;
|
||||
uint32_t actualSingleProcessSInnerSize;
|
||||
uint32_t actualSingleProcessSInnerSizeAlign;
|
||||
|
||||
uint64_t tensorQueryOffset;
|
||||
uint64_t tensorKeyOffset;
|
||||
uint64_t tensorKeyScaleOffset;
|
||||
uint64_t tensorWeightsOffset;
|
||||
uint64_t indiceOutOffset;
|
||||
|
||||
bool isFirstS2InnerLoop;
|
||||
bool isLastS2InnerLoop;
|
||||
bool isAllLoopEnd = false;
|
||||
bool isValid = false;
|
||||
};
|
||||
|
||||
struct ConstInfo {
|
||||
// CUBE与VEC核间同步的模式
|
||||
static constexpr uint32_t QLI_SYNC_MODE4 = 4;
|
||||
static constexpr uint32_t AIV0_AIV1_OFFSET = 16;
|
||||
static constexpr uint32_t CROSS_VC_EVENT = 0;
|
||||
static constexpr uint32_t CROSS_CV_EVENT = 2;
|
||||
// BUFFER的字节数
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32B = 32;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_64B = 64;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_256B = 256;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_512B = 512;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_1K = 1024;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_2K = 2048;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_4K = 4096;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_8K = 8192;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_16K = 16384;
|
||||
static constexpr uint32_t BUFFER_SIZE_BYTE_32K = 32768;
|
||||
// 无效索引
|
||||
static constexpr int INVALID_IDX = -1;
|
||||
|
||||
// CUBE和VEC的核间同步EventID
|
||||
uint32_t syncC1V1 = 0U;
|
||||
uint32_t syncC1V0 = 2U;
|
||||
uint32_t syncV1C1 = 0U;
|
||||
uint32_t syncV0C1 = 1U;
|
||||
|
||||
// 基本块大小
|
||||
uint32_t mBaseSize = 1ULL;
|
||||
uint32_t s1BaseSize = 1ULL;
|
||||
uint32_t s2BaseSize = 1ULL;
|
||||
|
||||
uint64_t batchSize = 0ULL;
|
||||
uint64_t gSize = 0ULL;
|
||||
uint64_t qHeadNum = 0ULL;
|
||||
uint64_t kHeadNum;
|
||||
uint64_t headDim;
|
||||
uint64_t sparseCount; // topK选取大小
|
||||
uint64_t kSeqSize = 0ULL; // kv最大S长度
|
||||
uint64_t qSeqSize = 1ULL; // q最大S长度
|
||||
uint32_t kCacheBlockSize = 0; // PA场景的block size
|
||||
uint32_t maxBlockNumPerBatch = 0; // PA场景的最大单batch block number
|
||||
LI_LAYOUT outputLayout; // 输出的格式
|
||||
bool attenMaskFlag = false;
|
||||
uint32_t cmpRatio = 1;
|
||||
bool batchSupperFlag = false; // Qactual_se长度是否为B+1
|
||||
int64_t stride = 1;
|
||||
int64_t scaleStride = 1;
|
||||
|
||||
uint32_t actualLenQDims = 0U; // query的actualSeqLength 的维度
|
||||
uint32_t actualLenDims = 0U; // KV 的actualSeqLength 的维度
|
||||
bool isAccumSeqS1 = false; // 是否累加模式
|
||||
bool isAccumSeqS2 = false; // 是否累加模式
|
||||
bool isLDOpen = false;
|
||||
};
|
||||
|
||||
struct SplitCoreInfo {
|
||||
uint32_t s2Start = 0U; // S2的起始位置
|
||||
uint32_t s2End = 0U; // S2循环index上限
|
||||
uint32_t bN2Start = 0U;
|
||||
uint32_t bN2End = 0U;
|
||||
uint32_t gS1Start = 0U;
|
||||
uint32_t gS1End = 0U;
|
||||
bool isLD = false; // 当前核是否需要进行Decode归约任务
|
||||
bool isCoreEnable = false;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T Align(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd) * (rnd)));
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Min(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (b) : (a);
|
||||
}
|
||||
|
||||
template <typename T1, typename T2>
|
||||
__aicore__ inline T1 Max(T1 a, T2 b)
|
||||
{
|
||||
return (a > b) ? (a) : (b);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline T CeilDiv(T num, T rnd)
|
||||
{
|
||||
return (((rnd) == 0) ? 0 : (((num) + (rnd)-1) / (rnd)));
|
||||
}
|
||||
} // namespace QLICommon
|
||||
|
||||
// bank冲突优化
|
||||
// david 256KB bank layout
|
||||
// shape ( bank_depth ( banks bank_groups block)) (512 ( 2 8 32))
|
||||
// stride (banks*bank_groups*block (bank_groups*block block 1)) (512 (256 32 1))
|
||||
#define UB_BLOCK 32 // 32B
|
||||
#define UB_BANK_GROUPS 8
|
||||
#define UB_BANKS 2
|
||||
#define UB_BANK_DEPTH 512
|
||||
|
||||
#define UB_BANK_GROUP_STRIDE UB_BLOCK // 32B
|
||||
#define UB_BANK_STRIDE (UB_BANK_GROUPS * UB_BLOCK) // 256B
|
||||
#define UB_BANK_DEPTH_STRIDE (UB_BANKS * UB_BANK_GROUPS * UB_BLOCK) // 512B
|
||||
|
||||
#endif // quant_lightning_indexer_COMMON_H
|
||||
@@ -0,0 +1,640 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_kernel.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef quant_lightning_indexer_KERNEL_H
|
||||
#define quant_lightning_indexer_KERNEL_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
#include "quant_lightning_indexer_service_vector.h"
|
||||
#include "quant_lightning_indexer_service_cube.h"
|
||||
#include "../vllm_quant_lightning_indexer_metadata.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
using namespace matmul;
|
||||
using namespace optiling;
|
||||
using namespace optiling::detail;
|
||||
using AscendC::CacheMode;
|
||||
using AscendC::CrossCoreSetFlag;
|
||||
using AscendC::CrossCoreWaitFlag;
|
||||
|
||||
// 由于S2循环前,RunInfo还没有赋值,使用TempLoopInfo临时存放B、N、S1轴相关的信息;同时减少重复计算
|
||||
struct TempLoopInfo {
|
||||
uint32_t bN2Idx = 0;
|
||||
uint32_t bIdx = 0U;
|
||||
uint32_t n2Idx = 0U;
|
||||
uint32_t gS1Idx = 0U;
|
||||
uint32_t gS1LoopEnd = 0U; // gS1方向循环的结束Idx
|
||||
uint32_t s2LoopEnd = 0U; // S2方向循环的结束Idx
|
||||
uint32_t actS1Size = 1ULL; // 当前Batch循环处理的S1轴的实际大小
|
||||
uint32_t actS2Size = 0ULL;
|
||||
uint32_t actS2SizeOrig = 0ULL;//压缩前s2
|
||||
bool curActSeqLenIsZero = false;
|
||||
bool needDealActS1LessThanS1 = false; // S1的实际长度小于shape的S1长度时,是否需要清理输出
|
||||
uint32_t actMBaseSize = 0U; // m轴(gS1)方向实际大小
|
||||
uint32_t mBasicSizeTail = 0U; // gS1方向循环的尾基本块大小
|
||||
uint32_t s2BasicSizeTail = 0U; // S2方向循环的尾基本块大小
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
class QLIPreload {
|
||||
public:
|
||||
__aicore__ inline QLIPreload(){};
|
||||
__aicore__ inline void Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights,
|
||||
__gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale, __gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengthsK, __gm__ uint8_t *blockTable,
|
||||
__gm__ uint8_t *metadata, __gm__ uint8_t *sparseIndices,
|
||||
__gm__ uint8_t *workspace, const QLITilingData *__restrict tiling, TPipe *tPipe);
|
||||
__aicore__ inline void Process();
|
||||
|
||||
// =================================类型定义区=================================
|
||||
using Q_T = typename QLIT::queryType;
|
||||
using K_T = typename QLIT::keyType;
|
||||
using OUT_T = typename QLIT::outputType;
|
||||
static constexpr bool PAGE_ATTENTION = QLIT::pageAttention;
|
||||
static constexpr LI_LAYOUT Q_LAYOUT_T = QLIT::layout;
|
||||
static constexpr LI_LAYOUT K_LAYOUT_T = QLIT::keyLayout;
|
||||
|
||||
using SCORE_T = typename QLIT::scoreType;
|
||||
|
||||
QLIMatmul<QLIT> matmulService;
|
||||
QLIVector<QLIT> vectorService;
|
||||
|
||||
// =================================常量区=================================
|
||||
static constexpr uint32_t SYNC_C1_V1_FLAG = 4;
|
||||
static constexpr uint32_t SYNC_V1_C1_FLAG = 5;
|
||||
|
||||
static constexpr uint32_t M_BASE_SIZE = 256;
|
||||
static constexpr uint32_t S2_BASE_SIZE = 128;
|
||||
static constexpr uint32_t HEAD_DIM = 128;
|
||||
static constexpr uint32_t K_HEAD_NUM = 1;
|
||||
static constexpr uint32_t GM_ALIGN_BYTES = 512;
|
||||
|
||||
static constexpr int64_t LD_PREFETCH_LEN = 2;
|
||||
// for workspace double
|
||||
static constexpr uint32_t WS_DOBULE = 2;
|
||||
|
||||
protected:
|
||||
TPipe *pipe = nullptr;
|
||||
|
||||
// offset
|
||||
uint64_t queryCoreOffset = 0ULL;
|
||||
uint64_t keyCoreOffset = 0ULL;
|
||||
uint64_t keyScaleCoreOffset = 0ULL;
|
||||
uint64_t weightsCoreOffset = 0ULL;
|
||||
uint64_t indiceOutCoreOffset = 0ULL;
|
||||
bool isUsedCoreEqZero = false;
|
||||
// ================================Global Buffer区=================================
|
||||
GlobalTensor<Q_T> queryGm;
|
||||
GlobalTensor<K_T> keyGm;
|
||||
GlobalTensor<float> weightsGm;
|
||||
GlobalTensor<float> qScaleGm;
|
||||
GlobalTensor<float> kScaleGm;
|
||||
GlobalTensor<uint32_t> metadataGm;
|
||||
|
||||
GlobalTensor<int32_t> indiceOutGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
|
||||
GlobalTensor<uint32_t> actualSeqLengthsGmQ;
|
||||
GlobalTensor<uint32_t> actualSeqLengthsGm;
|
||||
|
||||
// ================================类成员变量====================================
|
||||
// aic、aiv核信息
|
||||
uint32_t tmpBlockIdx = 0U;
|
||||
uint32_t aiCoreIdx = 0U;
|
||||
uint32_t usedCoreNum = 0U;
|
||||
|
||||
QLICommon::ConstInfo constInfo{};
|
||||
TempLoopInfo tempLoopInfo{};
|
||||
QLICommon::SplitCoreInfo splitCoreInfo{};
|
||||
|
||||
// ================================Init functions==================================
|
||||
__aicore__ inline void InitTilingData(const QLITilingData *__restrict tilingData);
|
||||
__aicore__ inline void InitBuffers();
|
||||
__aicore__ inline void InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsK);
|
||||
// ================================Split Core================================
|
||||
__aicore__ inline void SplitCoreByAICPU(uint32_t curCoreIdx, GlobalTensor<uint32_t> &metadataGm);
|
||||
__aicore__ inline uint32_t GetS2BaseBlockNumOnMask(uint32_t s1gIdx, uint32_t actS1Size, uint32_t actS2SizeOrig);
|
||||
// ================================Process functions================================
|
||||
__aicore__ inline void ProcessMain();
|
||||
__aicore__ inline void ProcessBaseBlock(uint32_t loop, uint64_t s2LoopIdx,
|
||||
QLICommon::RunInfo runInfo);
|
||||
__aicore__ inline void ProcessInvalid();
|
||||
// ================================Params Calc=====================================
|
||||
__aicore__ inline void CalcGS1LoopParams(uint32_t bN2Idx);
|
||||
__aicore__ inline void GetBN2Idx(uint32_t bN2Idx);
|
||||
__aicore__ inline uint32_t GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm, uint32_t defaultSeqLen);
|
||||
__aicore__ inline uint32_t GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm, uint32_t defaultSeqLen, uint32_t cmpRatio);
|
||||
__aicore__ inline void GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, uint32_t &actS2Size, uint32_t &actS2SizeOrig);
|
||||
__aicore__ inline void CalcS2LoopParams(uint32_t bN2LoopIdx, uint32_t gS1LoopIdx);
|
||||
__aicore__ inline void CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void DealActSeqLenIsZero(uint32_t bIdx, uint32_t n2Idx, uint32_t s1Start);
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitTilingData(const QLITilingData *__restrict tilingData)
|
||||
{
|
||||
usedCoreNum = tilingData->usedCoreNum;
|
||||
constInfo.batchSize = tilingData->bSize;
|
||||
constInfo.qHeadNum = constInfo.gSize = tilingData->gSize;
|
||||
constInfo.kSeqSize = tilingData->s2Size;
|
||||
constInfo.qSeqSize = tilingData->s1Size;
|
||||
constInfo.attenMaskFlag = (tilingData->sparseMode == 3);
|
||||
constInfo.kCacheBlockSize = tilingData->blockSize;
|
||||
constInfo.maxBlockNumPerBatch = tilingData->maxBlockNumPerBatch;
|
||||
constInfo.sparseCount = tilingData->sparseCount;
|
||||
constInfo.cmpRatio = tilingData->cmpRatio;
|
||||
constInfo.batchSupperFlag = tilingData->batchSupperFlag;
|
||||
constInfo.stride = tilingData->stride;
|
||||
constInfo.scaleStride = tilingData->scaleStride;
|
||||
constInfo.outputLayout = Q_LAYOUT_T; // 输出和输入形状一致
|
||||
if (Q_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
constInfo.isAccumSeqS1 = true;
|
||||
}
|
||||
if (K_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
constInfo.isAccumSeqS2 = true;
|
||||
}
|
||||
|
||||
constInfo.kHeadNum = K_HEAD_NUM;
|
||||
constInfo.headDim = HEAD_DIM;
|
||||
|
||||
constInfo.mBaseSize = M_BASE_SIZE;
|
||||
constInfo.s2BaseSize = S2_BASE_SIZE;
|
||||
constInfo.s1BaseSize = (constInfo.mBaseSize + constInfo.gSize - 1) / constInfo.gSize;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitBuffers()
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.InitBuffers(pipe);
|
||||
} else {
|
||||
matmulService.InitBuffers(pipe);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::InitActualSeqLen(__gm__ uint8_t *actualSeqLengthsQ,
|
||||
__gm__ uint8_t *actualSeqLengthsK)
|
||||
{
|
||||
if (actualSeqLengthsQ == nullptr) {
|
||||
constInfo.actualLenQDims = 0;
|
||||
} else {
|
||||
constInfo.actualLenQDims = (constInfo.batchSupperFlag) ? constInfo.batchSize + 1 : constInfo.batchSize;
|
||||
actualSeqLengthsGmQ.SetGlobalBuffer((__gm__ uint32_t *)actualSeqLengthsQ, constInfo.actualLenQDims);
|
||||
}
|
||||
if (actualSeqLengthsK == nullptr) {
|
||||
constInfo.actualLenDims = 0;
|
||||
} else {
|
||||
constInfo.actualLenDims = constInfo.batchSize;
|
||||
actualSeqLengthsGm.SetGlobalBuffer((__gm__ uint32_t *)actualSeqLengthsK, constInfo.actualLenDims);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetActualSeqLen(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm,
|
||||
uint32_t defaultSeqLen)
|
||||
{
|
||||
bIdx = (constInfo.batchSupperFlag)? bIdx + 1 : bIdx; // 如果为B+1情况,则向后移动一位
|
||||
if (actualLenDims == 0) {
|
||||
return defaultSeqLen;
|
||||
} else if (isAccumSeq && bIdx > 0) {
|
||||
return actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1);
|
||||
} else {
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetActualSeqLenKey(uint32_t bIdx, uint32_t actualLenDims, bool isAccumSeq,
|
||||
GlobalTensor<uint32_t> &actualSeqLengthsGm,
|
||||
uint32_t defaultSeqLen, uint32_t cmpRatio)
|
||||
{
|
||||
if (actualLenDims == 0) {
|
||||
return defaultSeqLen * cmpRatio;
|
||||
} else if (isAccumSeq && bIdx > 0) {
|
||||
return actualSeqLengthsGm.GetValue(bIdx) - actualSeqLengthsGm.GetValue(bIdx - 1);
|
||||
} else {
|
||||
return actualSeqLengthsGm.GetValue(bIdx);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::GetS1S2ActualSeqLen(uint32_t bIdx, uint32_t &actS1Size, uint32_t &actS2Size, uint32_t &actS2SizeOrig)
|
||||
{
|
||||
actS1Size = GetActualSeqLen(bIdx, constInfo.actualLenQDims, constInfo.isAccumSeqS1, actualSeqLengthsGmQ,
|
||||
constInfo.qSeqSize);
|
||||
actS2SizeOrig =
|
||||
GetActualSeqLenKey(bIdx, constInfo.actualLenDims, constInfo.isAccumSeqS2, actualSeqLengthsGm, constInfo.kSeqSize, constInfo.cmpRatio); // 压缩前的actS2Size
|
||||
actS2Size = actS2SizeOrig / constInfo.cmpRatio; // 真实使用的压缩后S2长度
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline uint32_t QLIPreload<QLIT>::GetS2BaseBlockNumOnMask(uint32_t s1gIdx, uint32_t actS1Size,
|
||||
uint32_t actS2SizeOrig)
|
||||
{
|
||||
if (actS2SizeOrig / constInfo.cmpRatio == 0) {
|
||||
return 0;
|
||||
}
|
||||
uint32_t s1Offset = constInfo.s1BaseSize * s1gIdx;
|
||||
int32_t validS2LenBase = static_cast<int32_t>(actS2SizeOrig) - static_cast<int32_t>(actS1Size); // 压缩前的validS2LenBase
|
||||
int32_t validS2Len = (static_cast<int32_t>(s1Offset) + validS2LenBase + static_cast<int32_t>(constInfo.s1BaseSize)) / static_cast<int32_t>(constInfo.cmpRatio);
|
||||
validS2Len = Min(validS2Len, static_cast<int32_t>(actS2SizeOrig) / constInfo.cmpRatio);
|
||||
validS2Len = Max(validS2Len, 1);
|
||||
return (validS2Len + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::SplitCoreByAICPU(uint32_t curCoreIdx, GlobalTensor<uint32_t> &metadataGm)
|
||||
{
|
||||
uint32_t liCoreEnableIndex = GetAttrAbsIndex(curCoreIdx, LI_CORE_ENABLE_INDEX);
|
||||
uint32_t bN2StartIndex = GetAttrAbsIndex(curCoreIdx, LI_BN2_START_INDEX);
|
||||
uint32_t mStartIndex = GetAttrAbsIndex(curCoreIdx, LI_M_START_INDEX);
|
||||
uint32_t s2StartIndex = GetAttrAbsIndex(curCoreIdx, LI_S2_START_INDEX);
|
||||
uint32_t bN2EndIndex = GetAttrAbsIndex(curCoreIdx, LI_BN2_END_INDEX);
|
||||
uint32_t mEndIndex = GetAttrAbsIndex(curCoreIdx, LI_M_END_INDEX);
|
||||
uint32_t s2EndIndex = GetAttrAbsIndex(curCoreIdx, LI_S2_END_INDEX);
|
||||
|
||||
uint32_t liZeroCoreEnableIndex = GetAttrAbsIndex(0, LI_CORE_ENABLE_INDEX);
|
||||
if (metadataGm.GetValue(liZeroCoreEnableIndex) == 0) {
|
||||
isUsedCoreEqZero = true;
|
||||
}
|
||||
if (metadataGm.GetValue(liCoreEnableIndex) == 0) {
|
||||
splitCoreInfo.isCoreEnable = false;
|
||||
return;
|
||||
} else {
|
||||
splitCoreInfo.isCoreEnable = true;
|
||||
}
|
||||
|
||||
splitCoreInfo.bN2Start = metadataGm.GetValue(bN2StartIndex);
|
||||
splitCoreInfo.gS1Start = metadataGm.GetValue(mStartIndex);
|
||||
splitCoreInfo.s2Start = metadataGm.GetValue(s2StartIndex);
|
||||
splitCoreInfo.bN2End = metadataGm.GetValue(bN2EndIndex);
|
||||
splitCoreInfo.gS1End = metadataGm.GetValue(mEndIndex);
|
||||
splitCoreInfo.s2End = metadataGm.GetValue(s2EndIndex);
|
||||
|
||||
if (splitCoreInfo.s2End != 0) {
|
||||
// 此时只需要s2End往前退一格,bN2End和gS1End都不变
|
||||
splitCoreInfo.s2End = splitCoreInfo.s2End - 1;
|
||||
} else {
|
||||
if (splitCoreInfo.gS1End != 0) {
|
||||
// splitCoreInfo.gS1End != 0 splitCoreInfo.s2End == 0 时,gS1End需要往前退一格, bN2End不变
|
||||
// 此时需要使用bIdx获取实际Actal S2来计算出 s2End
|
||||
splitCoreInfo.gS1End = splitCoreInfo.gS1End - 1;
|
||||
// 需要获取当前的Actaul S2
|
||||
uint32_t bIdx = splitCoreInfo.bN2End / constInfo.kHeadNum;
|
||||
uint32_t actS1Size, actS2Size, actS2SizeOrig;
|
||||
GetS1S2ActualSeqLen(bIdx, actS1Size, actS2Size, actS2SizeOrig);
|
||||
// s2的切块数量
|
||||
uint32_t s2BaseNum;
|
||||
if (constInfo.attenMaskFlag) {
|
||||
s2BaseNum = GetS2BaseBlockNumOnMask(splitCoreInfo.gS1End, actS1Size, actS2SizeOrig);
|
||||
} else {
|
||||
s2BaseNum = CeilDiv(actS2Size, constInfo.s2BaseSize);
|
||||
}
|
||||
splitCoreInfo.s2End = s2BaseNum - 1;
|
||||
} else {
|
||||
// splitCoreInfo.gS1End == 0 splitCoreInfo.s2End == 0 时,bN2End需要往前退一格
|
||||
// 此时需要使用bIdx获取实际Actal S1和S2来计算出 gS1End 和 s2End
|
||||
splitCoreInfo.bN2End = splitCoreInfo.bN2End - 1;
|
||||
|
||||
// 需要获取当前的Actaul S1 S2
|
||||
uint32_t bIdx = splitCoreInfo.bN2End / constInfo.kHeadNum;
|
||||
uint32_t actS1Size, actS2Size, actS2SizeOrig;
|
||||
GetS1S2ActualSeqLen(bIdx, actS1Size, actS2Size, actS2SizeOrig);
|
||||
|
||||
// s1的切块数量
|
||||
uint32_t s1GBaseNum = CeilDiv(actS1Size, constInfo.s1BaseSize);
|
||||
splitCoreInfo.gS1End = s1GBaseNum - 1;
|
||||
|
||||
// s2的切块数量
|
||||
uint32_t s2BaseNum;
|
||||
if (constInfo.attenMaskFlag) {
|
||||
s2BaseNum = GetS2BaseBlockNumOnMask(splitCoreInfo.gS1End, actS1Size, actS2SizeOrig);
|
||||
} else {
|
||||
s2BaseNum = CeilDiv(actS2Size, constInfo.s2BaseSize);
|
||||
}
|
||||
splitCoreInfo.s2End = s2BaseNum - 1;
|
||||
}
|
||||
}
|
||||
|
||||
splitCoreInfo.isLD = false;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::DealActSeqLenIsZero(uint32_t bIdx, uint32_t n2Idx, uint32_t s1Start)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
if (constInfo.outputLayout == LI_LAYOUT::TND) {
|
||||
uint32_t tSizeIdx = (constInfo.batchSupperFlag) ? constInfo.batchSize : constInfo.batchSize - 1;
|
||||
uint32_t tBaseIdx = (constInfo.batchSupperFlag) ? bIdx : bIdx - 1;
|
||||
uint32_t tSize = actualSeqLengthsGmQ.GetValue(tSizeIdx);
|
||||
uint32_t tBase = bIdx == 0 ? 0 : actualSeqLengthsGmQ.GetValue(tBaseIdx);
|
||||
uint32_t s1Count = tempLoopInfo.actS1Size;
|
||||
|
||||
for (uint32_t s1Idx = s1Start; s1Idx < s1Count; s1Idx++) {
|
||||
uint64_t indiceOutOffset =
|
||||
(tBase + s1Idx) * constInfo.kHeadNum * constInfo.sparseCount + // T轴、s1轴偏移
|
||||
n2Idx * constInfo.sparseCount; // N2轴偏移
|
||||
vectorService.CleanInvalidOutput(indiceOutOffset);
|
||||
}
|
||||
} else if (constInfo.outputLayout == LI_LAYOUT::BSND) {
|
||||
for (uint32_t s1Idx = s1Start; s1Idx < constInfo.qSeqSize; s1Idx++) {
|
||||
// B,S1,N2,K
|
||||
uint64_t indiceOutOffset = bIdx * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount +
|
||||
s1Idx * constInfo.kHeadNum * constInfo.sparseCount + // B轴、S1轴偏移
|
||||
n2Idx * constInfo.sparseCount; // N2轴偏移
|
||||
vectorService.CleanInvalidOutput(indiceOutOffset);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::Init(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights,
|
||||
__gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale,
|
||||
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsK,
|
||||
__gm__ uint8_t *blockTable, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *workspace,
|
||||
const QLITilingData *__restrict tiling, TPipe *tPipe)
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
tmpBlockIdx = GetBlockIdx(); // vec:0-47
|
||||
aiCoreIdx = tmpBlockIdx / 2;
|
||||
} else {
|
||||
tmpBlockIdx = GetBlockIdx(); // cube:0-23
|
||||
aiCoreIdx = tmpBlockIdx;
|
||||
}
|
||||
|
||||
InitTilingData(tiling);
|
||||
InitActualSeqLen(actualSeqLengthsQ, actualSeqLengthsK);
|
||||
|
||||
// 获取分核信息
|
||||
metadataGm.SetGlobalBuffer((__gm__ uint32_t *)metadata);
|
||||
SplitCoreByAICPU(aiCoreIdx, metadataGm);
|
||||
|
||||
pipe = tPipe;
|
||||
|
||||
uint64_t offset = 0;
|
||||
//vec 把整个s2的score存储在GM,大小为s1BaseSize * 16K * 4
|
||||
GlobalTensor<SCORE_T> scoreGm; //存放vec核写出的score
|
||||
if ASCEND_IS_AIV {
|
||||
uint64_t singleCoreScoreSize = constInfo.s1BaseSize * QLICommon::Align((uint64_t)constInfo.kSeqSize, (uint64_t)constInfo.s2BaseSize) * sizeof(SCORE_T);
|
||||
scoreGm.SetGlobalBuffer((__gm__ SCORE_T *)(workspace + aiCoreIdx * singleCoreScoreSize));
|
||||
offset += GetBlockNum() * singleCoreScoreSize;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.InitParams(constInfo, tiling);
|
||||
indiceOutGm.SetGlobalBuffer((__gm__ int32_t *)sparseIndices);
|
||||
weightsGm.SetGlobalBuffer((__gm__ float *)weights);
|
||||
qScaleGm.SetGlobalBuffer((__gm__ float *)queryScale);
|
||||
kScaleGm.SetGlobalBuffer((__gm__ float *)keyScale);
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
vectorService.InitVecInputTensor(weightsGm, qScaleGm, kScaleGm, indiceOutGm, blockTableGm);
|
||||
vectorService.InitVecWorkspaceTensor(scoreGm);
|
||||
} else {
|
||||
matmulService.InitParams(constInfo);
|
||||
queryGm.SetGlobalBuffer((__gm__ Q_T *)query);
|
||||
if constexpr (PAGE_ATTENTION) {
|
||||
blockTableGm.SetGlobalBuffer((__gm__ int32_t *)blockTable);
|
||||
}
|
||||
keyGm.SetGlobalBuffer((__gm__ K_T *)key);
|
||||
matmulService.InitMm1GlobalTensor(blockTableGm, keyGm, queryGm);
|
||||
}
|
||||
InitBuffers();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::GetBN2Idx(uint32_t bN2Idx)
|
||||
{
|
||||
tempLoopInfo.bN2Idx = bN2Idx;
|
||||
tempLoopInfo.bIdx = bN2Idx / constInfo.kHeadNum;
|
||||
tempLoopInfo.n2Idx = bN2Idx % constInfo.kHeadNum;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcS2LoopParams(uint32_t bN2LoopIdx, uint32_t gS1LoopIdx)
|
||||
{
|
||||
tempLoopInfo.gS1Idx = gS1LoopIdx;
|
||||
tempLoopInfo.actMBaseSize = constInfo.mBaseSize;
|
||||
uint32_t remainedGS1Size = tempLoopInfo.actS1Size * constInfo.gSize - tempLoopInfo.gS1Idx * constInfo.mBaseSize;
|
||||
if (remainedGS1Size <= constInfo.mBaseSize && remainedGS1Size > 0) {
|
||||
tempLoopInfo.actMBaseSize = tempLoopInfo.mBasicSizeTail;
|
||||
}
|
||||
|
||||
bool isEnd = (bN2LoopIdx == splitCoreInfo.bN2End) && (gS1LoopIdx == splitCoreInfo.gS1End);
|
||||
uint32_t s2BlockNum;
|
||||
if (constInfo.attenMaskFlag) {
|
||||
s2BlockNum = GetS2BaseBlockNumOnMask(gS1LoopIdx, tempLoopInfo.actS1Size, tempLoopInfo.actS2SizeOrig);
|
||||
} else {
|
||||
s2BlockNum = (tempLoopInfo.actS2Size + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
}
|
||||
tempLoopInfo.s2LoopEnd = isEnd ? splitCoreInfo.s2End : s2BlockNum - 1;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcGS1LoopParams(uint32_t bN2LoopIdx)
|
||||
{
|
||||
GetBN2Idx(bN2LoopIdx);
|
||||
GetS1S2ActualSeqLen(tempLoopInfo.bIdx, tempLoopInfo.actS1Size, tempLoopInfo.actS2Size, tempLoopInfo.actS2SizeOrig);
|
||||
if ((tempLoopInfo.actS2Size == 0) || (tempLoopInfo.actS1Size == 0)) {
|
||||
tempLoopInfo.curActSeqLenIsZero = true;
|
||||
return;
|
||||
}
|
||||
tempLoopInfo.curActSeqLenIsZero = false;
|
||||
tempLoopInfo.s2BasicSizeTail = tempLoopInfo.actS2Size % constInfo.s2BaseSize;
|
||||
tempLoopInfo.s2BasicSizeTail =
|
||||
(tempLoopInfo.s2BasicSizeTail == 0) ? constInfo.s2BaseSize : tempLoopInfo.s2BasicSizeTail;
|
||||
tempLoopInfo.mBasicSizeTail = (tempLoopInfo.actS1Size * constInfo.gSize) % constInfo.mBaseSize;
|
||||
tempLoopInfo.mBasicSizeTail =
|
||||
(tempLoopInfo.mBasicSizeTail == 0) ? constInfo.mBaseSize : tempLoopInfo.mBasicSizeTail;
|
||||
|
||||
uint32_t gS1SplitNum = (tempLoopInfo.actS1Size * constInfo.gSize + constInfo.mBaseSize - 1) / constInfo.mBaseSize;
|
||||
tempLoopInfo.gS1LoopEnd = (bN2LoopIdx == splitCoreInfo.bN2End) ? splitCoreInfo.gS1End : gS1SplitNum - 1;
|
||||
if constexpr (Q_LAYOUT_T == LI_LAYOUT::BSND) {
|
||||
if (tempLoopInfo.gS1LoopEnd == gS1SplitNum - 1 && constInfo.qSeqSize > tempLoopInfo.actS1Size) {
|
||||
tempLoopInfo.needDealActS1LessThanS1 = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::CalcRunInfo(uint32_t loop, uint32_t s2LoopIdx, QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
runInfo.loop = loop;
|
||||
runInfo.bIdx = tempLoopInfo.bIdx;
|
||||
runInfo.gS1Idx = tempLoopInfo.gS1Idx;
|
||||
runInfo.s2Idx = s2LoopIdx;
|
||||
runInfo.bN2Idx = tempLoopInfo.bN2Idx;
|
||||
runInfo.isValid = s2LoopIdx <= tempLoopInfo.s2LoopEnd;
|
||||
|
||||
if (!runInfo.isValid) {
|
||||
return; // 需要验证, v1 时候需要runInfo
|
||||
}
|
||||
|
||||
runInfo.actS1Size = tempLoopInfo.actS1Size;
|
||||
runInfo.actS2Size = tempLoopInfo.actS2Size;
|
||||
runInfo.actS2SizeOrig = tempLoopInfo.actS2SizeOrig;
|
||||
// 计算实际基本块size
|
||||
runInfo.actMBaseSize = tempLoopInfo.actMBaseSize;
|
||||
runInfo.actualSingleProcessSInnerSize = constInfo.s2BaseSize;
|
||||
uint32_t s2SplitNum = (tempLoopInfo.actS2Size + constInfo.s2BaseSize - 1) / constInfo.s2BaseSize;
|
||||
if (runInfo.s2Idx == s2SplitNum - 1) {
|
||||
runInfo.actualSingleProcessSInnerSize = tempLoopInfo.s2BasicSizeTail;
|
||||
}
|
||||
runInfo.actualSingleProcessSInnerSizeAlign =
|
||||
QLICommon::Align((uint32_t)runInfo.actualSingleProcessSInnerSize, QLICommon::ConstInfo::BUFFER_SIZE_BYTE_32B);
|
||||
|
||||
runInfo.isFirstS2InnerLoop = s2LoopIdx == splitCoreInfo.s2Start;
|
||||
runInfo.isLastS2InnerLoop = s2LoopIdx == tempLoopInfo.s2LoopEnd;
|
||||
runInfo.isAllLoopEnd = (runInfo.bN2Idx == splitCoreInfo.bN2End) && (runInfo.gS1Idx == splitCoreInfo.gS1End) &&
|
||||
(runInfo.s2Idx == splitCoreInfo.s2End);
|
||||
|
||||
if (runInfo.isFirstS2InnerLoop) {
|
||||
uint64_t actualSeqQPrefixSum;
|
||||
if constexpr (Q_LAYOUT_T == LI_LAYOUT::TND) {
|
||||
uint32_t actualSeqLengthsGmQIdx = (constInfo.batchSupperFlag) ? runInfo.bIdx : runInfo.bIdx - 1;
|
||||
actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : actualSeqLengthsGmQ.GetValue(actualSeqLengthsGmQIdx);
|
||||
} else { // BSND
|
||||
actualSeqQPrefixSum = (runInfo.bIdx <= 0) ? 0 : runInfo.bIdx * constInfo.qSeqSize;
|
||||
}
|
||||
uint64_t tndBIdxOffset = actualSeqQPrefixSum * constInfo.qHeadNum * constInfo.headDim;
|
||||
// B,S1,N1(N2,G),D
|
||||
queryCoreOffset = tndBIdxOffset + runInfo.gS1Idx * constInfo.mBaseSize * constInfo.headDim;
|
||||
// B,S1,N1(N2,G)/T,N1(N2,G)
|
||||
weightsCoreOffset = actualSeqQPrefixSum * constInfo.qHeadNum + runInfo.n2Idx * constInfo.gSize;
|
||||
// B,S1,N2,k/T,N2,k
|
||||
indiceOutCoreOffset =
|
||||
actualSeqQPrefixSum * constInfo.kHeadNum * constInfo.sparseCount + runInfo.n2Idx * constInfo.sparseCount;
|
||||
}
|
||||
uint64_t actualSeqKPrefixSum;
|
||||
if constexpr (K_LAYOUT_T == LI_LAYOUT::TND) { // T N2 D
|
||||
actualSeqKPrefixSum = (runInfo.bIdx <= 0) ? 0 : actualSeqLengthsGm.GetValue(runInfo.bIdx - 1);
|
||||
actualSeqKPrefixSum = actualSeqKPrefixSum / constInfo.cmpRatio;
|
||||
} else {
|
||||
actualSeqKPrefixSum = (runInfo.bIdx <= 0) ? 0 : runInfo.bIdx * constInfo.kSeqSize;
|
||||
}
|
||||
uint64_t tndBIdxOffsetForK = actualSeqKPrefixSum * constInfo.kHeadNum * constInfo.headDim;
|
||||
keyCoreOffset = tndBIdxOffsetForK + runInfo.s2Idx * constInfo.s2BaseSize * constInfo.kHeadNum * constInfo.headDim;
|
||||
keyScaleCoreOffset = (actualSeqKPrefixSum + runInfo.s2Idx * constInfo.s2BaseSize) * constInfo.kHeadNum;
|
||||
runInfo.tensorQueryOffset = queryCoreOffset;
|
||||
runInfo.tensorKeyOffset = keyCoreOffset;
|
||||
runInfo.tensorKeyScaleOffset = keyScaleCoreOffset;
|
||||
runInfo.tensorWeightsOffset = weightsCoreOffset;
|
||||
runInfo.indiceOutOffset = indiceOutCoreOffset;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::Process()
|
||||
{
|
||||
if (isUsedCoreEqZero) {
|
||||
// 没有计算任务,直接清理输出
|
||||
ProcessInvalid();
|
||||
return;
|
||||
}
|
||||
|
||||
ProcessMain();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessInvalid()
|
||||
{
|
||||
if ASCEND_IS_AIV {
|
||||
uint32_t aivCoreNum = GetBlockNum() * 2; // 2 means c:v = 1:2
|
||||
uint64_t totalOutputSize =
|
||||
constInfo.batchSize * constInfo.qSeqSize * constInfo.kHeadNum * constInfo.sparseCount;
|
||||
uint64_t singleCoreSize =
|
||||
QLICommon::Align((totalOutputSize + aivCoreNum - 1) / aivCoreNum, GM_ALIGN_BYTES / sizeof(OUT_T));
|
||||
uint64_t baseSize = tmpBlockIdx * singleCoreSize;
|
||||
if (baseSize < totalOutputSize) {
|
||||
uint64_t dealSize =
|
||||
(baseSize + singleCoreSize <= totalOutputSize) ? singleCoreSize : totalOutputSize - baseSize;
|
||||
GlobalTensor<OUT_T> output = indiceOutGm[baseSize];
|
||||
AscendC::InitGlobalMemory(output, dealSize, constInfo.INVALID_IDX);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessMain()
|
||||
{
|
||||
if(!splitCoreInfo.isCoreEnable){
|
||||
return;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.AllocEventID();
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_VC_EVENT + 0);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_VC_EVENT + 1);
|
||||
} else {
|
||||
matmulService.AllocEventID();
|
||||
}
|
||||
|
||||
QLICommon::RunInfo runInfo;
|
||||
uint32_t gloop = 0;
|
||||
for (uint32_t bN2LoopIdx = splitCoreInfo.bN2Start; bN2LoopIdx <= splitCoreInfo.bN2End; bN2LoopIdx++) {
|
||||
CalcGS1LoopParams(bN2LoopIdx);
|
||||
if (tempLoopInfo.curActSeqLenIsZero) {
|
||||
DealActSeqLenIsZero(tempLoopInfo.bIdx, tempLoopInfo.n2Idx, 0U);
|
||||
continue;
|
||||
}
|
||||
for (uint32_t gS1LoopIdx = splitCoreInfo.gS1Start; gS1LoopIdx <= tempLoopInfo.gS1LoopEnd; gS1LoopIdx++) {
|
||||
CalcS2LoopParams(bN2LoopIdx, gS1LoopIdx);
|
||||
for (int s2LoopIdx = splitCoreInfo.s2Start; s2LoopIdx <= tempLoopInfo.s2LoopEnd; s2LoopIdx++) {
|
||||
ProcessBaseBlock(gloop, s2LoopIdx, runInfo);
|
||||
++gloop;
|
||||
}
|
||||
splitCoreInfo.s2Start = 0;
|
||||
}
|
||||
if (tempLoopInfo.needDealActS1LessThanS1) {
|
||||
DealActSeqLenIsZero(tempLoopInfo.bIdx, tempLoopInfo.n2Idx, tempLoopInfo.actS1Size);
|
||||
}
|
||||
splitCoreInfo.gS1Start = 0;
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vectorService.FreeEventID();
|
||||
} else {
|
||||
matmulService.FreeEventID();
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_VC_EVENT + 0);
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_VC_EVENT + 1);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIPreload<QLIT>::ProcessBaseBlock(uint32_t loop, uint64_t s2LoopIdx, QLICommon::RunInfo runInfo)
|
||||
{
|
||||
CalcRunInfo(loop, s2LoopIdx, runInfo);
|
||||
if ASCEND_IS_AIC {
|
||||
matmulService.ComputeMm1(runInfo);
|
||||
} else {
|
||||
vectorService.ProcessVec1(runInfo);
|
||||
if (runInfo.isLastS2InnerLoop) { //本核s2last
|
||||
vectorService.ProcessTopK(runInfo);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace QLIKernel
|
||||
#endif // quant_lightning_indexer_KERNEL_H
|
||||
@@ -0,0 +1,438 @@
|
||||
/**
|
||||
* 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 lightning_indexer_service_cube.h
|
||||
* \brief use 5 buffer for matmul l1, better pipeline
|
||||
*/
|
||||
#ifndef quant_lightning_indexer_SERVICE_CUBE_H
|
||||
#define quant_lightning_indexer_SERVICE_CUBE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
template <typename QLIT>
|
||||
class QLIMatmul {
|
||||
public:
|
||||
using Q_T = typename QLIT::queryType;
|
||||
using K_T = typename QLIT::keyType;
|
||||
using QK_T = typename QLIT::queryKeyType;
|
||||
|
||||
__aicore__ inline QLIMatmul(){};
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void InitMm1GlobalTensor(const GlobalTensor<int32_t> &blkTableGm, const GlobalTensor<K_T> &keyGm,
|
||||
const GlobalTensor<Q_T> &queryGm);
|
||||
__aicore__ inline void InitParams(const ConstInfo &constInfo);
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
__aicore__ inline void ComputeMm1(const QLICommon::RunInfo &runInfo);
|
||||
|
||||
static constexpr IsResetLoad3dConfig LOAD3DV2_CONFIG = {true, true}; // isSetFMatrix isSetPadding;
|
||||
static constexpr uint64_t KEY_BUF_NUM = 3;
|
||||
static constexpr uint64_t QUERY_BUF_NUM = 2;
|
||||
static constexpr uint64_t L0_BUF_NUM = 2;
|
||||
|
||||
static constexpr uint32_t KEY_MTE1_MTE2_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t QUERY_MTE1_MTE2_EVENT = EVENT_ID5; // KEY_MTE1_MTE2_EVENT + KEY_BUF_NUM;
|
||||
static constexpr uint32_t M_MTE1_EVENT = EVENT_ID3;
|
||||
|
||||
static constexpr uint32_t MTE2_MTE1_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t MTE1_M_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t FIX_M_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t M_FIX_EVENT = EVENT_ID3;
|
||||
|
||||
static constexpr uint64_t M_BASIC_BLOCK = 256;
|
||||
static constexpr uint64_t D_BASIC_BLOCK = 128;
|
||||
static constexpr uint64_t S2_BASIC_BLOCK = 128;
|
||||
|
||||
static constexpr uint64_t M_BASIC_BLOCK_L0 = 256;
|
||||
static constexpr uint64_t D_BASIC_BLOCK_L0 = 128;
|
||||
static constexpr uint64_t S2_BASIC_BLOCK_L0 = 128;
|
||||
|
||||
static constexpr uint64_t FP8_BLOCK_CUBE = 32;
|
||||
static constexpr FixpipeConfig QLI_CFG_ROW_MAJOR_UB = {CO2Layout::ROW_MAJOR, true}; // ROW_MAJOR: 使能NZ2ND,输出数据格式为ND格式; true: 用于用户指定目的地址的位置是否是UB
|
||||
|
||||
static constexpr uint64_t QUERY_BUFFER_OFFSET = M_BASIC_BLOCK * D_BASIC_BLOCK;
|
||||
static constexpr uint64_t KEY_BUFFER_OFFSET = S2_BASIC_BLOCK * D_BASIC_BLOCK;
|
||||
static constexpr uint64_t L0AB_BUFFER_OFFSET = M_BASIC_BLOCK_L0 * D_BASIC_BLOCK_L0;
|
||||
static constexpr uint64_t L0C_BUFFER_OFFSET = M_BASIC_BLOCK_L0 * S2_BASIC_BLOCK_L0;
|
||||
|
||||
protected:
|
||||
__aicore__ inline void Fixp(uint64_t s1gGmOffset, uint64_t s2GmOffset, uint64_t s1gL0RealSize,
|
||||
uint64_t s2L0RealSize, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void ComuteL0c(uint64_t s1gL0RealSize, uint64_t s2L0RealSize, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void LoadKeyToL0b(uint64_t s2L0Offset, uint64_t s2L1RealSize, uint64_t s2L0RealSize,
|
||||
const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void LoadQueryToL0a(uint64_t s1gL1Offset, uint64_t s1gL0Offset, uint64_t s1gL1RealSize,
|
||||
uint64_t s1gL0RealSize, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void QueryNd2Nz(uint64_t s1gL1RealSize, uint64_t s1gL1Offset, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void KeyNd2Nz(uint64_t s2L1RealSize, uint64_t s2GmOffset, const QLICommon::RunInfo &runInfo);
|
||||
__aicore__ inline void KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset, const QLICommon::RunInfo &runInfo);
|
||||
GlobalTensor<int32_t> blkTableGm_;
|
||||
GlobalTensor<K_T> keyGm_;
|
||||
GlobalTensor<Q_T> queryGm_;
|
||||
|
||||
TBuf<TPosition::A1> bufQL1_;
|
||||
LocalTensor<Q_T> queryL1_;
|
||||
TBuf<TPosition::B1> bufKeyL1_;
|
||||
LocalTensor<K_T> keyL1_;
|
||||
|
||||
TBuf<TPosition::A2> bufQL0_;
|
||||
LocalTensor<Q_T> queryL0_;
|
||||
TBuf<TPosition::B2> bufKeyL0_;
|
||||
LocalTensor<K_T> keyL0_;
|
||||
|
||||
TBuf<TPosition::CO1> bufL0C_;
|
||||
LocalTensor<float> cL0_;
|
||||
|
||||
TBuf<TPosition::VECCALC> bufUB_;
|
||||
LocalTensor<QK_T> mm1ResUB_;
|
||||
|
||||
uint64_t keyL1BufIdx_ = 0;
|
||||
uint64_t queryL1Mte2BufIdx_ = 0;
|
||||
uint64_t queryL1Mte1BufIdx_ = 0;
|
||||
uint64_t l0BufIdx_ = 0;
|
||||
|
||||
ConstInfo constInfo_;
|
||||
|
||||
private:
|
||||
static constexpr bool PAGE_ATTENTION = QLIT::pageAttention;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::InitParams(const ConstInfo &constInfo)
|
||||
{
|
||||
constInfo_ = constInfo;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(bufUB_, 2 * CeilDiv(constInfo_.mBaseSize, 2) * constInfo_.s2BaseSize * sizeof(QK_T)); //大小:2(开dB) * 2 * 64 * 128 * 4 = 128KB
|
||||
mm1ResUB_ = bufUB_.Get<QK_T>();
|
||||
pipe->InitBuffer(bufQL1_, QUERY_BUF_NUM * M_BASIC_BLOCK * D_BASIC_BLOCK * sizeof(Q_T));
|
||||
queryL1_ = bufQL1_.Get<Q_T>();
|
||||
pipe->InitBuffer(bufKeyL1_, KEY_BUF_NUM * S2_BASIC_BLOCK * D_BASIC_BLOCK * sizeof(K_T));
|
||||
keyL1_ = bufKeyL1_.Get<K_T>();
|
||||
|
||||
pipe->InitBuffer(bufQL0_, L0_BUF_NUM * M_BASIC_BLOCK_L0 * D_BASIC_BLOCK_L0 * sizeof(Q_T));
|
||||
queryL0_ = bufQL0_.Get<Q_T>();
|
||||
pipe->InitBuffer(bufKeyL0_, L0_BUF_NUM * D_BASIC_BLOCK_L0 * S2_BASIC_BLOCK_L0 * sizeof(K_T));
|
||||
keyL0_ = bufKeyL0_.Get<K_T>();
|
||||
|
||||
pipe->InitBuffer(bufL0C_, L0_BUF_NUM * M_BASIC_BLOCK_L0 * S2_BASIC_BLOCK_L0 * sizeof(float));
|
||||
cL0_ = bufL0C_.Get<float>();
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void
|
||||
QLIMatmul<QLIT>::InitMm1GlobalTensor(const GlobalTensor<int32_t> &blkTableGm, const GlobalTensor<K_T> &keyGm,
|
||||
const GlobalTensor<Q_T> &queryGm)
|
||||
{
|
||||
blkTableGm_ = blkTableGm;
|
||||
keyGm_ = keyGm;
|
||||
queryGm_ = queryGm;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ComputeMm1(const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_VC_EVENT + runInfo.loop % 2);
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_VC_EVENT + runInfo.loop % 2 + QLICommon::ConstInfo::AIV0_AIV1_OFFSET);
|
||||
uint64_t s2GmBaseOffset = runInfo.s2Idx * constInfo_.s2BaseSize;
|
||||
uint64_t s1gProcessSize = runInfo.actMBaseSize;
|
||||
uint64_t s2ProcessSize = runInfo.actualSingleProcessSInnerSize;
|
||||
for (uint64_t s2GmOffset = 0; s2GmOffset < s2ProcessSize; s2GmOffset += S2_BASIC_BLOCK) {
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + keyL1BufIdx_ % KEY_BUF_NUM);
|
||||
uint64_t s2L1RealSize =
|
||||
s2GmOffset + S2_BASIC_BLOCK > s2ProcessSize ? s2ProcessSize - s2GmOffset : S2_BASIC_BLOCK;
|
||||
if (PAGE_ATTENTION) {
|
||||
KeyNd2NzForPA(s2L1RealSize, s2GmBaseOffset + s2GmOffset, runInfo);
|
||||
}else {
|
||||
KeyNd2Nz(s2L1RealSize, s2GmOffset, runInfo);
|
||||
}
|
||||
|
||||
SetFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
// s1gProcessSize当前必定不会超过2倍的s1g basic block
|
||||
for (uint64_t s1gGmOffset = 0; s1gGmOffset < s1gProcessSize; s1gGmOffset += M_BASIC_BLOCK) {
|
||||
uint64_t s1gL1RealSize =
|
||||
s1gGmOffset + M_BASIC_BLOCK > s1gProcessSize ? s1gProcessSize - s1gGmOffset : M_BASIC_BLOCK;
|
||||
uint64_t s1gL1SizeAlign2G = CeilAlign(s1gL1RealSize, 2 * constInfo_.gSize);
|
||||
if (runInfo.isFirstS2InnerLoop && s2GmOffset == 0) {
|
||||
queryL1Mte2BufIdx_++;
|
||||
queryL1Mte1BufIdx_ = queryL1Mte2BufIdx_;
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + queryL1Mte2BufIdx_ % QUERY_BUF_NUM);
|
||||
QueryNd2Nz(s1gL1SizeAlign2G, s1gGmOffset, runInfo);
|
||||
SetFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_MTE1>(MTE2_MTE1_EVENT);
|
||||
} else {
|
||||
queryL1Mte1BufIdx_ =
|
||||
queryL1Mte2BufIdx_ - (CeilDiv(s1gProcessSize, M_BASIC_BLOCK) - 1 - (s1gGmOffset > 0));
|
||||
}
|
||||
for (uint64_t s2L1Offset = 0; s2L1Offset < s2L1RealSize; s2L1Offset += S2_BASIC_BLOCK_L0) {
|
||||
uint64_t s2L0RealSize =
|
||||
s2L1Offset + S2_BASIC_BLOCK_L0 > s2L1RealSize ? s2L1RealSize - s2L1Offset : S2_BASIC_BLOCK_L0;
|
||||
for (uint64_t s1gL1Offset = 0; s1gL1Offset < s1gL1SizeAlign2G; s1gL1Offset += M_BASIC_BLOCK_L0) {
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
uint64_t s1gL0RealSize =
|
||||
s1gL1Offset + M_BASIC_BLOCK_L0 > s1gL1SizeAlign2G ? s1gL1SizeAlign2G - s1gL1Offset : M_BASIC_BLOCK_L0;
|
||||
LoadQueryToL0a(s1gGmOffset, s1gL1Offset, s1gL1SizeAlign2G, s1gL0RealSize, runInfo);
|
||||
LoadKeyToL0b(s2L1Offset, s2L1RealSize, s2L0RealSize, runInfo);
|
||||
|
||||
SetFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
WaitFlag<HardEvent::MTE1_M>(MTE1_M_EVENT);
|
||||
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
ComuteL0c(s1gL0RealSize, s2L0RealSize, runInfo);
|
||||
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
|
||||
Fixp(s1gGmOffset + s1gL1Offset, s2GmOffset + s2L1Offset, s1gL0RealSize, s2L0RealSize, runInfo);
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
l0BufIdx_++;
|
||||
}
|
||||
}
|
||||
if (s2GmOffset + S2_BASIC_BLOCK >= s2ProcessSize && runInfo.isLastS2InnerLoop) {
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + queryL1Mte1BufIdx_ % QUERY_BUF_NUM);
|
||||
}
|
||||
}
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + keyL1BufIdx_ % KEY_BUF_NUM);
|
||||
keyL1BufIdx_++;
|
||||
}
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_CV_EVENT + runInfo.loop % 2);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_FIX>(QLICommon::ConstInfo::CROSS_CV_EVENT + runInfo.loop % 2 + QLICommon::ConstInfo::AIV0_AIV1_OFFSET);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::KeyNd2Nz(uint64_t s2L1RealSize, uint64_t s2GmOffset,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s2L1RealSize; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = constInfo_.headDim;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
// 默认一块buf最多放两份
|
||||
DataCopy(keyL1_[(keyL1BufIdx_ % KEY_BUF_NUM) * KEY_BUFFER_OFFSET],
|
||||
keyGm_[runInfo.tensorKeyOffset + s2GmOffset * constInfo_.headDim], nd2nzPara);
|
||||
}
|
||||
|
||||
// blkNum, blkSize, N2, D
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::KeyNd2NzForPA(uint64_t s2L1RealSize, uint64_t s2GmOffset,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
uint64_t s2L1Offset = 0;
|
||||
while (s2L1Offset < s2L1RealSize) {
|
||||
uint64_t s2BlkId = (s2L1Offset + s2GmOffset) / constInfo_.kCacheBlockSize;
|
||||
uint64_t s2BlkOffset = (s2L1Offset + s2GmOffset) % constInfo_.kCacheBlockSize;
|
||||
uint64_t keyGmOffset = blkTableGm_.GetValue(runInfo.bIdx * constInfo_.maxBlockNumPerBatch + s2BlkId) *
|
||||
constInfo_.stride +
|
||||
s2BlkOffset * constInfo_.headDim;
|
||||
|
||||
uint64_t s2Mte2Size = s2L1RealSize - s2L1Offset;
|
||||
s2Mte2Size = s2BlkOffset + s2Mte2Size >= constInfo_.kCacheBlockSize ? constInfo_.kCacheBlockSize - s2BlkOffset
|
||||
: s2Mte2Size;
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s2Mte2Size; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = constInfo_.headDim;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s2L1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
DataCopy(keyL1_[(keyL1BufIdx_ % KEY_BUF_NUM) * KEY_BUFFER_OFFSET + s2L1Offset * FP8_BLOCK_CUBE],
|
||||
keyGm_[keyGmOffset], nd2nzPara);
|
||||
|
||||
s2L1Offset += s2Mte2Size;
|
||||
}
|
||||
}
|
||||
|
||||
// batch, s1, n2, g, d
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::QueryNd2Nz(uint64_t s1gL1RealSize, uint64_t s1gGmOffset,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
Nd2NzParams nd2nzPara;
|
||||
nd2nzPara.ndNum = 1;
|
||||
nd2nzPara.nValue = s1gL1RealSize; // 行数
|
||||
nd2nzPara.dValue = constInfo_.headDim;
|
||||
nd2nzPara.srcDValue = constInfo_.headDim;
|
||||
nd2nzPara.dstNzC0Stride = CeilAlign(s1gL1RealSize, (uint64_t)BLOCK_CUBE); // 对齐到16 单位block
|
||||
nd2nzPara.dstNzNStride = 1;
|
||||
nd2nzPara.srcNdMatrixStride = 0;
|
||||
nd2nzPara.dstNzMatrixStride = 0;
|
||||
// 默认一块buf最多放两份
|
||||
DataCopy(queryL1_[(queryL1Mte2BufIdx_ % QUERY_BUF_NUM) * QUERY_BUFFER_OFFSET],
|
||||
queryGm_[runInfo.tensorQueryOffset + s1gGmOffset * constInfo_.headDim], nd2nzPara);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadQueryToL0a(uint64_t s1gGmOffset, uint64_t s1gL1Offset, uint64_t s1gL1RealSize,
|
||||
uint64_t s1gL0RealSize, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
LoadData2DParamsV2 loadData2DParamsV2;
|
||||
loadData2DParamsV2.mStartPosition = CeilDiv(s1gL1Offset, BLOCK_CUBE);
|
||||
loadData2DParamsV2.kStartPosition = 0;
|
||||
loadData2DParamsV2.mStep = CeilDiv(s1gL0RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.kStep = CeilDiv(constInfo_.headDim, FP8_BLOCK_CUBE);
|
||||
loadData2DParamsV2.srcStride = CeilDiv(s1gL1RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.dstStride = CeilDiv(s1gL0RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.ifTranspose = false;
|
||||
|
||||
LoadData(queryL0_[(l0BufIdx_ % L0_BUF_NUM) * L0AB_BUFFER_OFFSET],
|
||||
queryL1_[(queryL1Mte1BufIdx_ % QUERY_BUF_NUM) * QUERY_BUFFER_OFFSET], loadData2DParamsV2);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::LoadKeyToL0b(uint64_t s2L1Offset, uint64_t s2L1RealSize, uint64_t s2L0RealSize,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
LoadData2DParamsV2 loadData2DParamsV2;
|
||||
loadData2DParamsV2.mStartPosition = CeilDiv(s2L1Offset, BLOCK_CUBE);
|
||||
loadData2DParamsV2.kStartPosition = 0;
|
||||
loadData2DParamsV2.mStep = CeilDiv(s2L0RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.kStep = CeilDiv(constInfo_.headDim, FP8_BLOCK_CUBE);
|
||||
loadData2DParamsV2.srcStride = CeilDiv(s2L1RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.dstStride = CeilDiv(s2L0RealSize, BLOCK_CUBE);
|
||||
loadData2DParamsV2.ifTranspose = false;
|
||||
|
||||
LoadData(keyL0_[(l0BufIdx_ % L0_BUF_NUM) * L0AB_BUFFER_OFFSET],
|
||||
keyL1_[(keyL1BufIdx_ % KEY_BUF_NUM) * KEY_BUFFER_OFFSET], loadData2DParamsV2);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::ComuteL0c(uint64_t s1gL0RealSize, uint64_t s2L0RealSize,
|
||||
const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
MmadParams mmadParams;
|
||||
mmadParams.m = CeilAlign(s1gL0RealSize, BLOCK_CUBE);
|
||||
mmadParams.n = s2L0RealSize;
|
||||
mmadParams.k = constInfo_.headDim;
|
||||
mmadParams.cmatrixInitVal = true;
|
||||
mmadParams.cmatrixSource = false;
|
||||
Mmad(cL0_[(l0BufIdx_ % L0_BUF_NUM) * L0C_BUFFER_OFFSET], queryL0_[(l0BufIdx_ % L0_BUF_NUM) * L0AB_BUFFER_OFFSET],
|
||||
keyL0_[(l0BufIdx_ % L0_BUF_NUM) * L0AB_BUFFER_OFFSET], mmadParams);
|
||||
if ((mmadParams.m / 16) * (mmadParams.n / 16) < 10) {
|
||||
PipeBarrier<PIPE_M>();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::Fixp(uint64_t s1gGmOffset, uint64_t s2GmOffset, uint64_t s1gL0RealSize,
|
||||
uint64_t s2L0RealSize, const QLICommon::RunInfo &runInfo)
|
||||
{
|
||||
SetFlag<HardEvent::M_FIX>(M_FIX_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
WaitFlag<HardEvent::M_FIX>(M_FIX_EVENT + l0BufIdx_ % L0_BUF_NUM);
|
||||
|
||||
static_assert(S2_BASIC_BLOCK == S2_BASIC_BLOCK_L0 && S2_BASIC_BLOCK_L0 == 128);
|
||||
if constexpr (std::is_same_v<QK_T, float>) {
|
||||
// s1gL0RealSize:2*gSize(128)对齐, 最大256
|
||||
// s2L0RealSize <= S2_BASIC_BLOCK_L0, 未约束
|
||||
uint32_t nSize = (s2L0RealSize + 7) >> 3 << 3; // 32B对齐
|
||||
uint32_t mSize = (s1gL0RealSize + 1) >> 1 << 1;
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams;
|
||||
// 固定参数
|
||||
fixpipeParams.mSize = mSize;
|
||||
fixpipeParams.srcStride = mSize; // 已16对齐
|
||||
fixpipeParams.dstStride = UB_BANK_DEPTH_STRIDE / sizeof(QK_T); // 落到同一个bank
|
||||
fixpipeParams.dualDstCtl = 1; // 双目标模式,按M维度拆分, M / 2 * N写入每个UB,M必须为2的倍数
|
||||
|
||||
// nSize已保证N方向32B对齐
|
||||
if (nSize <= (256 / sizeof(float))) {
|
||||
// N方向小于一个bank(256B), 只需搬一个ND块, 且不用补齐
|
||||
fixpipeParams.nSize = nSize;
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
} else {
|
||||
// N方向在(256B, 512B]范围, 直接按512B搬, 注意此时不能开unitflag
|
||||
fixpipeParams.nSize = S2_BASIC_BLOCK_L0 / 2; // 分2个ND搬, S2_BASIC_BLOCK_L0不为128会有问题
|
||||
fixpipeParams.params.ndNum = 2;
|
||||
fixpipeParams.params.srcNdStride = ((fixpipeParams.mSize + 15) / 16) * fixpipeParams.nSize;
|
||||
fixpipeParams.params.dstNdStride = constInfo_.s2BaseSize * constInfo_.mBaseSize / 2; // S2_BASIC_BLOCK * M_BASE_SIZE / 2
|
||||
}
|
||||
Fixpipe<QK_T, float, QLI_CFG_ROW_MAJOR_UB>(mm1ResUB_[(runInfo.loop % 2) * constInfo_.s2BaseSize / 2], // 未考虑s1gGmOffset和s2GmOffset
|
||||
cL0_[(l0BufIdx_ % L0_BUF_NUM) * L0C_BUFFER_OFFSET], fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
} else {
|
||||
// nSize * sizeof(QT) <= 256B, 小于一个UB bank大小(VL)
|
||||
uint32_t nSize = (s2L0RealSize + 7) >> 3 << 3; // 8个元素(32B)对齐
|
||||
uint32_t mSize = (s1gL0RealSize + 1) >> 1 << 1; // 有效数据不足16行,只需输出部分行即可;L0C上的bmm1结果矩阵M方向的size大小必须是偶数
|
||||
uint32_t srcStride = ((mSize + 15) / 16) * 16; // L0C上matmul结果相邻连续数据片断间隔(前面一个数据块的头与后面数据块的头的间隔),单位为16 *sizeof(T) //源NZ矩阵中相邻Z排布的起始地址偏移
|
||||
FixpipeParamsC310<CO2Layout::ROW_MAJOR> fixpipeParams; // L0C->UB
|
||||
fixpipeParams.nSize = nSize; // N方向全部输出
|
||||
fixpipeParams.mSize = mSize / 2; // M方向每个AIV一半
|
||||
fixpipeParams.srcStride = srcStride;
|
||||
fixpipeParams.dstStride = UB_BANK_DEPTH_STRIDE / sizeof(QK_T); // 落到同一个bank
|
||||
fixpipeParams.params.ndNum = 1;
|
||||
fixpipeParams.params.srcNdStride = 0;
|
||||
fixpipeParams.params.dstNdStride = 0;
|
||||
fixpipeParams.dualDstCtl = 0;
|
||||
fixpipeParams.quantPre = F322BF16;
|
||||
fixpipeParams.reluEn = true; // ReLU激活
|
||||
fixpipeParams.subBlockId = 0;
|
||||
Fixpipe<QK_T, float, QLI_CFG_ROW_MAJOR_UB>(mm1ResUB_[(runInfo.loop % 2) * (UB_BANK_STRIDE / sizeof(QK_T))], // 未考虑s1gGmOffset和s2GmOffset
|
||||
cL0_[(l0BufIdx_ % L0_BUF_NUM) * L0C_BUFFER_OFFSET], fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
|
||||
fixpipeParams.subBlockId = 1;
|
||||
Fixpipe<QK_T, float, QLI_CFG_ROW_MAJOR_UB>(mm1ResUB_[(runInfo.loop % 2) * (UB_BANK_STRIDE / sizeof(QK_T))], // 未考虑s1gGmOffset和s2GmOffset
|
||||
cL0_[(l0BufIdx_ % L0_BUF_NUM) * L0C_BUFFER_OFFSET + mSize / 2 * 16], fixpipeParams); // 将matmul结果从L0C搬运到UB
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::AllocEventID()
|
||||
{
|
||||
SetMMLayoutTransform(true);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 1);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 2);
|
||||
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + 0);
|
||||
SetFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + 1);
|
||||
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 0);
|
||||
SetFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 1);
|
||||
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + 0);
|
||||
SetFlag<HardEvent::FIX_M>(FIX_M_EVENT + 1);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIMatmul<QLIT>::FreeEventID()
|
||||
{
|
||||
SetMMLayoutTransform(false);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 1);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(KEY_MTE1_MTE2_EVENT + 2);
|
||||
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + 0);
|
||||
WaitFlag<HardEvent::MTE1_MTE2>(QUERY_MTE1_MTE2_EVENT + 1);
|
||||
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 0);
|
||||
WaitFlag<HardEvent::M_MTE1>(M_MTE1_EVENT + 1);
|
||||
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + 0);
|
||||
WaitFlag<HardEvent::FIX_M>(FIX_M_EVENT + 1);
|
||||
}
|
||||
} // namespace QLIKernel
|
||||
#endif
|
||||
@@ -0,0 +1,519 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_service_vector.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef quant_lightning_indexer_SERVICE_VECTOR_H
|
||||
#define quant_lightning_indexer_SERVICE_VECTOR_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_operator_list_tensor_intf.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#include "lib/matrix/matmul/tiling.h"
|
||||
#include "quant_lightning_indexer_common.h"
|
||||
#include "../arch35/vf/quant_lightning_indexer_vector1.h"
|
||||
#include "../arch35/vf/quant_lightning_indexer_topk.h"
|
||||
|
||||
namespace QLIKernel {
|
||||
using namespace QLICommon;
|
||||
constexpr uint32_t TRUNK_LEN_16K = 16384;
|
||||
template <typename QLIT>
|
||||
class QLIVector {
|
||||
public:
|
||||
// =================================类型定义区=================================
|
||||
static constexpr LI_LAYOUT Q_LAYOUT_T = QLIT::layout;
|
||||
static constexpr LI_LAYOUT K_LAYOUT_T = QLIT::keyLayout;
|
||||
static constexpr bool PAGE_ATTENTION = QLIT::pageAttention;
|
||||
|
||||
using QK_T = typename QLIT::queryKeyType;
|
||||
using SCORE_T = typename QLIT::scoreType;
|
||||
|
||||
__aicore__ inline QLIVector(){};
|
||||
__aicore__ inline void ProcessVec1(const QLICommon::RunInfo &info);
|
||||
__aicore__ inline void ProcessTopK(const QLICommon::RunInfo &info);
|
||||
__aicore__ inline void InitBuffers(TPipe *pipe);
|
||||
__aicore__ inline void InitParams(const struct QLICommon::ConstInfo &constInfo,
|
||||
const QLITilingData *__restrict tilingData);
|
||||
__aicore__ inline void InitVecWorkspaceTensor(GlobalTensor<SCORE_T> scoreGm);
|
||||
__aicore__ inline void InitVecInputTensor(GlobalTensor<float> weightsGm, GlobalTensor<float> qScaleGm,
|
||||
GlobalTensor<float> kScaleGm, GlobalTensor<int32_t> indiceOutGm,
|
||||
GlobalTensor<int32_t> blockTableGm);
|
||||
__aicore__ inline void CleanInvalidOutput(int64_t invalidS1offset);
|
||||
__aicore__ inline void AllocEventID();
|
||||
__aicore__ inline void FreeEventID();
|
||||
|
||||
protected:
|
||||
GlobalTensor<SCORE_T> scoreGm;
|
||||
GlobalTensor<float> weightsGm;
|
||||
GlobalTensor<float> qScaleGm;
|
||||
GlobalTensor<float> kScaleGm;
|
||||
GlobalTensor<int32_t> indiceOutGm;
|
||||
GlobalTensor<int32_t> blockTableGm;
|
||||
// =================================常量区=================================
|
||||
static constexpr uint32_t VEC1_V_MTE2_EVENT = EVENT_ID0;
|
||||
static constexpr uint32_t VEC1_MTE2_V_EVENT = EVENT_ID1;
|
||||
static constexpr uint32_t VEC1_V_MTE3_EVENT = EVENT_ID2;
|
||||
static constexpr uint32_t VEC1_MTE3_V_EVENT = EVENT_ID3;
|
||||
|
||||
static constexpr uint32_t TOPK_V_MTE2_EVENT = EVENT_ID4;
|
||||
static constexpr uint32_t TOPK_MTE2_V_EVENT = EVENT_ID5;
|
||||
static constexpr uint32_t TOPK_V_MTE3_EVENT = EVENT_ID6;
|
||||
static constexpr uint32_t TOPK_MTE3_V_EVENT = EVENT_ID7;
|
||||
|
||||
static constexpr uint32_t KSCALE_S_MTE2_EVENT = EVENT_ID7;
|
||||
static constexpr uint32_t MTE3_MTE2_EVENT = EVENT_ID0;
|
||||
static constexpr uint32_t V_MTE2_EVENT = EVENT_ID7;
|
||||
static constexpr uint32_t V_MTE2_EVENT1 = EVENT_ID2;
|
||||
static constexpr uint32_t V_MTE2_EVENT2 = EVENT_ID3;
|
||||
static constexpr uint32_t V_MTE2_EVENT3 = EVENT_ID5;
|
||||
|
||||
private:
|
||||
__aicore__ inline void GetKeyScale(const QLICommon::RunInfo &runInfo, LocalTensor<float> &kScaleUB,
|
||||
int64_t batchId, int64_t startS2, int64_t getLen);
|
||||
// ================================Local Buffer区====================================
|
||||
|
||||
// tmp buff for vector
|
||||
TBuf<TPosition::VECCALC> resMm1Buf_;
|
||||
LocalTensor<QK_T> resMm1UB_;
|
||||
//tmp buff for weight
|
||||
TBuf<TPosition::VECCALC> weightBuf_;
|
||||
LocalTensor<float> weightUB_;
|
||||
//tmp buff for kScale
|
||||
TBuf<TPosition::VECCALC> kScaleBuf_;
|
||||
LocalTensor<float> kScaleUB_;
|
||||
//tmp buff for qScale
|
||||
TBuf<TPosition::VECCALC> qScaleBuf_;
|
||||
LocalTensor<float> qScaleUB_;
|
||||
//tmp buff for out
|
||||
TBuf<TPosition::VECCALC> outBuf_;
|
||||
LocalTensor<SCORE_T> vec1OutUB_;
|
||||
// tmp buff for LD
|
||||
|
||||
// tmp buff for topk
|
||||
TBuf<TPosition::VECCALC> mrgValueBuf_;
|
||||
LocalTensor<SCORE_T> mrgValueLocal_;
|
||||
|
||||
TBuf<TPosition::VECCALC> indicesOutBuf_;
|
||||
LocalTensor<uint32_t> indicesOutLocal_;
|
||||
|
||||
TBuf<TPosition::VECCALC> scoreOutBuf_;
|
||||
LocalTensor<SCORE_T> scoreOutLocal_;
|
||||
|
||||
TBuf<TPosition::VECCALC> topkSharedTmpBuf_;
|
||||
LocalTensor<uint32_t> topkSharedTmpLocal_;
|
||||
|
||||
TBuf<TPosition::VECCALC> outInvalidBuf_;
|
||||
LocalTensor<int32_t> outInvalidLocal_;
|
||||
|
||||
int32_t blockId_ = -1;
|
||||
// para for vector
|
||||
int32_t groupInner_ = 0;
|
||||
int32_t globalTopkNum_ = 0;
|
||||
int64_t blockS2StartIdx_ = 0;
|
||||
int32_t gSize_ = 0;
|
||||
int32_t kSeqSize_ = 0;
|
||||
int32_t kHeadNum_ = 0;
|
||||
int32_t qHeadNum_ = 0;
|
||||
int32_t s1BaseSize_ = 0;
|
||||
int32_t s2BaseSize_ = 0;
|
||||
int32_t kCacheBlockSize_ = 0;
|
||||
int32_t maxBlockNumPerBatch_ = 0;
|
||||
uint32_t topkCount_ = 0;
|
||||
uint32_t topkCountAlign256_ = 0; // topkCount对齐到256(直方图需要),支持topk泛化
|
||||
uint32_t trunkLen_ = 0;
|
||||
|
||||
struct QLICommon::ConstInfo constInfo_;
|
||||
topk::LITopk<SCORE_T> topkOp_;
|
||||
};
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitBuffers(TPipe *pipe)
|
||||
{
|
||||
pipe->InitBuffer(resMm1Buf_, 2 * CeilDiv(constInfo_.mBaseSize, 2) * s2BaseSize_ * sizeof(QK_T)); //大小:2(开dB) * 2 * 64 * 128 * 4 = 128KB
|
||||
resMm1UB_ = resMm1Buf_.Get<QK_T>();//qk
|
||||
pipe->InitBuffer(weightBuf_, 2 * CeilDiv(s1BaseSize_, 2) * gSize_* sizeof(float)); // 大小:2(开dB) * 2 * 64 * 2 = 0.5KB
|
||||
weightUB_ = weightBuf_.Get<float>();//weight
|
||||
pipe->InitBuffer(kScaleBuf_, 2 * s2BaseSize_ * sizeof(float)); // 大小:2(开dB) * 128 * 4 = 1KB
|
||||
kScaleUB_ = kScaleBuf_.Get<float>();//kScale
|
||||
pipe->InitBuffer(qScaleBuf_, 2 * CeilDiv(s1BaseSize_, 2) * gSize_* sizeof(float)); // 大小:2(开dB) * 2 * 64 * 4 = 1KB
|
||||
qScaleUB_ = qScaleBuf_.Get<float>();//qScale
|
||||
pipe->InitBuffer(outBuf_, 2 * CeilDiv(s1BaseSize_, 2) * s2BaseSize_ * sizeof(SCORE_T)); // 大小:2(开dB) * 2 * 128 * 4 = 2KB
|
||||
vec1OutUB_ = outBuf_.Get<SCORE_T>();//out
|
||||
|
||||
// Topk
|
||||
pipe->InitBuffer(mrgValueBuf_, (topkCountAlign256_ + trunkLen_) * sizeof(SCORE_T)); // 大小:(topkCountAlign256_ + 每次排序长度) * sizeof(SCORE_T)
|
||||
mrgValueLocal_ = mrgValueBuf_.Get<SCORE_T>();
|
||||
|
||||
pipe->InitBuffer(indicesOutBuf_, (topkCountAlign256_ + 64) * sizeof(uint32_t)); // 大小:(topkCountAlign256_ + 64) * 4 64:duplicate刷-1需要额外空间
|
||||
indicesOutLocal_ = indicesOutBuf_.Get<uint32_t>();
|
||||
|
||||
pipe->InitBuffer(scoreOutBuf_, topkCountAlign256_ * sizeof(SCORE_T)); // 大小:topkCountAlign256_ * sizeof(SCORE_T)
|
||||
scoreOutLocal_ = scoreOutBuf_.Get<SCORE_T>();
|
||||
|
||||
uint64_t topkSharedTmpSize = topkOp_.GetSharedTmpBufferSize();
|
||||
pipe->InitBuffer(topkSharedTmpBuf_, topkSharedTmpSize);
|
||||
topkSharedTmpLocal_ = topkSharedTmpBuf_.Get<uint32_t>();
|
||||
topkOp_.InitBuffers(topkSharedTmpLocal_);
|
||||
|
||||
//刷-1
|
||||
pipe->InitBuffer(outInvalidBuf_, topkCount_ * sizeof(int32_t));
|
||||
outInvalidLocal_ = outInvalidBuf_.Get<int32_t>();
|
||||
Duplicate(kScaleUB_, float(0), 2 * s2BaseSize_);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitParams(const struct QLICommon::ConstInfo &constInfo,
|
||||
const QLITilingData *__restrict tilingData)
|
||||
{
|
||||
this->constInfo_ = constInfo;
|
||||
blockS2StartIdx_ = 0;
|
||||
gSize_ = constInfo.gSize;
|
||||
kSeqSize_ = constInfo.kSeqSize;
|
||||
// define N2 para
|
||||
kHeadNum_ = constInfo.kHeadNum;
|
||||
qHeadNum_ = constInfo.qHeadNum;
|
||||
// define MMBase para
|
||||
s1BaseSize_ = constInfo.s1BaseSize; // 4
|
||||
s2BaseSize_ = constInfo.s2BaseSize; // 128
|
||||
kCacheBlockSize_ = constInfo.kCacheBlockSize;
|
||||
maxBlockNumPerBatch_ = constInfo.maxBlockNumPerBatch;
|
||||
blockId_ = GetBlockIdx();
|
||||
trunkLen_ = TRUNK_LEN_16K;
|
||||
topkCount_ = constInfo.sparseCount;
|
||||
topkCountAlign256_ = QLICommon::Align(constInfo.sparseCount, (uint64_t)256); // topkCount对齐到256
|
||||
topkOp_.Init(topkCount_, trunkLen_);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitVecInputTensor(GlobalTensor<float> weightsGm, GlobalTensor<float> qScaleGm,
|
||||
GlobalTensor<float> kScaleGm,
|
||||
GlobalTensor<int32_t> indiceOutGm,
|
||||
GlobalTensor<int32_t> blockTableGm)
|
||||
{
|
||||
this->weightsGm = weightsGm;
|
||||
this->qScaleGm = qScaleGm;
|
||||
this->kScaleGm = kScaleGm;
|
||||
this->indiceOutGm = indiceOutGm;
|
||||
this->blockTableGm = blockTableGm;
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::InitVecWorkspaceTensor(GlobalTensor<SCORE_T> scoreGm)
|
||||
{
|
||||
this->scoreGm = scoreGm;//resucesum*k
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::AllocEventID()
|
||||
{
|
||||
SetFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + 0);
|
||||
SetFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + 1);
|
||||
SetFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + 0);
|
||||
SetFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + 1);
|
||||
|
||||
SetFlag<HardEvent::V_MTE2>(TOPK_V_MTE2_EVENT);
|
||||
SetFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT1);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::FreeEventID()
|
||||
{
|
||||
WaitFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + 0);
|
||||
WaitFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + 1);
|
||||
WaitFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + 0);
|
||||
WaitFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + 1);
|
||||
|
||||
WaitFlag<HardEvent::V_MTE2>(TOPK_V_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT1);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::CleanInvalidOutput(int64_t invalidS1Offset)
|
||||
{
|
||||
// init -1 and copy to output
|
||||
Duplicate(outInvalidLocal_, constInfo_.INVALID_IDX, constInfo_.sparseCount);
|
||||
|
||||
SetFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
|
||||
AscendC::DataCopyParams dataCopyOutParams;
|
||||
dataCopyOutParams.blockCount = 1;
|
||||
dataCopyOutParams.blockLen = constInfo_.sparseCount * sizeof(int32_t);
|
||||
dataCopyOutParams.srcStride = 0;
|
||||
dataCopyOutParams.dstStride = 0;
|
||||
AscendC::DataCopyPad(indiceOutGm[invalidS1Offset], outInvalidLocal_, dataCopyOutParams);
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::GetKeyScale(const QLICommon::RunInfo &runInfo, LocalTensor<float> &kScaleUB,
|
||||
int64_t batchId, int64_t startS2, int64_t getLen)
|
||||
{
|
||||
// startS2一定能整除kCacheBlockSize_
|
||||
AscendC::DataCopyPadExtParams<float> padParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyExtParams copyInParams;
|
||||
if constexpr (PAGE_ATTENTION) {
|
||||
int32_t startBlockTableIdx = startS2 / kCacheBlockSize_;
|
||||
int32_t startBlockTableOffset = startS2 % kCacheBlockSize_;
|
||||
int32_t blockTableBatchOffset = batchId * maxBlockNumPerBatch_;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
int32_t resUbBaseOffset = 0;
|
||||
if (startBlockTableOffset > 0) {
|
||||
int32_t firstPartLen =
|
||||
kCacheBlockSize_ - startBlockTableOffset > getLen ? getLen : kCacheBlockSize_ - startBlockTableOffset;
|
||||
copyInParams.blockLen = firstPartLen * sizeof(float);
|
||||
int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx);
|
||||
SetFlag<HardEvent::S_MTE2>(KSCALE_S_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::S_MTE2>(KSCALE_S_MTE2_EVENT);
|
||||
AscendC::DataCopyPad(kScaleUB[(runInfo.loop % 2) * s2BaseSize_],
|
||||
kScaleGm[blockId * constInfo_.scaleStride + startBlockTableOffset],
|
||||
copyInParams, padParams);
|
||||
startBlockTableIdx++;
|
||||
getLen = getLen - firstPartLen;
|
||||
resUbBaseOffset = firstPartLen;
|
||||
}
|
||||
int32_t getLoopNum = CeilDiv(getLen, kCacheBlockSize_);
|
||||
copyInParams.blockLen = kCacheBlockSize_ * sizeof(float);
|
||||
for (int32_t i = 0; i < getLoopNum; i++) {
|
||||
if (i == getLoopNum - 1) {
|
||||
copyInParams.blockLen = (getLen - i * kCacheBlockSize_) * sizeof(float);
|
||||
}
|
||||
int32_t blockId = blockTableGm.GetValue(blockTableBatchOffset + startBlockTableIdx + i);
|
||||
SetFlag<HardEvent::S_MTE2>(KSCALE_S_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::S_MTE2>(KSCALE_S_MTE2_EVENT);
|
||||
AscendC::DataCopyPad(kScaleUB[(runInfo.loop % 2) * s2BaseSize_ + resUbBaseOffset + i * kCacheBlockSize_],
|
||||
kScaleGm[blockId * constInfo_.scaleStride],
|
||||
copyInParams, padParams);
|
||||
}
|
||||
} else {
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.blockLen = getLen * sizeof(float);
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
AscendC::DataCopyPad(kScaleUB[(runInfo.loop % 2) * s2BaseSize_], kScaleGm[runInfo.tensorKeyScaleOffset], copyInParams, padParams);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::ProcessVec1(const QLICommon::RunInfo &info)
|
||||
{
|
||||
auto pingpong = (info.loop % 2);
|
||||
auto s1BaseSizePerAIV = CeilDiv(s1BaseSize_, 2);
|
||||
int64_t curS1Idx = info.gS1Idx * s1BaseSize_;
|
||||
int64_t curS2Idx = info.s2Idx * s2BaseSize_;
|
||||
int64_t curS1ProcNum = curS1Idx + s1BaseSize_ > info.actS1Size ? info.actS1Size % s1BaseSize_ : s1BaseSize_;
|
||||
int64_t curAivS1Idx = curS1Idx + (blockId_ % 2) * CeilDiv(curS1ProcNum, 2);
|
||||
int64_t curAivS1ProcNum = (blockId_ % 2 == 0) ? CeilDiv(curS1ProcNum, 2) : curS1ProcNum / 2;
|
||||
|
||||
if (curAivS1ProcNum == 0) {
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_CV_EVENT + pingpong); // V核等C核计算完mm1,mm1Res已搬运到UB
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_VC_EVENT + pingpong); // V核处理完,通知C核可以把mm1Res搬运到UB
|
||||
return;
|
||||
}
|
||||
WaitFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + pingpong);
|
||||
//weightsGm --> weightUB_
|
||||
int64_t weightGmOffset = info.tensorWeightsOffset + curAivS1Idx * kHeadNum_ * gSize_;
|
||||
DataCopyPadExtParams<float> padWeightsParams{false, 0, 0, 0};
|
||||
DataCopyExtParams qwDataCopyExtParams;
|
||||
qwDataCopyExtParams.blockCount = curAivS1ProcNum;
|
||||
qwDataCopyExtParams.blockLen = gSize_ * sizeof(float);
|
||||
qwDataCopyExtParams.srcStride = 0;
|
||||
qwDataCopyExtParams.dstStride = (UB_BANK_DEPTH_STRIDE - UB_BANK_STRIDE) / 32;
|
||||
DataCopyPad(weightUB_[pingpong * (UB_BANK_STRIDE / sizeof(float))],
|
||||
weightsGm[weightGmOffset], qwDataCopyExtParams, padWeightsParams);
|
||||
|
||||
//qScaleGm --> qScaleUB_
|
||||
DataCopyPadExtParams<float> padQScaleParams{false, 0, 0, 0};
|
||||
DataCopyPad(qScaleUB_[pingpong * (UB_BANK_STRIDE / sizeof(float))],
|
||||
qScaleGm[weightGmOffset], qwDataCopyExtParams, padQScaleParams);
|
||||
|
||||
//kScaleGm --> kScaleUB_
|
||||
GetKeyScale(info, kScaleUB_, info.bIdx, curS2Idx, info.actualSingleProcessSInnerSize);
|
||||
SetFlag<HardEvent::MTE2_V>(VEC1_MTE2_V_EVENT + pingpong);
|
||||
WaitFlag<HardEvent::MTE2_V>(VEC1_MTE2_V_EVENT + pingpong);
|
||||
WaitFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + pingpong);
|
||||
|
||||
//CV同步
|
||||
CrossCoreWaitFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_CV_EVENT + info.loop % 2); //V核等C核计算完mm1,mm1Res已搬运到UB
|
||||
|
||||
static_assert(std::is_same_v<SCORE_T, uint16_t>);
|
||||
auto outBase = vec1OutUB_[pingpong * (UB_BANK_STRIDE / sizeof(SCORE_T))];
|
||||
auto weightBase = weightUB_[pingpong * (UB_BANK_STRIDE / sizeof(float))];
|
||||
auto qScaleBase = qScaleUB_[pingpong * (UB_BANK_STRIDE / sizeof(float))];
|
||||
auto kScaleBase = kScaleUB_[pingpong * s2BaseSize_];
|
||||
|
||||
auto qkBase = resMm1UB_[pingpong * (UB_BANK_STRIDE / sizeof(QK_T))];
|
||||
auto qkVLstride = (UB_BANK_DEPTH_STRIDE / sizeof(QK_T)) / 2 * constInfo_.mBaseSize;
|
||||
vector1::BatchMulWeightAndReduceSum(outBase, UB_BANK_DEPTH_STRIDE / sizeof(SCORE_T),
|
||||
qkBase, qkVLstride, (uint32_t)(gSize_ * UB_BANK_DEPTH_STRIDE / sizeof(QK_T)),
|
||||
weightBase, UB_BANK_DEPTH_STRIDE / sizeof(float),
|
||||
kScaleBase, (uint32_t)0,
|
||||
qScaleBase, UB_BANK_DEPTH_STRIDE / sizeof(float),
|
||||
gSize_, curAivS1ProcNum);
|
||||
SetFlag<HardEvent::V_MTE2>(VEC1_V_MTE2_EVENT + pingpong);
|
||||
SetFlag<HardEvent::V_MTE3>(VEC1_V_MTE3_EVENT + pingpong);
|
||||
WaitFlag<HardEvent::V_MTE3>(VEC1_V_MTE3_EVENT + pingpong);
|
||||
//outUB_ ---> scoreGm
|
||||
int64_t vec1OutGmOffset = blockId_ % 2 == 0 ? curS2Idx :
|
||||
s1BaseSizePerAIV * QLICommon::Align((uint64_t)constInfo_.kSeqSize, (uint64_t)s2BaseSize_) + curS2Idx;
|
||||
DataCopyExtParams copyOutParams;
|
||||
copyOutParams.blockCount = curAivS1ProcNum;
|
||||
copyOutParams.blockLen = s2BaseSize_ * sizeof(SCORE_T);
|
||||
copyOutParams.srcStride = (UB_BANK_DEPTH_STRIDE - UB_BANK_STRIDE) / 32;
|
||||
copyOutParams.dstStride = (QLICommon::Align((uint64_t)constInfo_.kSeqSize, (uint64_t)s2BaseSize_) - s2BaseSize_) * sizeof(SCORE_T);
|
||||
DataCopyPad(scoreGm[vec1OutGmOffset], outBase, copyOutParams);
|
||||
SetFlag<HardEvent::MTE3_V>(VEC1_MTE3_V_EVENT + pingpong);
|
||||
CrossCoreSetFlag<QLICommon::ConstInfo::QLI_SYNC_MODE4, PIPE_V>(QLICommon::ConstInfo::CROSS_VC_EVENT + pingpong); //V核处理完,通知C核可以把mm1Res搬运到UB
|
||||
}
|
||||
|
||||
template <typename QLIT>
|
||||
__aicore__ inline void QLIVector<QLIT>::ProcessTopK(const QLICommon::RunInfo &info)
|
||||
{
|
||||
SetFlag<HardEvent::MTE3_MTE2>(MTE3_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(MTE3_MTE2_EVENT);
|
||||
|
||||
int64_t curS1Idx = info.gS1Idx * s1BaseSize_;
|
||||
int64_t curS2Idx = info.s2Idx * s2BaseSize_;
|
||||
int64_t curS1ProcNum = curS1Idx + s1BaseSize_ > info.actS1Size ? info.actS1Size % s1BaseSize_ : s1BaseSize_;
|
||||
int64_t curAivS1Idx = curS1Idx + (blockId_ % 2) * CeilDiv(curS1ProcNum, 2);
|
||||
int64_t curAivS1ProcNum = (blockId_ % 2 == 0) ? CeilDiv(curS1ProcNum, 2) : curS1ProcNum / 2;
|
||||
|
||||
AscendC::DataCopyExtParams copyInParams;
|
||||
copyInParams.blockCount = 1;
|
||||
copyInParams.srcStride = 0;
|
||||
copyInParams.dstStride = 0;
|
||||
copyInParams.rsv = 0;
|
||||
|
||||
AscendC::DataCopyParams copyOutParams;
|
||||
copyOutParams.blockCount = 1;
|
||||
copyOutParams.blockLen = topkCount_ * sizeof(uint32_t); // bytes
|
||||
copyOutParams.srcStride = 0;
|
||||
copyOutParams.dstStride = 0;
|
||||
|
||||
int32_t cuRealAcSeq = info.actS2Size;
|
||||
if (constInfo_.attenMaskFlag) {
|
||||
cuRealAcSeq = info.actS2SizeOrig - info.actS1Size + curAivS1Idx + 1;
|
||||
}
|
||||
|
||||
int32_t validS2Len = cuRealAcSeq;
|
||||
for (uint32_t i = 0; i < curAivS1ProcNum; i++) {
|
||||
uint32_t rowIdx = blockId_ % 2 * CeilDiv(curS1ProcNum, 2) + i;
|
||||
uint32_t vecOffset = blockId_ % 2 * CeilDiv(s1BaseSize_, 2) + i;
|
||||
|
||||
SCORE_T zero = 0;
|
||||
int32_t neg = -1;
|
||||
if (constInfo_.attenMaskFlag) {
|
||||
validS2Len = ((int32_t)i + cuRealAcSeq) / static_cast<int32_t>(constInfo_.cmpRatio);
|
||||
}
|
||||
if (validS2Len <= 0) {
|
||||
WaitFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
Duplicate(indicesOutLocal_.ReinterpretCast<int32_t>(), neg, topkCount_);
|
||||
SetFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
AscendC::DataCopyPad(indiceOutGm[info.indiceOutOffset + (curS1Idx + rowIdx) * topkCount_], indicesOutLocal_.ReinterpretCast<int32_t>(), copyOutParams);
|
||||
SetFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
continue;
|
||||
}
|
||||
|
||||
WaitFlag<HardEvent::V_MTE2>(TOPK_V_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
|
||||
AscendC::DataCopyPadExtParams<SCORE_T> padParams{true, 0, 0, 0};
|
||||
if (validS2Len >= topkCount_) {
|
||||
uint32_t s2LoopNum = (validS2Len + trunkLen_ - 1) / trunkLen_;
|
||||
if (s2LoopNum == 1) {
|
||||
uint32_t validS2LenAlign = QLICommon::Align(validS2Len, (int32_t)256);
|
||||
Duplicate(mrgValueLocal_[validS2Len / 256 * 256], zero, validS2LenAlign - validS2Len / 256 * 256);
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT);
|
||||
copyInParams.blockLen = validS2Len * sizeof(SCORE_T); // byte
|
||||
AscendC::DataCopyPadExtParams<SCORE_T> padParams{true, 0, 0, 0};
|
||||
AscendC::DataCopyPad(mrgValueLocal_, scoreGm[vecOffset * QLICommon::Align((uint64_t)constInfo_.kSeqSize, (uint64_t)s2BaseSize_)], copyInParams, padParams);
|
||||
SetFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
topkOp_(mrgValueLocal_, indicesOutLocal_, scoreOutLocal_, validS2LenAlign, 0, 1);
|
||||
} else {
|
||||
for (uint32_t loopIdx = 0; loopIdx < s2LoopNum; loopIdx++) {
|
||||
if (loopIdx == 0) {
|
||||
copyInParams.blockLen = trunkLen_ * sizeof(SCORE_T); // byte
|
||||
AscendC::DataCopyPad(mrgValueLocal_, scoreGm[vecOffset * QLICommon::Align((uint64_t)constInfo_.kSeqSize, (uint64_t)s2BaseSize_)], copyInParams, padParams);
|
||||
SetFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
topkOp_(mrgValueLocal_, indicesOutLocal_, scoreOutLocal_, trunkLen_, loopIdx, s2LoopNum);
|
||||
continue;
|
||||
}
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT2);
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT2);
|
||||
uint32_t validTrunkLen = (loopIdx * trunkLen_ + trunkLen_) > validS2Len ? validS2Len % trunkLen_ : trunkLen_;
|
||||
uint32_t offset = vecOffset * QLICommon::Align((uint64_t)constInfo_.kSeqSize, (uint64_t)s2BaseSize_) + loopIdx * trunkLen_;
|
||||
AscendC::DataCopy(mrgValueLocal_, scoreOutLocal_, topkCountAlign256_);
|
||||
// topk如果没有对齐到256,则把topkCountAlign256_ - topkCount_部分刷0
|
||||
if (topkCountAlign256_ != topkCount_) {
|
||||
uint64_t mask[1];
|
||||
mask[0] = ~0;
|
||||
mask[0] = mask[0] << (topkCount_ % 64);
|
||||
PipeBarrier<PIPE_V>();
|
||||
// 把topkCount_对齐到64刷0,此处由于duplicate的限制mask[0]刷64个数
|
||||
Duplicate(mrgValueLocal_[topkCount_ / 64 * 64], zero, mask, 1, 1, 0);
|
||||
PipeBarrier<PIPE_V>();
|
||||
// 把topk剩余对齐到256的部分刷0
|
||||
Duplicate(mrgValueLocal_[topkCount_ / 64 * 64 + 64], zero, topkCountAlign256_ - (topkCount_ / 64 * 64 + 64));
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT3);
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT3);
|
||||
}
|
||||
copyInParams.blockLen = validTrunkLen * sizeof(SCORE_T); // byte
|
||||
// TOPK 直方图一次必须计算256,输入处理数据需要和256对齐
|
||||
if ((topkCountAlign256_ + validTrunkLen) % 256 != 0) {
|
||||
Duplicate(mrgValueLocal_[topkCountAlign256_ + validTrunkLen / 256 * 256], zero, QLICommon::Align(validTrunkLen, (uint32_t)256) - validTrunkLen / 256 * 256);
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT);
|
||||
}
|
||||
WaitFlag<HardEvent::V_MTE2>(V_MTE2_EVENT1);
|
||||
AscendC::DataCopyPad(mrgValueLocal_[topkCountAlign256_], scoreGm[offset], copyInParams, padParams);
|
||||
SetFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
WaitFlag<HardEvent::MTE2_V>(TOPK_MTE2_V_EVENT);
|
||||
topkOp_(mrgValueLocal_, indicesOutLocal_, scoreOutLocal_, QLICommon::Align(topkCountAlign256_ + validTrunkLen, (uint32_t)256), loopIdx, s2LoopNum);
|
||||
SetFlag<HardEvent::V_MTE2>(V_MTE2_EVENT1);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
AscendC::CreateVecIndex(indicesOutLocal_.ReinterpretCast<int32_t>(), (int32_t)zero, validS2Len);
|
||||
}
|
||||
|
||||
if (validS2Len < topkCount_) {
|
||||
uint64_t mask[1];
|
||||
mask[0] = ~0;
|
||||
mask[0] = mask[0] << (validS2Len % 8);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Duplicate(indicesOutLocal_.ReinterpretCast<int32_t>()[validS2Len / 8 * 8], neg, mask, 1, 1, 0);
|
||||
}
|
||||
|
||||
if (validS2Len / 8 * 8 + 64 < topkCount_) {
|
||||
PipeBarrier<PIPE_V>();
|
||||
Duplicate(indicesOutLocal_.ReinterpretCast<int32_t>()[validS2Len / 8 * 8 + 64], neg, topkCount_ - (validS2Len / 8 * 8 + 64));
|
||||
}
|
||||
|
||||
SetFlag<HardEvent::V_MTE2>(TOPK_V_MTE2_EVENT);
|
||||
SetFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
WaitFlag<HardEvent::V_MTE3>(TOPK_V_MTE3_EVENT);
|
||||
AscendC::DataCopyPad(indiceOutGm[info.indiceOutOffset + (curS1Idx + rowIdx) * topkCount_], indicesOutLocal_.ReinterpretCast<int32_t>(), copyOutParams);
|
||||
SetFlag<HardEvent::MTE3_V>(TOPK_MTE3_V_EVENT);
|
||||
}
|
||||
}
|
||||
} // namespace QLIKernel
|
||||
#endif
|
||||
@@ -0,0 +1,165 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_topk.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef quant_lightning_indexer_TOPK_H
|
||||
#define quant_lightning_indexer_TOPK_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "vf_topk.h"
|
||||
#include "vf_topk_16_gather.h"
|
||||
|
||||
namespace topk {
|
||||
template<typename T>
|
||||
class LITopk
|
||||
{
|
||||
public:
|
||||
__aicore__ inline void operator()(LocalTensor<uint32_t>& outputIdxLocal,
|
||||
LocalTensor<T>& inputLocal,
|
||||
uint32_t s2SeqLen)
|
||||
{
|
||||
}
|
||||
};
|
||||
|
||||
template<>
|
||||
class LITopk<uint32_t> {
|
||||
public:
|
||||
static __aicore__ inline uint32_t GetSharedTmpBufferSize(uint32_t topK)
|
||||
{
|
||||
return 2 * topK * sizeof(uint32_t) + 5 * 256 * sizeof(uint32_t) + 64 * sizeof(uint32_t) +
|
||||
(topK + 64) * sizeof(uint32_t); // for output value tensor
|
||||
}
|
||||
|
||||
static __aicore__ inline uint32_t GetIndexBufferSize(uint32_t topK)
|
||||
{
|
||||
return (topK + 64) * sizeof(uint32_t);
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(uint32_t topK)
|
||||
{
|
||||
this->topK = topK;
|
||||
}
|
||||
|
||||
__aicore__ inline void InitBuffers(LocalTensor<uint32_t>& sharedTmpBuffer)
|
||||
{
|
||||
tmpIdxLocal = sharedTmpBuffer[0];
|
||||
tmpValueLocal = tmpIdxLocal[topK];
|
||||
histogramsLocal = tmpValueLocal[topK];
|
||||
idx0Local = histogramsLocal[256];
|
||||
idx1Local = idx0Local[256];
|
||||
idx2Local = idx1Local[256];
|
||||
idx3Local = idx2Local[256];
|
||||
nkValueLocal = idx3Local[256];
|
||||
outputValueLocal = nkValueLocal[64];
|
||||
}
|
||||
|
||||
__aicore__ inline void operator()(LocalTensor<uint32_t>& outputIdxLocal,
|
||||
LocalTensor<uint32_t>& inputLocal,
|
||||
uint32_t s2SeqLen)
|
||||
{
|
||||
topkb32::LiTopKVF(outputIdxLocal, // filter阶段使用输出value Buf topK * 4B
|
||||
outputValueLocal, // filter阶段使用输出 Idx Buf topK * 4B
|
||||
inputLocal, // 输入 s2SeqLen * 4B
|
||||
tmpIdxLocal, // filter阶段使用暂存index Buf topK * 4B
|
||||
tmpValueLocal, // filter阶段使用暂存value Buf topK * 4B
|
||||
histogramsLocal, // 直方图的临时Buf 256 * 4B
|
||||
idx0Local, // 输入数据第1个8位Buf 256 * 4B
|
||||
idx1Local, // 输入数据第2个8位Buf 256 * 4B
|
||||
idx2Local, // 输入数据第3个8位Buf 256 * 4B
|
||||
idx3Local, // 输入数据第4个8位Buf 256 * 4B
|
||||
nkValueLocal, // next_k 暂存Buf 64 * 4B
|
||||
topK, // topk数量
|
||||
s2SeqLen); // 输入元素总数
|
||||
}
|
||||
private:
|
||||
LocalTensor<uint32_t> tmpIdxLocal; // filter阶段使用暂存index Buf topK * 4B
|
||||
LocalTensor<uint32_t> tmpValueLocal; // filter阶段使用暂存value Buf topK * 4B
|
||||
LocalTensor<uint32_t> histogramsLocal; // 直方图的临时Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idx0Local; // 输入数据第1个8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idx1Local; // 输入数据第2个8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idx2Local; // 输入数据第3个8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idx3Local; // 输入数据第4个8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> nkValueLocal; // next_k 暂存Buf 64 * 4B
|
||||
LocalTensor<uint32_t> outputValueLocal; // 输出value tensor
|
||||
uint32_t topK;
|
||||
};
|
||||
|
||||
template<>
|
||||
class LITopk<uint16_t> {
|
||||
public:
|
||||
__aicore__ inline uint32_t GetSharedTmpBufferSize()
|
||||
{
|
||||
// 2 * QLICommon::Align(topK, (uint32_t)256):两块hisIndexLocal;3 * 256:histogramsLocal idxHighLocal idxLowLocal;64:nkValueLocal
|
||||
uint64_t bufferSize1 = (2 * QLICommon::Align(topK, (uint32_t)256) + 3 * 256 + 64) * sizeof(uint32_t);
|
||||
// QLICommon::Align(topK, (uint32_t)256) + trunkLen:tmpIndexLocal
|
||||
uint64_t bufferSize2 = (QLICommon::Align(topK, (uint32_t)256) + trunkLen) * sizeof(uint16_t);
|
||||
return bufferSize1 + bufferSize2;
|
||||
}
|
||||
|
||||
__aicore__ inline void Init(uint32_t topK, uint32_t trunkLen)
|
||||
{
|
||||
this->topK = topK;
|
||||
this->trunkLen = trunkLen;
|
||||
}
|
||||
|
||||
__aicore__ inline void InitBuffers(LocalTensor<uint32_t>& sharedTmpBuffer)
|
||||
{
|
||||
LocalTensor<uint32_t> hisIndexLocal1 = sharedTmpBuffer[0];
|
||||
LocalTensor<uint32_t> hisIndexLocal2 = hisIndexLocal1[QLICommon::Align(topK, (uint32_t)256)];
|
||||
hisIndexLocal[0] = hisIndexLocal1;
|
||||
hisIndexLocal[1] = hisIndexLocal2;
|
||||
histogramsLocal = hisIndexLocal2[QLICommon::Align(topK, (uint32_t)256)];
|
||||
idxHighLocal = histogramsLocal[256];
|
||||
idxLowLocal = idxHighLocal[256];
|
||||
nkValueLocal = idxLowLocal[256];
|
||||
LocalTensor<uint32_t> tmpIndexLocalTmp = nkValueLocal[64];
|
||||
tmpIndexLocal = tmpIndexLocalTmp.template ReinterpretCast<uint16_t>();
|
||||
}
|
||||
|
||||
__aicore__ inline void operator()(LocalTensor<uint16_t>& mrgValueLocal, LocalTensor<uint32_t>& indicesOutLocal,
|
||||
LocalTensor<uint16_t>& hisValueLocal, uint32_t s2SeqLen, uint32_t loopIdx, uint32_t s2LoopNum)
|
||||
{
|
||||
if (s2LoopNum == 1) {
|
||||
topkb16gather::LiTopKVF<false>(tmpIndexLocal, hisValueLocal, mrgValueLocal, histogramsLocal, idxHighLocal, idxLowLocal, nkValueLocal, topK, s2SeqLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(indicesOutLocal, tmpIndexLocal, RoundMode::CAST_NONE, topK);
|
||||
return;
|
||||
}
|
||||
|
||||
if (loopIdx == 0) {
|
||||
topkb16gather::LiTopKVF<true>(tmpIndexLocal, hisValueLocal, mrgValueLocal, histogramsLocal, idxHighLocal, idxLowLocal, nkValueLocal, topK, s2SeqLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Cast(hisIndexLocal[(loopIdx + 1) % 2], tmpIndexLocal, RoundMode::CAST_NONE, topK);
|
||||
} else {
|
||||
topkb16gather::LiTopKVF<true>(tmpIndexLocal, hisValueLocal, mrgValueLocal, histogramsLocal, idxHighLocal, idxLowLocal, nkValueLocal, topK, s2SeqLen);
|
||||
PipeBarrier<PIPE_V>();
|
||||
topkb16gather::LiTopKGatherVF(hisIndexLocal[(loopIdx + 1) % 2], hisValueLocal, mrgValueLocal, tmpIndexLocal, hisIndexLocal[loopIdx % 2],
|
||||
topK, loopIdx * trunkLen - QLICommon::Align(topK, (uint32_t)256), s2SeqLen);
|
||||
if (loopIdx == s2LoopNum - 1) {
|
||||
PipeBarrier<PIPE_V>();
|
||||
AscendC::DataCopy(indicesOutLocal, hisIndexLocal[(loopIdx + 1) % 2], QLICommon::Align(topK, (uint32_t)256));
|
||||
}
|
||||
}
|
||||
}
|
||||
private:
|
||||
LocalTensor<uint32_t> hisIndexLocal[2]; // 每trunkLen长度的s2选出的topK个索引
|
||||
LocalTensor<uint32_t> histogramsLocal; // 直方图的临时Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idxHighLocal; // 输入数据高8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> idxLowLocal; // 输入数据低8位Buf 256 * 4B
|
||||
LocalTensor<uint32_t> nkValueLocal; // next_k 暂存Buf 64 * 4B
|
||||
LocalTensor<uint16_t> tmpIndexLocal; // 每trunkLen + topK的临时index
|
||||
uint32_t topK = 512;
|
||||
uint32_t trunkLen = 16384;
|
||||
};
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,614 @@
|
||||
/**
|
||||
* 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 quant_lightning_indexer_vector1.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef quant_lightning_indexer_VECTOR1_H
|
||||
#define quant_lightning_indexer_VECTOR1_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace vector1 {
|
||||
|
||||
template <typename T>
|
||||
struct FloatSortTraits;
|
||||
|
||||
// fp32
|
||||
template <>
|
||||
struct FloatSortTraits<float> {
|
||||
using UInt = uint32_t;
|
||||
static constexpr UInt ZERO = 0x00000000;
|
||||
static constexpr UInt SIGN_MASK = 0x80000000;
|
||||
static constexpr UInt NAN_MASK = 0x7FC00000;
|
||||
static constexpr UInt ALL_ONE = 0xFFFFFFFF;
|
||||
};
|
||||
|
||||
// bf16
|
||||
template <>
|
||||
struct FloatSortTraits<bfloat16_t> {
|
||||
using UInt = uint16_t;
|
||||
static constexpr UInt ZERO = 0x0000;
|
||||
static constexpr UInt SIGN_MASK = 0x8000;
|
||||
static constexpr UInt NAN_MASK = 0x7FC0;
|
||||
static constexpr UInt ALL_ONE = 0xFFFF;
|
||||
};
|
||||
|
||||
|
||||
template <typename FloatT>
|
||||
struct FloatSortConstCtx {
|
||||
using Traits = FloatSortTraits<FloatT>;
|
||||
using UInt = typename Traits::UInt;
|
||||
AscendC::MicroAPI::RegTensor<UInt> zeros;
|
||||
AscendC::MicroAPI::RegTensor<UInt> allOnes;
|
||||
AscendC::MicroAPI::RegTensor<UInt> signMask;
|
||||
AscendC::MicroAPI::RegTensor<UInt> nan;
|
||||
};
|
||||
|
||||
|
||||
template <typename FloatT>
|
||||
__simd_callee__ inline void InitFloatSortConstCtx(FloatSortConstCtx<FloatT>& ctx, AscendC::MicroAPI::MaskReg& maskAll)
|
||||
{
|
||||
using Traits = FloatSortTraits<FloatT>;
|
||||
AscendC::MicroAPI::Duplicate(ctx.zeros, Traits::ZERO, maskAll);
|
||||
AscendC::MicroAPI::Duplicate(ctx.allOnes, Traits::ALL_ONE, maskAll);
|
||||
AscendC::MicroAPI::Duplicate(ctx.signMask, Traits::SIGN_MASK, maskAll);
|
||||
AscendC::MicroAPI::Duplicate(ctx.nan, Traits::NAN_MASK, maskAll);
|
||||
}
|
||||
|
||||
|
||||
template <typename FloatT>
|
||||
__simd_callee__ inline void FloatToSortableKey(AscendC::MicroAPI::RegTensor<typename FloatSortTraits<FloatT>::UInt>& outKey,
|
||||
AscendC::MicroAPI::RegTensor<FloatT>& inVal,
|
||||
FloatSortConstCtx<FloatT>& ctx,
|
||||
AscendC::MicroAPI::MaskReg& maskAll)
|
||||
{
|
||||
using Traits = FloatSortTraits<FloatT>;
|
||||
using UInt = typename Traits::UInt;
|
||||
|
||||
AscendC::MicroAPI::RegTensor<UInt> regTemp;
|
||||
AscendC::MicroAPI::RegTensor<UInt> regMask;
|
||||
AscendC::MicroAPI::MaskReg regSelectNan;
|
||||
AscendC::MicroAPI::MaskReg regSelectSign;
|
||||
|
||||
auto& inBits = (AscendC::MicroAPI::RegTensor<UInt>&)inVal;
|
||||
|
||||
// 1. NaN check
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::EQ>(regSelectNan, inBits, ctx.nan, maskAll);
|
||||
|
||||
// 2. NaN -> ALL_ONE
|
||||
AscendC::MicroAPI::Select(outKey, ctx.allOnes, inBits, regSelectNan);
|
||||
|
||||
// 3. sign bit
|
||||
AscendC::MicroAPI::And(regTemp, outKey, ctx.signMask, maskAll);
|
||||
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::GT>(regSelectSign, regTemp, ctx.zeros, maskAll);
|
||||
|
||||
// 4. xor mask
|
||||
AscendC::MicroAPI::Select(regMask, ctx.allOnes, ctx.signMask, regSelectSign);
|
||||
AscendC::MicroAPI::Xor(outKey, outKey, regMask, maskAll);
|
||||
}
|
||||
|
||||
template <typename FloatT>
|
||||
__simd_callee__ inline void FloatX2ToSortableKey(AscendC::MicroAPI::RegTensor<typename FloatSortTraits<FloatT>::UInt>& outKey0,
|
||||
AscendC::MicroAPI::RegTensor<typename FloatSortTraits<FloatT>::UInt>& outKey1,
|
||||
AscendC::MicroAPI::RegTensor<FloatT>& inVal0,
|
||||
AscendC::MicroAPI::RegTensor<FloatT>& inVal1,
|
||||
FloatSortConstCtx<FloatT>& ctx,
|
||||
AscendC::MicroAPI::MaskReg& maskAll)
|
||||
{
|
||||
using Traits = FloatSortTraits<FloatT>;
|
||||
using UInt = typename Traits::UInt;
|
||||
|
||||
AscendC::MicroAPI::RegTensor<UInt> regTemp[2];
|
||||
AscendC::MicroAPI::RegTensor<UInt> regMask[2];
|
||||
AscendC::MicroAPI::MaskReg regSelectNan[2];
|
||||
AscendC::MicroAPI::MaskReg regSelectSign[2];
|
||||
|
||||
auto& inBits0 = (AscendC::MicroAPI::RegTensor<UInt>&)inVal0;
|
||||
auto& inBits1 = (AscendC::MicroAPI::RegTensor<UInt>&)inVal1;
|
||||
|
||||
// 1. NaN check
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::EQ>(regSelectNan[0], inBits0, ctx.nan, maskAll);
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::EQ>(regSelectNan[1], inBits1, ctx.nan, maskAll);
|
||||
|
||||
// 2. NaN -> ALL_ONE
|
||||
AscendC::MicroAPI::Select(outKey0, ctx.allOnes, inBits0, regSelectNan[0]);
|
||||
AscendC::MicroAPI::Select(outKey1, ctx.allOnes, inBits1, regSelectNan[1]);
|
||||
|
||||
// 3. sign bit
|
||||
AscendC::MicroAPI::And(regTemp[0], outKey0, ctx.signMask, maskAll);
|
||||
AscendC::MicroAPI::And(regTemp[1], outKey1, ctx.signMask, maskAll);
|
||||
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::GT>(regSelectSign[0], regTemp[0], ctx.zeros, maskAll);
|
||||
AscendC::MicroAPI::Compare<UInt, CMPMODE::GT>(regSelectSign[1], regTemp[1], ctx.zeros, maskAll);
|
||||
|
||||
// 4. xor mask
|
||||
AscendC::MicroAPI::Select(regMask[0], ctx.allOnes, ctx.signMask, regSelectSign[0]);
|
||||
AscendC::MicroAPI::Select(regMask[1], ctx.allOnes, ctx.signMask, regSelectSign[1]);
|
||||
AscendC::MicroAPI::Xor(outKey0, outKey0, regMask[0], maskAll);
|
||||
AscendC::MicroAPI::Xor(outKey1, outKey1, regMask[1], maskAll);
|
||||
}
|
||||
|
||||
|
||||
template <typename T, size_t N>
|
||||
__simd_callee__ inline void DuplicateZero(AscendC::MicroAPI::RegTensor<T> (®Array)[N],
|
||||
AscendC::MicroAPI::MaskReg& mask)
|
||||
{
|
||||
static_assert(N <= 4, "N must be <= 4");
|
||||
// 不能用循环, 会导致fatal error: error in backend: Unsupported Inst must be hoisted.
|
||||
if constexpr (N >= 1) {
|
||||
AscendC::MicroAPI::Duplicate(regArray[0], static_cast<T>(0), mask);
|
||||
}
|
||||
if constexpr (N >= 2) {
|
||||
AscendC::MicroAPI::Duplicate(regArray[1], static_cast<T>(0), mask);
|
||||
}
|
||||
if constexpr (N >= 3) {
|
||||
AscendC::MicroAPI::Duplicate(regArray[2], static_cast<T>(0), mask);
|
||||
}
|
||||
if constexpr (N >= 4) {
|
||||
AscendC::MicroAPI::Duplicate(regArray[3], static_cast<T>(0), mask);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <typename T, size_t N, bool ApplyRelu = true>
|
||||
__simd_callee__ inline void WeightedAccum(AscendC::MicroAPI::RegTensor<T> (&accum)[N],
|
||||
AscendC::MicroAPI::RegTensor<T> (&input)[N],
|
||||
AscendC::MicroAPI::RegTensor<T>& weight,
|
||||
AscendC::MicroAPI::MaskReg& mask)
|
||||
{
|
||||
static_assert(N <= 2, "N must be <= 2");
|
||||
// ---- Relu block ----
|
||||
if constexpr (ApplyRelu) {
|
||||
if constexpr (N >= 1) {
|
||||
AscendC::MicroAPI::Relu(input[0], input[0], mask);
|
||||
}
|
||||
if constexpr (N >= 2) {
|
||||
AscendC::MicroAPI::Relu(input[1], input[1], mask);
|
||||
}
|
||||
}
|
||||
// ---- MulAdd block ----
|
||||
if constexpr (N >= 1) {
|
||||
AscendC::MicroAPI::MulAddDst(accum[0], input[0], weight, mask);
|
||||
}
|
||||
if constexpr (N >= 2) {
|
||||
AscendC::MicroAPI::MulAddDst(accum[1], input[1], weight, mask);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
__simd_callee__ inline void BroadcastLane(AscendC::MicroAPI::RegTensor<float>& dst,
|
||||
AscendC::MicroAPI::RegTensor<float>& src,
|
||||
uint16_t laneIdx)
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<uint32_t> brcGatherIndex;
|
||||
AscendC::MicroAPI::Duplicate(brcGatherIndex, laneIdx);
|
||||
AscendC::MicroAPI::Gather(dst, src, brcGatherIndex);
|
||||
}
|
||||
|
||||
__simd_callee__ inline void BroadcastLane(AscendC::MicroAPI::RegTensor<float>& dst,
|
||||
__local_mem__ float* src,
|
||||
uint16_t laneIdx)
|
||||
{
|
||||
AscendC::MicroAPI::LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(dst, src + laneIdx);
|
||||
}
|
||||
|
||||
// float in uint16 out
|
||||
__aicore__ inline void MulWeightAndReduceSum(const LocalTensor<uint16_t> &out_, // out [S2Base] [128 ]
|
||||
const LocalTensor<float> &qk_, // q*k^t [G, S2Base] [64 128]
|
||||
const uint32_t qkVLStride,
|
||||
const LocalTensor<float> &weight_, // w [G] [64 ]
|
||||
const LocalTensor<float> &kScale_, // kScale [S2Base] [128 ]
|
||||
const LocalTensor<float> &qScale_, // qScale [G] [64 ]
|
||||
const int gSize) // G 64
|
||||
{
|
||||
auto weight = (__local_mem__ float*)weight_.GetPhyAddr();
|
||||
auto qScale = (__local_mem__ float*)qScale_.GetPhyAddr();
|
||||
auto kScale = (__local_mem__ float*)kScale_.GetPhyAddr();
|
||||
auto qk = (__local_mem__ float*)qk_.GetPhyAddr();
|
||||
auto out = (__local_mem__ uint16_t*)out_.GetPhyAddr();
|
||||
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> regwBrc;
|
||||
AscendC::MicroAPI::RegTensor<float> regQK[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regW;
|
||||
|
||||
AscendC::MicroAPI::RegTensor<float> regQScale;
|
||||
AscendC::MicroAPI::RegTensor<float> regKScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum0[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum1[2];
|
||||
AscendC::MicroAPI::MaskReg maskAllB32 = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg maskAllB16 = AscendC::MicroAPI::CreateMask<bfloat16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
FloatSortConstCtx<bfloat16_t> bf16Ctx;
|
||||
InitFloatSortConstCtx(bf16Ctx, maskAllB16);
|
||||
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND};
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND};
|
||||
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW, weight);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale, qScale);
|
||||
AscendC::MicroAPI::Mul(regW, regW, regQScale, maskAllB32);
|
||||
|
||||
DuplicateZero(regSum0, maskAllB32);
|
||||
DuplicateZero(regSum1, maskAllB32);
|
||||
|
||||
MicroAPI::LoadAlign<float>(regKScale[0], kScale);
|
||||
MicroAPI::LoadAlign<float>(regKScale[1], kScale + 64);
|
||||
|
||||
// unroll2
|
||||
for (uint16_t i = (uint16_t)(0); i < (uint16_t)(gSize); i += 2) {
|
||||
MicroAPI::LoadAlign<float>(regQK[0], qk + 128 * i); // RowStride是128, 行都落在一个bank上
|
||||
MicroAPI::LoadAlign<float>(regQK[1], qk + 128 * i + qkVLStride);
|
||||
BroadcastLane(regwBrc, regW, i);
|
||||
WeightedAccum(regSum0, regQK, regwBrc, maskAllB32);
|
||||
|
||||
MicroAPI::LoadAlign<float>(regQK[0], qk + 128 * i + 128);
|
||||
MicroAPI::LoadAlign<float>(regQK[1], qk + 128 * i + 128 + qkVLStride);
|
||||
BroadcastLane(regwBrc, regW, i + 1);
|
||||
WeightedAccum(regSum1, regQK, regwBrc, maskAllB32);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::Add(regSum0[0], regSum0[0], regSum1[0], maskAllB32);
|
||||
AscendC::MicroAPI::Add(regSum0[1], regSum0[1], regSum1[1], maskAllB32);
|
||||
|
||||
AscendC::MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32);
|
||||
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regSumBF16;
|
||||
// interleave cast ==> regSum[1] high regSum[0] low
|
||||
AscendC::MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16, regSum0[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16, regSum0[0], maskAllB32);
|
||||
|
||||
AscendC::MicroAPI::RegTensor<uint16_t> regOut;
|
||||
FloatToSortableKey<bfloat16_t>(regOut, regSumBF16, bf16Ctx, maskAllB16);
|
||||
// normal store
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out, regOut, maskAllB16);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// bfloat16_t in uint16 out
|
||||
__aicore__ inline void MulWeightAndReduceSum(const LocalTensor<uint16_t> &out_, // out [S2Base] [128 ]
|
||||
const LocalTensor<bfloat16_t> &qk_, // q*k^t [G, S2Base] [64 128]
|
||||
const uint32_t qkVLStride, // unused for bfloat16
|
||||
const LocalTensor<float> &weight_, // w [G] [64 ]
|
||||
const LocalTensor<float> &kScale_, // kScale [S2Base] [128 ]
|
||||
const LocalTensor<float> &qScale_, // qScale [G] [64 ]
|
||||
const int gSize) // G 64
|
||||
{
|
||||
auto weight = (__local_mem__ float*)weight_.GetPhyAddr();
|
||||
auto qScale = (__local_mem__ float*)qScale_.GetPhyAddr();
|
||||
auto qk = (__local_mem__ bfloat16_t*)qk_.GetPhyAddr();
|
||||
auto kScale = (__local_mem__ float*)kScale_.GetPhyAddr();
|
||||
auto out = (__local_mem__ uint16_t*)out_.GetPhyAddr();
|
||||
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> regQK[4];
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regQKB16[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regW;
|
||||
AscendC::MicroAPI::RegTensor<float> regwBrc[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regQScale;
|
||||
AscendC::MicroAPI::RegTensor<float> regKScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum[2];
|
||||
|
||||
AscendC::MicroAPI::MaskReg maskAllB32 = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg maskAllB16 = AscendC::MicroAPI::CreateMask<bfloat16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regSumBF16;
|
||||
|
||||
FloatSortConstCtx<bfloat16_t> bf16Ctx;
|
||||
InitFloatSortConstCtx(bf16Ctx, maskAllB16);
|
||||
|
||||
|
||||
using CastTrait = AscendC::MicroAPI::CastTrait;
|
||||
static constexpr CastTrait castTraitB162B32_EVEN = {AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
static constexpr CastTrait castTraitB162B32_ODD = {AscendC::MicroAPI::RegLayout::ONE, AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
constexpr static CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND};
|
||||
constexpr static CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND};
|
||||
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW, weight);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale, qScale);
|
||||
AscendC::MicroAPI::Mul(regW, regW, regQScale, maskAllB32);
|
||||
AscendC::MicroAPI::StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_NORM>(weight, regW, maskAllB32);
|
||||
AscendC::MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
DuplicateZero(regSum, maskAllB32);
|
||||
|
||||
// interleave load
|
||||
MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>(regKScale[0], regKScale[1], kScale);
|
||||
|
||||
// Duplicate + Gather方法劣化
|
||||
// Relu在cube随路做
|
||||
for (uint16_t i = (uint16_t)(0); i < (uint16_t)(gSize); i++) {
|
||||
AscendC::MicroAPI::LoadAlign<bfloat16_t>(regQKB16[0], qk + 256 * i); // RowStride是256, 行都落在一个bank上
|
||||
AscendC::MicroAPI::LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(regwBrc[0], weight + i);
|
||||
// interleave cast
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_EVEN>(regQK[0], regQKB16[0], maskAllB16);
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_ODD>(regQK[1], regQKB16[0], maskAllB16);
|
||||
AscendC::MicroAPI::MulAddDst(regSum[0], regQK[0], regwBrc[0], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum[1], regQK[1], regwBrc[0], maskAllB32);
|
||||
}
|
||||
|
||||
AscendC::MicroAPI::Mul(regSum[0], regSum[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum[1], regSum[1], regKScale[1], maskAllB32);
|
||||
// interleave cast back
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16, regSum[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16, regSum[0], maskAllB32);
|
||||
|
||||
AscendC::MicroAPI::RegTensor<uint16_t> regOut;
|
||||
FloatToSortableKey<bfloat16_t>(regOut, regSumBF16, bf16Ctx, maskAllB16);
|
||||
// norm load
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out, regOut, maskAllB16);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// 计算S1=2
|
||||
// float in uint16 out
|
||||
__aicore__ inline void MulWeightAndReduceSum2(const LocalTensor<uint16_t> &out_, // out [2, S2Base] [128 ]
|
||||
uint32_t outStride,
|
||||
const LocalTensor<float> &qk_, // q*k^t [2, G, S2Base] [64 128]
|
||||
uint32_t qkVLStride,
|
||||
uint32_t qkStride,
|
||||
const LocalTensor<float> &weight_, // w [2, G] [64 ]
|
||||
uint32_t weightStride,
|
||||
const LocalTensor<float> &kScale_, // kScale [S2Base] [128 ]
|
||||
uint32_t kScaleStride,
|
||||
const LocalTensor<float> &qScale_, // qScale [2, G] [64 ]
|
||||
uint32_t qScaleStride,
|
||||
const int gSize) // G 64
|
||||
{
|
||||
auto weight0 = (__local_mem__ float*)weight_.GetPhyAddr();
|
||||
auto qScale0 = (__local_mem__ float*)qScale_.GetPhyAddr();
|
||||
auto kScale0 = (__local_mem__ float*)kScale_.GetPhyAddr();
|
||||
auto qk0 = (__local_mem__ float*)qk_.GetPhyAddr();
|
||||
auto out0 = (__local_mem__ uint16_t*)out_.GetPhyAddr();
|
||||
|
||||
auto weight1 = weight0 + weightStride;
|
||||
auto qScale1 = qScale0 + qScaleStride;
|
||||
auto qk1 = qk0 + qkStride;
|
||||
// kScaleStride is zero
|
||||
auto out1 = out0 + outStride;
|
||||
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> regwBrc[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regQK0[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regQK1[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regW[2];
|
||||
|
||||
AscendC::MicroAPI::RegTensor<float> regQScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regKScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum0[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum1[2];
|
||||
AscendC::MicroAPI::MaskReg maskAllB32 = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg maskAllB16 = AscendC::MicroAPI::CreateMask<bfloat16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
FloatSortConstCtx<bfloat16_t> bf16Ctx;
|
||||
InitFloatSortConstCtx(bf16Ctx, maskAllB16);
|
||||
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND};
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND};
|
||||
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW[0], weight0);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW[1], weight1);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale[0], qScale0);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale[1], qScale1);
|
||||
AscendC::MicroAPI::Mul(regW[0], regW[0], regQScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regW[1], regW[1], regQScale[1], maskAllB32);
|
||||
// regW[0]与weight1混合使用
|
||||
AscendC::MicroAPI::StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_NORM>(weight1, regW[1], maskAllB32);
|
||||
AscendC::MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
DuplicateZero(regSum0, maskAllB32);
|
||||
DuplicateZero(regSum1, maskAllB32);
|
||||
|
||||
MicroAPI::LoadAlign<float>(regKScale[0], kScale0);
|
||||
MicroAPI::LoadAlign<float>(regKScale[1], kScale0 + 64);
|
||||
|
||||
for (uint16_t i = (uint16_t)(0); i < (uint16_t)(gSize); i++) {
|
||||
MicroAPI::LoadAlign<float>(regQK0[0], qk0 + 128 * i);
|
||||
MicroAPI::LoadAlign<float>(regQK0[1], qk0 + 128 * i + qkVLStride);
|
||||
MicroAPI::LoadAlign<float>(regQK1[0], qk1 + 128 * i);
|
||||
MicroAPI::LoadAlign<float>(regQK1[1], qk1 + 128 * i + qkVLStride);
|
||||
// 混合使用对整体性能更好
|
||||
BroadcastLane(regwBrc[0], regW[0], i);
|
||||
// Weight无bank冲突,用LoadAlign来提取weight标量
|
||||
BroadcastLane(regwBrc[1], weight1, i);
|
||||
AscendC::MicroAPI::Relu(regQK0[0], regQK0[0], maskAllB32);
|
||||
AscendC::MicroAPI::Relu(regQK0[1], regQK0[1], maskAllB32);
|
||||
AscendC::MicroAPI::Relu(regQK1[0], regQK1[0], maskAllB32);
|
||||
AscendC::MicroAPI::Relu(regQK1[1], regQK1[1], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32);
|
||||
}
|
||||
|
||||
// Apply kScale scaling
|
||||
AscendC::MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32);
|
||||
|
||||
|
||||
// Convert to bfloat16 and store output channel
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regSumBF16[2];
|
||||
AscendC::MicroAPI::RegTensor<uint16_t> regOut[2];
|
||||
AscendC::MicroAPI::DeInterleave(regSum0[0], regSum0[1], regSum0[0], regSum0[1]);
|
||||
AscendC::MicroAPI::DeInterleave(regSum1[0], regSum1[1], regSum1[0], regSum1[1]);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16[0], regSum0[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16[1], regSum1[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16[0], regSum0[0], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16[1], regSum1[0], maskAllB32);
|
||||
|
||||
FloatX2ToSortableKey<bfloat16_t>(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16);
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out0, regOut[0], maskAllB16);
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out1, regOut[1], maskAllB16);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// 计算S1=2
|
||||
// bfloat16 in uint16 out
|
||||
__aicore__ inline void MulWeightAndReduceSum2(const LocalTensor<uint16_t> &out_, // out [2, S2Base] [128 ]
|
||||
uint32_t outStride,
|
||||
const LocalTensor<bfloat16_t> &qk_, // q*k^t [2, G, S2Base] [64 128]
|
||||
uint32_t qkVLStride,
|
||||
uint32_t qkStride, // gSize * 256
|
||||
const LocalTensor<float> &weight_, // w [2, G] [64 ]
|
||||
uint32_t weightStride,
|
||||
const LocalTensor<float> &kScale_, // kScale [S2Base] [128 ]
|
||||
uint32_t kScaleStride,
|
||||
const LocalTensor<float> &qScale_, // qScale [2, G] [64 ]
|
||||
uint32_t qScaleStride,
|
||||
const int gSize) // G 64
|
||||
{
|
||||
auto weight0 = (__local_mem__ float*)weight_.GetPhyAddr();
|
||||
auto qScale0 = (__local_mem__ float*)qScale_.GetPhyAddr();
|
||||
auto kScale0 = (__local_mem__ float*)kScale_.GetPhyAddr();
|
||||
auto qk0 = (__local_mem__ bfloat16_t*)qk_.GetPhyAddr();
|
||||
auto out0 = (__local_mem__ uint16_t*)out_.GetPhyAddr();
|
||||
|
||||
auto weight1 = weight0 + weightStride;
|
||||
auto qScale1 = qScale0 + qScaleStride;
|
||||
auto qk1 = qk0 + qkStride;
|
||||
// kScaleStride is zero
|
||||
auto out1 = out0 + outStride;
|
||||
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
AscendC::MicroAPI::RegTensor<float> regwBrc[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regQK0[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regQK1[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regW[2];
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regQKB16[2];
|
||||
|
||||
AscendC::MicroAPI::RegTensor<float> regQScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regKScale[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum0[2];
|
||||
AscendC::MicroAPI::RegTensor<float> regSum1[2];
|
||||
AscendC::MicroAPI::MaskReg maskAllB32 = AscendC::MicroAPI::CreateMask<float, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
AscendC::MicroAPI::MaskReg maskAllB16 = AscendC::MicroAPI::CreateMask<bfloat16_t, AscendC::MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
FloatSortConstCtx<bfloat16_t> bf16Ctx;
|
||||
InitFloatSortConstCtx(bf16Ctx, maskAllB16);
|
||||
|
||||
using CastTrait = AscendC::MicroAPI::CastTrait;
|
||||
static constexpr CastTrait castTraitB162B32_EVEN = {AscendC::MicroAPI::RegLayout::ZERO, AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
static constexpr CastTrait castTraitB162B32_ODD = {AscendC::MicroAPI::RegLayout::ONE, AscendC::MicroAPI::SatMode::UNKNOWN,
|
||||
AscendC::MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_EVEN = {MicroAPI::RegLayout::ZERO, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::MERGING, RoundMode::CAST_ROUND};
|
||||
constexpr static MicroAPI::CastTrait castTraitF32ToF16_ODD = {MicroAPI::RegLayout::ONE, MicroAPI::SatMode::NO_SAT,
|
||||
MicroAPI::MaskMergeMode::ZEROING, RoundMode::CAST_ROUND};
|
||||
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW[0], weight0);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regW[1], weight1);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale[0], qScale0);
|
||||
AscendC::MicroAPI::LoadAlign<float>(regQScale[1], qScale1);
|
||||
AscendC::MicroAPI::Mul(regW[0], regW[0], regQScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regW[1], regW[1], regQScale[1], maskAllB32);
|
||||
// 读写依赖,寄存器可以保序
|
||||
AscendC::MicroAPI::StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_NORM>(weight0, regW[0], maskAllB32);
|
||||
AscendC::MicroAPI::StoreAlign<float, AscendC::MicroAPI::StoreDist::DIST_NORM>(weight1, regW[1], maskAllB32);
|
||||
DuplicateZero(regSum0, maskAllB32);
|
||||
DuplicateZero(regSum1, maskAllB32);
|
||||
|
||||
// interleave load
|
||||
MicroAPI::LoadAlign<float, MicroAPI::LoadDist::DIST_DINTLV_B32>(regKScale[0], regKScale[1], kScale0);
|
||||
|
||||
for (uint16_t i = (uint16_t)(0); i < (uint16_t)(gSize); i++) {
|
||||
AscendC::MicroAPI::LoadAlign<bfloat16_t>(regQKB16[0], qk0 + 256 * i); // RowStride是256, 行都落在一个bank上
|
||||
AscendC::MicroAPI::LoadAlign<bfloat16_t>(regQKB16[1], qk1 + 256 * i); // RowStride是256, 行都落在一个bank上
|
||||
AscendC::MicroAPI::LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(regwBrc[0], weight0 + i);
|
||||
AscendC::MicroAPI::LoadAlign<float, AscendC::MicroAPI::LoadDist::DIST_BRC_B32>(regwBrc[1], weight1 + i);
|
||||
// interleave cast
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_EVEN>(regQK0[0], regQKB16[0], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_ODD>(regQK0[1], regQKB16[0], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_EVEN>(regQK1[0], regQKB16[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<float, bfloat16_t, castTraitB162B32_ODD>(regQK1[1], regQKB16[1], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum0[0], regQK0[0], regwBrc[0], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum0[1], regQK0[1], regwBrc[0], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum1[0], regQK1[0], regwBrc[1], maskAllB32);
|
||||
AscendC::MicroAPI::MulAddDst(regSum1[1], regQK1[1], regwBrc[1], maskAllB32);
|
||||
}
|
||||
|
||||
// Apply kScale scaling
|
||||
AscendC::MicroAPI::Mul(regSum0[0], regSum0[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum0[1], regSum0[1], regKScale[1], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum1[0], regSum1[0], regKScale[0], maskAllB32);
|
||||
AscendC::MicroAPI::Mul(regSum1[1], regSum1[1], regKScale[1], maskAllB32);
|
||||
|
||||
// Convert to bfloat16 and store output channel
|
||||
AscendC::MicroAPI::RegTensor<bfloat16_t> regSumBF16[2];
|
||||
AscendC::MicroAPI::RegTensor<uint16_t> regOut[2];
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16[0], regSum0[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_ODD>(regSumBF16[1], regSum1[1], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16[0], regSum0[0], maskAllB32);
|
||||
AscendC::MicroAPI::Cast<bfloat16_t, float, castTraitF32ToF16_EVEN>(regSumBF16[1], regSum1[0], maskAllB32);
|
||||
|
||||
FloatX2ToSortableKey<bfloat16_t>(regOut[0], regOut[1], regSumBF16[0], regSumBF16[1], bf16Ctx, maskAllB16);
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out0, regOut[0], maskAllB16);
|
||||
AscendC::MicroAPI::StoreAlign<uint16_t, AscendC::MicroAPI::StoreDist::DIST_NORM>(out1, regOut[1], maskAllB16);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template<typename QK_T, typename SCORE_T>
|
||||
__aicore__ inline void BatchMulWeightAndReduceSum(const LocalTensor<SCORE_T> &out_, // out [S2Base] [128 ]
|
||||
uint32_t outStride,
|
||||
const LocalTensor<QK_T> &qk_, // q*k^t [G, S2Base] [64 128]
|
||||
uint32_t qkVLStride,
|
||||
uint32_t qkStride,
|
||||
const LocalTensor<float> &weight_, // w [G] [64 ]
|
||||
uint32_t weightStride,
|
||||
const LocalTensor<float> &kScale_, // kScale [S2Base] [128 ]
|
||||
uint32_t kScaleStride,
|
||||
const LocalTensor<float> &qScale_, // qScale [G] [64 ]
|
||||
uint32_t qScaleStride,
|
||||
const int gSize, // G 64
|
||||
const int batch)
|
||||
{
|
||||
// 暂只支持这两种情况, 后续改成循环
|
||||
if (batch != 2 && batch != 1) {
|
||||
return;
|
||||
}
|
||||
if (batch == 2) {
|
||||
MulWeightAndReduceSum2(out_, outStride,
|
||||
qk_, qkVLStride, qkStride,
|
||||
weight_, weightStride,
|
||||
kScale_, kScaleStride,
|
||||
qScale_, qScaleStride,
|
||||
gSize);
|
||||
} else {
|
||||
MulWeightAndReduceSum(out_, qk_, qkVLStride, weight_, kScale_, qScale_, gSize);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,678 @@
|
||||
/**
|
||||
* 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 vf_top_k.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef VF_TOP_K_H
|
||||
#define VF_TOP_K_H
|
||||
|
||||
namespace topkb32 {
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsFirstVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* inputBuf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// 计算直方图cout0 0-127 cout1 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
// 32bit 高16bit
|
||||
MicroAPI::RegTensor<uint32_t> vreg0U16;
|
||||
// 32bit 低16bit
|
||||
MicroAPI::RegTensor<uint32_t> vreg1U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg2U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg3U16;
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> vreg0;
|
||||
MicroAPI::RegTensor<uint8_t> vreg1;
|
||||
MicroAPI::RegTensor<uint8_t> vreg2;
|
||||
MicroAPI::RegTensor<uint8_t> vreg3;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg1U16, vreg0U16, inputBuf + i * 256);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128);
|
||||
|
||||
MicroAPI::DeInterleave(vreg1, vreg0, (MicroAPI::RegTensor<uint8_t>&)vreg0U16, (MicroAPI::RegTensor<uint8_t>&)vreg2U16);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, vreg0, pregB8);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, vreg0, pregB8);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindFirstTargetBinVFImpl(__ubuf__ uint32_t* idx0Buf, __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* histogramsBuf, uint32_t bottomK)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx0;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK;
|
||||
MicroAPI::Duplicate(btmK, bottomK);
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdx0;
|
||||
|
||||
MicroAPI::MaskReg pregGE = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK, pregB32);
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdx0, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idx0Buf, sqzIdx0, alignIdx0);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idx0Buf, alignIdx0);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx0;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx0, idx0Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> idxAll1;
|
||||
MicroAPI::RegTensor<uint32_t> idxPrev0;
|
||||
MicroAPI::RegTensor<uint32_t> prevBinValue;
|
||||
MicroAPI::Duplicate(idxAll1, 1);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> zeroAll;
|
||||
MicroAPI::Duplicate(zeroAll, 0);
|
||||
|
||||
MicroAPI::MaskReg preg0 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(preg0, idx0, zeroAll, pregB32);
|
||||
MicroAPI::Sub(idxPrev0, idx0, (MicroAPI::RegTensor<uint32_t>&)idxAll1, pregB32);
|
||||
MicroAPI::ShiftRights(idxPrev0, idxPrev0, (int16_t)24, pregB32);
|
||||
|
||||
MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev0, pregB32);
|
||||
MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, preg0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> nextK;
|
||||
MicroAPI::Sub(nextK, btmK, prevBinValue, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(nkValueBuf, nextK, pregB32);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsSecondVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// 计算直方图0-127 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx0;
|
||||
// 0x000000fc -> 0xfcfcfcfc
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx0, idx0Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vreg0U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg1U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg2U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg3U16;
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> vreg0;
|
||||
MicroAPI::RegTensor<uint8_t> vreg1;
|
||||
MicroAPI::RegTensor<uint8_t> vreg2;
|
||||
MicroAPI::RegTensor<uint8_t> vreg3;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg1U16, vreg0U16, inputBuf + i * 256);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128);
|
||||
|
||||
MicroAPI::DeInterleave(vreg1, vreg0, (MicroAPI::RegTensor<uint8_t>&)vreg0U16, (MicroAPI::RegTensor<uint8_t>&)vreg2U16);
|
||||
|
||||
MicroAPI::MaskReg pregEQ = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ, vreg0, (MicroAPI::RegTensor<uint8_t>&)idx0, pregB8);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, vreg1, pregEQ);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, vreg1, pregEQ);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
// kValue新的bottomK
|
||||
__simd_vf__ void FindSecondTargetBinVFImpl(__ubuf__ uint32_t* idx1Buf, __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* kValue, __ubuf__ uint32_t* histogramsBuf)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx1;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK1;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(btmK1, kValue);
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdx1;
|
||||
|
||||
MicroAPI::MaskReg pregGE = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK1, pregB32);
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdx1, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idx1Buf, sqzIdx1, alignIdx1);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idx1Buf, alignIdx1);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx1;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx1, idx1Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> idxAll1;
|
||||
MicroAPI::RegTensor<uint32_t> idxPrev1;
|
||||
MicroAPI::RegTensor<uint32_t> prevBinValue;
|
||||
MicroAPI::Duplicate(idxAll1, 1);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> zeroAll;
|
||||
MicroAPI::Duplicate(zeroAll, 0);
|
||||
|
||||
MicroAPI::MaskReg preg1 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(preg1, idx1, zeroAll, pregB32);
|
||||
MicroAPI::Sub(idxPrev1, idx1, (MicroAPI::RegTensor<uint32_t>&)idxAll1, pregB32);
|
||||
MicroAPI::ShiftRights(idxPrev1, idxPrev1, (int16_t)24, pregB32);
|
||||
|
||||
MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev1, pregB32);
|
||||
MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, preg1);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> nextK;
|
||||
MicroAPI::Sub(nextK, btmK1, prevBinValue, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(nkValueBuf, nextK, pregB32);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsThirdVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, __ubuf__ uint32_t* idx1Buf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// 计算直方图0-127 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx0;
|
||||
MicroAPI::RegTensor<uint32_t> idx1;
|
||||
// 0x000000fc -> 0xfcfcfcfc
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx0, idx0Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx1, idx1Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vreg0U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg1U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg2U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg3U16;
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> vreg0;
|
||||
MicroAPI::RegTensor<uint8_t> vreg1;
|
||||
MicroAPI::RegTensor<uint8_t> vreg2;
|
||||
MicroAPI::RegTensor<uint8_t> vreg3;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg1U16, vreg0U16, inputBuf + i * 256);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128);
|
||||
|
||||
MicroAPI::DeInterleave(vreg1, vreg0, (MicroAPI::RegTensor<uint8_t>&)vreg0U16, (MicroAPI::RegTensor<uint8_t>&)vreg2U16);
|
||||
MicroAPI::DeInterleave(vreg3, vreg2, (MicroAPI::RegTensor<uint8_t>&)vreg1U16, (MicroAPI::RegTensor<uint8_t>&)vreg3U16);
|
||||
|
||||
MicroAPI::MaskReg pregEQ0 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregEQ1 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ0, vreg0, (MicroAPI::RegTensor<uint8_t>&)idx0, pregB8);
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ1, vreg1, (MicroAPI::RegTensor<uint8_t>&)idx1, pregB8);
|
||||
|
||||
MicroAPI::MaskReg pregEQ = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::And(pregEQ, pregEQ0, pregEQ1, pregB8);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, vreg2, pregEQ);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, vreg2, pregEQ);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindThirdTargetBinVFImpl(__ubuf__ uint32_t* idx2Buf, __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* kValue, __ubuf__ uint32_t* histogramsBuf)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx2;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK2;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(btmK2, kValue);
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdx2;
|
||||
|
||||
MicroAPI::MaskReg pregGE = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK2, pregB32);
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdx2, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idx2Buf, sqzIdx2, alignIdx2);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idx2Buf, alignIdx2);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx2;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx2, idx2Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> idxAll1;
|
||||
MicroAPI::RegTensor<uint32_t> idxPrev2;
|
||||
MicroAPI::RegTensor<uint32_t> prevBinValue;
|
||||
MicroAPI::Duplicate(idxAll1, 1);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> zeroAll;
|
||||
MicroAPI::Duplicate(zeroAll, 0);
|
||||
|
||||
MicroAPI::MaskReg preg2 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(preg2, idx2, zeroAll, pregB32);
|
||||
MicroAPI::Sub(idxPrev2, idx2, (MicroAPI::RegTensor<uint32_t>&)idxAll1, pregB32);
|
||||
MicroAPI::ShiftRights(idxPrev2, idxPrev2, (int16_t)24, pregB32);
|
||||
|
||||
MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev2, pregB32);
|
||||
MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, preg2);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> nextK;
|
||||
MicroAPI::Sub(nextK, btmK2, prevBinValue, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(nkValueBuf, nextK, pregB32);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsLastVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* idx0Buf, __ubuf__ uint32_t* idx1Buf, __ubuf__ uint32_t* idx2Buf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// 计算直方图0-127 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx0;
|
||||
MicroAPI::RegTensor<uint32_t> idx1;
|
||||
MicroAPI::RegTensor<uint32_t> idx2;
|
||||
// 0x000000fc -> 0xfcfcfcfc
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx0, idx0Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx1, idx1Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idx2, idx2Buf);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vreg0U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg1U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg2U16;
|
||||
MicroAPI::RegTensor<uint32_t> vreg3U16;
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> vreg0;
|
||||
MicroAPI::RegTensor<uint8_t> vreg1;
|
||||
MicroAPI::RegTensor<uint8_t> vreg2;
|
||||
MicroAPI::RegTensor<uint8_t> vreg3;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg1U16, vreg0U16, inputBuf + i * 256);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_DINTLV_B16>(vreg3U16, vreg2U16, inputBuf + (i * 256) + 128);
|
||||
|
||||
MicroAPI::DeInterleave(vreg1, vreg0, (MicroAPI::RegTensor<uint8_t>&)vreg0U16, (MicroAPI::RegTensor<uint8_t>&)vreg2U16);
|
||||
MicroAPI::DeInterleave(vreg3, vreg2, (MicroAPI::RegTensor<uint8_t>&)vreg1U16, (MicroAPI::RegTensor<uint8_t>&)vreg3U16);
|
||||
|
||||
MicroAPI::MaskReg pregEQ0 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregEQ1 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregEQ2 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ0, vreg0, (MicroAPI::RegTensor<uint8_t>&)idx0, pregB8);
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ1, vreg1, (MicroAPI::RegTensor<uint8_t>&)idx1, pregB8);
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ2, vreg2, (MicroAPI::RegTensor<uint8_t>&)idx2, pregB8);
|
||||
|
||||
MicroAPI::MaskReg pregEQ0And1 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregEQAll = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::And(pregEQ0And1, pregEQ0, pregEQ1, pregB8);
|
||||
MicroAPI::And(pregEQAll, pregEQ0And1, pregEQ2, pregB8);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, vreg3, pregEQAll);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, vreg3, pregEQAll);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t* kValue, __ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* idx0Buf, __ubuf__ uint32_t* idx1Buf, __ubuf__ uint32_t* idx2Buf, __ubuf__ uint32_t* idx3Buf)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx3;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK3;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(btmK3, kValue);
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdx3;
|
||||
|
||||
MicroAPI::MaskReg pregGE = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK3, pregB32);
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdx3, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idx3Buf, sqzIdx3, alignIdx3);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idx3Buf, alignIdx3);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idx0;
|
||||
MicroAPI::RegTensor<uint32_t> idx1;
|
||||
MicroAPI::RegTensor<uint32_t> idx2;
|
||||
MicroAPI::RegTensor<uint32_t> idx3;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B32>(idx0, idx0Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B32>(idx1, idx1Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B32>(idx2, idx2Buf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B32>(idx3, idx3Buf);
|
||||
|
||||
MicroAPI::ShiftLefts(idx0, idx0, (int16_t)24, pregB32);
|
||||
MicroAPI::ShiftLefts(idx1, idx1, (int16_t)16, pregB32);
|
||||
MicroAPI::ShiftLefts(idx2, idx2, (int16_t)8, pregB32);
|
||||
|
||||
// ADD
|
||||
MicroAPI::Add(idx0, idx0, idx1, pregB32);
|
||||
MicroAPI::Add(idx0, idx0, idx2, pregB32);
|
||||
MicroAPI::Add(idx0, idx0, idx3, pregB32);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(kValue, idx0, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ uint32_t* inputBuf, uint32_t beginIdx, __ubuf__ uint32_t* kValue, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vregInput;
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::Arange(idxC, beginIdx + i * 64);
|
||||
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputBuf + i * 64);
|
||||
|
||||
MicroAPI::MaskReg poutGT = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdxOut;
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GT>(poutGT, vregInput, kthValue, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxOut, (MicroAPI::RegTensor<uint32_t>&)idxC, poutGT);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputIdxBuf, sqzIdxOut, alignIdx);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ uint32_t* inputBuf, uint32_t beginIdx, __ubuf__ uint32_t* kValue)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vregInput;
|
||||
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::Arange(idxC, beginIdx);
|
||||
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputBuf);
|
||||
|
||||
MicroAPI::MaskReg poutEQ = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdxOut;
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(poutEQ, vregInput, kthValue, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxOut, (MicroAPI::RegTensor<uint32_t>&)idxC, poutEQ);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputIdxBuf, sqzIdxOut, alignIdx);
|
||||
MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindValueGTOutputVFImpl(__ubuf__ uint32_t* outputValueBuf, __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* kValue, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignValue;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vregInput;
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputBuf + i * 64);
|
||||
|
||||
MicroAPI::MaskReg poutGT = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> sqzValueOut;
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GT>(poutGT, vregInput, kthValue, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzValueOut, vregInput, poutGT);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputValueBuf, sqzValueOut, alignValue);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(outputValueBuf, alignValue);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindValueEQOutputVFImpl(__ubuf__ uint32_t* outputValueBuf, __ubuf__ uint32_t* inputBuf, __ubuf__ uint32_t* kValue)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignValue;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> vregInput;
|
||||
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputBuf);
|
||||
|
||||
MicroAPI::MaskReg poutEQ = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> sqzValueOut;
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(poutEQ, vregInput, kthValue, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzValueOut, vregInput, poutEQ);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputValueBuf, sqzValueOut, alignValue);
|
||||
MicroAPI::StoreUnAlignPost(outputValueBuf, alignValue);
|
||||
}
|
||||
|
||||
__aicore__ inline void LiTopKVF(const LocalTensor<uint32_t>& outputIdxLocal,
|
||||
const LocalTensor<uint32_t>& outputValueLocal,
|
||||
const LocalTensor<uint32_t>& inputLocal,
|
||||
const LocalTensor<uint32_t>& tmpIdxLocal,
|
||||
const LocalTensor<uint32_t>& tmpValueLocal,
|
||||
const LocalTensor<uint32_t>& histogramsLocal,
|
||||
const LocalTensor<uint32_t>& idx0Local,
|
||||
const LocalTensor<uint32_t>& idx1Local,
|
||||
const LocalTensor<uint32_t>& idx2Local,
|
||||
const LocalTensor<uint32_t>& idx3Local,
|
||||
const LocalTensor<uint32_t>& nkValueLocal,
|
||||
uint32_t topK,
|
||||
uint32_t s2SeqLen)
|
||||
{
|
||||
__ubuf__ uint32_t* outputIdxBuf = (__ubuf__ uint32_t*)outputIdxLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* outputValueBuf = (__ubuf__ uint32_t*)outputValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* inputBuf = (__ubuf__ uint32_t*)inputLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* tmpIdxBuf = (__ubuf__ uint32_t*)tmpIdxLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* tmpValueBuf = (__ubuf__ uint32_t*)tmpValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* histogramsBuf = (__ubuf__ uint32_t*)histogramsLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idx0Buf = (__ubuf__ uint32_t*)idx0Local.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idx1Buf = (__ubuf__ uint32_t*)idx1Local.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idx2Buf = (__ubuf__ uint32_t*)idx2Local.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idx3Buf = (__ubuf__ uint32_t*)idx3Local.GetPhyAddr();
|
||||
__ubuf__ uint32_t* nkValueBuf = (__ubuf__ uint32_t*)nkValueLocal.GetPhyAddr();
|
||||
|
||||
uint32_t bottomK = s2SeqLen - topK + 1;
|
||||
uint32_t beginIdx = 0;
|
||||
bool flag = true;
|
||||
|
||||
const uint16_t repeatSize8 = 256;
|
||||
const uint16_t repeatSize32 = 64;
|
||||
|
||||
uint16_t histogramsLoopNum = (s2SeqLen + repeatSize8 - 1) / repeatSize8;
|
||||
uint16_t inputLoopNum = (s2SeqLen + repeatSize32 - 1) / repeatSize32;
|
||||
uint16_t topkLoopNum = (topK + 64 - 1) / 64;
|
||||
|
||||
// find kth-value
|
||||
HistogramsFirstVFImpl<uint32_t>(histogramsBuf, inputBuf, histogramsLoopNum, flag);
|
||||
FindFirstTargetBinVFImpl(idx0Buf, nkValueBuf, histogramsBuf, bottomK);
|
||||
HistogramsSecondVFImpl<uint32_t>(histogramsBuf, inputBuf, idx0Buf, histogramsLoopNum, flag);
|
||||
FindSecondTargetBinVFImpl(idx1Buf, nkValueBuf, nkValueBuf, histogramsBuf);
|
||||
HistogramsThirdVFImpl<uint32_t>(histogramsBuf, inputBuf, idx0Buf, idx1Buf, histogramsLoopNum, flag);
|
||||
FindThirdTargetBinVFImpl(idx2Buf, nkValueBuf, nkValueBuf, histogramsBuf);
|
||||
HistogramsLastVFImpl<uint32_t>(histogramsBuf, inputBuf, idx0Buf, idx1Buf, idx2Buf, histogramsLoopNum, flag);
|
||||
FindKthVFImpl(nkValueBuf, histogramsBuf, idx0Buf, idx1Buf, idx2Buf, idx3Buf);
|
||||
|
||||
// filter
|
||||
// 输出大于k-value的值value
|
||||
FindValueGTOutputVFImpl(outputValueBuf, inputBuf, nkValueBuf, inputLoopNum);
|
||||
// value-当前偏移大于k-value的值在AR特殊寄存器中的有效字节数
|
||||
int64_t arValueNum = AscendC::GetSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
// value-剩余需要输出等于k-value的数量
|
||||
int64_t remainValueNum = topK - (arValueNum / sizeof(uint32_t));
|
||||
for(uint16_t i = 0; i < inputLoopNum; ++i) {
|
||||
int64_t arValueNumPerLoop = AscendC::GetSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
if (((arValueNumPerLoop - arValueNum) / sizeof(uint32_t)) < remainValueNum) {
|
||||
// 调用一次查找等于k-value情况的过程
|
||||
FindValueEQOutputVFImpl(outputValueBuf, inputBuf + i * 64, nkValueBuf);
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// 输出大于k-value的值idx
|
||||
FindIdxGTOutputVFImpl(outputIdxBuf, inputBuf, (uint32_t)(0), nkValueBuf, inputLoopNum);
|
||||
// idx-当前偏移大于k-value的值在AR特殊寄存器中的有效字节数
|
||||
int64_t arIdxNum = AscendC::GetSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
int64_t remainIdxNum = topK - (arIdxNum / sizeof(uint32_t));
|
||||
for(uint16_t i = 0; i < inputLoopNum; ++i) {
|
||||
int64_t arIdxNumPerLoop = AscendC::GetSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
if (((arIdxNumPerLoop - arIdxNum) / sizeof(uint32_t)) < remainIdxNum) {
|
||||
// 调用一次查找等于k-value情况的过程
|
||||
beginIdx = i * 64;
|
||||
FindIdxEQOutputVFImpl(outputIdxBuf, inputBuf + i * 64, beginIdx, nkValueBuf);
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,430 @@
|
||||
/**
|
||||
* 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 vf_top_k_16_gather.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef VF_TOP_K_16_GATHER_H
|
||||
#define VF_TOP_K_16_GATHER_H
|
||||
|
||||
namespace topkb16gather {
|
||||
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsHighVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint16_t* inputBuf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
// 计算直方图cout0 0-127 cout1 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> vregHigh;
|
||||
MicroAPI::RegTensor<uint16_t> vregLow;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_DINTLV_B8>(vregLow, vregHigh, inputBuf + i * 256);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, (MicroAPI::RegTensor<uint8_t>&)vregHigh, pregB8);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, (MicroAPI::RegTensor<uint8_t>&)vregHigh, pregB8);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindHighTargetBinVFImpl(__ubuf__ uint32_t* idxHighBuf, __ubuf__ uint32_t* nkValueBuf, __ubuf__ uint32_t* histogramsBuf, uint32_t bottomK)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg pregGE;
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdxHigh;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK;
|
||||
MicroAPI::Duplicate(btmK, bottomK);
|
||||
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdxHigh;
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxHigh, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idxHighBuf, sqzIdxHigh, alignIdxHigh);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idxHighBuf, alignIdxHigh);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idxHigh;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idxHigh, idxHighBuf);
|
||||
|
||||
MicroAPI::RegTensor<uint8_t> idxAll1;
|
||||
MicroAPI::RegTensor<uint32_t> idxPrev0;
|
||||
MicroAPI::RegTensor<uint32_t> prevBinValue;
|
||||
MicroAPI::Duplicate(idxAll1, 1);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> zeroAll;
|
||||
MicroAPI::Duplicate(zeroAll, 0);
|
||||
|
||||
MicroAPI::MaskReg preg0 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::EQ>(preg0, idxHigh, zeroAll, pregB32);
|
||||
MicroAPI::Sub(idxPrev0, idxHigh, (MicroAPI::RegTensor<uint32_t>&)idxAll1, pregB32);
|
||||
MicroAPI::ShiftRights(idxPrev0, idxPrev0, (int16_t)24, pregB32);
|
||||
|
||||
MicroAPI::Gather(prevBinValue, histogramsBuf, idxPrev0, pregB32);
|
||||
MicroAPI::Select(prevBinValue, zeroAll, prevBinValue, preg0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> nextK;
|
||||
MicroAPI::Sub(nextK, btmK, prevBinValue, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(nkValueBuf, nextK, pregB32);
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
__simd_vf__ void HistogramsLowVFImpl(__ubuf__ uint32_t* histogramsBuf, __ubuf__ uint16_t* inputBuf, __ubuf__ uint32_t* idxHighBuf, uint16_t vfLoop, bool init)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB8 = MicroAPI::CreateMask<uint8_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg pregEQ;
|
||||
|
||||
// 计算直方图0-127 128-255
|
||||
MicroAPI::RegTensor<uint16_t> cout0;
|
||||
MicroAPI::RegTensor<uint16_t> cout1;
|
||||
MicroAPI::Duplicate(cout0, 0);
|
||||
MicroAPI::Duplicate(cout1, 0);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout0U32Odd;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Even;
|
||||
MicroAPI::RegTensor<uint32_t> cout1U32Odd;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idxHigh;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idxHigh, idxHighBuf);
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> vregHigh;
|
||||
MicroAPI::RegTensor<uint16_t> vregLow;
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_EVEN = {MicroAPI::RegLayout::ZERO,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
static constexpr MicroAPI::CastTrait CAST_TRAIT_UINT16_TOUINT32_ODD = {MicroAPI::RegLayout::ONE,
|
||||
MicroAPI::SatMode::UNKNOWN, MicroAPI::MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
for (uint16_t i = 0; i < vfLoop; ++i) {
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_DINTLV_B8>(vregLow, vregHigh, inputBuf + i * 256);
|
||||
|
||||
MicroAPI::Compare<uint8_t, CMPMODE::EQ>(pregEQ, (MicroAPI::RegTensor<uint8_t>&)vregHigh, (MicroAPI::RegTensor<uint8_t>&)idxHigh, pregB8);
|
||||
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN0,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout0, (MicroAPI::RegTensor<uint8_t>&)vregLow, pregEQ);
|
||||
MicroAPI::Histograms<uint8_t, uint16_t, MicroAPI::HistogramsBinType::BIN1,
|
||||
MicroAPI::HistogramsType::ACCUMULATE>(cout1, (MicroAPI::RegTensor<uint8_t>&)vregLow, pregEQ);
|
||||
}
|
||||
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout0U32Even, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout0U32Odd, cout0, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_EVEN>(cout1U32Even, cout1, pregB16);
|
||||
MicroAPI::Cast<uint32_t, uint16_t, CAST_TRAIT_UINT16_TOUINT32_ODD>(cout1U32Odd, cout1, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf, cout0U32Even, cout0U32Odd, pregB32);
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_INTLV_B32>(histogramsBuf + 128, cout1U32Even, cout1U32Odd, pregB32);
|
||||
}
|
||||
|
||||
__simd_vf__ void FindKthVFImpl(__ubuf__ uint32_t* kValue, __ubuf__ uint32_t* histogramsBuf, __ubuf__ uint32_t* idxHighBuf, __ubuf__ uint32_t* idxLowBuf)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg pregGE;
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdxLow;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> btmK;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(btmK, kValue);
|
||||
|
||||
MicroAPI::RegTensor<int32_t> idxC;
|
||||
MicroAPI::RegTensor<uint32_t> cout;
|
||||
MicroAPI::RegTensor<uint32_t> sqzIdxLow;
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(4); ++i) {
|
||||
MicroAPI::Arange(idxC, i * 64);
|
||||
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_NORM>(cout, histogramsBuf + i * 64);
|
||||
|
||||
MicroAPI::Compare<uint32_t, CMPMODE::GE>(pregGE, cout, btmK, pregB32);
|
||||
|
||||
MicroAPI::Squeeze<uint32_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxLow, (MicroAPI::RegTensor<uint32_t>&)idxC, pregGE);
|
||||
MicroAPI::StoreUnAlign<uint32_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(idxLowBuf, sqzIdxLow, alignIdxLow);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(idxLowBuf, alignIdxLow);
|
||||
|
||||
MicroAPI::LocalMemBar<AscendC::MicroAPI::MemType::VEC_STORE, AscendC::MicroAPI::MemType::VEC_LOAD>();
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idxHigh;
|
||||
MicroAPI::RegTensor<uint32_t> idxLow;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B8>(idxHigh, idxHighBuf);
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B16>(idxLow, idxLowBuf);
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> idxTmp;
|
||||
MicroAPI::Duplicate(idxTmp, 0xff00);
|
||||
|
||||
MicroAPI::And(idxHigh, idxHigh, (MicroAPI::RegTensor<uint32_t>&)idxTmp, pregB32);
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> idxK;
|
||||
MicroAPI::Add(idxK, idxHigh, idxLow, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM_B16>(kValue, idxK, pregB32);
|
||||
}
|
||||
|
||||
/**
|
||||
输出所有大于的kth-value的Index
|
||||
*/
|
||||
__simd_vf__ void FindIdxGTOutputVFImpl(__ubuf__ uint16_t* outputIdxBuf, __ubuf__ uint16_t* inputValueBuf, uint16_t beginIdx, __ubuf__ uint32_t* kValue, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg poutGT;
|
||||
|
||||
MicroAPI::ClearSpr<AscendC::SpecialPurposeReg::AR>();
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B16>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> vregInput;
|
||||
MicroAPI::RegTensor<int16_t> idxC;
|
||||
MicroAPI::RegTensor<uint16_t> sqzIdxOut;
|
||||
|
||||
for (uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::Arange(idxC, beginIdx + i * 128);
|
||||
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputValueBuf + i * 128);
|
||||
|
||||
MicroAPI::Compare<uint16_t, CMPMODE::GT>(poutGT, vregInput, (MicroAPI::RegTensor<uint16_t>&)kthValue, pregB16);
|
||||
|
||||
MicroAPI::Squeeze<uint16_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxOut, (MicroAPI::RegTensor<uint16_t>&)idxC, poutGT);
|
||||
MicroAPI::StoreUnAlign<uint16_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputIdxBuf, sqzIdxOut, alignIdx);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx);
|
||||
}
|
||||
|
||||
/**
|
||||
输出所有等于的kth-value的Index
|
||||
*/
|
||||
__simd_vf__ void FindIdxEQOutputVFImpl(__ubuf__ uint16_t* outputIdxBuf, __ubuf__ uint16_t* inputValueBuf, uint16_t beginIdx, __ubuf__ uint32_t* kValue, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg poutEQ;
|
||||
|
||||
MicroAPI::UnalignRegForStore alignIdx;
|
||||
|
||||
MicroAPI::RegTensor<uint32_t> kthValue;
|
||||
MicroAPI::LoadAlign<uint32_t, MicroAPI::LoadDist::DIST_BRC_B16>(kthValue, kValue);
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> vregInput;
|
||||
MicroAPI::RegTensor<int16_t> idxC;
|
||||
MicroAPI::RegTensor<uint16_t> sqzIdxOut;
|
||||
|
||||
for(uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::Arange(idxC, beginIdx + i * 128);
|
||||
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_NORM>(vregInput, inputValueBuf + i * 128);
|
||||
|
||||
MicroAPI::Compare<uint16_t, CMPMODE::EQ>(poutEQ, vregInput, (MicroAPI::RegTensor<uint16_t>&)kthValue, pregB16);
|
||||
|
||||
MicroAPI::Squeeze<uint16_t, MicroAPI::GatherMaskMode::STORE_REG>(sqzIdxOut, (MicroAPI::RegTensor<uint16_t>&)idxC, poutEQ);
|
||||
MicroAPI::StoreUnAlign<uint16_t, MicroAPI::PostLiteral::POST_MODE_UPDATE>(outputIdxBuf, sqzIdxOut, alignIdx);
|
||||
}
|
||||
MicroAPI::StoreUnAlignPost(outputIdxBuf, alignIdx);
|
||||
}
|
||||
|
||||
/**
|
||||
输出最终的Value
|
||||
*/
|
||||
__simd_vf__ void FindValueOutputVFImpl(__ubuf__ uint16_t* outputValueBuf, __ubuf__ uint16_t* inputValueBuf, __ubuf__ uint16_t* tmpIdxBuf, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB16 = MicroAPI::CreateMask<uint16_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> tmpIdx;
|
||||
MicroAPI::RegTensor<uint16_t> outputValue;
|
||||
|
||||
for(uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_NORM>(tmpIdx, tmpIdxBuf + i * 128);
|
||||
|
||||
MicroAPI::Gather(outputValue, inputValueBuf, tmpIdx, pregB16);
|
||||
|
||||
MicroAPI::StoreAlign<uint16_t, MicroAPI::StoreDist::DIST_NORM>(outputValueBuf + i * 128, outputValue, pregB16);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
输出最终的Idx
|
||||
*/
|
||||
__simd_vf__ void FindRealIndexVFImpl(__ubuf__ uint32_t* outputIdxBuf, __ubuf__ uint16_t* tmpIdxBuf, __ubuf__ uint32_t* hisIdxBuf, uint32_t topK, uint32_t loopIndex, uint16_t vfLoop)
|
||||
{
|
||||
MicroAPI::MaskReg pregB32 = MicroAPI::CreateMask<uint32_t, MicroAPI::MaskPattern::ALL>();
|
||||
|
||||
MicroAPI::MaskReg pregNow;
|
||||
MicroAPI::MaskReg pregHis;
|
||||
|
||||
MicroAPI::RegTensor<uint16_t> tmpIdx;
|
||||
MicroAPI::RegTensor<uint32_t> outputGatherIdx;
|
||||
MicroAPI::RegTensor<uint32_t> outputAddsIdx;
|
||||
|
||||
for(uint16_t i = 0; i < (uint16_t)(vfLoop); ++i) {
|
||||
MicroAPI::LoadAlign<uint16_t, MicroAPI::LoadDist::DIST_UNPACK_B16>(tmpIdx, tmpIdxBuf + i * 64);
|
||||
|
||||
MicroAPI::Compares<uint32_t, CMPMODE::GT>(pregNow, (MicroAPI::RegTensor<uint32_t>&)tmpIdx, topK - 1, pregB32);
|
||||
MicroAPI::Xor(pregHis, pregNow, pregB32, pregB32);
|
||||
|
||||
MicroAPI::Gather(outputGatherIdx, hisIdxBuf, (MicroAPI::RegTensor<uint32_t>&)tmpIdx, pregHis);
|
||||
MicroAPI::Adds(outputAddsIdx, (MicroAPI::RegTensor<uint32_t>&)tmpIdx, loopIndex, pregNow);
|
||||
|
||||
MicroAPI::Add(outputGatherIdx, outputGatherIdx, outputAddsIdx, pregB32);
|
||||
|
||||
MicroAPI::StoreAlign<uint32_t, MicroAPI::StoreDist::DIST_NORM>(outputIdxBuf + i * 64, outputGatherIdx, pregB32);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief LiTopKVF 对一个validLen的输入进行topk算法,输出idx_tmp
|
||||
* @param tmpIdxLocal Temp阶段输出的TopKIndex;如果s2SeqLen < 16K作为最终输出 validLen * 2B
|
||||
* @param outputValueLocal 如果s2SeqLen > 16K并且是首轮输出Value topK * 2B
|
||||
* @param inputValueLocal 输入Value validLen * 2B
|
||||
* @param histogramsLocal 直方图 256 * 4B
|
||||
* @param idxHighLocal 目标桶高八位 256 * 4B
|
||||
* @param idxLowLocal 目标桶低八位 256 * 4B
|
||||
* @param nkValueLocal 存储next_k的值 64 * 4B
|
||||
* @param topK topK元素
|
||||
* @param validLen 有效元素个数:QLICommon::Align(topkCountAlign256_ + validTrunkLen, (uint32_t)256)
|
||||
*/
|
||||
template<bool ISOUTVALUE> // 是否输出VALUE
|
||||
__aicore__ inline void LiTopKVF(const LocalTensor<uint16_t>& tmpIdxLocal,
|
||||
const LocalTensor<uint16_t>& outputValueLocal,
|
||||
const LocalTensor<uint16_t>& inputValueLocal,
|
||||
const LocalTensor<uint32_t>& histogramsLocal,
|
||||
const LocalTensor<uint32_t>& idxHighLocal,
|
||||
const LocalTensor<uint32_t>& idxLowLocal,
|
||||
const LocalTensor<uint32_t>& nkValueLocal,
|
||||
uint32_t topK,
|
||||
uint32_t validLen)
|
||||
{
|
||||
__ubuf__ uint16_t* tmpIdxBuf = (__ubuf__ uint16_t*)tmpIdxLocal.GetPhyAddr();
|
||||
__ubuf__ uint16_t* outputValueBuf = (__ubuf__ uint16_t*)outputValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint16_t* inputValueBuf = (__ubuf__ uint16_t*)inputValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* histogramsBuf = (__ubuf__ uint32_t*)histogramsLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idxHighBuf = (__ubuf__ uint32_t*)idxHighLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* idxLowBuf = (__ubuf__ uint32_t*)idxLowLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* nkValueBuf = (__ubuf__ uint32_t*)nkValueLocal.GetPhyAddr();
|
||||
|
||||
uint32_t bottomK = validLen - topK + 1;
|
||||
uint32_t beginIdx = 0;
|
||||
bool flag = true;
|
||||
|
||||
const uint16_t repeatSize8 = 256;
|
||||
const uint16_t repeatSize16 = 128;
|
||||
const uint16_t repeatSize32 = 64;
|
||||
|
||||
uint16_t histogramsLoopNum = (validLen + repeatSize8 - 1) / repeatSize8;
|
||||
uint16_t inputLoopNum = (validLen + repeatSize16 - 1) / repeatSize16;
|
||||
uint16_t topkLoopNum = (topK + repeatSize32 - 1) / repeatSize32;
|
||||
uint16_t topkLoopNum16 = (topK + repeatSize16 - 1) / repeatSize16;
|
||||
|
||||
// find kth-value
|
||||
HistogramsHighVFImpl<uint16_t>(histogramsBuf, inputValueBuf, histogramsLoopNum, flag);
|
||||
FindHighTargetBinVFImpl(idxHighBuf, nkValueBuf, histogramsBuf, bottomK);
|
||||
|
||||
HistogramsLowVFImpl<uint16_t>(histogramsBuf, inputValueBuf, idxHighBuf, histogramsLoopNum, flag);
|
||||
FindKthVFImpl(nkValueBuf, histogramsBuf, idxHighBuf, idxLowBuf);
|
||||
|
||||
// filter
|
||||
// 输出大于k-value的值idx
|
||||
FindIdxGTOutputVFImpl(tmpIdxBuf, inputValueBuf, (uint32_t)(0), nkValueBuf, inputLoopNum);
|
||||
// 输出等于k-value的值idx
|
||||
FindIdxEQOutputVFImpl(tmpIdxBuf, inputValueBuf, (uint32_t)(0), nkValueBuf, inputLoopNum);
|
||||
|
||||
// 是否输出Value
|
||||
if constexpr (ISOUTVALUE) {
|
||||
FindValueOutputVFImpl(outputValueBuf, inputValueBuf, tmpIdxBuf, topkLoopNum16);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief 通过idx_tmp gather出实际的TopKIndex,s2SeqLen > 16K才会执行
|
||||
* @param outputIdxLocal 输出Idx 有效:topK * 2B
|
||||
* @param outputValueLocal 输出Value topK * 2B(以后需要输出实际value使用)
|
||||
* @param inputValueLocal 输入Value validLen * 2B
|
||||
* @param tmpIdxLocal 本轮tmpIdx输入 validLen * 2B (0 ~ validLen - 1)
|
||||
* @param hisIdxLocal 上一轮实际Idx输入 有效:topK * 4B
|
||||
* @param topK topK元素个数
|
||||
* @param loopBasicIdx 当前循环需要加上得基准Index
|
||||
* @param validLen 有效元素个数
|
||||
*/
|
||||
__aicore__ inline void LiTopKGatherVF(const LocalTensor<uint32_t>& outputIdxLocal,
|
||||
const LocalTensor<uint16_t>& outputValueLocal,
|
||||
const LocalTensor<uint16_t>& inputValueLocal,
|
||||
const LocalTensor<uint16_t>& tmpIdxLocal,
|
||||
const LocalTensor<uint32_t>& hisIdxLocal,
|
||||
uint32_t topK,
|
||||
uint32_t loopBasicIdx,
|
||||
uint32_t validLen)
|
||||
{
|
||||
__ubuf__ uint32_t* outputIdxBuf = (__ubuf__ uint32_t*)outputIdxLocal.GetPhyAddr();
|
||||
__ubuf__ uint16_t* outputValueBuf = (__ubuf__ uint16_t*)outputValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint16_t* inputValueBuf = (__ubuf__ uint16_t*)inputValueLocal.GetPhyAddr();
|
||||
__ubuf__ uint16_t* tmpIdxBuf = (__ubuf__ uint16_t*)tmpIdxLocal.GetPhyAddr();
|
||||
__ubuf__ uint32_t* hisIdxBuf = (__ubuf__ uint32_t*)hisIdxLocal.GetPhyAddr();
|
||||
|
||||
const uint16_t repeatSize32 = 64;
|
||||
const uint16_t repeatSize16 = 128;
|
||||
uint16_t topkLoopNum16 = (topK + repeatSize16 - 1) / repeatSize16;
|
||||
uint16_t topkLoopNum32 = (topK + repeatSize32 - 1) / repeatSize32;
|
||||
|
||||
FindRealIndexVFImpl(outputIdxBuf, tmpIdxBuf, hisIdxBuf, topK, loopBasicIdx, topkLoopNum32);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,55 @@
|
||||
/**
|
||||
* 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 vllm_quant_lightning_indexer.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "lib/matmul_intf.h"
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
#include "arch35/quant_lightning_indexer_kernel.h"
|
||||
#else
|
||||
#include "arch32/quant_lightning_indexer_kernel.h"
|
||||
#endif
|
||||
#include "vllm_quant_lightning_indexer_template_tiling_key.h"
|
||||
using namespace QLIKernel;
|
||||
using namespace optiling::detail;
|
||||
|
||||
#define INVOKE_LI_NO_KFC_OP_IMPL(templateClass, ...) \
|
||||
do { \
|
||||
templateClass<QLIType<__VA_ARGS__>> op; \
|
||||
GET_TILING_DATA_WITH_STRUCT(QLITilingData, tiling_data_in, tiling); \
|
||||
const QLITilingData *__restrict tiling_data = &tiling_data_in; \
|
||||
op.Init(query, key, weights, queryScale, keyScale, actualSeqLengthsQ, actualSeqLengthsK, blocktable, \
|
||||
metadata, sparseIndices, user, tiling_data, &tPipe); \
|
||||
op.Process(); \
|
||||
} while (0)
|
||||
|
||||
template <int DT_Q, int DT_K, int DT_OUT, int PAGE_ATTENTION, int Q_LAYOUT_T, int K_LAYOUT_T>
|
||||
__global__ __aicore__ void vllm_quant_lightning_indexer(__gm__ uint8_t *query, __gm__ uint8_t *key, __gm__ uint8_t *weights,
|
||||
__gm__ uint8_t *queryScale, __gm__ uint8_t *keyScale,
|
||||
__gm__ uint8_t *actualSeqLengthsQ, __gm__ uint8_t *actualSeqLengthsK,
|
||||
__gm__ uint8_t *blocktable, __gm__ uint8_t *metadata,
|
||||
__gm__ uint8_t *sparseIndices, __gm__ uint8_t *sparseValues,
|
||||
__gm__ uint8_t *workspace, __gm__ uint8_t *tiling)
|
||||
{
|
||||
TPipe tPipe;
|
||||
__gm__ uint8_t *user = GetUserWorkspace(workspace);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
INVOKE_LI_NO_KFC_OP_IMPL(QLIPreload, fp8_e4m3fn_t, fp8_e4m3fn_t, float, uint16_t, int32_t,
|
||||
PAGE_ATTENTION, LI_LAYOUT(Q_LAYOUT_T), LI_LAYOUT(K_LAYOUT_T));
|
||||
#else
|
||||
INVOKE_LI_NO_KFC_OP_IMPL(QLIPreload, int8_t, int8_t, int32_t,
|
||||
PAGE_ATTENTION, LI_LAYOUT(Q_LAYOUT_T), LI_LAYOUT(K_LAYOUT_T));
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vllm_quant_lightning_indexer_metadata.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_METADATA_H
|
||||
#define QUANT_LIGHTNING_INDEXER_METADATA_H
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
// Constants
|
||||
inline constexpr uint32_t AIC_CORE_NUM = 36;
|
||||
inline constexpr uint32_t AIV_CORE_NUM = 72;
|
||||
constexpr uint32_t QLI_META_SIZE = 1024;
|
||||
using QLI_METADATA_T = int32_t;
|
||||
|
||||
inline constexpr uint32_t LI_METADATA_SIZE = 8;
|
||||
inline constexpr uint32_t LD_METADATA_SIZE = 8;
|
||||
|
||||
// LI Metadata Index Definitions
|
||||
inline constexpr uint32_t LI_CORE_ENABLE_INDEX = 0;
|
||||
inline constexpr uint32_t LI_BN2_START_INDEX = 1;
|
||||
inline constexpr uint32_t LI_M_START_INDEX = 2;
|
||||
inline constexpr uint32_t LI_S2_START_INDEX = 3;
|
||||
inline constexpr uint32_t LI_BN2_END_INDEX = 4;
|
||||
inline constexpr uint32_t LI_M_END_INDEX = 5;
|
||||
inline constexpr uint32_t LI_S2_END_INDEX = 6;
|
||||
inline constexpr uint32_t LI_FIRST_LD_DATA_WORKSPACE_IDX_INDEX = 7;
|
||||
|
||||
// LD Metadata Index Definitions
|
||||
inline constexpr uint32_t LD_CORE_ENABLE_INDEX = 0;
|
||||
inline constexpr uint32_t LD_BN2_IDX_INDEX = 1;
|
||||
inline constexpr uint32_t LD_M_IDX_INDEX = 2;
|
||||
inline constexpr uint32_t LD_WORKSPACE_IDX_INDEX = 3;
|
||||
inline constexpr uint32_t LD_WORKSPACE_NUM_INDEX = 4;
|
||||
inline constexpr uint32_t LD_M_START_INDEX = 5;
|
||||
inline constexpr uint32_t LD_M_NUM_INDEX = 6;
|
||||
|
||||
/**
|
||||
* @brief 获取属性的绝对索引
|
||||
* @param coreIdx 核索引
|
||||
* @param metaIdx 元数据索引
|
||||
* @param isAIV 是否为AIV数据,默认为false
|
||||
* @return 返回属性的绝对索引
|
||||
*/
|
||||
#ifdef __CCE_AICORE__
|
||||
__aicore__ inline uint32_t GetAttrAbsIndex(uint32_t coreIdx, uint32_t metaIdx, bool isAIV=false)
|
||||
{
|
||||
if (isAIV) {
|
||||
return LI_METADATA_SIZE * AIC_CORE_NUM + LD_METADATA_SIZE * coreIdx + metaIdx;
|
||||
} else {
|
||||
return LI_METADATA_SIZE * coreIdx + metaIdx;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
namespace detail {
|
||||
struct QliMetaData {
|
||||
uint32_t LIMetadata[AIC_CORE_NUM][LI_METADATA_SIZE];
|
||||
uint32_t LDMetadata[AIV_CORE_NUM][LD_METADATA_SIZE];
|
||||
};
|
||||
};
|
||||
|
||||
static_assert(QLI_META_SIZE * sizeof(QLI_METADATA_T) >= sizeof(detail::QliMetaData));
|
||||
};
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,79 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vllm_quant_lightning_indexer_template_tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_TEMPLATE_TILING_KEY_H
|
||||
#define QUANT_LIGHTNING_INDEXER_TEMPLATE_TILING_KEY_H
|
||||
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define QLI_TPL_INT8 2
|
||||
#define QLI_TPL_INT32 3
|
||||
#define QLI_TPL_FLOAT32_E4M3FN 36
|
||||
#define QLI_LAYOUT_BSND 0
|
||||
#define QLI_LAYOUT_TND 1
|
||||
#define QLI_LAYOUT_PA_BSND 2
|
||||
|
||||
#define ASCENDC_TPL_4_BW 4
|
||||
|
||||
// 模板参数支持的范围定义
|
||||
#if (__CCE_AICORE__ == 310)
|
||||
ASCENDC_TPL_ARGS_DECL(VllmQuantLightningIndexer, // 算子OpType
|
||||
ASCENDC_TPL_DTYPE_DECL(DT_Q, QLI_TPL_FLOAT32_E4M3FN), ASCENDC_TPL_DTYPE_DECL(DT_K, QLI_TPL_FLOAT32_E4M3FN),
|
||||
ASCENDC_TPL_DTYPE_DECL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 1, 0),
|
||||
ASCENDC_TPL_UINT_DECL(Q_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND,
|
||||
QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_DECL(K_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QLI_LAYOUT_BSND, QLI_LAYOUT_TND, QLI_LAYOUT_PA_BSND), );
|
||||
// 支持的模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_FLOAT32_E4M3FN), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_FLOAT32_E4M3FN),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND, QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_PA_BSND), ),
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_FLOAT32_E4M3FN), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_FLOAT32_E4M3FN),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND), ),
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_FLOAT32_E4M3FN), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_FLOAT32_E4M3FN),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_TND), ), );
|
||||
#else
|
||||
ASCENDC_TPL_ARGS_DECL(VllmQuantLightningIndexer, // 算子OpType
|
||||
ASCENDC_TPL_DTYPE_DECL(DT_Q, QLI_TPL_INT8), ASCENDC_TPL_DTYPE_DECL(DT_K, QLI_TPL_INT8),
|
||||
ASCENDC_TPL_DTYPE_DECL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_DECL(PAGE_ATTENTION, 1, 0),
|
||||
ASCENDC_TPL_UINT_DECL(Q_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND,
|
||||
QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_DECL(K_LAYOUT_T, ASCENDC_TPL_4_BW, ASCENDC_TPL_UI_LIST,
|
||||
QLI_LAYOUT_BSND, QLI_LAYOUT_TND, QLI_LAYOUT_PA_BSND), );
|
||||
// 支持的模板参数组合
|
||||
// 用于调用GET_TPL_TILING_KEY获取TilingKey时,接口内部校验TilingKey是否合法
|
||||
ASCENDC_TPL_SEL(
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_INT8),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 1),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND, QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_PA_BSND), ),
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_INT8),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_BSND), ),
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DTYPE_SEL(DT_Q, QLI_TPL_INT8), ASCENDC_TPL_DTYPE_SEL(DT_K, QLI_TPL_INT8),
|
||||
ASCENDC_TPL_DTYPE_SEL(DT_OUT, QLI_TPL_INT32), ASCENDC_TPL_BOOL_SEL(PAGE_ATTENTION, 0),
|
||||
ASCENDC_TPL_UINT_SEL(Q_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_TND),
|
||||
ASCENDC_TPL_UINT_SEL(K_LAYOUT_T, ASCENDC_TPL_UI_LIST, QLI_LAYOUT_TND), ), );
|
||||
#endif
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user