19
csrc/attention/load_index_kv_cache/CMakeLists.txt
Normal file
19
csrc/attention/load_index_kv_cache/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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(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()
|
||||
62
csrc/attention/load_index_kv_cache/op_host/CMakeLists.txt
Normal file
62
csrc/attention/load_index_kv_cache/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,62 @@
|
||||
# ----------------------------------------------------------------------------
|
||||
# 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.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
# add_ops_compile_options(
|
||||
# OP_NAME LoadIndexKvCache
|
||||
# OPTIONS --cce-auto-sync=off
|
||||
# -Wno-deprecated-declarations
|
||||
# -Werror
|
||||
# -mllvm -cce-aicore-hoist-movemask=false
|
||||
# --op_relocatable_kernel_binary=true
|
||||
# )
|
||||
|
||||
# set(load_index_kv_cache_depends transformer/attention/load_index_kv_cache PARENT_SCOPE)
|
||||
|
||||
# target_sources(op_host_aclnn PRIVATE
|
||||
# op_host/load_index_kv_cache_def.cpp
|
||||
# )
|
||||
|
||||
# target_sources(optiling PRIVATE
|
||||
# op_host/load_index_kv_cache_tiling.cpp
|
||||
# )
|
||||
|
||||
# if (NOT BUILD_OPEN_PROJECT)
|
||||
# target_sources(opmaster_ct PRIVATE
|
||||
# op_host/load_index_kv_cache_tiling.cpp
|
||||
# )
|
||||
# endif ()
|
||||
|
||||
# target_include_directories(optiling PRIVATE
|
||||
# ${CMAKE_CURRENT_SOURCE_DIR}/op_host
|
||||
# )
|
||||
|
||||
# target_sources(opsproto PRIVATE
|
||||
# op_host/load_index_kv_cache_proto.cpp
|
||||
# )
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
load_index_kv_cache_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME LoadIndexKvCache
|
||||
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 load_index_kv_cache ACLNNTYPE aclnn)
|
||||
endif()
|
||||
@@ -0,0 +1,49 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class LoadIndexKvCache : public OpDef {
|
||||
public:
|
||||
explicit LoadIndexKvCache(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("kv_cache")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND})
|
||||
.IgnoreContiguous();
|
||||
this->Input("slot_mapping")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("kv")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Output("kv_scale")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND});
|
||||
this->Attr("block_stride").AttrType(OPTIONAL).Int(0);
|
||||
this->AICore().AddConfig("ascend950");
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(LoadIndexKvCache);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,36 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_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 {
|
||||
|
||||
graphStatus InferShape4LoadIndexKvCache(gert::InferShapeContext* context)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
graphStatus InferDtype4LoadIndexKvCache(gert::InferDataTypeContext* context)
|
||||
{
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(LoadIndexKvCache)
|
||||
.InferShape(InferShape4LoadIndexKvCache)
|
||||
.InferDataType(InferDtype4LoadIndexKvCache);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,204 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <sstream>
|
||||
#include "load_index_kv_cache_tiling.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace optiling {
|
||||
namespace {
|
||||
constexpr uint64_t WORKSPACE_SIZE = 32;
|
||||
int64_t CeilDiv(int64_t x, int64_t y)
|
||||
{
|
||||
if (y != 0) {
|
||||
return (x + y - 1) / y;
|
||||
}
|
||||
return x;
|
||||
}
|
||||
int64_t DownAlign(int64_t x, int64_t y) {
|
||||
if (y == 0) {
|
||||
return x;
|
||||
}
|
||||
return (x / y) * y;
|
||||
}
|
||||
int64_t RoundUp(int64_t x, int64_t y) {
|
||||
return CeilDiv(x, y) * y;
|
||||
}
|
||||
|
||||
constexpr int64_t INPUT_KV_CACHE_IDX = 0;
|
||||
constexpr int64_t INPUT_SLOT_MAPPING_IDX = 1;
|
||||
constexpr int64_t ATTR_BLOCK_STRIDE_INDEX = 0;
|
||||
constexpr int64_t DIM_0 = 0;
|
||||
constexpr int64_t DIM_1 = 1;
|
||||
constexpr int64_t DIM_2 = 2;
|
||||
constexpr int64_t DIM_3 = 3;
|
||||
constexpr int64_t INPUTS_DIM_LIMIT = 4;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::GetPlatformInfo()
|
||||
{
|
||||
auto platformInfo = context_->GetPlatformInfo();
|
||||
if (platformInfo == nullptr) {
|
||||
auto compileInfoPtr = context_->GetCompileInfo<LoadIndexKvCacheCompileInfo>();
|
||||
OPS_ERR_IF(compileInfoPtr == nullptr, OPS_LOG_E(context_, "compile info is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
coreNum_ = compileInfoPtr->coreNum;
|
||||
ubSize_ = compileInfoPtr->ubSize;
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
|
||||
coreNum_ = ascendcPlatform.GetCoreNumAiv();
|
||||
uint64_t ubSizePlatForm;
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
|
||||
ubSize_ = ubSizePlatForm;
|
||||
socVersion_ = ascendcPlatform.GetSocVersion();
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::GetAttr()
|
||||
{
|
||||
auto* attrs = context_->GetAttrs();
|
||||
OPS_LOG_E_IF_NULL(context_, attrs, return ge::GRAPH_FAILED);
|
||||
|
||||
auto blockStrideAttr = attrs->GetAttrPointer<int64_t>(ATTR_BLOCK_STRIDE_INDEX);
|
||||
blockStride_ = blockStrideAttr != nullptr ? *blockStrideAttr : 0;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::GetShapeAttrsInfoInner()
|
||||
{
|
||||
auto shapeKvCache = context_->GetInputShape(INPUT_KV_CACHE_IDX);
|
||||
OPS_LOG_E_IF_NULL(context_, shapeKvCache, return ge::GRAPH_FAILED);
|
||||
auto kvCacheStorageShape = shapeKvCache->GetStorageShape();
|
||||
OPS_ERR_IF(
|
||||
(kvCacheStorageShape.GetDimNum() != INPUTS_DIM_LIMIT),
|
||||
OPS_LOG_E(context_, "the dim of kv_cache only support %d, please check.", INPUTS_DIM_LIMIT),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
bn_ = kvCacheStorageShape.GetDim(DIM_0);
|
||||
bs_ = kvCacheStorageShape.GetDim(DIM_1);
|
||||
d_ = kvCacheStorageShape.GetDim(DIM_3);
|
||||
|
||||
auto shapeSlotMapping = context_->GetInputShape(INPUT_SLOT_MAPPING_IDX);
|
||||
OPS_LOG_E_IF_NULL(context_, shapeSlotMapping, return ge::GRAPH_FAILED);
|
||||
auto slotMappingStorageShape = shapeSlotMapping->GetStorageShape();
|
||||
OPS_ERR_IF(
|
||||
(slotMappingStorageShape.GetDimNum() != 1),
|
||||
OPS_LOG_E(context_, "the dim of slot_mapping must be 1, please check."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
n_ = slotMappingStorageShape.GetDim(0);
|
||||
|
||||
OPS_ERR_IF(GetAttr() != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(context_->GetNodeName(), "get attr failed."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::CheckShapesAndAttrs()
|
||||
{
|
||||
blockStride_ = blockStride_ == 0 ? bs_ * d_ : blockStride_;
|
||||
OPS_ERR_IF(blockStride_ < bs_ * d_,
|
||||
OPS_LOG_E(context_, "stride_kvcache must be greater than last dim of kv_cache %ld.", d_),
|
||||
return ge::GRAPH_FAILED);
|
||||
// 如果blockStride_为0,则认为连续
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::CalcOpTiling()
|
||||
{
|
||||
// 默认尾轴在ub内全载
|
||||
rowOfFormerBlock_ = CeilDiv(n_, static_cast<int64_t>(coreNum_));
|
||||
usedCoreNums_ = std::min(CeilDiv(n_, rowOfFormerBlock_), static_cast<int64_t>(coreNum_));
|
||||
rowOfTailBlock_ = n_ - (usedCoreNums_ - 1) * rowOfFormerBlock_;
|
||||
|
||||
tilingData_.set_bn(bn_);
|
||||
tilingData_.set_bs(bs_);
|
||||
tilingData_.set_d(d_);
|
||||
tilingData_.set_n(n_);
|
||||
tilingData_.set_rowOfFormerBlock(rowOfFormerBlock_);
|
||||
tilingData_.set_rowOfTailBlock(rowOfTailBlock_);
|
||||
tilingData_.set_blockStride(blockStride_);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::DoOpTiling()
|
||||
{
|
||||
if (GetPlatformInfo() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetShapeAttrsInfoInner() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (CheckShapesAndAttrs() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (CalcOpTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (GetWorkspaceSize() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (PostTiling() == ge::GRAPH_FAILED) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
context_->SetTilingKey(0);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::GetWorkspaceSize()
|
||||
{
|
||||
workspaceSize_ = WORKSPACE_SIZE;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus LoadIndexKvCacheTiling::PostTiling()
|
||||
{
|
||||
context_->SetBlockDim(usedCoreNums_);
|
||||
size_t* workspaces = context_->GetWorkspaceSizes(1);
|
||||
workspaces[0] = workspaceSize_;
|
||||
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
|
||||
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepareForLoadIndexKvCache(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingForLoadIndexKvCache(gert::TilingContext *context)
|
||||
{
|
||||
OPS_ERR_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("LoadIndexKvCache", "Tiling context is null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
LoadIndexKvCacheTiling LoadIndexKvCacheTiling(context);
|
||||
return LoadIndexKvCacheTiling.DoOpTiling();
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(LoadIndexKvCache)
|
||||
.Tiling(TilingForLoadIndexKvCache)
|
||||
.TilingParse<LoadIndexKvCacheCompileInfo>(TilingPrepareForLoadIndexKvCache);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,98 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef LOAD_INDEX_KV_CACHE_TILING_H
|
||||
#define LOAD_INDEX_KV_CACHE_TILING_H
|
||||
|
||||
|
||||
#include <vector>
|
||||
#include <iostream>
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "platform/platform_infos_def.h"
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "error/ops_error.h"
|
||||
#include "platform/platform_info.h"
|
||||
|
||||
namespace optiling {
|
||||
// ----------公共定义----------
|
||||
struct TilingRequiredParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::StorageShape *shape;
|
||||
};
|
||||
|
||||
struct TilingOptionalParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::Tensor *tensor;
|
||||
};
|
||||
|
||||
// ----------算子TilingData定义----------
|
||||
BEGIN_TILING_DATA_DEF(LoadIndexKvCacheTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, bn); // block_num
|
||||
TILING_DATA_FIELD_DEF(int64_t, bs); // block_size
|
||||
TILING_DATA_FIELD_DEF(int64_t, d); // head_dim
|
||||
TILING_DATA_FIELD_DEF(int64_t, n);
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfFormerBlock); // 头核共需要处理多少行
|
||||
TILING_DATA_FIELD_DEF(int64_t, rowOfTailBlock); // 尾核共需要处理多少行
|
||||
TILING_DATA_FIELD_DEF(int64_t, blockStride);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(LoadIndexKvCache, LoadIndexKvCacheTilingData)
|
||||
|
||||
// ----------算子CompileInfo定义----------
|
||||
struct LoadIndexKvCacheCompileInfo {
|
||||
uint64_t coreNum = 0;
|
||||
uint64_t ubSize = 0;
|
||||
};
|
||||
|
||||
// ----------算子Tiling入参信息解析及check类----------
|
||||
class LoadIndexKvCacheTiling {
|
||||
public:
|
||||
explicit LoadIndexKvCacheTiling(gert::TilingContext* tilingContext) : context_(tilingContext)
|
||||
{
|
||||
}
|
||||
~LoadIndexKvCacheTiling() = default;
|
||||
|
||||
ge::graphStatus GetPlatformInfo();
|
||||
ge::graphStatus DoOpTiling();
|
||||
ge::graphStatus GetWorkspaceSize();
|
||||
ge::graphStatus PostTiling();
|
||||
ge::graphStatus GetAttr();
|
||||
ge::graphStatus GetShapeAttrsInfoInner();
|
||||
ge::graphStatus CheckShapesAndAttrs();
|
||||
ge::graphStatus CalcOpTiling();
|
||||
private:
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
LoadIndexKvCacheTilingData tilingData_;
|
||||
uint64_t coreNum_ = 0;
|
||||
uint64_t workspaceSize_ = 0;
|
||||
uint64_t usedCoreNums_ = 0;
|
||||
uint64_t ubSize_ = 0;
|
||||
int64_t bn_ = 0;
|
||||
int64_t bs_ = 0;
|
||||
int64_t d_ = 0;
|
||||
int64_t n_ = 0;
|
||||
int64_t rowOfFormerBlock_ = 0;
|
||||
int64_t rowOfTailBlock_ = 0;
|
||||
int64_t blockStride_ = 0;
|
||||
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
|
||||
int64_t tilingKey_ = 0;
|
||||
};
|
||||
|
||||
} // namespace optiling
|
||||
#endif // LOAD_INDEX_KV_CACHE_TILING_H
|
||||
@@ -0,0 +1,46 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "load_index_kv_cache_perf.h"
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
extern "C" __global__ __aicore__ void load_index_kv_cache(
|
||||
GM_ADDR kv_cache,
|
||||
GM_ADDR slot_mapping,
|
||||
GM_ADDR kv,
|
||||
GM_ADDR kv_scale,
|
||||
GM_ADDR workspace,
|
||||
GM_ADDR tiling)
|
||||
{
|
||||
if (workspace == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
GM_ADDR userWs = GetUserWorkspace(workspace);
|
||||
if (userWs == nullptr) {
|
||||
return;
|
||||
}
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
TPipe pipe;
|
||||
|
||||
if (TILING_KEY_IS(0)) {
|
||||
LoadIndexKvCache::LoadIndexKvCachePerf<DTYPE_KV_CACHE> op;
|
||||
op.Init(kv_cache, slot_mapping, kv, kv_scale, userWs, &tilingData, &pipe);
|
||||
op.Process();
|
||||
return;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef LOAD_INDEX_LV_CACHE_BASE_H
|
||||
#define LOAD_INDEX_LV_CACHE_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace LoadIndexKvCache {
|
||||
using namespace AscendC;
|
||||
constexpr int32_t BLOCK_SIZE = 32;
|
||||
constexpr int32_t VL_FP32 = 64;
|
||||
constexpr int32_t KV_LAST_DIM = 128;
|
||||
constexpr int32_t KV_SCALE_LAST_DIM = 4;
|
||||
|
||||
__aicore__ inline int32_t CeilDiv(int32_t a, int b)
|
||||
{
|
||||
if (b == 0) {
|
||||
return a;
|
||||
}
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t CeilAlign(int32_t a, int b)
|
||||
{
|
||||
return CeilDiv(a, b) * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline int32_t RoundUp(int32_t num)
|
||||
{
|
||||
int32_t elemNum = BLOCK_SIZE / sizeof(T);
|
||||
return CeilAlign(num, elemNum);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline int32_t RoundUp(int32_t num, int32_t elemNum)
|
||||
{
|
||||
return CeilAlign(num, elemNum);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void CopyIn(
|
||||
const GlobalTensor<T>& inputGm, const LocalTensor<T>& inputTensor, const uint16_t nBurst, const uint32_t copyLen,
|
||||
uint32_t srcStride = 0)
|
||||
{
|
||||
DataCopyPadExtParams<T> dataCopyPadExtParams;
|
||||
dataCopyPadExtParams.isPad = false;
|
||||
dataCopyPadExtParams.leftPadding = 0;
|
||||
dataCopyPadExtParams.rightPadding = 0;
|
||||
dataCopyPadExtParams.paddingValue = 0;
|
||||
|
||||
DataCopyExtParams dataCoptExtParams;
|
||||
dataCoptExtParams.blockCount = nBurst;
|
||||
dataCoptExtParams.blockLen = copyLen * sizeof(T);
|
||||
dataCoptExtParams.srcStride = srcStride * sizeof(T);
|
||||
dataCoptExtParams.dstStride = 0;
|
||||
DataCopyPad(inputTensor, inputGm, dataCoptExtParams, dataCopyPadExtParams);
|
||||
}
|
||||
|
||||
|
||||
template <typename T, AscendC::PaddingMode mode = AscendC::PaddingMode::Normal>
|
||||
__aicore__ inline void CopyOut(
|
||||
const LocalTensor<T>& outputTensor, const GlobalTensor<T>& outputGm, const uint16_t nBurst, const uint32_t copyLen,
|
||||
uint32_t dstStride = 0)
|
||||
{
|
||||
DataCopyExtParams dataCopyParams;
|
||||
dataCopyParams.blockCount = nBurst;
|
||||
dataCopyParams.blockLen = copyLen * sizeof(T);
|
||||
dataCopyParams.srcStride = 0;
|
||||
dataCopyParams.dstStride = dstStride * sizeof(T);
|
||||
DataCopyPad<T, mode>(outputGm, outputTensor, dataCopyParams);
|
||||
}
|
||||
|
||||
} // namespace IndexerCompressEpilogV2
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,110 @@
|
||||
/**
|
||||
* 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 load_index_kv_cache_single_row.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef LOAD_INDEX_KV_CACH_PERF_H
|
||||
#define LOAD_INDEX_KV_CACH_PERF_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "load_index_kv_cache_base.h"
|
||||
|
||||
namespace LoadIndexKvCache {
|
||||
using namespace AscendC;
|
||||
template <typename T>
|
||||
class LoadIndexKvCachePerf {
|
||||
public:
|
||||
__aicore__ inline LoadIndexKvCachePerf()
|
||||
{}
|
||||
|
||||
__aicore__ inline void Init(
|
||||
GM_ADDR kvCache, GM_ADDR slotMapping, GM_ADDR kv, GM_ADDR kvScale, GM_ADDR workspace,
|
||||
const LoadIndexKvCacheTilingData* tilingDataPtr, TPipe* pipePtr)
|
||||
{
|
||||
pipe = pipePtr;
|
||||
tilingData = tilingDataPtr;
|
||||
|
||||
kvCacheGm.SetGlobalBuffer((__gm__ T*)kvCache);
|
||||
slotMappingGm.SetGlobalBuffer((__gm__ int32_t*)slotMapping);
|
||||
|
||||
kvGm.SetGlobalBuffer((__gm__ T*)kv);
|
||||
kvScaleGm.SetGlobalBuffer((__gm__ float*)kvScale);
|
||||
|
||||
pipe->InitBuffer(kvCacheQue, 2, RoundUp<T>(tilingData->d) * sizeof(T));
|
||||
pipe->InitBuffer(kvAndKvScaleQue, 2, RoundUp<T>(tilingData->d) * sizeof(T));
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
int64_t curBlockIdx = GetBlockIdx();
|
||||
int64_t rowOuterLoop =
|
||||
(curBlockIdx == GetBlockNum() - 1) ? tilingData->rowOfTailBlock : tilingData->rowOfFormerBlock;
|
||||
|
||||
int64_t baseSlotMappingOffset = curBlockIdx * tilingData->rowOfFormerBlock;
|
||||
int64_t kvBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock * KV_LAST_DIM;
|
||||
int64_t kvScaleBaseOffset = curBlockIdx * tilingData->rowOfFormerBlock;
|
||||
for (int64_t rowOuterIdx = 0; rowOuterIdx < rowOuterLoop; rowOuterIdx++) {
|
||||
int64_t curSlotIdx = baseSlotMappingOffset + rowOuterIdx;
|
||||
int64_t slot = slotMappingGm.GetValue(curSlotIdx);
|
||||
if (slot == -1) {
|
||||
// slot为-1时,需要将对应位置的kv和kvscale全部刷成0
|
||||
kvAndKvScaleLocal = kvAndKvScaleQue.template AllocTensor<T>();
|
||||
Duplicate(kvAndKvScaleLocal.template ReinterpretCast<uint8_t>(), uint8_t(0), KV_LAST_DIM);
|
||||
Duplicate(kvAndKvScaleLocal[KV_LAST_DIM].template ReinterpretCast<float>(), float(0), 1);
|
||||
kvAndKvScaleQue.template EnQue(kvAndKvScaleLocal);
|
||||
kvAndKvScaleLocal = kvAndKvScaleQue.template DeQue<T>();
|
||||
CopyOut(kvAndKvScaleLocal, kvGm[kvBaseOffset + rowOuterIdx * KV_LAST_DIM], 1, KV_LAST_DIM);
|
||||
CopyOut(kvAndKvScaleLocal[KV_LAST_DIM].template ReinterpretCast<float>(), kvScaleGm[kvScaleBaseOffset + rowOuterIdx], 1, 1);
|
||||
kvAndKvScaleQue.template FreeTensor(kvAndKvScaleLocal);
|
||||
continue;
|
||||
}
|
||||
|
||||
// 处理kvCacheLocal
|
||||
kvCacheLocal = kvCacheQue.template AllocTensor<T>();
|
||||
int64_t kvCacheGmOffset = (slot / tilingData->bs) * tilingData->blockStride + (slot % tilingData->bs) * KV_LAST_DIM;
|
||||
int64_t kvCacheScaleGmOffset = (slot / tilingData->bs) * tilingData->blockStride + tilingData->bs * KV_LAST_DIM + (slot % tilingData->bs) * KV_SCALE_LAST_DIM;
|
||||
CopyIn(kvCacheGm[kvCacheGmOffset].template ReinterpretCast<uint8_t>(), kvCacheLocal.template ReinterpretCast<uint8_t>(), 1, KV_LAST_DIM);
|
||||
CopyIn(kvCacheGm[kvCacheScaleGmOffset].template ReinterpretCast<uint8_t>(), kvCacheLocal[KV_LAST_DIM].template ReinterpretCast<uint8_t>(), 1, KV_SCALE_LAST_DIM);
|
||||
kvCacheQue.template EnQue(kvCacheLocal);
|
||||
kvCacheLocal = kvCacheQue.template DeQue<T>();
|
||||
|
||||
kvAndKvScaleLocal = kvAndKvScaleQue.AllocTensor<T>();
|
||||
DataCopy(kvAndKvScaleLocal, kvCacheLocal, RoundUp<T>(tilingData->d));
|
||||
kvCacheQue.template FreeTensor(kvCacheLocal);
|
||||
|
||||
kvAndKvScaleQue.template EnQue(kvAndKvScaleLocal);
|
||||
kvAndKvScaleLocal = kvAndKvScaleQue.template DeQue<T>();
|
||||
CopyOut(kvAndKvScaleLocal, kvGm[kvBaseOffset + rowOuterIdx * KV_LAST_DIM], 1, KV_LAST_DIM);
|
||||
CopyOut(kvAndKvScaleLocal[KV_LAST_DIM].template ReinterpretCast<float>(), kvScaleGm[kvScaleBaseOffset + rowOuterIdx * 1], 1, 1);
|
||||
kvAndKvScaleQue.template FreeTensor(kvAndKvScaleLocal);
|
||||
}
|
||||
}
|
||||
private:
|
||||
TPipe* pipe;
|
||||
const LoadIndexKvCacheTilingData* tilingData;
|
||||
GlobalTensor<T> kvCacheGm;
|
||||
GlobalTensor<int32_t> slotMappingGm;
|
||||
GlobalTensor<T> kvGm;
|
||||
GlobalTensor<float> kvScaleGm;
|
||||
|
||||
TQue<QuePosition::VECIN, 1> kvCacheQue;
|
||||
TQue<QuePosition::VECOUT, 1> kvAndKvScaleQue;
|
||||
TQue<QuePosition::VECOUT, 1> kvScaleQue;
|
||||
|
||||
LocalTensor<T> kvCacheLocal;
|
||||
LocalTensor<T> kvAndKvScaleLocal;
|
||||
};
|
||||
|
||||
} // namespace LoadIndexKvCache
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user