10
csrc/attention/sparse_attn_sharedkv_metadata/CMakeLists.txt
Normal file
10
csrc/attention/sparse_attn_sharedkv_metadata/CMakeLists.txt
Normal file
@@ -0,0 +1,10 @@
|
||||
# ---------------------------------------------------------------------------------------------------------
|
||||
# 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_modules_sources_aicpu(DEPENDENCIES sparse_attn_sharedkv)
|
||||
208
csrc/attention/sparse_attn_sharedkv_metadata/README.md
Normal file
208
csrc/attention/sparse_attn_sharedkv_metadata/README.md
Normal file
@@ -0,0 +1,208 @@
|
||||
# SparseAttnSharedkvMetadata
|
||||
|
||||
## 产品支持情况
|
||||
|
||||
| 产品 | 是否支持 |
|
||||
| ------------------------------------------------------------ | :------: |
|
||||
|<term>Atlas A2 推理系列产品</term> | √ |
|
||||
|<term>Atlas A3 推理系列产品</term> | √ |
|
||||
|
||||
## 功能说明
|
||||
|
||||
- API功能:`SparseAttnSharedkvMetadata`算子旨在生成一个任务列表,包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及 Q 和 K 的分块的索引,供后续`SparseAttnSharedkv`算子使用。
|
||||
|
||||
## 参数说明
|
||||
|
||||
<table style="undefined;table-layout: fixed; width: 1576px">
|
||||
<colgroup>
|
||||
<col style="width: 170px">
|
||||
<col style="width: 170px">
|
||||
<col style="width: 310px">
|
||||
<col style="width: 212px">
|
||||
<col style="width: 100px">
|
||||
</colgroup>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>参数名</th>
|
||||
<th>输入/输出/属性</th>
|
||||
<th>描述</th>
|
||||
<th>数据类型</th>
|
||||
<th>数据格式</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td>num_heads_q</td>
|
||||
<td>属性</td>
|
||||
<td>Q的多头数,目前仅支持64。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>num_heads_kv</td>
|
||||
<td>属性</td>
|
||||
<td>K和V的多头数,目前仅支持1。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>head_dim</td>
|
||||
<td>属性</td>
|
||||
<td>注意力头的维度,目前仅支持512。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cu_seqlens_q</td>
|
||||
<td>可选输入</td>
|
||||
<td>当layout_query为TND时,表示不同Batch中q的有效token数,维度为B+1,大小为参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cu_seqlens_ori_kv</td>
|
||||
<td>可选输入</td>
|
||||
<td>当layout_kv为TND时,表示不同Batch中ori_kv的有效token数,维度为B+1,大小为参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和。目前layout_kv仅支持PA_ND,故设置此参数无效。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cu_seqlens_cmp_kv</td>
|
||||
<td>可选输入</td>
|
||||
<td>当layout_kv为TND时,表示不同Batch中cmp_kv的有效token数,维度为B+1,大小为参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和。目前layout_kv仅支持PA_ND,故设置此参数无效。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>seqused_q</td>
|
||||
<td>可选输入</td>
|
||||
<td>表示不同Batch中q实际参与运算的token数,维度为B。目前暂不支持指定该参数。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>seqused_kv</td>
|
||||
<td>可选输入</td>
|
||||
<td>表示不同Batch中ori_kv实际参与运算的token数,维度为B。</td>
|
||||
<td>INT32</td>
|
||||
<td>ND</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>batch_size</td>
|
||||
<td>可选属性</td>
|
||||
<td>输入样本批量大小。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>max_seqlen_q</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示所有batch中`q`的最大有效token数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>max_seqlen_kv</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示所有batch中`ori_kv`的最大有效token数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ori_topk</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示通过QLI算法从`ori_kv`中筛选出的关键稀疏token的个数。目前暂不支持指定该参数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cmp_topk</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示通过QLI算法从`cmp_kv`中筛选出的关键稀疏token的个数,目前仅支持512。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cmp_ratio</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示对`ori_kv`的压缩率,数据范围支持4/128,</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ori_mask_mode</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示q和ori_kv计算的mask模式,仅支持输入默认值4,代表band模式的mask。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cmp_mask_mode</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示q和cmp_kv计算的mask模式,仅支持输入默认值3,代表rightDownCausal模式的mask。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ori_win_left</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示q和ori_kv计算中q对过去token计算的数量,仅支持默认值127。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>ori_win_right</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示q和ori_kv计算中q对未来token计算的数量,仅支持默认值0。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_q</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于标识输入q的数据排布格式,默认值为BSND,目前支持传入BSND和TND。</td>
|
||||
<td>STRING</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_kv</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于标识输入ori_kv和cmp_kv的数据排布格式,目前仅支持传入默认值PA_ND。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>has_ori_kv</td>
|
||||
<td>可选属性</td>
|
||||
<td>是否传入ori_kv。</td>
|
||||
<td>BOOL</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>has_cmp_kv</td>
|
||||
<td>可选属性</td>
|
||||
<td>是否传入cmp_kv。</td>
|
||||
<td>BOOL</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>device</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于获取设备信息。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>metadata</td>
|
||||
<td>输出</td>
|
||||
<td>每个cube核上FlashAttention计算任务的Batch、Head、以及 Q 和 K 的分块的索引,以及每个vector核上FlashDecode的规约任务索引。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## 约束说明
|
||||
|
||||
- 该接口支持推理场景下使用。
|
||||
- 该接口支持aclgraph模式。
|
||||
- Tensor不能全传None。
|
||||
@@ -0,0 +1,143 @@
|
||||
/**
|
||||
* 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_sparse_attn_sharedkv_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "aclnn_sparse_attn_sharedkv_metadata.h"
|
||||
#include "l0_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 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 aclnnSparseAttnSharedkvMetadataGetWorkspaceSize(
|
||||
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 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(aclnnSparseAttnSharedkvMetadata,
|
||||
DFX_IN(cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional, sequsedQOptional, sequsedKvOptional, numHeadsQ, numHeadsKv, headDim, batchSizeOptional,
|
||||
maxSeqlenQOptional, maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, 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, 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::SparseAttnSharedkvMetadata(
|
||||
cuSeqLensQOptionalContiguous, cuSeqLensOriKvOptionalContiguous, cuSeqLensCmpKvOptionalContiguous,
|
||||
sequsedQOptionalContiguous, sequsedKvOptionalContiguous, numHeadsQ, numHeadsKv, headDim, batchSizeOptional,
|
||||
maxSeqlenQOptional, maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, 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
|
||||
aclnnSparseAttnSharedkvMetadata(void *workspace, uint64_t workspaceSize,
|
||||
aclOpExecutor *executor, aclrtStream stream) {
|
||||
L2_DFX_PHASE_2(aclnnSparseAttnSharedkvMetadata);
|
||||
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,58 @@
|
||||
/**
|
||||
* 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_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define ACLNN_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnSparseAttnSharedkvMetadataGetWorkspaceSize(
|
||||
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 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
|
||||
aclnnSparseAttnSharedkvMetadata(void* workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor* executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif // ACLNN_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
@@ -0,0 +1,86 @@
|
||||
/**
|
||||
* 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_sparse_attn_sharedkv_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "l0_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(SparseAttnSharedkvMetadata);
|
||||
|
||||
const aclTensor* SparseAttnSharedkvMetadata(
|
||||
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 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(SparseAttnSharedkvMetadata, cuSeqLensQOptional, cuSeqLensOriKvOptional, cuSeqLensCmpKvOptional, sequsedQOptional, sequsedKvOptional, numHeadsQ, numHeadsKv, headDim, batchSizeOptional,
|
||||
maxSeqlenQOptional, maxSeqlenKvOptional, oriTopKOptional, cmpTopKOptional, cmpRatioOptional, oriMaskModeOptional,
|
||||
cmpMaskModeOptional, oriWinLeftOptional, oriWinRightOptional, layoutQOptional, layoutKvOptional,
|
||||
hasOriKvOptional, hasCmpKvOptional, socVersion, aicCoreNum, aivCoreNum, metaData);
|
||||
|
||||
static internal::AicpuTaskSpace space(
|
||||
"SparseAttnSharedkvMetadata");
|
||||
|
||||
auto ret = ADD_TO_LAUNCHER_LIST_AICPU(
|
||||
SparseAttnSharedkvMetadata,
|
||||
OP_ATTR_NAMES({"num_heads_q", "num_heads_kv", "head_dim", "batch_size", "max_seqlen_q", "max_seqlen_kv",
|
||||
"ori_topk", "cmp_topk", "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, cmpRatioOptional, oriMaskModeOptional,
|
||||
cmpMaskModeOptional, oriWinLeftOptional, oriWinRightOptional, layoutQOptional, layoutKvOptional,
|
||||
hasOriKvOptional, hasCmpKvOptional, socVersion,
|
||||
aicCoreNum, aivCoreNum));
|
||||
OP_CHECK(ret == ACL_SUCCESS,
|
||||
OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
|
||||
"SparseAttnSharedkvMetadata"
|
||||
" ADD_TO_LAUNCHER_LIST_AICPU failed."),
|
||||
return nullptr);
|
||||
return metaData;
|
||||
}
|
||||
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,47 @@
|
||||
/**
|
||||
* 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_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define L0_SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
const aclTensor* SparseAttnSharedkvMetadata(
|
||||
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 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,54 @@
|
||||
/**
|
||||
* 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 sparse_attn_sharedkv_metadata_proto.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
#define SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
|
||||
#include "graph/operator_reg.h"
|
||||
#include "graph/types.h"
|
||||
|
||||
namespace ge {
|
||||
|
||||
REG_OP(SparseAttnSharedkvMetadata)
|
||||
.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)
|
||||
.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(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(SparseAttnSharedkvMetadata)
|
||||
|
||||
} // namespace ge
|
||||
|
||||
#endif // SPARSE_ATTN_SHAREDKV_METADATA_PROTO_H
|
||||
@@ -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 sparse_attn_sharedkv_metadata_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "../../sparse_attn_sharedkv/op_kernel/sparse_attn_sharedkv_metadata.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
static ge::graphStatus InferShapeSparseAttnSharedkvMetadata(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 InferDtypeSparseAttnSharedkvMetadata(gert::InferDataTypeContext* context)
|
||||
{
|
||||
context->SetOutputDataType(0, DT_INT32);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(SparseAttnSharedkvMetadata)
|
||||
.InferShape(InferShapeSparseAttnSharedkvMetadata)
|
||||
.InferDataType(InferDtypeSparseAttnSharedkvMetadata);
|
||||
} // namespace ops
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,350 @@
|
||||
/**
|
||||
* 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 sparse_attn_sharedkv_metadata_aicpu.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef SPARSE_ATTN_SHAREDKV_METADATA_AICPU_H
|
||||
#define 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;
|
||||
}
|
||||
|
||||
// 分核功能模块输出: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 }; // 单个归约任务最大分核数量
|
||||
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 curFdDataNum { 1U };
|
||||
|
||||
int64_t bN2Cost { 0 };
|
||||
uint32_t bN2Block { 0U };
|
||||
bool isFinished { false };
|
||||
BatchCache batchCache {};
|
||||
S1GCache s1GCache {};
|
||||
CoreCache coreCache {};
|
||||
};
|
||||
|
||||
class SparseAttnSharedkvMetadataCpuKernel : public CpuKernel {
|
||||
public:
|
||||
SparseAttnSharedkvMetadataCpuKernel() = default;
|
||||
~SparseAttnSharedkvMetadataCpuKernel() = default;
|
||||
uint32_t Compute(CpuKernelContext &ctx) override;
|
||||
|
||||
private:
|
||||
bool Prepare(CpuKernelContext &ctx);
|
||||
bool ParamsCheck();
|
||||
int32_t GetQueryBatchSize();
|
||||
int32_t GetKvBatchSize();
|
||||
bool CheckSingleParam();
|
||||
bool CheckExistence();
|
||||
bool CheckConsistency();
|
||||
bool CheckFeature();
|
||||
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:
|
||||
// context for log use
|
||||
CpuKernelContext *context_ = nullptr;
|
||||
|
||||
// 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_;
|
||||
|
||||
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 @@
|
||||
{
|
||||
"SparseAttnSharedkvMetadata":{
|
||||
"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