@@ -0,0 +1,17 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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}/*)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
@@ -0,0 +1,54 @@
|
||||
# KvQuantSparseAttnSharedkvMetadata
|
||||
|
||||
## 产品支持情况
|
||||
|
||||
| 产品 | 是否支持 |
|
||||
| ------------------------------------------------------------ | :------: |
|
||||
|<term>Ascend 950PR/Ascend 950DT</term>| √ |
|
||||
|<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>| × |
|
||||
|<term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>| × |
|
||||
|<term>Atlas 200I/500 A2 推理产品</term>| × |
|
||||
|<term>Atlas 推理系列加速卡产品</term>| × |
|
||||
|<term>Atlas 训练系列产品</term>| × |
|
||||
|
||||
## 功能说明
|
||||
|
||||
- API功能:`KvQuantSparseAttnSharedkvMetadata`算子旨在生成一个任务列表,包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及 Q 和 K 的分块的索引,供后续`KvQuantSparseAttnSharedkv`算子使用。
|
||||
|
||||
## 参数说明
|
||||
|
||||
| 参数名 |输入/输出/属性| 描述 | 数据类型 |数据格式|
|
||||
|-------------|------------|------|-----|-----|
|
||||
|num_heads_q|属性|query对应的多头数,目前支持64/128。|INT32|-|
|
||||
|num_heads_kv|属性|key和value对应的多头数,目前仅支持1。|INT32|-|
|
||||
|head_dim|属性|注意力头的维度,目前仅支持512。|INT32|-|
|
||||
|kv_quant_mode|属性|kv nope的量化模式,仅支持1,表示K、V nope为per_tile量化,量化后的KV数据类型为FLOAT8_E4M3FN。|INT32|-|
|
||||
|cu_seqlens_q|可选输入|表示不同Batch中q的有效token数,维度为B+1。|INT32|ND|
|
||||
|cu_seqlens_ori_kv|可选输入|预留参数,当前不生效,表示不同Batch中ori_kv的有效token数,维度为B+1。|INT32|ND|
|
||||
|cu_seqlens_cmp_kv|可选输入|预留参数,当前不生效,表示不同Batch中cmp_kv的有效token数,维度为B+1。|INT32|ND|
|
||||
|seqused_q|可选输入|预留参数,当前不生效,表示不同Batch中q的有效token数,维度为B。|INT32|ND|
|
||||
|seqused_kv|可选输入|表示不同Batch中ori_kv的有效token数,维度为B。|INT32|ND|
|
||||
|batch_size|可选属性|输入样本批量大小。|INT32|-|
|
||||
|max_seqlen_q|可选属性|表示所有Batch中q的最大有效token数。|INT32|-|
|
||||
|max_seqlen_kv|可选属性|表示所有Batch中ori_kv的最大有效token数。|INT32|-|
|
||||
|ori_topk|可选属性|预留参数,当前不生效,表示通过QLI算法从ori_kv中筛选出的关键稀疏token的个数。|INT32|-|
|
||||
|cmp_topk|可选属性|表示通过QLI算法从cmp_kv中筛选出的关键稀疏token的个数,目前支持512/1024。|INT32|-|
|
||||
|tile_size|可选属性|表示量化粒度,必须能被rope_head_dim整除,默认值为None,当前仅支持64。|INT32|-|
|
||||
|rope_head_dim|可选属性|默认值为0,当前仅支持64。|INT32|-|
|
||||
|cmp_ratio|可选属性|表示对ori_kv的压缩率,数据范围支持4/128,默认值为None。|INT32|-|
|
||||
|ori_mask_mode|可选属性|表示q和ori_kv计算的mask模式,仅支持输入默认值4,代表band模式的mask。|INT32|-|
|
||||
|cmp_mask_mode|可选属性|表示q和cmp_kv计算的mask模式,仅支持输入默认值3,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景。|INT32|-|
|
||||
|ori_win_left|可选属性|表示q和ori_kv计算中q对过去token计算的数量,仅支持默认值127。|INT32|-|
|
||||
|ori_win_right|可选属性|表示q和ori_kv计算中q对未来token计算的数量,仅支持默认值0。|INT32|-|
|
||||
|layout_q|可选属性|用于标识输入q的数据排布格式,支持BSND和TND,默认值为BSND。|STRING|-|
|
||||
|layout_kv|可选属性|用于标识输入ori_kv和cmp_kv的数据排布格式,仅支持传入默认值PA_ND(PageAttention)。|STRING|-|
|
||||
|has_ori_kv|可选属性|用于标识是否含有ori_kv。|BOOL|-|
|
||||
|has_cmp_kv|可选属性|用于标识是否含有cmp_kv。|BOOL|-|
|
||||
|device|可选属性|用于获取设备信息。|STRING|-|
|
||||
|metadata|输出|包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及 Q 和 K 的分块的索引的列表,shape固定为1024。|INT32|-|
|
||||
|
||||
## 约束说明
|
||||
|
||||
- 该接口支持推理场景下使用。
|
||||
- 该接口支持aclgraph模式。
|
||||
- Tensor不能全传None。
|
||||
@@ -0,0 +1,56 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_metadata_proto.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
|
||||
#include "graph/operator_reg.h"
|
||||
#include "graph/types.h"
|
||||
|
||||
namespace ge {
|
||||
|
||||
REG_OP(KvQuantSparseAttnSharedkvMetadata)
|
||||
.OPTIONAL_INPUT(cu_seqlens_q, TensorType({DT_INT32}))
|
||||
.OPTIONAL_INPUT(cu_seqlens_ori_kv, TensorType({DT_INT32}))
|
||||
.OPTIONAL_INPUT(cu_seqlens_cmp_kv, TensorType({DT_INT32}))
|
||||
.OPTIONAL_INPUT(seqused_q, TensorType({DT_INT32}))
|
||||
.OPTIONAL_INPUT(seqused_kv, TensorType({DT_INT32}))
|
||||
.OUTPUT(metadata, TensorType({DT_INT32}))
|
||||
.REQUIRED_ATTR(num_heads_q, Int)
|
||||
.REQUIRED_ATTR(num_heads_kv, Int)
|
||||
.REQUIRED_ATTR(head_dim, Int)
|
||||
.REQUIRED_ATTR(kv_quant_mode, Int)
|
||||
.ATTR(batch_size, Int, 0)
|
||||
.ATTR(max_seqlen_q, Int, 0)
|
||||
.ATTR(max_seqlen_kv, Int, 0)
|
||||
.ATTR(ori_topk, Int, 0)
|
||||
.ATTR(cmp_topk, Int, 0)
|
||||
.ATTR(tile_size, Int, 0)
|
||||
.ATTR(rope_head_dim, Int, 0)
|
||||
.ATTR(cmp_ratio, Int, -1)
|
||||
.ATTR(ori_mask_mode, Int, 4)
|
||||
.ATTR(cmp_mask_mode, Int, 3)
|
||||
.ATTR(ori_win_left, Int, 127)
|
||||
.ATTR(ori_win_right, Int, 0)
|
||||
.ATTR(layout_q, String, "BSND")
|
||||
.ATTR(layout_kv, String, "PA_ND")
|
||||
.ATTR(has_ori_kv, Bool, true)
|
||||
.ATTR(has_cmp_kv, Bool, true)
|
||||
.REQUIRED_ATTR(soc_version, String)
|
||||
.REQUIRED_ATTR(aic_core_num, Int)
|
||||
.REQUIRED_ATTR(aiv_core_num, Int)
|
||||
.OP_END_FACTORY_REG(KvQuantSparseAttnSharedkvMetadata)
|
||||
} // namespace ge
|
||||
|
||||
#endif // KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
@@ -0,0 +1,11 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
add_modules_sources(OPTYPE kv_quant_sparse_attn_sharedkv ACLNNTYPE aclnn)
|
||||
@@ -0,0 +1,39 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_metadata_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "../../kv_quant_sparse_attn_sharedkv/op_kernel/kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
static ge::graphStatus InferShapeKvQuantSparseAttnSharedkvMetadata(gert::InferShapeContext* context)
|
||||
{
|
||||
gert::Shape* oShape = context->GetOutputShape(0);
|
||||
// output shape (SAS_METADATA_T, )
|
||||
oShape->SetDimNum(1);
|
||||
oShape->SetDim(0, optiling::SAS_META_SIZE);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDtypeKvQuantSparseAttnSharedkvMetadata(gert::InferDataTypeContext* context)
|
||||
{
|
||||
context->SetOutputDataType(0, DT_INT32);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(KvQuantSparseAttnSharedkvMetadata)
|
||||
.InferShape(InferShapeKvQuantSparseAttnSharedkvMetadata)
|
||||
.InferDataType(InferDtypeKvQuantSparseAttnSharedkvMetadata);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,152 @@
|
||||
/**
|
||||
* 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 aclnn_kv_quant_sparse_attn_sharedkv_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "aclnn_kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
#include "l0_kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
#include "aclnn_kernels/contiguous.h"
|
||||
#include "aclnn_kernels/reshape.h"
|
||||
#include "aclnn/aclnn_base.h"
|
||||
#include "aclnn_kernels/common/op_error_check.h"
|
||||
#include "opdev/common_types.h"
|
||||
#include "opdev/data_type_utils.h"
|
||||
#include "opdev/format_utils.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/tensor_view_utils.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
static aclnnStatus ParamsCheck(const aclTensor* cuSeqLensQOptional,
|
||||
const aclTensor* cuSeqLensOriKvOptional,
|
||||
const aclTensor* cuSeqLensCmpKvOptional,
|
||||
const aclTensor* sequsedQOptional,
|
||||
const aclTensor* sequsedKvOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsKv,
|
||||
int64_t headDim,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKvOptional,
|
||||
int64_t oriTopKOptional,
|
||||
int64_t cmpTopKOptional,
|
||||
int64_t kvQuantMode,
|
||||
int64_t tileSizeOptional,
|
||||
int64_t ropeHeadDimOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
int64_t oriMaskModeOptional,
|
||||
int64_t cmpMaskModeOptional,
|
||||
int64_t oriWinLeftOptional,
|
||||
int64_t oriWinRightOptional,
|
||||
char *layoutQOptional,
|
||||
char *layoutKvOptional,
|
||||
bool hasOriKvOptional,
|
||||
bool hasCmpKvOptional,
|
||||
const aclTensor* metaData) {
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
aclnnStatus aclnnKvQuantSparseAttnSharedkvMetadataGetWorkspaceSize(
|
||||
const aclTensor* cuSeqLensQOptional,
|
||||
const aclTensor* cuSeqLensOriKvOptional,
|
||||
const aclTensor* cuSeqLensCmpKvOptional,
|
||||
const aclTensor* sequsedQOptional,
|
||||
const aclTensor* sequsedKvOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsKv,
|
||||
int64_t headDim,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKvOptional,
|
||||
int64_t oriTopKOptional,
|
||||
int64_t cmpTopKOptional,
|
||||
int64_t kvQuantMode,
|
||||
int64_t tileSizeOptional,
|
||||
int64_t ropeHeadDimOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
int64_t oriMaskModeOptional,
|
||||
int64_t cmpMaskModeOptional,
|
||||
int64_t oriWinLeftOptional,
|
||||
int64_t oriWinRightOptional,
|
||||
char *layoutQOptional,
|
||||
char *layoutKvOptional,
|
||||
bool hasOriKvOptional,
|
||||
bool hasCmpKvOptional,
|
||||
const aclTensor* metaData,
|
||||
uint64_t* workspaceSize,
|
||||
aclOpExecutor** executor) {
|
||||
L2_DFX_PHASE_1(aclnnKvQuantSparseAttnSharedkvMetadata,
|
||||
DFX_IN(cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional, sequsedQOptional,
|
||||
sequsedKvOptional, numHeadsQ, numHeadsKv, headDim, batchSizeOptional, maxSeqlenQOptional,
|
||||
maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, kvQuantMode, tileSizeOptional,
|
||||
ropeHeadDimOptional, cmpRatioOptional, oriMaskModeOptional, cmpMaskModeOptional,
|
||||
oriWinLeftOptional, oriWinRightOptional, layoutQOptional, layoutKvOptional,
|
||||
hasOriKvOptional, hasCmpKvOptional),
|
||||
DFX_OUT(metaData));
|
||||
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
|
||||
auto ret = ParamsCheck(cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional, sequsedQOptional,
|
||||
sequsedKvOptional, numHeadsQ, numHeadsKv, headDim, batchSizeOptional, maxSeqlenQOptional,
|
||||
maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, kvQuantMode, tileSizeOptional,
|
||||
ropeHeadDimOptional, cmpRatioOptional, oriMaskModeOptional, cmpMaskModeOptional,
|
||||
oriWinLeftOptional, oriWinRightOptional, layoutQOptional, layoutKvOptional,
|
||||
hasOriKvOptional, hasCmpKvOptional, metaData);
|
||||
CHECK_RET(ret == ACLNN_SUCCESS, ret);
|
||||
|
||||
const op::PlatformInfo &npuInfo = op::GetCurrentPlatformInfo();
|
||||
uint32_t aicCoreNum = npuInfo.GetCubeCoreNum();
|
||||
uint32_t aivCoreNum = npuInfo.GetVectorCoreNum();
|
||||
const char *socVersion = npuInfo.GetSocLongVersion().c_str();
|
||||
|
||||
auto cuSeqLensQOptionalContiguous = l0op::Contiguous(cuSeqLensQOptional, uniqueExecutor.get());
|
||||
CHECK_RET(cuSeqLensQOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto cuSeqLensOriKvOptionalContiguous = l0op::Contiguous(cuSeqLensOriKvOptional, uniqueExecutor.get());
|
||||
CHECK_RET(cuSeqLensOriKvOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto cuSeqLensCmpKvOptionalContiguous = l0op::Contiguous(cuSeqLensCmpKvOptional, uniqueExecutor.get());
|
||||
CHECK_RET(cuSeqLensCmpKvOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto sequsedQOptionalContiguous = l0op::Contiguous(sequsedQOptional, uniqueExecutor.get());
|
||||
CHECK_RET(sequsedQOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto sequsedKvOptionalContiguous = l0op::Contiguous(sequsedKvOptional, uniqueExecutor.get());
|
||||
CHECK_RET(sequsedKvOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
auto output = l0op::KvQuantSparseAttnSharedkvMetadata(
|
||||
cuSeqLensQOptionalContiguous, cuSeqLensOriKvOptionalContiguous, cuSeqLensCmpKvOptionalContiguous,
|
||||
sequsedQOptionalContiguous, sequsedKvOptionalContiguous, numHeadsQ, numHeadsKv, headDim, batchSizeOptional,
|
||||
maxSeqlenQOptional, maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, kvQuantMode, tileSizeOptional,
|
||||
ropeHeadDimOptional, cmpRatioOptional, oriMaskModeOptional, cmpMaskModeOptional, oriWinLeftOptional,
|
||||
oriWinRightOptional, layoutQOptional, layoutKvOptional, hasOriKvOptional, hasCmpKvOptional, socVersion,
|
||||
aicCoreNum, aivCoreNum, metaData, uniqueExecutor.get());
|
||||
CHECK_RET(output != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
*workspaceSize = 0;
|
||||
uniqueExecutor.ReleaseTo(executor);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnKvQuantSparseAttnSharedkvMetadata(void *workspace, uint64_t workspaceSize,
|
||||
aclOpExecutor *executor, aclrtStream stream) {
|
||||
L2_DFX_PHASE_2(aclnnKvQuantSparseAttnSharedkvMetadata);
|
||||
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,61 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef ACLNN_KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define ACLNN_KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnKvQuantSparseAttnSharedkvMetadataGetWorkspaceSize(
|
||||
const aclTensor* cuSeqLensQOptional,
|
||||
const aclTensor* cuSeqLensOriKvOptional,
|
||||
const aclTensor* cuSeqLensCmpKvOptional,
|
||||
const aclTensor* sequsedQOptional,
|
||||
const aclTensor* sequsedKvOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsKv,
|
||||
int64_t headDim,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKvOptional,
|
||||
int64_t oriTopKOptional,
|
||||
int64_t cmpTopKOptional,
|
||||
int64_t kvQuantMode,
|
||||
int64_t tileSizeOptional,
|
||||
int64_t ropeHeadDimOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
int64_t oriMaskModeOptional,
|
||||
int64_t cmpMaskModeOptional,
|
||||
int64_t oriWinLeftOptional,
|
||||
int64_t oriWinRightOptional,
|
||||
char *layoutQOptional,
|
||||
char *layoutKvOptional,
|
||||
bool hasOriKvOptional,
|
||||
bool hasCmpKvOptional,
|
||||
const aclTensor* metaData,
|
||||
uint64_t* workspaceSize,
|
||||
aclOpExecutor** executor);
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnKvQuantSparseAttnSharedkvMetadata(void* workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor* executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif // ACLNN_SPARSE_FLASH_ATTENTION_ANTIQUANT_METADATA_AICPU_H
|
||||
@@ -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 l0_kv_quant_sparse_attn_sharedkv_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "l0_kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
#include "opdev/aicpu/aicpu_task.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
#include "opdev/op_def.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/shape_utils.h"
|
||||
|
||||
using namespace op;
|
||||
namespace l0op {
|
||||
OP_TYPE_REGISTER(KvQuantSparseAttnSharedkvMetadata);
|
||||
|
||||
const aclTensor* KvQuantSparseAttnSharedkvMetadata(
|
||||
const aclTensor* cuSeqLensQOptional,
|
||||
const aclTensor* cuSeqLensOriKvOptional,
|
||||
const aclTensor* cuSeqLensCmpKvOptional,
|
||||
const aclTensor* sequsedQOptional,
|
||||
const aclTensor* sequsedKvOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsKv,
|
||||
int64_t headDim,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKvOptional,
|
||||
int64_t oriTopKOptional,
|
||||
int64_t cmpTopKOptional,
|
||||
int64_t kvQuantMode,
|
||||
int64_t tileSizeOptional,
|
||||
int64_t ropeHeadDimOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
int64_t oriMaskModeOptional,
|
||||
int64_t cmpMaskModeOptional,
|
||||
int64_t oriWinLeftOptional,
|
||||
int64_t oriWinRightOptional,
|
||||
char *layoutQOptional,
|
||||
char *layoutKvOptional,
|
||||
bool hasOriKvOptional,
|
||||
bool hasCmpKvOptional,
|
||||
const char *socVersion,
|
||||
int64_t aicCoreNum,
|
||||
int64_t aivCoreNum,
|
||||
const aclTensor* metaData,
|
||||
aclOpExecutor* executor) {
|
||||
L0_DFX(KvQuantSparseAttnSharedkvMetadata, cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional,
|
||||
sequsedQOptional, sequsedKvOptional, numHeadsQ, numHeadsKv, headDim, batchSizeOptional, maxSeqlenQOptional,
|
||||
maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, kvQuantMode, tileSizeOptional, ropeHeadDimOptional,
|
||||
cmpRatioOptional, oriMaskModeOptional, cmpMaskModeOptional, oriWinLeftOptional, oriWinRightOptional,
|
||||
layoutQOptional, layoutKvOptional, hasOriKvOptional, hasCmpKvOptional, socVersion, aicCoreNum, aivCoreNum,
|
||||
metaData);
|
||||
|
||||
static internal::AicpuTaskSpace space("KvQuantSparseAttnSharedkvMetadata");
|
||||
|
||||
auto ret = ADD_TO_LAUNCHER_LIST_AICPU(
|
||||
KvQuantSparseAttnSharedkvMetadata,
|
||||
OP_ATTR_NAMES({"num_heads_q", "num_heads_kv", "head_dim", "batch_size", "max_seqlen_q", "max_seqlen_kv",
|
||||
"ori_topk", "cmp_topk", "kv_quant_mode", "tile_size", "rope_head_dim", "cmp_ratio", "ori_mask_mode",
|
||||
"cmp_mask_mode", "ori_win_left", "ori_win_right", "layout_q", "layout_kv", "has_ori_kv",
|
||||
"has_cmp_kv", "soc_version", "aic_core_num", "aiv_core_num"}),
|
||||
OP_INPUT(cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional, sequsedQOptional, sequsedKvOptional),
|
||||
OP_OUTPUT(metaData),
|
||||
OP_ATTR(numHeadsQ, numHeadsKv, headDim, batchSizeOptional, maxSeqlenQOptional, maxSeqlenKvOptional, oriTopKOptional,
|
||||
cmpTopKOptional, kvQuantMode, tileSizeOptional, ropeHeadDimOptional, cmpRatioOptional, oriMaskModeOptional,
|
||||
cmpMaskModeOptional, oriWinLeftOptional, oriWinRightOptional, layoutQOptional, layoutKvOptional,
|
||||
hasOriKvOptional, hasCmpKvOptional, socVersion, aicCoreNum, aivCoreNum));
|
||||
OP_CHECK(ret == ACL_SUCCESS,
|
||||
OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
|
||||
"KvQuantSparseAttnSharedkvMetadata"
|
||||
" ADD_TO_LAUNCHER_LIST_AICPU failed."),
|
||||
return nullptr);
|
||||
return metaData;
|
||||
}
|
||||
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,50 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef L0_KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define L0_KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
const aclTensor* KvQuantSparseAttnSharedkvMetadata(
|
||||
const aclTensor* cuSeqLensQOptional,
|
||||
const aclTensor* cuSeqLensOriKvOptional,
|
||||
const aclTensor* cuSeqLensCmpKvOptional,
|
||||
const aclTensor* sequsedQOptional,
|
||||
const aclTensor* sequsedKvOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsKv,
|
||||
int64_t headDim,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKvOptional,
|
||||
int64_t oriTopKOptional,
|
||||
int64_t cmpTopKOptional,
|
||||
int64_t kvQuantMode,
|
||||
int64_t tileSizeOptional,
|
||||
int64_t ropeHeadDimOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
int64_t oriMaskModeOptional,
|
||||
int64_t cmpMaskModeOptional,
|
||||
int64_t oriWinLeftOptional,
|
||||
int64_t oriWinRightOptional,
|
||||
char *layoutQOptional,
|
||||
char *layoutKvOptional,
|
||||
bool hasOriKvOptional,
|
||||
bool hasCmpKvOptional,
|
||||
const char *socVersion,
|
||||
int64_t aicCoreNum,
|
||||
int64_t aivCoreNum,
|
||||
const aclTensor* metaData,
|
||||
aclOpExecutor* executor);
|
||||
} // namespace l0op
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,26 @@
|
||||
# ---------------------------------------------------------------------------------------------------------
|
||||
# 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.
|
||||
# ---------------------------------------------------------------------------------------------------------
|
||||
|
||||
if (BUILD_WITH_INSTALLED_DEPENDENCY_CANN_PKG)
|
||||
if (NOT (UT_TEST_ALL OR OP_KERNEL_AICPU_UT))
|
||||
add_definitions(-D_GLIBCXX_USE_CXX11_ABI=1)
|
||||
set(CMAKE_CXX_COMPILER ${ASCEND_DIR}/toolkit/toolchain/hcc/bin/aarch64-target-linux-gnu-g++)
|
||||
endif()
|
||||
|
||||
file(GLOB_RECURSE JSON_FILE ${CMAKE_CURRENT_SOURCE_DIR}/*.json)
|
||||
file(GLOB AICPU_SRC ${CMAKE_CURRENT_SOURCE_DIR}/*_aicpu*.cpp)
|
||||
message(STATUS "[kv_quant_sparse_attn_sharedkv_metadata] Found aicpu sources: ${AICPU_SRC}, ascend dir: ${ASCEND_DIR}, ophost name: ${OPHOST_NAME}")
|
||||
|
||||
add_aicpu_cust_kernel_modules(kv_quant_sparse_attn_sharedkv_metadata ${AICPU_SRC} ${JSON_FILE})
|
||||
endif()
|
||||
|
||||
if(UT_TEST_ALL OR OP_KERNEL_AICPU_UT)
|
||||
AddAicpuOpTestCase(kv_quant_sparse_attn_sharedkv_metadata)
|
||||
endif()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,409 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_metadata_aicpu.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
|
||||
#include <array>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
#include "cpu_context.h"
|
||||
#include "cpu_kernel.h"
|
||||
#include "cpu_tensor.h"
|
||||
|
||||
namespace aicpu {
|
||||
constexpr int64_t FA_TOLERANCE_RATIO = 2;
|
||||
|
||||
enum BlockType : uint32_t {
|
||||
WIN_NORMAL_BLOCK = 0,
|
||||
WIN_TAIL_BLOCK,
|
||||
CMP_NORMAL_BLOCK,
|
||||
CMP_TAIL_BLOCK,
|
||||
BLOCK_MAX_TYPE
|
||||
};
|
||||
|
||||
enum class SparseMode : uint8_t {
|
||||
DEFAULT_MASK = 0,
|
||||
ALL_MASK,
|
||||
LEFT_UP_CAUSAL,
|
||||
RIGHT_DOWN_CAUSAL,
|
||||
BAND,
|
||||
SPARSE_BUTT,
|
||||
};
|
||||
|
||||
enum class ValidSocVersion {
|
||||
ASCEND910 = 0,
|
||||
ASCEND950,
|
||||
RESERVED_VERSION = 99999
|
||||
};
|
||||
|
||||
template<class T>
|
||||
using Range = std::pair<T, T>;
|
||||
|
||||
template<class T>
|
||||
using BlockCost = std::array<std::array<T, static_cast<size_t>(BLOCK_MAX_TYPE)>, static_cast<size_t>(BLOCK_MAX_TYPE)>;
|
||||
|
||||
template<typename T>
|
||||
T Clip(T value, T minValue, T maxValue)
|
||||
{
|
||||
if (value < minValue) {
|
||||
return minValue;
|
||||
}
|
||||
if (value > maxValue) {
|
||||
return maxValue;
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
template<typename T>
|
||||
inline bool IsWithinTolerance(T limit, T tolerance, T value)
|
||||
{
|
||||
return limit + tolerance >= value;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
inline typename std::enable_if<std::is_integral_v<T>, bool>::type GetAttrValue(CpuKernelContext &ctx,
|
||||
const std::string &name, T &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (!attr) {
|
||||
KERNEL_LOG_ERROR("attr is null: %s", name.c_str());
|
||||
return false;
|
||||
}
|
||||
value = static_cast<T>(attr->GetInt());
|
||||
return true;
|
||||
}
|
||||
|
||||
inline bool GetAttrValue(CpuKernelContext &ctx, const std::string &name, std::string &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (!attr) {
|
||||
KERNEL_LOG_ERROR("attr is null: %s", name.c_str());
|
||||
return false;
|
||||
}
|
||||
value = attr->GetString();
|
||||
return true;
|
||||
}
|
||||
|
||||
inline bool GetAttrValue(CpuKernelContext &ctx, const std::string &name, bool &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (!attr) {
|
||||
KERNEL_LOG_ERROR("attr is null: %s", name.c_str());
|
||||
return false;
|
||||
}
|
||||
value = attr->GetBool();
|
||||
return true;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
inline typename std::enable_if<std::is_integral_v<T>, void>::type GetAttrValueOpt(CpuKernelContext &ctx,
|
||||
const std::string &name, T &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (attr != nullptr) {
|
||||
value = static_cast<T>(attr->GetInt());
|
||||
}
|
||||
}
|
||||
|
||||
inline void GetAttrValueOpt(CpuKernelContext &ctx, const std::string &name, std::string &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (attr != nullptr) {
|
||||
value = attr->GetString();
|
||||
}
|
||||
}
|
||||
|
||||
inline void GetAttrValueOpt(CpuKernelContext &ctx, const std::string &name, bool &value)
|
||||
{
|
||||
auto attr = ctx.GetAttr(name);
|
||||
if (attr != nullptr) {
|
||||
value = attr->GetBool();
|
||||
}
|
||||
}
|
||||
|
||||
// 分核功能模块输出:FD信息,包含需要归约的数据索引及其分核信息
|
||||
struct FlashDecodeResult {
|
||||
uint32_t fdUsedVecNum { 0U }; // 归约过程使用的vector数量
|
||||
// 1、归约任务的索引信息
|
||||
std::vector<uint32_t> fdBN2Idx {}; // 每个归约任务的BN2索引,脚标为归约任务的序号,最大为核数-1
|
||||
std::vector<uint32_t> fdMIdx {}; // 每个归约任务的GS1索引,脚标为归约任务的序号
|
||||
std::vector<uint32_t> fdWorkspaceIdx {}; // 每个归约任务在workspace中的存放位置
|
||||
std::vector<uint32_t> fdS2SplitNum {}; // 每个归约任务的S2核间切分份数,脚标为归约任务的序号
|
||||
std::vector<uint32_t> fdMSize {}; // 每个归约任务m轴大小,脚标为归约任务的序号
|
||||
// 2、FD负载均衡阶段,归约任务的分核(vec)信息
|
||||
std::vector<uint32_t> fdIdx {}; // FD负载均衡阶段,每个vector处理的归约任务对应ID
|
||||
std::vector<uint32_t> fdMStart {}; // FD负载均衡阶段,每个vector处理的归约任务的m轴起点
|
||||
std::vector<uint32_t> fdMNum {}; // FD负载均衡阶段,每个vector处理的归约任务的m轴行数
|
||||
|
||||
FlashDecodeResult(uint32_t aicNum, uint32_t aivNum) :
|
||||
fdBN2Idx(aicNum),
|
||||
fdMIdx(aicNum),
|
||||
fdWorkspaceIdx(aicNum),
|
||||
fdS2SplitNum(aicNum),
|
||||
fdMSize(aicNum),
|
||||
fdIdx(aivNum),
|
||||
fdMStart(aivNum),
|
||||
fdMNum(aivNum) {}
|
||||
};
|
||||
|
||||
// 分核功能模块输出:FA阶段的核间分核信息
|
||||
struct SplitResult {
|
||||
uint32_t usedCoreNum { 0U }; // 使用的核数量
|
||||
std::vector<uint32_t> bN2End {}; // 每个核处理数据的BN2结束点
|
||||
std::vector<uint32_t> gS1End {}; // 每个核处理数据的GS1结束点
|
||||
std::vector<uint32_t> s2End {}; // 每个核处理数据的S2结束点
|
||||
std::vector<uint32_t> firstFdDataWorkspaceIdx {}; // 每个核第一份归约任务的存放位置
|
||||
int64_t maxCost { 0 }; // 慢核开销
|
||||
uint32_t numOfFdHead { 0U }; // 归约任务数量
|
||||
uint32_t maxS2SplitNum { 0U }; // 单个归约任务最大分核数量
|
||||
uint32_t maxS2GBaseNum { 0U }; // 单个核最大s2基本块数量
|
||||
FlashDecodeResult fdRes { 0U, 0U }; // FD信息
|
||||
|
||||
SplitResult(uint32_t aicNum, uint32_t aivNum) :
|
||||
bN2End(aicNum),
|
||||
gS1End(aicNum),
|
||||
s2End(aicNum),
|
||||
firstFdDataWorkspaceIdx(aicNum),
|
||||
fdRes(aicNum, aivNum) {};
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录切分信息
|
||||
struct SplitInfo {
|
||||
std::vector<uint32_t> s1GBaseNum {}; // S1G方向,切了多少个基本块
|
||||
std::vector<uint32_t> s2BaseNum {}; // S2方向,切了多少个基本块
|
||||
std::vector<uint32_t> s1GTailSize {}; // S1G方向,尾块size
|
||||
std::vector<uint32_t> s2TailSize {}; // S2方向,尾块size
|
||||
bool isKvSeqAllZero { true };
|
||||
|
||||
explicit SplitInfo(uint32_t batchSize) :
|
||||
s1GBaseNum(batchSize),
|
||||
s2BaseNum(batchSize),
|
||||
s1GTailSize(batchSize),
|
||||
s2TailSize(batchSize) {}
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录batch的开销信息
|
||||
struct CostInfo {
|
||||
std::vector<int64_t> bN2CostOfEachBatch {}; // 整个batch的开销
|
||||
std::vector<uint32_t> bN2BlockOfEachBatch {}; // 整个batch的开销
|
||||
std::vector<int64_t> bN2LastBlockCostOfEachBatch {}; // batch最后一块的开销
|
||||
uint32_t totalBlockNum { 0U };
|
||||
int64_t totalCost { 0 };
|
||||
int64_t maxS1GCost { 0 }; // 记录所有S1G行中的最大开销
|
||||
|
||||
explicit CostInfo(uint32_t batchSize) :
|
||||
bN2CostOfEachBatch(batchSize),
|
||||
bN2BlockOfEachBatch(batchSize),
|
||||
bN2LastBlockCostOfEachBatch(batchSize) {}
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:分核过程中,case基本信息的上下文信息,组合以减少接口传参数量
|
||||
struct SplitContext {
|
||||
SplitInfo splitInfo { 0U };
|
||||
CostInfo costInfo { 0U };
|
||||
|
||||
explicit SplitContext(uint32_t batchSize) :
|
||||
splitInfo(batchSize),
|
||||
costInfo(batchSize) {}
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录batch相关的临时信息
|
||||
struct BatchCache {
|
||||
uint32_t bIdx { 0U };
|
||||
uint32_t s1Size { 0U };
|
||||
uint32_t s2Size { 0U };
|
||||
int64_t preTokenLeftUp { 0 };
|
||||
int64_t nextTokenLeftUp { 0 };
|
||||
BlockCost<int64_t> typeCost {};
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录当前行(S1G)的临时信息
|
||||
struct S1GCache {
|
||||
uint32_t bIdx { 0U };
|
||||
uint32_t s1GIdx { 0U };
|
||||
uint32_t s2Start { 0U };
|
||||
uint32_t s2End { 0U };
|
||||
uint32_t winS2Start { 0U };
|
||||
uint32_t winS2End { 0U };
|
||||
uint32_t cmpS2Start { 0U }; // win部分与cmp部分的切分点
|
||||
uint32_t cmpS2End { 0U };
|
||||
int64_t s1GCost { 0 };
|
||||
int64_t s1GLastBlockCost { 0 };
|
||||
uint32_t s1GBlock { 0U };
|
||||
int64_t s1GNormalBlockCost { 0 };
|
||||
uint32_t winS1GBlock { 0U };
|
||||
int64_t winS1GCost { 0 };
|
||||
int64_t winS1GLastBlockCost { 0 };
|
||||
int64_t winS1GNormalBlockCost { 0 };
|
||||
uint32_t cmpS1GBlock { 0U };
|
||||
int64_t cmpS1GCost { 0 };
|
||||
int64_t cmpS1GLastBlockCost { 0 };
|
||||
int64_t cmpS1GNormalBlockCost { 0 };
|
||||
int64_t cmpS2TailSize {0};
|
||||
int64_t winS2TailSize {0};
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录分配过程中,当前核的负载信息
|
||||
struct CoreCache {
|
||||
int64_t costLimit { 0 }; // 负载上限
|
||||
int64_t cost { 0 }; // 已分配负载
|
||||
uint32_t block { 0U }; // 已分配块数
|
||||
};
|
||||
|
||||
// 分核功能模块内部使用:记录分配过程中的上下文信息
|
||||
struct AssignContext {
|
||||
uint32_t curBIdx { 0U };
|
||||
uint32_t curBN2Idx { 0U };
|
||||
uint32_t curS1GIdx { 0U };
|
||||
uint32_t curS2Idx { 0U };
|
||||
uint32_t curCoreIdx { 0U };
|
||||
int64_t unassignedCost { 0 };
|
||||
uint32_t usedCoreNum { 0U };
|
||||
uint32_t curKvSplitPart { 1U };
|
||||
uint32_t preFdDataNum { 0U };
|
||||
|
||||
int64_t bN2Cost { 0 };
|
||||
uint32_t bN2Block { 0U };
|
||||
bool isFinished { false };
|
||||
BatchCache batchCache {};
|
||||
S1GCache s1GCache {};
|
||||
CoreCache coreCache {};
|
||||
};
|
||||
|
||||
class KvQuantSparseAttnSharedkvMetadataCpuKernel : public CpuKernel {
|
||||
public:
|
||||
KvQuantSparseAttnSharedkvMetadataCpuKernel() = default;
|
||||
~KvQuantSparseAttnSharedkvMetadataCpuKernel() = default;
|
||||
uint32_t Compute(CpuKernelContext &ctx) override;
|
||||
|
||||
private:
|
||||
bool Prepare(CpuKernelContext &ctx);
|
||||
int32_t GetQueryBatchSize();
|
||||
int32_t GetKvBatchSize();
|
||||
bool CheckSingleParam();
|
||||
bool CheckExistence();
|
||||
bool CheckConsistency();
|
||||
bool CheckFeature();
|
||||
bool ParamsCheck();
|
||||
bool ParamsInit();
|
||||
bool BalanceSchedule(SplitResult &splitRes);
|
||||
bool GenMetaData(SplitResult &splitRes);
|
||||
ValidSocVersion ProcessSocVersion();
|
||||
// util
|
||||
uint32_t GetS1SeqSize(uint32_t bIdx);
|
||||
uint32_t GetS2SeqSize(uint32_t bIdx);
|
||||
int64_t CalcPreTokenLeftUp(uint32_t s1Size, uint32_t s2Size);
|
||||
int64_t CalcNextTokenLeftUp(uint32_t s1Size, uint32_t s2Size);
|
||||
Range<int64_t> CalcS2TokenRange(uint32_t s1GIdx, const BatchCache &batchCache);
|
||||
int64_t WinCalcCost(uint32_t basicM, uint32_t basicS2);
|
||||
int64_t CmpCalcCost(uint32_t basicM, uint32_t basicS2);
|
||||
void CalcCostTable(uint32_t s1NormalSize, uint32_t s2NormalSize, uint32_t s1GTailSize,
|
||||
uint32_t winS2TailSize, uint32_t cmpS2TailSize);
|
||||
|
||||
// cache calculation
|
||||
void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache);
|
||||
void CalcBlockRangeAndTailSize(Range<int64_t> &oriS2TokenRange, const BatchCache &batchCache, S1GCache &s1GCache);
|
||||
void CalcWinS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo);
|
||||
void CalcCmpS1GCache(S1GCache &s1GCache, const SplitInfo &splitInfo);
|
||||
void GatherWinAndCmpCache(S1GCache &s1GCache);
|
||||
void CalcS1GCache(uint32_t s1GIdx, const SplitContext &splitContext, const BatchCache &batchCache, S1GCache &s1GCache);
|
||||
|
||||
// preprocess
|
||||
void CalcSplitInfo(SplitContext &splitContext);
|
||||
void CalcBatchCost(uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo);
|
||||
void CalcCostInfo(SplitContext &splitContext);
|
||||
|
||||
// assign
|
||||
void UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void AssignByRow(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
int64_t CalcCurBlockCost(AssignContext &assignContext);
|
||||
void AssignByBlock(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void ForceAssign(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void AssignBlocksToCore(const SplitContext &splitContext, AssignContext &assignContext, SplitResult &result);
|
||||
|
||||
// FD
|
||||
bool IsNeedRecordFDInfo(const AssignContext &assignContext, const SplitResult &splitRes);
|
||||
void RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, SplitResult &result);
|
||||
|
||||
// main
|
||||
void SplitFD(SplitResult &splitRes);
|
||||
void CalcSplitPlan(int64_t costLimit, const SplitContext &splitContext, SplitResult &result);
|
||||
void SplitCore();
|
||||
|
||||
private:
|
||||
// input
|
||||
Tensor *actSeqLenQ_ = nullptr;
|
||||
Tensor *actSeqLenOriKv_ = nullptr;
|
||||
Tensor *actSeqLenCmpKv_ = nullptr;
|
||||
Tensor *seqUsedQ_ = nullptr;
|
||||
Tensor *seqUsedKv_ = nullptr;
|
||||
|
||||
// output
|
||||
Tensor *metaData_ = nullptr;
|
||||
|
||||
// attributes
|
||||
int32_t batchSize_ = 0;
|
||||
int32_t querySeqSize_ = 0;
|
||||
int32_t queryHeadNum_ = 0;
|
||||
int32_t kvSeqSize_ = 0;
|
||||
int32_t kvHeadNum_ = 0;
|
||||
int32_t headDim_ = 0;
|
||||
int32_t oriTopK_ = 0;
|
||||
int32_t cmpTopK_ = 0;
|
||||
int32_t cmpRatio_ = -1;
|
||||
int32_t oriMaskMode_ = 4;
|
||||
int32_t cmpMaskMode_ = 3;
|
||||
int64_t winLeft_ = 127;
|
||||
int64_t winRight_ = 0;
|
||||
std::string layoutQuery_ = "BSND";
|
||||
std::string layoutKv_ = "PA_ND";
|
||||
bool hasOriKv_ = true;
|
||||
bool hasCmpKv_ = true;
|
||||
uint32_t aicCoreNum_ = 24U;
|
||||
uint32_t aivCoreNum_ = 48U;
|
||||
|
||||
// attr
|
||||
std::string socVersion_ = "ascend910B";
|
||||
int64_t preToken_ = 0; // new
|
||||
int64_t nextToken_ = 0; // new
|
||||
uint32_t groupSize_ = 0;
|
||||
uint32_t mBaseSize_ = 0;
|
||||
uint32_t s2BaseSize_ = 0;
|
||||
bool isS1G_ = true;
|
||||
bool isCFA = false;
|
||||
bool isSCFA = false;
|
||||
bool supportFd = false;
|
||||
uint32_t sparseMode_ = 0;
|
||||
uint32_t attentionMode_ = 1;
|
||||
BlockCost<int64_t> typeCost_;
|
||||
bool isN128 = false;
|
||||
|
||||
private:
|
||||
enum class ParamId : uint32_t {
|
||||
// input
|
||||
actSeqLenQ = 0,
|
||||
actSeqLenOriKv = 1,
|
||||
actSeqLenCmpKv = 2,
|
||||
seqUsedQ = 3,
|
||||
seqUsedKv = 4,
|
||||
// output
|
||||
metaData = 0,
|
||||
};
|
||||
};
|
||||
} // namespace aicpu
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,15 @@
|
||||
{
|
||||
"KvQuantSparseAttnSharedkvMetadata":{
|
||||
"opInfo":{
|
||||
"computeCost":"100",
|
||||
"engine":"DNN_VM_AICPU",
|
||||
"flagAsync":"False",
|
||||
"flagPartial":"False",
|
||||
"functionName":"RunCpuKernel",
|
||||
"kernelSo":"libtransformer_aicpu_kernels.so",
|
||||
"opKernelLib":"CUSTAICPUKernel",
|
||||
"userDefined":"True",
|
||||
"workspaceSize":"100"
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user