19
csrc/moe/transpose_kv_cache_by_block/CMakeLists.txt
Normal file
19
csrc/moe/transpose_kv_cache_by_block/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
24
csrc/moe/transpose_kv_cache_by_block/op_host/CMakeLists.txt
Normal file
24
csrc/moe/transpose_kv_cache_by_block/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,24 @@
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
transpose_kv_cache_by_block_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME TransposeKvCacheByBlock
|
||||
OPTIONS
|
||||
--cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-mllvm -cce-aicore-hoist-movemask=false
|
||||
--op_relocatable_kernel_binary=true
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE transpose_kv_cache_by_block ACLNNTYPE aclnn)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class TransposeKvCacheByBlock : public OpDef {
|
||||
public:
|
||||
explicit TransposeKvCacheByBlock(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("KCache")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("VCache")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("blockIDs")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Attr("blockSize").Int();
|
||||
this->Attr("headNum").Int();
|
||||
this->Attr("headDim").Int();
|
||||
this->Attr("splitNum").Int();
|
||||
this->Attr("layerNum").Int();
|
||||
|
||||
this->AICore().AddConfig("ascend910b");
|
||||
this->AICore().AddConfig("ascend910_93");
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(TransposeKvCacheByBlock);
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under 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 transpose_kv_cache_by_block_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
#include "error/ops_error.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
|
||||
static ge::graphStatus InferShapeTransposeKvCacheByBlock(gert::InferShapeContext* context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDataTypeTransposeKvCacheByBlock(gert::InferDataTypeContext *context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(TransposeKvCacheByBlock)
|
||||
.InferShape(InferShapeTransposeKvCacheByBlock)
|
||||
.InferDataType(InferDataTypeTransposeKvCacheByBlock);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,182 @@
|
||||
#include "transpose_kv_cache_by_block_tiling.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "log/ops_log.h"
|
||||
#include <algorithm>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
constexpr uint64_t DATA_SIZE = 2;
|
||||
constexpr uint64_t BLOCK_SIZE = 32;
|
||||
constexpr uint64_t DB_ON = 2;
|
||||
|
||||
constexpr uint32_t FULL_LOAD = 0;
|
||||
constexpr uint32_t SPLIT_BLOCK_SIZE_ALIGNED_AND_DB = 1;
|
||||
constexpr uint32_t SPLIT_BLOCK_SIZE_UNALIGNED_AND_DB = 3;
|
||||
constexpr uint32_t SPLIT_BLOCK_SIZE_ALIGNED_AND_NOT_DB = 2;
|
||||
constexpr uint32_t SPLIT_BLOCK_SIZE_UNALIGNED_AND_NOT_DB = 4;
|
||||
|
||||
void findFactorsOptimized(std::vector<int64_t> &factors, int64_t n) {
|
||||
|
||||
for (int64_t i = 1; i * i <= n; i++) {
|
||||
if (n % i == 0) {
|
||||
factors.push_back(i);
|
||||
|
||||
if (i != n / i) {
|
||||
factors.push_back(n / i);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
sort(factors.begin(), factors.end());
|
||||
}
|
||||
|
||||
ge::graphStatus CalTiling(gert::TilingContext* context, TransposeKvCacheByBlockTilingData &tiling)
|
||||
{
|
||||
fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
|
||||
OPS_LOG_E_IF_NULL(context, platformInfoPtr, return ge::GRAPH_FAILED);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
int64_t useCoreNum = ascendcPlatform.GetCoreNumAiv();
|
||||
uint64_t ubSize;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
|
||||
|
||||
auto attr = context->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context, attr, return ge::GRAPH_FAILED);
|
||||
const int64_t* blockSizePtr = attr->GetAttrPointer<int64_t>(0);
|
||||
const int64_t* headNumPtr = attr->GetAttrPointer<int64_t>(1);
|
||||
const int64_t* headDimPtr = attr->GetAttrPointer<int64_t>(2);
|
||||
const int64_t* splitNumPtr = attr->GetAttrPointer<int64_t>(3);
|
||||
const int64_t* layerNumPtr = attr->GetAttrPointer<int64_t>(4);
|
||||
OPS_CHECK(blockSizePtr == nullptr || headNumPtr == nullptr || headDimPtr == nullptr ||
|
||||
splitNumPtr == nullptr || layerNumPtr == nullptr,
|
||||
OPS_LOG_E(context->GetNodeName(), "Get attr failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto blockIDsTensor = context->GetDynamicInputTensor(2, 0);
|
||||
OPS_LOG_E_IF_NULL(context, blockIDsTensor, return ge::GRAPH_FAILED);
|
||||
|
||||
gert::Shape blockIDsTensorShape = blockIDsTensor->GetStorageShape();
|
||||
int64_t calBlockNum = static_cast<int64_t>(blockIDsTensorShape.GetDim(0));
|
||||
|
||||
tiling.set_calBlockNum(static_cast<uint32_t>(calBlockNum));
|
||||
|
||||
int64_t blockSize = *blockSizePtr;
|
||||
int64_t headNum = *headNumPtr;
|
||||
int64_t headDim = *headDimPtr;
|
||||
int64_t splitNum = *splitNumPtr;
|
||||
int64_t layerNum = *layerNumPtr;
|
||||
uint32_t tilingKey = FULL_LOAD;
|
||||
|
||||
if (headDim * DATA_SIZE % BLOCK_SIZE != 0) {
|
||||
OPS_LOG_E(context, "headDim * DATA_SIZE must be a multiple of 32 bytes.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
std::vector<int64_t> factors;
|
||||
findFactorsOptimized(factors, useCoreNum);
|
||||
|
||||
uint32_t factorIndex = 0;
|
||||
bool findSplitNum = true;
|
||||
int64_t blockSizeSplitNum = factors[factorIndex];
|
||||
uint64_t dataSizeloadOnce = blockSize * headNum * headDim * DATA_SIZE;
|
||||
// if can full load, not split blockSize and db
|
||||
if (dataSizeloadOnce > ubSize) {
|
||||
tilingKey = SPLIT_BLOCK_SIZE_ALIGNED_AND_DB;
|
||||
// split blockSize and db
|
||||
while (dataSizeloadOnce > (ubSize / DB_ON)) {
|
||||
factorIndex += 1;
|
||||
if (factorIndex == factors.size()) {
|
||||
tilingKey = FULL_LOAD;
|
||||
findSplitNum = false;
|
||||
break;
|
||||
}
|
||||
blockSizeSplitNum = factors[factorIndex];
|
||||
dataSizeloadOnce = ((blockSize + blockSizeSplitNum - 1) / blockSizeSplitNum) * headNum * headDim * DATA_SIZE;
|
||||
}
|
||||
if (tilingKey == SPLIT_BLOCK_SIZE_ALIGNED_AND_DB && (blockSize % blockSizeSplitNum != 0)) {
|
||||
tilingKey = SPLIT_BLOCK_SIZE_UNALIGNED_AND_DB;
|
||||
}
|
||||
}
|
||||
|
||||
if (!findSplitNum) {
|
||||
tilingKey = SPLIT_BLOCK_SIZE_ALIGNED_AND_NOT_DB;
|
||||
// split blockSize but not db
|
||||
findSplitNum = true;
|
||||
factorIndex = 0;
|
||||
blockSizeSplitNum = factors[factorIndex];
|
||||
dataSizeloadOnce = blockSize * headNum * headDim * DATA_SIZE;
|
||||
while (dataSizeloadOnce > ubSize) {
|
||||
factorIndex += 1;
|
||||
if (factorIndex == factors.size()) {
|
||||
tilingKey = FULL_LOAD;
|
||||
findSplitNum = false;
|
||||
break;
|
||||
}
|
||||
blockSizeSplitNum = factors[factorIndex];
|
||||
dataSizeloadOnce = ((blockSize + blockSizeSplitNum - 1) / blockSizeSplitNum) * headNum * headDim * DATA_SIZE;
|
||||
}
|
||||
if (tilingKey == SPLIT_BLOCK_SIZE_ALIGNED_AND_NOT_DB && (blockSize % blockSizeSplitNum != 0)) {
|
||||
tilingKey = SPLIT_BLOCK_SIZE_UNALIGNED_AND_NOT_DB;
|
||||
}
|
||||
}
|
||||
|
||||
// headNum * headDim too large
|
||||
if (!findSplitNum) {
|
||||
OPS_LOG_E(context, "headNum * headDim * sizeof(half) > ubSize "
|
||||
"or blockSize * headNum * headDim * sizeof(half) > ubSize * vectorCoreNum. "
|
||||
"Currently, splitting headNum or headDim is not supported.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
tiling.set_blockSizePerTime(static_cast<uint32_t>((blockSize + blockSizeSplitNum - 1) / blockSizeSplitNum));
|
||||
tiling.set_blockSizePerTimeTail(static_cast<uint32_t>(blockSize % blockSizeSplitNum));
|
||||
tiling.set_blockSizeSplitNum(static_cast<uint32_t>(blockSizeSplitNum));
|
||||
|
||||
tiling.set_blockSize(static_cast<uint32_t>(blockSize));
|
||||
tiling.set_headNum(static_cast<uint32_t>(headNum));
|
||||
tiling.set_headDim(static_cast<uint32_t>(headDim));
|
||||
tiling.set_splitNum(static_cast<uint32_t>(splitNum));
|
||||
tiling.set_layerNum(static_cast<uint32_t>(layerNum));
|
||||
|
||||
int64_t totalRound = layerNum * calBlockNum;
|
||||
|
||||
if ((totalRound * blockSizeSplitNum) < useCoreNum) {
|
||||
useCoreNum = totalRound * blockSizeSplitNum;
|
||||
}
|
||||
int64_t blockPerCore = totalRound / (useCoreNum / blockSizeSplitNum);
|
||||
int64_t tailCoreNum = totalRound % (useCoreNum / blockSizeSplitNum);
|
||||
|
||||
tiling.set_useCoreNum(static_cast<uint32_t>(useCoreNum));
|
||||
tiling.set_blockPerCore(static_cast<uint32_t>(blockPerCore));
|
||||
tiling.set_tailCoreNum(static_cast<uint32_t>(tailCoreNum));
|
||||
context->SetBlockDim(useCoreNum);
|
||||
context->SetTilingKey(tilingKey);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
static ge::graphStatus TransposeKvCacheByBlockTilingFunc(gert::TilingContext* context)
|
||||
{
|
||||
|
||||
TransposeKvCacheByBlockTilingData tiling;
|
||||
auto status = CalTiling(context, tiling);
|
||||
OP_CHECK(status != ge::GRAPH_SUCCESS, OPS_LOG_E(context->GetNodeName(), "Cal tiling failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
|
||||
context->GetRawTilingData()->SetDataSize(tiling.GetDataSize());
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
struct TransposeKvCacheByBlockCompileInfo {};
|
||||
ge::graphStatus TilingParseForTransposeKvCacheByBlock(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(TransposeKvCacheByBlock)
|
||||
.Tiling(TransposeKvCacheByBlockTilingFunc)
|
||||
.TilingParse<TransposeKvCacheByBlockCompileInfo>(TilingParseForTransposeKvCacheByBlock);
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
#include "register/tilingdata_base.h"
|
||||
|
||||
namespace optiling {
|
||||
BEGIN_TILING_DATA_DEF(TransposeKvCacheByBlockTilingData)
|
||||
// shape info
|
||||
// TILING_DATA_FIELD_DEF(uint32_t, blockNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockSize);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, headNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, headDim);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, splitNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, layerNum);
|
||||
// tiling info
|
||||
TILING_DATA_FIELD_DEF(uint32_t, useCoreNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockPerCore);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tailCoreNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, calBlockNum);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockSizePerTime);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockSizePerTimeTail);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, blockSizeSplitNum);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(TransposeKvCacheByBlock, TransposeKvCacheByBlockTilingData)
|
||||
}
|
||||
16
csrc/moe/transpose_kv_cache_by_block/op_kernel/common.h
Normal file
16
csrc/moe/transpose_kv_cache_by_block/op_kernel/common.h
Normal file
@@ -0,0 +1,16 @@
|
||||
#include "kernel_operator.h"
|
||||
using namespace AscendC;
|
||||
|
||||
#ifndef __OP_KERNEL_KV_CACHE_TRANSPOSE_H__
|
||||
#define __OP_KERNEL_KV_CACHE_TRANSPOSE_H__
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline __gm__ T* GetTensorAddr(uint16_t index, GM_ADDR tensorPtr) {
|
||||
__gm__ uint64_t* dataAddr = reinterpret_cast<__gm__ uint64_t*>(tensorPtr);
|
||||
// The offset of the data address from the first address.
|
||||
uint64_t tensorPtrOffset = *dataAddr;
|
||||
// Moving 3 bits to the right means dividing by sizeof(uint64 t).
|
||||
__gm__ uint64_t* retPtr = dataAddr + (tensorPtrOffset >> 3);
|
||||
return reinterpret_cast<__gm__ T*>(*(retPtr + index));
|
||||
}
|
||||
#endif
|
||||
141
csrc/moe/transpose_kv_cache_by_block/op_kernel/full_load.h
Normal file
141
csrc/moe/transpose_kv_cache_by_block/op_kernel/full_load.h
Normal file
@@ -0,0 +1,141 @@
|
||||
#include "common.h"
|
||||
|
||||
template <typename T>
|
||||
class TransposeKvCacheByBlockKernelFullLoad {
|
||||
protected:
|
||||
TQueBind<TPosition::VECIN, TPosition::VECOUT, 1> vecInQueue_;
|
||||
GlobalTensor<T> kCacheGm_;
|
||||
GlobalTensor<T> vCacheGm_;
|
||||
GlobalTensor<int64_t> blockIDsGm_;
|
||||
|
||||
GM_ADDR kCachePtr_;
|
||||
GM_ADDR vCachePtr_;
|
||||
|
||||
// shape info
|
||||
uint32_t blockNum_;
|
||||
uint32_t blockSize_;
|
||||
uint32_t headNum_;
|
||||
uint32_t headDim_;
|
||||
uint32_t splitNum_;
|
||||
uint32_t layerNum_;
|
||||
// tiling info
|
||||
uint32_t useCoreNum_;
|
||||
uint32_t blockPerCore_;
|
||||
uint32_t tailCoreNum_;
|
||||
uint32_t calBlockNum_;
|
||||
|
||||
uint32_t srcFactor_;
|
||||
uint32_t dstFactor_;
|
||||
uint32_t copyOutLength_;
|
||||
uint32_t dataBlockSize_;
|
||||
|
||||
__aicore__ inline void CopyIn(GlobalTensor<T> &cacheGm, uint32_t offsetBlock, DataCopyParams &repeatParams) {
|
||||
LocalTensor<T> cacheLocal = vecInQueue_.AllocTensor<T>();
|
||||
for (uint32_t i = 0; i < splitNum_; ++i) {
|
||||
DataCopy(cacheLocal[i * dstFactor_], cacheGm[i * srcFactor_ + offsetBlock], repeatParams);
|
||||
}
|
||||
vecInQueue_.EnQue(cacheLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOut(GlobalTensor<T> &cacheGm, uint32_t offsetBlock) {
|
||||
LocalTensor<T> cacheLocal = vecInQueue_.DeQue<T>();
|
||||
DataCopy(cacheGm[offsetBlock], cacheLocal, copyOutLength_);
|
||||
vecInQueue_.FreeTensor(cacheLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetGlobalBuffers(uint32_t layerId) {
|
||||
kCacheGm_.SetGlobalBuffer(GetTensorAddr<T>(layerId, kCachePtr_));
|
||||
vCacheGm_.SetGlobalBuffer(GetTensorAddr<T>(layerId, vCachePtr_));
|
||||
}
|
||||
|
||||
__aicore__ inline void Caloffset(uint32_t &startBlock, uint32_t &endBlock, uint32_t &startLayer, uint32_t &endLayer) {
|
||||
uint32_t blockIdx = GetBlockIdx();
|
||||
uint32_t curBlockStart;
|
||||
uint32_t curBlocknum;
|
||||
|
||||
if (blockIdx < tailCoreNum_) {
|
||||
curBlockStart = blockIdx * (blockPerCore_ + 1);
|
||||
curBlocknum = blockPerCore_ + 1;
|
||||
} else {
|
||||
curBlockStart = blockIdx * blockPerCore_ + tailCoreNum_;
|
||||
curBlocknum = blockPerCore_;
|
||||
}
|
||||
uint32_t curBlockEnd = curBlockStart + curBlocknum;
|
||||
startBlock = curBlockStart / layerNum_;
|
||||
startLayer = curBlockStart % layerNum_;
|
||||
endBlock = (curBlockEnd + layerNum_ - 1) / layerNum_;
|
||||
endLayer = curBlockEnd % layerNum_;
|
||||
if (endLayer == 0) {
|
||||
endLayer = layerNum_;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR KCache, GM_ADDR VCache, GM_ADDR blockIDs,
|
||||
TransposeKvCacheByBlockTilingData* tilingData, TPipe* tPipe) {
|
||||
kCachePtr_ = KCache;
|
||||
vCachePtr_ = VCache;
|
||||
blockIDsGm_.SetGlobalBuffer((__gm__ int64_t*)blockIDs);
|
||||
|
||||
// shape info
|
||||
blockSize_ = tilingData->blockSize;
|
||||
headNum_ = tilingData->headNum;
|
||||
headDim_ = tilingData->headDim;
|
||||
splitNum_ = tilingData->splitNum;
|
||||
layerNum_ = tilingData->layerNum;
|
||||
// tiling info
|
||||
useCoreNum_ = tilingData->useCoreNum;
|
||||
blockPerCore_ = tilingData->blockPerCore;
|
||||
tailCoreNum_ = tilingData->tailCoreNum;
|
||||
calBlockNum_ = tilingData->calBlockNum;
|
||||
|
||||
tPipe->InitBuffer(vecInQueue_, 1, TOTAL_UB_SIZE);
|
||||
srcFactor_ = blockSize_ * headNum_ / splitNum_ * headDim_;
|
||||
dstFactor_ = headNum_ / splitNum_ * headDim_;
|
||||
copyOutLength_ = blockSize_ * headNum_ * headDim_;
|
||||
dataBlockSize_ = static_cast<uint32_t>(AscendC::GetDataBlockSizeInBytes());
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
DataCopyParams repeatParams;
|
||||
repeatParams.blockCount = blockSize_;
|
||||
repeatParams.blockLen = headNum_ / splitNum_ * headDim_ * sizeof(T) / dataBlockSize_;
|
||||
repeatParams.srcStride = 0;
|
||||
repeatParams.dstStride = (headNum_ * headDim_ - headNum_ / splitNum_ * headDim_) * sizeof(T) / dataBlockSize_;
|
||||
|
||||
uint32_t startBlock;
|
||||
uint32_t endBlock;
|
||||
uint32_t startLayer;
|
||||
uint32_t endLayer;
|
||||
|
||||
Caloffset(startBlock, endBlock, startLayer, endLayer);
|
||||
for (uint32_t i = startBlock; i < endBlock; ++i) {
|
||||
int64_t blockId = blockIDsGm_.GetValue(i);
|
||||
uint32_t offsetBlock = blockId * blockSize_ * headNum_ * headDim_;
|
||||
uint32_t realStartLayer;
|
||||
uint32_t realEndLayer;
|
||||
if (i == startBlock) {
|
||||
realStartLayer = startLayer;
|
||||
} else {
|
||||
realStartLayer = 0;
|
||||
}
|
||||
|
||||
if (i == (endBlock - 1)) {
|
||||
realEndLayer = endLayer;
|
||||
} else {
|
||||
realEndLayer = layerNum_;
|
||||
}
|
||||
for (uint32_t layerId = realStartLayer; layerId < realEndLayer; ++layerId) {
|
||||
SetGlobalBuffers(layerId);
|
||||
|
||||
CopyIn(kCacheGm_, offsetBlock, repeatParams);
|
||||
CopyOut(kCacheGm_, offsetBlock);
|
||||
|
||||
CopyIn(vCacheGm_, offsetBlock, repeatParams);
|
||||
CopyOut(vCacheGm_, offsetBlock);
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
};
|
||||
190
csrc/moe/transpose_kv_cache_by_block/op_kernel/general.h
Normal file
190
csrc/moe/transpose_kv_cache_by_block/op_kernel/general.h
Normal file
@@ -0,0 +1,190 @@
|
||||
#include "common.h"
|
||||
|
||||
template <typename T, uint32_t DB, bool needHandleUnFactorSplit>
|
||||
class TransposeKvCacheByBlockKernelGeneral {
|
||||
protected:
|
||||
TQueBind<TPosition::VECIN, TPosition::VECOUT, 1> queBind_;
|
||||
GlobalTensor<T> kCacheGm_;
|
||||
GlobalTensor<T> vCacheGm_;
|
||||
GlobalTensor<int64_t> blockIDsGm_;
|
||||
|
||||
GM_ADDR kCachePtr_;
|
||||
GM_ADDR vCachePtr_;
|
||||
|
||||
// shape info
|
||||
uint32_t blockNum_;
|
||||
uint32_t blockSize_;
|
||||
uint32_t headNum_;
|
||||
uint32_t headDim_;
|
||||
uint32_t splitNum_;
|
||||
uint32_t layerNum_;
|
||||
uint32_t headNumSplited_;
|
||||
uint32_t blockSizeSplitNum_;
|
||||
|
||||
// tiling info
|
||||
uint32_t useCoreNum_;
|
||||
uint32_t blockPerCore_;
|
||||
uint32_t tailCoreNum_;
|
||||
uint32_t calBlockNum_;
|
||||
|
||||
uint32_t srcFactor_;
|
||||
uint32_t dstFactor_;
|
||||
uint32_t copyOutLength_;
|
||||
|
||||
uint32_t blockSizePerTime_;
|
||||
uint32_t blockSizePerTimeTail_;
|
||||
|
||||
uint32_t blockIdx_;
|
||||
uint32_t dataBlockSize_;
|
||||
bool needSync_;
|
||||
|
||||
__aicore__ inline void CopyIn(GlobalTensor<T> &cacheGm, uint32_t offsetBlock, DataCopyParams &repeatParams) {
|
||||
LocalTensor<T> cacheLocal = queBind_.AllocTensor<T>();
|
||||
for (uint32_t i = 0; i < splitNum_; ++i) {
|
||||
DataCopy(cacheLocal[i * dstFactor_], cacheGm[i * srcFactor_ + offsetBlock], repeatParams);
|
||||
}
|
||||
queBind_.EnQue(cacheLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void CopyOut(GlobalTensor<T> &cacheGm, uint32_t offsetBlock) {
|
||||
LocalTensor<T> cacheLocal = queBind_.DeQue<T>();
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE2>(0x8);
|
||||
AscendC::CrossCoreWaitFlag(0x8);
|
||||
DataCopy(cacheGm[offsetBlock], cacheLocal, copyOutLength_);
|
||||
queBind_.FreeTensor(cacheLocal);
|
||||
}
|
||||
|
||||
__aicore__ inline void SetGlobalBuffers(uint32_t layerId) {
|
||||
kCacheGm_.SetGlobalBuffer(GetTensorAddr<T>(layerId, kCachePtr_));
|
||||
vCacheGm_.SetGlobalBuffer(GetTensorAddr<T>(layerId, vCachePtr_));
|
||||
}
|
||||
|
||||
__aicore__ inline void Caloffset(uint32_t &startBlock, uint32_t &endBlock, uint32_t &startLayer, uint32_t &endLayer) {
|
||||
|
||||
uint32_t curBlockStart;
|
||||
uint32_t curBlocknum;
|
||||
uint32_t groupBlockIdx = blockIdx_ / blockSizeSplitNum_;
|
||||
if (groupBlockIdx < tailCoreNum_) {
|
||||
needSync_ = false;
|
||||
curBlockStart = groupBlockIdx * (blockPerCore_ + 1);
|
||||
curBlocknum = blockPerCore_ + 1;
|
||||
} else {
|
||||
needSync_ = true;
|
||||
curBlockStart = groupBlockIdx * blockPerCore_ + tailCoreNum_;
|
||||
curBlocknum = blockPerCore_;
|
||||
}
|
||||
uint32_t curBlockEnd = curBlockStart + curBlocknum;
|
||||
startBlock = curBlockStart / layerNum_;
|
||||
startLayer = curBlockStart % layerNum_;
|
||||
endBlock = (curBlockEnd + layerNum_ - 1) / layerNum_;
|
||||
endLayer = curBlockEnd % layerNum_;
|
||||
if (endLayer == 0) {
|
||||
endLayer = layerNum_;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR KCache, GM_ADDR VCache, GM_ADDR blockIDs,
|
||||
TransposeKvCacheByBlockTilingData* tilingData, TPipe* tPipe) {
|
||||
kCachePtr_ = KCache;
|
||||
vCachePtr_ = VCache;
|
||||
blockIDsGm_.SetGlobalBuffer((__gm__ int64_t*)blockIDs);
|
||||
blockIdx_ = GetBlockIdx();
|
||||
// shape info
|
||||
blockSize_ = tilingData->blockSize;
|
||||
headNum_ = tilingData->headNum;
|
||||
headDim_ = tilingData->headDim;
|
||||
splitNum_ = tilingData->splitNum;
|
||||
layerNum_ = tilingData->layerNum;
|
||||
// tiling info
|
||||
useCoreNum_ = tilingData->useCoreNum;
|
||||
blockPerCore_ = tilingData->blockPerCore;
|
||||
tailCoreNum_ = tilingData->tailCoreNum;
|
||||
calBlockNum_ = tilingData->calBlockNum;
|
||||
blockSizeSplitNum_ = tilingData->blockSizeSplitNum;
|
||||
blockSizePerTime_ = tilingData->blockSizePerTime;
|
||||
blockSizePerTimeTail_ = tilingData->blockSizePerTimeTail;
|
||||
headNumSplited_ = headNum_ / splitNum_;
|
||||
|
||||
tPipe->InitBuffer(queBind_, DB, TOTAL_UB_SIZE / DB);
|
||||
srcFactor_ = blockSize_ * headNumSplited_ * headDim_;
|
||||
dstFactor_ = headNumSplited_ * headDim_;
|
||||
copyOutLength_ = blockSizePerTime_ * headNum_ * headDim_;
|
||||
dataBlockSize_ = static_cast<uint32_t>(AscendC::GetDataBlockSizeInBytes());
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
DataCopyParams repeatParams;
|
||||
repeatParams.blockCount = blockSizePerTime_;
|
||||
repeatParams.blockLen = headNumSplited_ * headDim_ * sizeof(T) / dataBlockSize_;
|
||||
repeatParams.srcStride = 0;
|
||||
repeatParams.dstStride = (headNum_ * headDim_ - headNumSplited_ * headDim_) * sizeof(T) / dataBlockSize_;
|
||||
|
||||
uint32_t startBlock;
|
||||
uint32_t endBlock;
|
||||
uint32_t startLayer;
|
||||
uint32_t endLayer;
|
||||
|
||||
Caloffset(startBlock, endBlock, startLayer, endLayer);
|
||||
|
||||
for (uint32_t i = startBlock; i < endBlock; ++i) {
|
||||
int64_t blockId = blockIDsGm_.GetValue(i);
|
||||
uint32_t offsetBlock = blockId * blockSize_ * headNum_ * headDim_;
|
||||
uint32_t realStartLayer;
|
||||
uint32_t realEndLayer;
|
||||
if (i == startBlock) {
|
||||
realStartLayer = startLayer;
|
||||
} else {
|
||||
realStartLayer = 0;
|
||||
}
|
||||
|
||||
if (i == (endBlock - 1)) {
|
||||
realEndLayer = endLayer;
|
||||
} else {
|
||||
realEndLayer = layerNum_;
|
||||
}
|
||||
for (uint32_t layerId = realStartLayer; layerId < realEndLayer; ++layerId) {
|
||||
SetGlobalBuffers(layerId);
|
||||
uint32_t blockSizeIndex = blockIdx_ % blockSizeSplitNum_;
|
||||
uint32_t srcOffset;
|
||||
uint32_t dstOffset;
|
||||
if constexpr (needHandleUnFactorSplit) {
|
||||
// handle tail
|
||||
if (blockSizeIndex >= blockSizePerTimeTail_) {
|
||||
repeatParams.blockCount = (blockSizePerTime_ - 1);
|
||||
copyOutLength_ = (blockSizePerTime_ - 1) * headNum_ * headDim_;
|
||||
srcOffset = (blockSizeIndex * blockSizePerTime_ - (blockSizeIndex - blockSizePerTimeTail_)) * headNumSplited_ * headDim_;
|
||||
dstOffset = (blockSizeIndex * blockSizePerTime_ - (blockSizeIndex - blockSizePerTimeTail_)) * headNum_ * headDim_;
|
||||
} else {
|
||||
repeatParams.blockCount = blockSizePerTime_;
|
||||
copyOutLength_ = blockSizePerTime_ * headNum_ * headDim_;
|
||||
srcOffset = blockSizeIndex * blockSizePerTime_ * headNumSplited_ * headDim_;
|
||||
dstOffset = blockSizeIndex * blockSizePerTime_ * headNum_ * headDim_;
|
||||
}
|
||||
} else {
|
||||
repeatParams.blockCount = blockSizePerTime_;
|
||||
copyOutLength_ = blockSizePerTime_ * headNum_ * headDim_;
|
||||
srcOffset = blockSizeIndex * blockSizePerTime_ * headNumSplited_ * headDim_;
|
||||
dstOffset = blockSizeIndex * blockSizePerTime_ * headNum_ * headDim_;
|
||||
}
|
||||
|
||||
CopyIn(kCacheGm_, offsetBlock + srcOffset, repeatParams);
|
||||
CopyOut(kCacheGm_, offsetBlock + dstOffset);
|
||||
|
||||
CopyIn(vCacheGm_, offsetBlock + srcOffset, repeatParams);
|
||||
CopyOut(vCacheGm_, offsetBlock + dstOffset);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
if (needSync_) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE2>(0x8);
|
||||
AscendC::CrossCoreWaitFlag(0x8);
|
||||
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE2>(0x8);
|
||||
AscendC::CrossCoreWaitFlag(0x8);
|
||||
}
|
||||
|
||||
}
|
||||
};
|
||||
@@ -0,0 +1,36 @@
|
||||
#include "kernel_operator.h"
|
||||
#include "full_load.h"
|
||||
#include "general.h"
|
||||
|
||||
|
||||
extern "C" __global__ __aicore__ void transpose_kv_cache_by_block(GM_ADDR KCache, GM_ADDR VCache, GM_ADDR blockIDs, GM_ADDR workspace, GM_ADDR tiling) {
|
||||
GET_TILING_DATA(tiling_data, tiling);
|
||||
TPipe tPipe;
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0);
|
||||
if (TILING_KEY_IS(0)) {
|
||||
// full load not db
|
||||
TransposeKvCacheByBlockKernelFullLoad<DTYPE_KCACHE> kernel;
|
||||
kernel.Init(KCache, VCache, blockIDs, &tiling_data, &tPipe);
|
||||
kernel.Process();
|
||||
} else if (TILING_KEY_IS(1)) {
|
||||
// db \ align split blockSize
|
||||
TransposeKvCacheByBlockKernelGeneral<DTYPE_KCACHE, uint32_t(2), false> kernel;
|
||||
kernel.Init(KCache, VCache, blockIDs, &tiling_data, &tPipe);
|
||||
kernel.Process();
|
||||
} else if (TILING_KEY_IS(2)) {
|
||||
// not db \ align split blockSize
|
||||
TransposeKvCacheByBlockKernelGeneral<DTYPE_KCACHE, uint32_t(1), false> kernel;
|
||||
kernel.Init(KCache, VCache, blockIDs, &tiling_data, &tPipe);
|
||||
kernel.Process();
|
||||
} else if (TILING_KEY_IS(3)) {
|
||||
// db \ unalign split blockSize
|
||||
TransposeKvCacheByBlockKernelGeneral<DTYPE_KCACHE, uint32_t(2), true> kernel;
|
||||
kernel.Init(KCache, VCache, blockIDs, &tiling_data, &tPipe);
|
||||
kernel.Process();
|
||||
} else if (TILING_KEY_IS(4)) {
|
||||
// not db \ unalign split blockSize
|
||||
TransposeKvCacheByBlockKernelGeneral<DTYPE_KCACHE, uint32_t(1), true> kernel;
|
||||
kernel.Init(KCache, VCache, blockIDs, &tiling_data, &tPipe);
|
||||
kernel.Process();
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user