@@ -0,0 +1,11 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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(OPTYPE vllm_quant_lightning_indexer_metadata ACLNNTYPE aclnn)
|
||||
181
csrc/attention/vllm_quant_lightning_indexer_metadata/README.md
Normal file
181
csrc/attention/vllm_quant_lightning_indexer_metadata/README.md
Normal file
@@ -0,0 +1,181 @@
|
||||
# VllmQuantLightningIndexerMetadata
|
||||
|
||||
## 产品支持情况
|
||||
|
||||
| 产品 | 是否支持 |
|
||||
| ------------------------------------------------------------ | :------: |
|
||||
|<term>Ascend 950PR/Ascend 950DT</term> | √ |
|
||||
|<term>Atlas A3 推理系列产品</term> | √ |
|
||||
|<term>Atlas A2 推理系列产品</term> | √ |
|
||||
|
||||
## 功能说明
|
||||
|
||||
- API功能:VllmQuantLightningIndexerMetadata是VllmQuantLightningIndexer的前置算子,通过AICPU为VllmQuantLightningIndexer算子生成分核结果,包括每个核需要处理的数据的起始点、结束点等内容,随后,VllmQuantLightningIndexer根据该分核结果进行实际计算。
|
||||
|
||||
- 主要计算过程为:
|
||||
1. 获取每个`batch`的基本块大小,并计算负载。
|
||||
2. 计算所有`batch`的总负载和总的基本块个数。
|
||||
3. 为每个核分配负载,并记录分核结果,分核结果包括每个核需要处理的数据的起始点、结束点等内容。
|
||||
|
||||
## 参数说明
|
||||
|
||||
>- 参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示hidden层的大小、N(Head Num)表示多头数、D(Head Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。
|
||||
>- 使用S1和S2分别表示query和key的输入样本序列长度,N1和N2分别表示query和key对应的多头数,k表示最后选取的索引个数。
|
||||
|
||||
<table style="undefined;table-layout: fixed; width: 1000px">
|
||||
<colgroup>
|
||||
<col style="width: 100px">
|
||||
<col style="width: 120px">
|
||||
<col style="width: 500px">
|
||||
<col style="width: 80px">
|
||||
<col style="width: 80px">
|
||||
</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的多头数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>num_heads_k</td>
|
||||
<td>属性</td>
|
||||
<td>K的多头数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>head_dim</td>
|
||||
<td>属性</td>
|
||||
<td>注意力头的维度。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>query_quant_mode</td>
|
||||
<td>属性</td>
|
||||
<td>用于标识query的量化模式,当前支持Per-Token-Head量化模式,当前仅支持传入0</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>key_quant_mode</td>
|
||||
<td>属性</td>
|
||||
<td>用于标识输入key的量化模式,当前支持Per-Token-Head量化模式,当前仅支持传入0。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>actual_seq_lengths_query</td>
|
||||
<td>可选输入</td>
|
||||
<td>表示不同Batch中`query`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和`query`的shape的S长度相同。该入参中每个Batch的有效token数不超过`query`中的维度S大小且不小于0。支持长度为B的一维tensor。<br>当`layout_query`为TND时,该入参必须传入,且以该入参元素的数量作为B值,该入参中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。不能出现负值。特殊说明:actual\_seq\_lengths\_query和actual\_seq\_lengths\_key至少传入一个。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>actual_seq_lengths_key</td>
|
||||
<td>可选输入</td>
|
||||
<td>表示不同Batch中压缩前原始`key`的有效token数,数据类型支持`int32`。如果不指定seqlen可传入None,表示和key的shape的S长度相同。该参数中每个Batch的原始有效token数除以压缩率后不超过`key`中的维度S大小且不小于0,支持长度为B的一维tensor。<br>当`layout_kv`为TND或PA_BSND时,该入参必须传入,`layout_kv`为TND,该参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须大于等于前一个元素的值。特殊说明:actual\_seq\_lengths\_query和actual\_seq\_lengths\_key至少传入一个。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</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>当layout_query为BSND时,表示每个Batch中的q的有效token数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>max_seqlen_k</td>
|
||||
<td>可选属性</td>
|
||||
<td>当layout_kv为BSND时,表示每个Batch中的k的有效token数。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>cmp_ratio</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于稀疏计算,表示key的压缩倍数。Atlas A3 推理系列产品支持1/2/4/8/16/32/64/128,Ascend 950PR/Ascend 950DT支持1/4/128。数据类型支持int32,默认值1。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_query</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于标识query的数据排布格式,当前支持BSND、TND,默认值"BSND"。</td>
|
||||
<td>STRING</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>layout_key</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于标识key的数据排布格式,当前仅支持PA_BSND。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>sparse_count</td>
|
||||
<td>可选属性</td>
|
||||
<td>代表topK阶段需要保留的block数量,Atlas A3推理系列产品支持[1, 2048],Ascend 950PR/Ascend 950DT支持512。</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>sparse_mode</td>
|
||||
<td>可选属性</td>
|
||||
<td>表示sparse的模式,支持0/3,数据类型支持int32。为0时,代表defaultMask模式。为3时,代表rightDownCausal模式的mask,对应以右顶点为划分的下三角场景.</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>pre_token</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于稀疏计算,表示attention需要和前几个Token计算关联,仅支持默认值2^63-1。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>next_token</td>
|
||||
<td>可选属性</td>
|
||||
<td>用于稀疏计算,表示attention需要和前几个Token计算关联,仅支持默认值2^63-1。</td>
|
||||
<td>INT64</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>device</td>
|
||||
<td>可选属性</td>
|
||||
<td>npu的ID。</td>
|
||||
<td>STRING</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>metadata</td>
|
||||
<td>输出</td>
|
||||
<td>VllmQuantLightningIndexerMetadata算子传入的分核信息,包括每个Cube核上FlashAttention计算任务的Batch、Head以及Q和K分块的索引,以及每个Vector核上FlashDecode的规约任务索引。数据类型支持`int32`,shape大小为[1024]</td>
|
||||
<td>INT32</td>
|
||||
<td>-</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
## 约束说明
|
||||
|
||||
- 该接口支持推理场景下使用。
|
||||
- 该接口支持aclgraph模式。
|
||||
@@ -0,0 +1,128 @@
|
||||
/**
|
||||
* 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_vllm_quant_lightning_indexer_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "aclnn_vllm_quant_lightning_indexer_metadata.h"
|
||||
#include "aclnn/aclnn_base.h"
|
||||
#include "aclnn_kernels/common/op_error_check.h"
|
||||
#include "aclnn_kernels/contiguous.h"
|
||||
#include "aclnn_kernels/reshape.h"
|
||||
#include "l0_vllm_quant_lightning_indexer_metadata.h"
|
||||
#include "opdev/common_types.h"
|
||||
#include "opdev/data_type_utils.h"
|
||||
#include "opdev/format_utils.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/tensor_view_utils.h"
|
||||
#include "opdev/platform.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
static aclnnStatus ParamsCheck(
|
||||
const aclTensor* actualSeqLengthsQueryOptional,
|
||||
const aclTensor* actualSeqLengthsKeyOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsK,
|
||||
int64_t headDim,
|
||||
int64_t queryQuantMode,
|
||||
int64_t keyQuantMode,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKOptional,
|
||||
char* layoutQueryOptional,
|
||||
char* layoutKeyOptional,
|
||||
int64_t sparseCountOptional,
|
||||
int64_t sparseModeOptional,
|
||||
int64_t preTokensOptional,
|
||||
int64_t nextTokensOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
const aclTensor* metaData) {
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
__attribute__((visibility("default")))
|
||||
aclnnStatus aclnnVllmQuantLightningIndexerMetadataGetWorkspaceSize(
|
||||
const aclTensor* actualSeqLengthsQueryOptional,
|
||||
const aclTensor* actualSeqLengthsKeyOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsK,
|
||||
int64_t headDim,
|
||||
int64_t queryQuantMode,
|
||||
int64_t keyQuantMode,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKOptional,
|
||||
char* layoutQueryOptional,
|
||||
char* layoutKeyOptional,
|
||||
int64_t sparseCountOptional,
|
||||
int64_t sparseModeOptional,
|
||||
int64_t preTokensOptional,
|
||||
int64_t nextTokensOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
const aclTensor* metaData,
|
||||
uint64_t* workspaceSize,
|
||||
aclOpExecutor** executor) {
|
||||
L2_DFX_PHASE_1(
|
||||
aclnnVllmQuantLightningIndexerMetadata,
|
||||
DFX_IN(actualSeqLengthsQueryOptional, actualSeqLengthsKeyOptional, numHeadsQ, numHeadsK, headDim, queryQuantMode,
|
||||
keyQuantMode, batchSizeOptional, maxSeqlenQOptional, maxSeqlenKOptional, layoutQueryOptional, layoutKeyOptional,
|
||||
sparseCountOptional, sparseModeOptional, preTokensOptional, nextTokensOptional, cmpRatioOptional),
|
||||
DFX_OUT(metaData));
|
||||
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
|
||||
auto ret = ParamsCheck(actualSeqLengthsQueryOptional, actualSeqLengthsKeyOptional, numHeadsQ, numHeadsK, headDim, queryQuantMode,
|
||||
keyQuantMode, batchSizeOptional, maxSeqlenQOptional, maxSeqlenKOptional, layoutQueryOptional, layoutKeyOptional,
|
||||
sparseCountOptional, sparseModeOptional, preTokensOptional, nextTokensOptional, cmpRatioOptional, 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 actualSeqLengthsQueryOptionalContiguous = l0op::Contiguous(actualSeqLengthsQueryOptional, uniqueExecutor.get());
|
||||
CHECK_RET(actualSeqLengthsQueryOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto actualSeqLengthsKeyOptionalContiguous = l0op::Contiguous(actualSeqLengthsKeyOptional, uniqueExecutor.get());
|
||||
CHECK_RET(actualSeqLengthsKeyOptionalContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
auto output = l0op::VllmQuantLightningIndexerMetadata(
|
||||
actualSeqLengthsQueryOptionalContiguous, actualSeqLengthsKeyOptionalContiguous, aicCoreNum, aivCoreNum, socVersion,
|
||||
numHeadsQ, numHeadsK, headDim, queryQuantMode, keyQuantMode, batchSizeOptional, maxSeqlenQOptional,
|
||||
maxSeqlenKOptional, layoutQueryOptional, layoutKeyOptional, sparseCountOptional, sparseModeOptional,
|
||||
preTokensOptional, nextTokensOptional, cmpRatioOptional, metaData, uniqueExecutor.get());
|
||||
CHECK_RET(output != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
|
||||
*workspaceSize = 0;
|
||||
uniqueExecutor.ReleaseTo(executor);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnVllmQuantLightningIndexerMetadata(void* workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor* executor,
|
||||
aclrtStream stream) {
|
||||
L2_DFX_PHASE_2(aclnnVllmQuantLightningIndexerMetadata);
|
||||
return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,83 @@
|
||||
/**
|
||||
* 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_QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
#define ACLNN_QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
/* function: aclnnVllmQuantLightningIndexerMetadataGetWorkspaceSize
|
||||
* parameters :
|
||||
* actualSeqLengthsQuery : optional
|
||||
* actualSeqLengthsKey : optional
|
||||
* numHeadsQ : required
|
||||
* numHeadsK : required
|
||||
* headDim : required
|
||||
* queryQuantMode : required
|
||||
* keyQuantMode : required
|
||||
* int64_t batchSize : optional
|
||||
* int64_t maxSeqlenQ : optional
|
||||
* int64_t maxSeqlenK : optional
|
||||
* char* layoutQuery : optional
|
||||
* char* layoutKey : optional
|
||||
* int64_t sparseCount : optional
|
||||
* int64_t sparseMode : optional
|
||||
* int64_t preTokens : optional
|
||||
* int64_t nextTokens : optional
|
||||
* cmpRatio : optional
|
||||
* out : required
|
||||
* workspaceSize : size of workspace(output).
|
||||
* executor : executor context(output).
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnVllmQuantLightningIndexerMetadataGetWorkspaceSize(
|
||||
const aclTensor* actualSeqLengthsQueryOptional,
|
||||
const aclTensor* actualSeqLengthsKeyOptional,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsK,
|
||||
int64_t headDim,
|
||||
int64_t queryQuantMode,
|
||||
int64_t keyQuantMode,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKOptional,
|
||||
char* layoutQueryOptional,
|
||||
char* layoutKeyOptional,
|
||||
int64_t sparseCountOptional,
|
||||
int64_t sparseModeOptional,
|
||||
int64_t preTokensOptional,
|
||||
int64_t nextTokensOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
const aclTensor* metaData,
|
||||
uint64_t* workspaceSize,
|
||||
aclOpExecutor** executor);
|
||||
|
||||
/* function: aclnnVllmQuantLightningIndexerMetadata
|
||||
* parameters :
|
||||
* workspace : workspace memory addr(input).
|
||||
* workspaceSize : size of workspace(input).
|
||||
* executor : executor context(input).
|
||||
* stream : acl stream.
|
||||
*/
|
||||
__attribute__((visibility("default"))) aclnnStatus
|
||||
aclnnVllmQuantLightningIndexerMetadata(void* workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor* executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif // ACLNN_QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
@@ -0,0 +1,75 @@
|
||||
/**
|
||||
* 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_vllm_quant_lightning_indexer_metadata.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "l0_vllm_quant_lightning_indexer_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(VllmQuantLightningIndexerMetadata);
|
||||
|
||||
const aclTensor* VllmQuantLightningIndexerMetadata(
|
||||
const aclTensor* actualSeqLengthsQueryOptional,
|
||||
const aclTensor* actualSeqLengthsKeyOptional,
|
||||
int64_t aicCoreNum,
|
||||
int64_t aivCoreNum,
|
||||
const char* socVersion,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsK,
|
||||
int64_t headDim,
|
||||
int64_t queryQuantMode,
|
||||
int64_t keyQuantMode,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKOptional,
|
||||
char* layoutQueryOptional,
|
||||
char* layoutKeyOptional,
|
||||
int64_t sparseCountOptional,
|
||||
int64_t sparseModeOptional,
|
||||
int64_t preTokensOptional,
|
||||
int64_t nextTokensOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
const aclTensor* metaData,
|
||||
aclOpExecutor* executor) {
|
||||
L0_DFX(VllmQuantLightningIndexerMetadata, actualSeqLengthsQueryOptional, actualSeqLengthsKeyOptional, aicCoreNum, aivCoreNum, socVersion,
|
||||
numHeadsQ, numHeadsK, headDim, queryQuantMode, keyQuantMode, batchSizeOptional, maxSeqlenQOptional,
|
||||
maxSeqlenKOptional, layoutQueryOptional, layoutKeyOptional, sparseCountOptional, sparseModeOptional,
|
||||
preTokensOptional, nextTokensOptional, cmpRatioOptional, metaData);
|
||||
|
||||
static internal::AicpuTaskSpace space("VllmQuantLightningIndexerMetadata");
|
||||
|
||||
auto ret = ADD_TO_LAUNCHER_LIST_AICPU(
|
||||
VllmQuantLightningIndexerMetadata,
|
||||
OP_ATTR_NAMES({"aic_core_num", "aiv_core_num", "soc_version", "num_heads_q", "num_heads_k", "head_dim", "query_quant_mode",
|
||||
"key_quant_mode", "batch_size", "max_seqlen_q", "max_seqlen_k", "layout_query", "layout_key", "sparse_count",
|
||||
"sparse_mode", "pre_tokens", "next_tokens", "cmp_ratio"}),
|
||||
OP_INPUT(actualSeqLengthsQueryOptional, actualSeqLengthsKeyOptional), OP_OUTPUT(metaData),
|
||||
OP_ATTR(aicCoreNum, aivCoreNum, socVersion, numHeadsQ, numHeadsK, headDim, queryQuantMode, keyQuantMode,
|
||||
batchSizeOptional, maxSeqlenQOptional, maxSeqlenKOptional, layoutQueryOptional, layoutKeyOptional,
|
||||
sparseCountOptional, sparseModeOptional, preTokensOptional, nextTokensOptional, cmpRatioOptional));
|
||||
OP_CHECK(ret == ACL_SUCCESS,
|
||||
OP_LOGE(ACLNN_ERR_INNER_NULLPTR,
|
||||
"VllmQuantLightningIndexerMetadata"
|
||||
" ADD_TO_LAUNCHER_LIST_AICPU failed."),
|
||||
return nullptr);
|
||||
return metaData;
|
||||
}
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,42 @@
|
||||
/**
|
||||
* 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_QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
#define L0_QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
const aclTensor* VllmQuantLightningIndexerMetadata(
|
||||
const aclTensor* actualSeqLengthsQueryOptional,
|
||||
const aclTensor* actualSeqLengthsKeyOptional,
|
||||
int64_t aicCoreNum,
|
||||
int64_t aivCoreNum,
|
||||
const char* socVersion,
|
||||
int64_t numHeadsQ,
|
||||
int64_t numHeadsK,
|
||||
int64_t headDim,
|
||||
int64_t queryQuantMode,
|
||||
int64_t keyQuantMode,
|
||||
int64_t batchSizeOptional,
|
||||
int64_t maxSeqlenQOptional,
|
||||
int64_t maxSeqlenKOptional,
|
||||
char* layoutQueryOptional,
|
||||
char* layoutKeyOptional,
|
||||
int64_t sparseCountOptional,
|
||||
int64_t sparseModeOptional,
|
||||
int64_t preTokensOptional,
|
||||
int64_t nextTokensOptional,
|
||||
int64_t cmpRatioOptional,
|
||||
const aclTensor* metaData,
|
||||
aclOpExecutor* executor);
|
||||
} // namespace l0op
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,83 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vllm_quant_lightning_indexer_metadata_proto.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_METADATA_PROTO_H
|
||||
#define QUANT_LIGHTNING_INDEXER_METADATA_PROTO_H
|
||||
|
||||
#include "graph/operator_reg.h"
|
||||
#include "graph/types.h"
|
||||
|
||||
namespace ge {
|
||||
|
||||
/**
|
||||
* @brief Function VllmQuantLightningIndexerMetadata.
|
||||
|
||||
* @par Inputs:
|
||||
* @li actual_seq_lengths_query: A matrix tensor. The type support int32.
|
||||
* Effective sequence length of q in different batches.
|
||||
* @li actual_seq_lengths_key: A matrix tensor. The type support int32.
|
||||
* Effective sequence length of key/value in different batches.
|
||||
|
||||
* @par Attributes:
|
||||
* @li aic_core_num: An int. Cube core num of device.
|
||||
* @li aiv_core_num: An int. Vector core num of device.
|
||||
* @li soc_version: A string. Version of SOC.
|
||||
* @li num_heads_q: An int. Heads num of query.
|
||||
* @li num_heads_k: An int. Heads num of key.
|
||||
* @li head_dim: An Int. Dim of head.
|
||||
* @li query_quant_mode: An int. Mode of query quant.
|
||||
* @li key_quant_mode: An int. Mode of query quant.
|
||||
* @li batch_size: An int. Size of batch.
|
||||
* @li max_seqlen_q: An int. Max sequence length of query.
|
||||
* @li max_seqlen_k: An int. Max sequence length of key.
|
||||
* @li layout_query: A string. Layout of query.
|
||||
* @li layout_key: A string. Layout of key.
|
||||
* @li sparse_count: An int. Sparse count.
|
||||
* @li sparse_mode: An int. Mode of sparse.
|
||||
* @li pre_tokens: An int. Num of pretokens.
|
||||
* @li next_tokens: An int. Num of nexttokens.
|
||||
* @li cmp_ratio: An int. Ratio of compressor.
|
||||
|
||||
* @par Outputs:
|
||||
* @li metadata: A matrix tensor. The type support int32.
|
||||
* The output of attention structure.
|
||||
*/
|
||||
REG_OP(VllmQuantLightningIndexerMetadata)
|
||||
.OPTIONAL_INPUT(actual_seq_lengths_query, TensorType({DT_INT32}))
|
||||
.OPTIONAL_INPUT(actual_seq_lengths_key, TensorType({DT_INT32}))
|
||||
.OUTPUT(metadata, TensorType({DT_INT32}))
|
||||
.REQUIRED_ATTR(aic_core_num, Int)
|
||||
.REQUIRED_ATTR(aiv_core_num, Int)
|
||||
.REQUIRED_ATTR(soc_version, String)
|
||||
.REQUIRED_ATTR(num_heads_q, Int)
|
||||
.REQUIRED_ATTR(num_heads_k, Int)
|
||||
.REQUIRED_ATTR(head_dim, Int)
|
||||
.REQUIRED_ATTR(query_quant_mode, Int)
|
||||
.REQUIRED_ATTR(key_quant_mode, Int)
|
||||
.ATTR(batch_size, Int, 0)
|
||||
.ATTR(max_seqlen_q, Int, 0)
|
||||
.ATTR(max_seqlen_k, Int, 0)
|
||||
.ATTR(layout_query, String, "BSND")
|
||||
.ATTR(layout_key, String, "BSND")
|
||||
.ATTR(sparse_count, Int, 2048)
|
||||
.ATTR(sparse_mode, Int, 3)
|
||||
.ATTR(pre_tokens, Int, 9223372036854775807)
|
||||
.ATTR(next_tokens, Int, 9223372036854775807)
|
||||
.ATTR(cmp_ratio, Int, 1)
|
||||
.OP_END_FACTORY_REG(VllmQuantLightningIndexerMetadata)
|
||||
} // namespace ge
|
||||
|
||||
#endif // QUANT_LIGHTNING_INDEXER_METADATA_PROTO_H
|
||||
|
||||
// FD
|
||||
@@ -0,0 +1,38 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vllm_quant_lightning_indexer_metadata_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "../../vllm_quant_lightning_indexer/op_kernel/vllm_quant_lightning_indexer_metadata.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
static ge::graphStatus InferShapeVllmQuantLightningIndexerMetaData(gert::InferShapeContext* context)
|
||||
{
|
||||
gert::Shape* oShape = context->GetOutputShape(0);
|
||||
oShape->SetDimNum(1);
|
||||
oShape->SetDim(0, optiling::QLI_META_SIZE);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDtypeVllmQuantLightningIndexerMetaData(gert::InferDataTypeContext* context)
|
||||
{
|
||||
context->SetOutputDataType(0, DT_INT32);
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(VllmQuantLightningIndexerMetadata)
|
||||
.InferShape(InferShapeVllmQuantLightningIndexerMetaData)
|
||||
.InferDataType(InferDtypeVllmQuantLightningIndexerMetaData);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,907 @@
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include "../../vllm_quant_lightning_indexer/op_kernel/vllm_quant_lightning_indexer_metadata.h"
|
||||
#include "../../common/aicpu/cpu_context_util.h"
|
||||
#include "vllm_quant_lightning_indexer_metadata_aicpu.h"
|
||||
|
||||
using namespace optiling;
|
||||
|
||||
namespace aicpu {
|
||||
uint32_t
|
||||
VllmQuantLightningIndexerMetadataCpuKernel::Compute(CpuKernelContext &ctx)
|
||||
{
|
||||
bool success = Prepare(ctx);
|
||||
if (!success) {
|
||||
return KERNEL_STATUS_PARAM_INVALID;
|
||||
}
|
||||
SplitResult splitRes {aicCoreNum_, aivCoreNum_};
|
||||
success = BalanceSchedule(splitRes) && GenMetaData(splitRes);
|
||||
return success ? KERNEL_STATUS_OK : KERNEL_STATUS_PARAM_INVALID;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::Prepare(CpuKernelContext &ctx)
|
||||
{
|
||||
// input
|
||||
actSeqLenQ_ = ctx.Input(static_cast<uint32_t>(ParamId::actSeqLenQ));
|
||||
actSeqLenKey_ = ctx.Input(static_cast<uint32_t>(ParamId::actSeqLenKV));
|
||||
// output
|
||||
metaData_ = ctx.Output(static_cast<uint32_t>(ParamId::metaData));
|
||||
|
||||
bool requiredAttrs = GetAttrValue(ctx, "aic_core_num", aicCoreNum_) &&
|
||||
GetAttrValue(ctx, "aiv_core_num", aivCoreNum_) &&
|
||||
GetAttrValue(ctx, "soc_version", socVersion_) &&
|
||||
GetAttrValue(ctx, "num_heads_q", numHeadsQ_) &&
|
||||
GetAttrValue(ctx, "num_heads_k", numHeadsK_) &&
|
||||
GetAttrValue(ctx, "head_dim", headDim_) &&
|
||||
GetAttrValue(ctx, "query_quant_mode", queryQuantMode_) &&
|
||||
GetAttrValue(ctx, "key_quant_mode", keyQuantMode_);
|
||||
if (!requiredAttrs) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// attributes optional
|
||||
GetAttrValueOpt(ctx, "batch_size", batchSize_);
|
||||
GetAttrValueOpt(ctx, "max_seqlen_q", maxSeqlenQ_);
|
||||
GetAttrValueOpt(ctx, "max_seqlen_k", maxSeqlenK_);
|
||||
GetAttrValueOpt(ctx, "layout_query", layoutQuery_);
|
||||
GetAttrValueOpt(ctx, "layout_key", layoutKey_);
|
||||
GetAttrValueOpt(ctx, "sparse_count", sparseCount_);
|
||||
GetAttrValueOpt(ctx, "sparse_mode", sparseMode_);
|
||||
GetAttrValueOpt(ctx, "pre_tokens", preToken_);
|
||||
GetAttrValueOpt(ctx, "next_tokens", nextToken_);
|
||||
GetAttrValueOpt(ctx, "cmp_ratio", cmpRatio_);
|
||||
|
||||
return (ParamsCheck() && ParamsInit());
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::ParamsCheck()
|
||||
{
|
||||
return (CheckSingleParam() && CheckExistence() && CheckConsistency() && CheckFeature());
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::CheckSingleParam()
|
||||
{
|
||||
// 基础输出校验
|
||||
KERNEL_CHECK_NULLPTR(metaData_, false, "metadata is null");
|
||||
auto metaShape = metaData_->GetTensorShape();
|
||||
KERNEL_CHECK_NULLPTR(metaShape, false, "shape of metadata is null");
|
||||
KERNEL_CHECK_NULLPTR(metaData_->GetData(), false, "data of metadata is null");
|
||||
// 核心数校验
|
||||
if (aicCoreNum_ == 0 || aivCoreNum_ == 0 || (aivCoreNum_ % aicCoreNum_ != 0)) {
|
||||
KERNEL_LOG_ERROR("Core num invalid: aic:%u, aiv:%u", aicCoreNum_, aivCoreNum_);
|
||||
return false;
|
||||
}
|
||||
// batch_size 非负校验
|
||||
if (batchSize_ < 0) {
|
||||
KERNEL_LOG_ERROR("batch_size should not be negative, but got %d", batchSize_);
|
||||
return false;
|
||||
}
|
||||
// max_seqlen_q 非负校验
|
||||
if (maxSeqlenQ_ < 0) {
|
||||
KERNEL_LOG_ERROR("max_seqlen_q should not be negative, but got %d", maxSeqlenQ_);
|
||||
return false;
|
||||
}
|
||||
// max_seqlen_k 非负校验
|
||||
if (maxSeqlenK_ < 0) {
|
||||
KERNEL_LOG_ERROR("max_seqlen_k should not be negative, but got %d", maxSeqlenK_);
|
||||
return false;
|
||||
}
|
||||
// num_heads_q 校验
|
||||
if (numHeadsQ_ != 64) {
|
||||
KERNEL_LOG_ERROR("num_heads_q should only be 64, but got %d", numHeadsQ_);
|
||||
return false;
|
||||
}
|
||||
// num_heads_k 校验
|
||||
if (numHeadsK_ != 1) {
|
||||
KERNEL_LOG_ERROR("num_heads_k should only be 1, but got %d", numHeadsK_);
|
||||
return false;
|
||||
}
|
||||
// layout_query 校验
|
||||
if (layoutQuery_ != "TND" && layoutQuery_ != "BSND") {
|
||||
KERNEL_LOG_ERROR("For layout_query, layout must be TND or BSND!");
|
||||
return false;
|
||||
}
|
||||
// layout_key 校验
|
||||
if (layoutKey_ != "PA_BSND" && layoutKey_ != "TND" && layoutKey_ != "BSND") {
|
||||
KERNEL_LOG_ERROR("For layout_key, layout must be PA_BSND/TND/BSND!");
|
||||
return false;
|
||||
}
|
||||
if (layoutQuery_ == "TND" && layoutKey_ == "BSND") {
|
||||
KERNEL_LOG_ERROR("For layout_query TND, layout_key should be PA_BSND/TND!");
|
||||
return false;
|
||||
}
|
||||
if (layoutQuery_ == "BSND" && layoutKey_ == "TND") {
|
||||
KERNEL_LOG_ERROR("For layout_query BSND, layout_key should be PA_BSND/BSND!");
|
||||
return false;
|
||||
}
|
||||
// sparse_mode 校验
|
||||
if (sparseMode_ != static_cast<uint32_t>(SparseMode::DEFAULT_MASK) &&
|
||||
sparseMode_ != static_cast<uint32_t>(SparseMode::RIGHT_DOWN_CAUSAL)) {
|
||||
KERNEL_LOG_ERROR("sparse_mode should be 0/3, but got %d", sparseMode_);
|
||||
return false;
|
||||
}
|
||||
// pre_tokens 校验
|
||||
if (preToken_ != INT64_MAX) {
|
||||
KERNEL_LOG_ERROR("pre_tokens should only be 2^63-1, but got %ld", preToken_);
|
||||
return false;
|
||||
}
|
||||
// next_tokens 校验
|
||||
if (nextToken_ != INT64_MAX) {
|
||||
KERNEL_LOG_ERROR("next_tokens should only be 2^63-1, but got %ld", nextToken_);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::CheckExistence()
|
||||
{
|
||||
auto isInvalid = [](Tensor* t) { return t == nullptr || t->GetData() == nullptr; };
|
||||
// Query 存在性逻辑
|
||||
if (layoutQuery_ == "TND") {
|
||||
if (isInvalid(actSeqLenQ_)) {
|
||||
KERNEL_LOG_ERROR("For layout_query TND, actual_seq_lengths_query must be provided!");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
// KV 存在性逻辑
|
||||
if (layoutKey_ == "PA_BSND") {
|
||||
if (isInvalid(actSeqLenKey_)) {
|
||||
KERNEL_LOG_ERROR("For layout_key PA_BSND, actual_seq_lengths_key must be provided!");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
if (layoutKey_ == "TND") {
|
||||
if (isInvalid(actSeqLenKey_)) {
|
||||
KERNEL_LOG_ERROR("For layout_key TND, actual_seq_lengths_key must be provided!");
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
int32_t VllmQuantLightningIndexerMetadataCpuKernel::GetQueryBatchSize()
|
||||
{
|
||||
if (actSeqLenQ_ != nullptr && actSeqLenQ_->GetData() != nullptr) {
|
||||
if (actSeqLenQ_->GetTensorShape() != nullptr) {
|
||||
return actSeqLenQ_->GetTensorShape()->GetDimSize(0);
|
||||
}
|
||||
}
|
||||
return batchSize_;
|
||||
}
|
||||
|
||||
int32_t VllmQuantLightningIndexerMetadataCpuKernel::GetKvBatchSize()
|
||||
{
|
||||
if (actSeqLenKey_ != nullptr && actSeqLenKey_->GetData() != nullptr) {
|
||||
if (actSeqLenKey_->GetTensorShape() != nullptr) {
|
||||
return actSeqLenKey_->GetTensorShape()->GetDimSize(0);
|
||||
}
|
||||
}
|
||||
return batchSize_;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::CheckConsistency()
|
||||
{
|
||||
int32_t queryBatchSize = GetQueryBatchSize();
|
||||
int32_t kvBatchSize = GetKvBatchSize();
|
||||
if ((layoutQuery_ == "BSND" || (layoutQuery_ == "TND" && layoutKey_ == "TND")) && queryBatchSize != kvBatchSize) {
|
||||
KERNEL_LOG_ERROR("For the layout_query is BSND or both layout_query and layout_key are TND, the dim of actual_seq_lengths_query and the dim of actual_seq_lengths_key should be equal.");
|
||||
return false;
|
||||
}
|
||||
if (std::abs(queryBatchSize - kvBatchSize) > 1) {
|
||||
KERNEL_LOG_ERROR("The difference between the dim of actual_seq_lengths_query and the dim of actual_seq_lengths_key should not be greater than 1.");
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::CheckFeature()
|
||||
{
|
||||
// 压缩率校验
|
||||
if (cmpRatio_ < 1 || cmpRatio_ > 128) {
|
||||
KERNEL_LOG_ERROR("cmp_ratio should be [1, 128], but got %d", cmpRatio_);
|
||||
return false;
|
||||
}
|
||||
validSocVersion_ = ProcessSocVersion();
|
||||
if (validSocVersion_ == ValidSocVersion::ASCEND910B) {
|
||||
// 校验 2 的幂次方: 1, 2, 4, ..., 128
|
||||
if ((cmpRatio_ & (cmpRatio_ - 1)) != 0) {
|
||||
KERNEL_LOG_ERROR("For Atlas A3, cmp_ratio should be 1/2/4/8/16/32/64/128, but got %d", cmpRatio_);
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
if (cmpRatio_ != 1 && cmpRatio_ != 4 && cmpRatio_ != 128) {
|
||||
KERNEL_LOG_ERROR("For Ascend950, cmp_ratio should be 1/4/128, but got %d", cmpRatio_);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
ValidSocVersion VllmQuantLightningIndexerMetadataCpuKernel::ProcessSocVersion()
|
||||
{
|
||||
const std::string ascend950 = "Ascend950";
|
||||
if (socVersion_.find(ascend950) != std::string::npos) {
|
||||
return ValidSocVersion::ASCEND950;
|
||||
} else {
|
||||
return ValidSocVersion::ASCEND910B;
|
||||
}
|
||||
|
||||
return ValidSocVersion::RESERVED_VERSION;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::ParamsInit()
|
||||
{
|
||||
int32_t qBatchSize_ = GetQueryBatchSize();
|
||||
int32_t kvBatchSize_ = GetKvBatchSize();
|
||||
isActQBatchPlus = (qBatchSize_ > kvBatchSize_);
|
||||
auto mode = static_cast<SparseMode>(sparseMode_);
|
||||
if (mode == SparseMode::RIGHT_DOWN_CAUSAL) {
|
||||
attentionMode_ = 1;
|
||||
preToken_ = INT64_MAX;
|
||||
} else if (mode == SparseMode::DEFAULT_MASK) {
|
||||
attentionMode_ = 0;
|
||||
} else if (mode == SparseMode::BAND) {
|
||||
attentionMode_ = 1;
|
||||
}
|
||||
groupSize_ = numHeadsQ_ / numHeadsK_;
|
||||
batchSize_ = std::min(qBatchSize_, kvBatchSize_);
|
||||
validSocVersion_ = ProcessSocVersion();
|
||||
if (validSocVersion_ == ValidSocVersion::ASCEND910B){
|
||||
s2BaseSize_ = 2048U; // 仅用于A3
|
||||
} else if (validSocVersion_ == ValidSocVersion::ASCEND950){
|
||||
s2BaseSize_ = 128U; // 仅用于A5
|
||||
} else {
|
||||
s2BaseSize_ = 128U; // 其他情况
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
uint32_t VllmQuantLightningIndexerMetadataCpuKernel::GetS1SeqSize(uint32_t bIdx)
|
||||
{
|
||||
if (actSeqLenQ_ == nullptr || actSeqLenQ_->GetData() == nullptr) {
|
||||
return static_cast<uint32_t>(maxSeqlenQ_);
|
||||
}
|
||||
const int32_t *s1Ptr = (int32_t*)actSeqLenQ_->GetData();
|
||||
if (layoutQuery_ == "TND") {
|
||||
if (isActQBatchPlus) {
|
||||
return static_cast<uint32_t>(s1Ptr[bIdx + 1U] - s1Ptr[bIdx]);
|
||||
}
|
||||
return (bIdx == 0) ? static_cast<uint32_t>(s1Ptr[bIdx]) :
|
||||
static_cast<uint32_t>(s1Ptr[bIdx] - s1Ptr[bIdx - 1U]);
|
||||
} else {
|
||||
return static_cast<uint32_t>(s1Ptr[bIdx]);
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t VllmQuantLightningIndexerMetadataCpuKernel::GetS2SeqSize(uint32_t bIdx)
|
||||
{
|
||||
if (actSeqLenKey_ == nullptr || actSeqLenKey_->GetData() == nullptr) {
|
||||
return static_cast<uint32_t>(maxSeqlenK_ * cmpRatio_);
|
||||
}
|
||||
const int32_t *s2Ptr = (int32_t*)actSeqLenKey_->GetData();
|
||||
if (layoutKey_ == "TND") {
|
||||
return (bIdx == 0) ? static_cast<uint32_t>(s2Ptr[bIdx]) :
|
||||
static_cast<uint32_t>(s2Ptr[bIdx] - s2Ptr[bIdx - 1U]);
|
||||
} else {
|
||||
return static_cast<uint32_t>(s2Ptr[bIdx]);
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcSplitInfo(SplitContext &splitContext)
|
||||
{
|
||||
// 计算每个batch的切分,统计是否为空batch,记录最后有效batch(每个batch的每个N2切分是一样的)
|
||||
SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) {
|
||||
uint32_t s1Size = GetS1SeqSize(bIdx);
|
||||
uint32_t s2Size = GetS2SeqSize(bIdx) / cmpRatio_;
|
||||
splitInfo.s1GBaseNum[bIdx] = (s1Size * groupSize_ + (mBaseSize_ - 1U)) / mBaseSize_;
|
||||
splitInfo.s1GTailSize[bIdx] = (s1Size * groupSize_) % mBaseSize_;
|
||||
splitInfo.s2BaseNum[bIdx] = (s2Size + s2BaseSize_ - 1U) / s2BaseSize_;
|
||||
splitInfo.s2TailSize[bIdx] = s2Size % s2BaseSize_;
|
||||
if (splitInfo.s1GBaseNum[bIdx] != 0U && splitInfo.s2BaseNum[bIdx] != 0U) {
|
||||
splitInfo.isKvSeqAllZero = false;
|
||||
}
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
int64_t VllmQuantLightningIndexerMetadataCpuKernel::CalcPreTokenLeftUp(
|
||||
uint32_t s1Size, uint32_t s2Size)
|
||||
{
|
||||
auto mode = static_cast<SparseMode>(sparseMode_);
|
||||
if (mode == SparseMode::BAND) {
|
||||
return static_cast<int64_t>(s1Size) - static_cast<int64_t>(s2Size) + preToken_;
|
||||
}
|
||||
return preToken_;
|
||||
}
|
||||
|
||||
int64_t VllmQuantLightningIndexerMetadataCpuKernel::CalcNextTokenLeftUp(
|
||||
uint32_t s1Size, uint32_t s2Size)
|
||||
{
|
||||
auto mode = static_cast<SparseMode>(sparseMode_);
|
||||
switch (mode) {
|
||||
case SparseMode::DEFAULT_MASK:
|
||||
case SparseMode::ALL_MASK:
|
||||
case SparseMode::LEFT_UP_CAUSAL:
|
||||
return nextToken_;
|
||||
case SparseMode::RIGHT_DOWN_CAUSAL:
|
||||
return static_cast<int64_t>(s2Size) - static_cast<int64_t>(s1Size);
|
||||
case SparseMode::BAND:
|
||||
return static_cast<int64_t>(s2Size) - static_cast<int64_t>(s1Size) + nextToken_;
|
||||
default:
|
||||
return nextToken_;
|
||||
}
|
||||
}
|
||||
|
||||
int64_t VllmQuantLightningIndexerMetadataCpuKernel::CalcCost(
|
||||
uint32_t basicM, uint32_t basicS2)
|
||||
{
|
||||
uint32_t alignBasicM = basicM / groupSize_;
|
||||
uint32_t alignBasicS2 = 1;
|
||||
if (validSocVersion_ == ValidSocVersion::ASCEND910B) {
|
||||
if (basicS2 <= 1024) {
|
||||
alignBasicS2 = (1536 + basicS2) / 2;
|
||||
} else {
|
||||
alignBasicS2 = (2560 + basicS2) / 2;
|
||||
}
|
||||
return static_cast<int64_t>(alignBasicM * alignBasicS2);
|
||||
} else {
|
||||
uint32_t alignCoefM = 16U;
|
||||
uint32_t alignCoefS2 = 64U;
|
||||
alignBasicM = (basicM + alignCoefM - 1U) >> 4U; // 按alignCoefM对齐,向上取整,4:移位操作实现除16
|
||||
alignBasicS2 = (basicS2 + alignCoefS2 - 1U) >> 6U; // 按alignCoefS2对齐,向上取整,6:移位操作实现除64
|
||||
return static_cast<int64_t>(6U * alignBasicM + 10U * alignBasicS2);
|
||||
}
|
||||
}
|
||||
|
||||
BlockCost<int64_t> VllmQuantLightningIndexerMetadataCpuKernel::CalcCostTable(uint32_t s1NormalSize,
|
||||
uint32_t s2NormalSize, uint32_t s1GTailSize, uint32_t s2TailSize)
|
||||
{
|
||||
BlockCost<int64_t> typeCost {};
|
||||
typeCost[NORMAL_BLOCK][NORMAL_BLOCK] = CalcCost(s1NormalSize, s2NormalSize);
|
||||
typeCost[TAIL_BLOCK][NORMAL_BLOCK] = (s1GTailSize == 0U) ? 0U : CalcCost(s1GTailSize, s2NormalSize);
|
||||
typeCost[NORMAL_BLOCK][TAIL_BLOCK] = (s2TailSize == 0U) ? 0U : CalcCost(s1NormalSize, s2TailSize);
|
||||
typeCost[TAIL_BLOCK][TAIL_BLOCK] = (s1GTailSize == 0U || s2TailSize == 0U) ? 0U : CalcCost(s1GTailSize, s2TailSize);
|
||||
return typeCost;
|
||||
}
|
||||
|
||||
Range<int64_t> VllmQuantLightningIndexerMetadataCpuKernel::CalcS2TokenRange(
|
||||
uint32_t s1GIdx, const BatchCache &batchCache) {
|
||||
|
||||
// no mask
|
||||
if (!attentionMode_) { //attentionMaskFlag ?
|
||||
return std::make_pair(0, static_cast<int64_t>(batchCache.s2Size));
|
||||
}
|
||||
// 1. calc index of s2FirstToken, s2LastToken by index of s1GFirstToken, s1GLastToken
|
||||
int64_t s1GFirstToken = static_cast<int64_t>(s1GIdx) * static_cast<int64_t>(mBaseSize_);
|
||||
int64_t s1GLastToken = std::min(s1GFirstToken + static_cast<int64_t>(mBaseSize_),
|
||||
static_cast<int64_t>(batchCache.s1Size) * static_cast<int64_t>(groupSize_)) - 1;
|
||||
int64_t s1FirstToken = 0;
|
||||
int64_t s1LastToken = 0;
|
||||
if (isS1G_) {
|
||||
s1FirstToken = s1GFirstToken / static_cast<int64_t>(groupSize_);
|
||||
s1LastToken = s1GLastToken / static_cast<int64_t>(groupSize_);
|
||||
} else {
|
||||
if (s1GFirstToken / batchCache.s1Size == s1GLastToken / batchCache.s1Size) {
|
||||
// start and end locate in one G
|
||||
s1FirstToken = s1GFirstToken % static_cast<int64_t>(batchCache.s1Size);
|
||||
s1LastToken = s1GLastToken % static_cast<int64_t>(batchCache.s1Size);
|
||||
} else {
|
||||
// start and end locate in tow or more G, but working same as crossing a complete block
|
||||
s1FirstToken = 0;
|
||||
s1LastToken = batchCache.s1Size;
|
||||
}
|
||||
}
|
||||
|
||||
int64_t s2FirstToken = s1FirstToken - batchCache.preTokenLeftUp;
|
||||
int64_t s2LastToken = s1LastToken + batchCache.nextTokenLeftUp;
|
||||
|
||||
return std::make_pair(s2FirstToken, s2LastToken);
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcBatchCache(
|
||||
uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
|
||||
batchCache.bIdx = bIdx;
|
||||
batchCache.s1Size = GetS1SeqSize(bIdx);
|
||||
batchCache.s2Size = GetS2SeqSize(bIdx);
|
||||
batchCache.preTokenLeftUp = CalcPreTokenLeftUp(batchCache.s1Size, batchCache.s2Size);
|
||||
batchCache.nextTokenLeftUp = CalcNextTokenLeftUp(batchCache.s1Size, batchCache.s2Size);
|
||||
batchCache.typeCost = CalcCostTable(mBaseSize_, s2BaseSize_, splitInfo.s1GTailSize[bIdx],
|
||||
splitInfo.s2TailSize[bIdx]);
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcS1GCache(uint32_t s1GIdx,
|
||||
const SplitContext &splitContext, const BatchCache &batchCache, S1GCache &s1GCache)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
|
||||
s1GCache.bIdx = batchCache.bIdx;
|
||||
s1GCache.s1GIdx = s1GIdx;
|
||||
|
||||
if (splitInfo.s1GBaseNum[batchCache.bIdx] == 0 || splitInfo.s2BaseNum[batchCache.bIdx] == 0) {
|
||||
s1GCache.s1GBlock = 0;
|
||||
s1GCache.s1GCost = 0;
|
||||
s1GCache.s1GLastBlockCost = 0;
|
||||
s1GCache.s1GNormalBlockCost = 0;
|
||||
return;
|
||||
}
|
||||
|
||||
auto s2TokenRange = CalcS2TokenRange(s1GIdx, batchCache);
|
||||
int64_t s2FirstToken = s2TokenRange.first;
|
||||
int64_t s2LastToken = s2TokenRange.second;
|
||||
// get valid range
|
||||
s2FirstToken = Clip(s2FirstToken, static_cast<int64_t>(0), static_cast<int64_t>(batchCache.s2Size - 1U));
|
||||
s2LastToken = Clip(s2LastToken, static_cast<int64_t>(0), static_cast<int64_t>(batchCache.s2Size - 1U));
|
||||
|
||||
// get block start & end
|
||||
uint32_t s2CmpLength = (s2LastToken - s2FirstToken + 1) / cmpRatio_;
|
||||
if (s2CmpLength == 0) {
|
||||
s1GCache.s2Start = 0U;
|
||||
s1GCache.s2End = 0U;
|
||||
} else {
|
||||
s1GCache.s2Start = 0U;
|
||||
s1GCache.s2End = (s2CmpLength + s2BaseSize_ - 1) / s2BaseSize_; // end of block index, Right-open interval
|
||||
}
|
||||
|
||||
if (s1GCache.s2Start >= s1GCache.s2End) {
|
||||
s1GCache.s1GBlock = 0;
|
||||
s1GCache.s1GCost = 0;
|
||||
s1GCache.s1GLastBlockCost = 0;
|
||||
s1GCache.s1GNormalBlockCost = 0;
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t s2TailSize = s2CmpLength % s2BaseSize_;
|
||||
|
||||
// 计算S2方向满块、尾块数量
|
||||
s1GCache.s1GBlock = s1GCache.s2End - s1GCache.s2Start;
|
||||
uint32_t curTailS2Num = s2TailSize != 0 ? 1U : 0U;
|
||||
uint32_t curNormalS2Num = s1GCache.s1GBlock - curTailS2Num;
|
||||
|
||||
BlockCost<int64_t> typeCost = CalcCostTable(mBaseSize_, s2BaseSize_,
|
||||
splitInfo.s1GTailSize[batchCache.bIdx], s2TailSize);
|
||||
|
||||
if (s1GIdx == (splitInfo.s1GBaseNum[batchCache.bIdx] - 1U) &&
|
||||
splitInfo.s1GTailSize[batchCache.bIdx] != 0U) {
|
||||
s1GCache.s1GCost = typeCost[TAIL_BLOCK][NORMAL_BLOCK] * curNormalS2Num +
|
||||
typeCost[TAIL_BLOCK][TAIL_BLOCK] * curTailS2Num;
|
||||
s1GCache.s1GLastBlockCost = curTailS2Num > 0U ? typeCost[TAIL_BLOCK][TAIL_BLOCK] :
|
||||
typeCost[TAIL_BLOCK][NORMAL_BLOCK];
|
||||
s1GCache.s1GNormalBlockCost = typeCost[TAIL_BLOCK][NORMAL_BLOCK];
|
||||
} else {
|
||||
s1GCache.s1GCost = typeCost[NORMAL_BLOCK][NORMAL_BLOCK] * curNormalS2Num +
|
||||
typeCost[NORMAL_BLOCK][TAIL_BLOCK] * curTailS2Num;
|
||||
s1GCache.s1GLastBlockCost = curTailS2Num > 0U ? typeCost[NORMAL_BLOCK][TAIL_BLOCK] :
|
||||
typeCost[NORMAL_BLOCK][NORMAL_BLOCK];
|
||||
s1GCache.s1GNormalBlockCost = typeCost[NORMAL_BLOCK][NORMAL_BLOCK];
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcBatchCost(
|
||||
uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
|
||||
costInfo.bN2CostOfEachBatch[bIdx] = 0;
|
||||
costInfo.bN2BlockOfEachBatch[bIdx] = 0U;
|
||||
costInfo.bN2LastBlockCostOfEachBatch[bIdx] = 0U;
|
||||
|
||||
if (GetS1SeqSize(bIdx) == 0U || GetS2SeqSize(bIdx) == 0U) {
|
||||
return;
|
||||
}
|
||||
|
||||
BatchCache bCache;
|
||||
S1GCache s1GCache;
|
||||
CalcBatchCache(bIdx, splitContext, bCache);
|
||||
for (uint32_t s1GIdx = 0; s1GIdx < splitInfo.s1GBaseNum[bIdx]; s1GIdx++) {
|
||||
CalcS1GCache(s1GIdx, splitContext, bCache, s1GCache);
|
||||
costInfo.bN2CostOfEachBatch[bIdx] += s1GCache.s1GCost;
|
||||
costInfo.bN2BlockOfEachBatch[bIdx] += s1GCache.s1GBlock;
|
||||
// 更新最大S1G行开销
|
||||
if (s1GCache.s1GCost > costInfo.maxS1GCost) {
|
||||
costInfo.maxS1GCost = s1GCache.s1GCost;
|
||||
}
|
||||
|
||||
if(s1GCache.s1GBlock > 0){
|
||||
costInfo.bN2LastBlockCostOfEachBatch[bIdx] = s1GCache.s1GLastBlockCost;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcCostInfo(SplitContext &splitContext)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
CostInfo &costInfo = splitContext.costInfo;
|
||||
|
||||
if (splitInfo.isKvSeqAllZero) {
|
||||
costInfo.totalCost = 0;
|
||||
costInfo.totalBlockNum = 0U;
|
||||
return;
|
||||
}
|
||||
|
||||
// 计算batch的负载并记录,用于按batch分配,需要按行计算起止点,统计块数、负载
|
||||
for (uint32_t bIdx = 0; bIdx < batchSize_; bIdx++) {
|
||||
CalcBatchCost(bIdx, splitContext, costInfo);
|
||||
costInfo.totalCost += costInfo.bN2CostOfEachBatch[bIdx] * numHeadsK_;
|
||||
costInfo.totalBlockNum += costInfo.bN2BlockOfEachBatch[bIdx] * numHeadsK_;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
const CostInfo &costInfo = splitContext.costInfo;
|
||||
|
||||
bool UpdateS1G = false;
|
||||
bool UpdateBatch = false;
|
||||
|
||||
// Update S2
|
||||
if (assignContext.curS2Idx >= assignContext.s1GCache.s2End) { // 边界assignInfo.s2End是取不到的开区间
|
||||
assignContext.curS2Idx = 0U;
|
||||
assignContext.curS1GIdx++;
|
||||
UpdateS1G = true;
|
||||
}
|
||||
|
||||
// Update S1G
|
||||
if (assignContext.curS1GIdx >= splitInfo.s1GBaseNum[assignContext.curBIdx]) {
|
||||
assignContext.curS1GIdx = 0U;
|
||||
assignContext.curBN2Idx++;
|
||||
}
|
||||
|
||||
// Update Batch
|
||||
if (assignContext.curBN2Idx == batchSize_ * numHeadsK_) { // 所有负载全部分配完,设置最后一个核的右开区间,返回
|
||||
assignContext.curS1GIdx = 0U;
|
||||
assignContext.curS2Idx = 0U;
|
||||
assignContext.isFinished = true;
|
||||
return;
|
||||
}
|
||||
|
||||
if (assignContext.curBN2Idx / numHeadsK_ != assignContext.curBIdx) {
|
||||
assignContext.curBIdx = assignContext.curBN2Idx / numHeadsK_;
|
||||
assignContext.curS1GIdx = 0U;
|
||||
UpdateBatch = true;
|
||||
UpdateS1G = true;
|
||||
}
|
||||
|
||||
// Update Cache
|
||||
if (UpdateBatch) {
|
||||
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
|
||||
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
|
||||
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
|
||||
}
|
||||
if (UpdateS1G) {
|
||||
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
|
||||
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext)
|
||||
{
|
||||
if (assignContext.isFinished) {
|
||||
return;
|
||||
}
|
||||
const CostInfo &costInfo = splitContext.costInfo;
|
||||
while (assignContext.bN2Cost == 0 || IsWithinTolerance(assignContext.coreCache.costLimit,
|
||||
costInfo.bN2LastBlockCostOfEachBatch[assignContext.curBIdx] / FA_TOLERANCE_RATIO,
|
||||
assignContext.coreCache.cost + assignContext.bN2Cost)) {
|
||||
assignContext.coreCache.cost += assignContext.bN2Cost;
|
||||
assignContext.coreCache.block += assignContext.bN2Block;
|
||||
assignContext.curBN2Idx++;
|
||||
|
||||
// to the end
|
||||
if (assignContext.curBN2Idx == batchSize_ * numHeadsK_) {
|
||||
assignContext.curS1GIdx = 0U;
|
||||
assignContext.curS2Idx = 0U;
|
||||
assignContext.isFinished = true;
|
||||
return;
|
||||
}
|
||||
|
||||
// next batch
|
||||
if (assignContext.curBN2Idx / numHeadsK_ != assignContext.curBIdx) {
|
||||
assignContext.curBIdx = assignContext.curBN2Idx / numHeadsK_;
|
||||
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
|
||||
}
|
||||
|
||||
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
|
||||
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
|
||||
assignContext.curS1GIdx = 0U;
|
||||
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
|
||||
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::AssignByRow(const SplitContext &splitContext, AssignContext &assignContext)
|
||||
{
|
||||
if (assignContext.isFinished) {
|
||||
return;
|
||||
}
|
||||
|
||||
while (IsWithinTolerance(assignContext.coreCache.costLimit,
|
||||
assignContext.s1GCache.s1GLastBlockCost / FA_TOLERANCE_RATIO,
|
||||
assignContext.coreCache.cost + assignContext.s1GCache.s1GCost)) {
|
||||
assignContext.coreCache.cost += assignContext.s1GCache.s1GCost;
|
||||
assignContext.coreCache.block += assignContext.s1GCache.s1GBlock;
|
||||
|
||||
// 当前batch被分配一行出去,更新剩余负载
|
||||
assignContext.bN2Cost = assignContext.bN2Cost > assignContext.s1GCache.s1GCost ?
|
||||
assignContext.bN2Cost - assignContext.s1GCache.s1GCost : 0;
|
||||
assignContext.bN2Block = assignContext.bN2Block > assignContext.s1GCache.s1GBlock ?
|
||||
assignContext.bN2Block - assignContext.s1GCache.s1GBlock : 0U;
|
||||
// 计算新一行的信息
|
||||
do{
|
||||
assignContext.curS1GIdx++;
|
||||
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
|
||||
}while(assignContext.s1GCache.s1GBlock == 0);
|
||||
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::AssignByBlock(const SplitContext &splitContext, AssignContext &assignContext)
|
||||
{
|
||||
if (assignContext.isFinished) {
|
||||
return;
|
||||
}
|
||||
|
||||
int64_t curCost = assignContext.s1GCache.s1GNormalBlockCost;
|
||||
if (assignContext.curS2Idx == (assignContext.s1GCache.s2End - 1U)) {
|
||||
curCost = assignContext.s1GCache.s1GLastBlockCost;
|
||||
}
|
||||
|
||||
while (IsWithinTolerance(assignContext.coreCache.costLimit, curCost / FA_TOLERANCE_RATIO,
|
||||
assignContext.coreCache.cost + curCost)) { // (costLimit - curCostOnCore) * FA_TOLERANCE_RATIO > curCost;至少分配1块
|
||||
assignContext.coreCache.cost += curCost;
|
||||
assignContext.coreCache.block++;
|
||||
assignContext.curS2Idx++;
|
||||
// 当前batch被分配一块出去,更新剩余负载
|
||||
assignContext.bN2Cost = assignContext.bN2Cost - curCost;
|
||||
// 当前行被分配一块出去,更新剩余负载
|
||||
assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost;
|
||||
assignContext.bN2Block--;
|
||||
assignContext.s1GCache.s1GBlock--;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::ForceAssign(const SplitContext &splitContext, AssignContext &assignContext)
|
||||
{
|
||||
if (assignContext.isFinished) {
|
||||
return;
|
||||
}
|
||||
|
||||
int64_t curCost = assignContext.s1GCache.s1GNormalBlockCost;
|
||||
if (assignContext.curS2Idx == (assignContext.s1GCache.s2End - 1U)) {
|
||||
curCost = assignContext.s1GCache.s1GLastBlockCost;
|
||||
}
|
||||
|
||||
assignContext.coreCache.cost += curCost;
|
||||
assignContext.coreCache.block++;
|
||||
assignContext.curS2Idx++;
|
||||
// 当前batch被分配一块出去,更新剩余负载
|
||||
assignContext.bN2Cost = assignContext.bN2Cost - curCost;
|
||||
assignContext.bN2Block--;
|
||||
// 当前行被分配一块出去,更新剩余负载
|
||||
assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost;
|
||||
assignContext.s1GCache.s1GBlock--;
|
||||
UpdateCursor(splitContext, assignContext);
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::IsNeedRecordFDInfo(const AssignContext &assignContext, const SplitResult &splitRes)
|
||||
{
|
||||
// 切分点大概率不会刚好在行尾,因此滞后处理归约信息的统计,到下一个切分点再判断是否需要归约
|
||||
// 核0无需处理
|
||||
if (assignContext.curCoreIdx == 0U) {
|
||||
return false;
|
||||
}
|
||||
// 无跨核行,无需处理
|
||||
if (assignContext.curKvSplitPart <= 1U) {
|
||||
return false;
|
||||
}
|
||||
// 需要归约的行还未处理完
|
||||
if (assignContext.curBN2Idx == splitRes.bN2End[assignContext.curCoreIdx - 1U] &&
|
||||
assignContext.curS1GIdx == splitRes.gS1End[assignContext.curCoreIdx - 1U]) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, SplitResult &result)
|
||||
{
|
||||
const SplitInfo &splitInfo = splitContext.splitInfo;
|
||||
// 需要规约的行是上一个核的切分点所在位置
|
||||
uint32_t splitBIdx = result.bN2End[assignContext.curCoreIdx - 1U] / numHeadsK_;
|
||||
uint32_t splitS1GIdx = result.gS1End[assignContext.curCoreIdx - 1U];
|
||||
uint32_t s1Size = GetS1SeqSize(splitBIdx);
|
||||
|
||||
// 计算归约数据的FD均衡划分信息
|
||||
uint32_t curFdS1gSize = (splitS1GIdx == splitInfo.s1GBaseNum[splitBIdx] - 1U) ?
|
||||
(s1Size * groupSize_ - splitS1GIdx * mBaseSize_) : mBaseSize_;
|
||||
// 记录
|
||||
result.maxS2SplitNum = std::max(result.maxS2SplitNum, assignContext.curKvSplitPart);
|
||||
// 若存在头归约,则切分点一定为上一个核结束的位置
|
||||
result.fdRes.fdBN2Idx[result.numOfFdHead] = result.bN2End[assignContext.curCoreIdx - 1U];
|
||||
result.fdRes.fdMIdx[result.numOfFdHead] = result.gS1End[assignContext.curCoreIdx - 1U];
|
||||
result.fdRes.fdWorkspaceIdx[result.numOfFdHead] = assignContext.preFdDataNum;
|
||||
result.fdRes.fdS2SplitNum[result.numOfFdHead] = assignContext.curKvSplitPart;
|
||||
result.fdRes.fdMSize[result.numOfFdHead] = curFdS1gSize / groupSize_;
|
||||
result.numOfFdHead++;
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::AssignBlockToCore(uint32_t coreNum, const SplitContext &splitContext,
|
||||
AssignContext &assignContext, SplitResult &result)
|
||||
{
|
||||
const CostInfo &costInfo = splitContext.costInfo;
|
||||
result.firstFdDataWorkspaceIdx[assignContext.curCoreIdx] = assignContext.preFdDataNum + assignContext.curKvSplitPart - 1U;
|
||||
assignContext.coreCache = {};
|
||||
assignContext.coreCache.costLimit = assignContext.unassignedCost / (coreNum - assignContext.curCoreIdx);
|
||||
if (!supportFd_){
|
||||
assignContext.coreCache.costLimit = costInfo.maxS1GCost > assignContext.coreCache.costLimit ? costInfo.maxS1GCost :assignContext.coreCache.costLimit;
|
||||
}
|
||||
// 1、按整batch分配
|
||||
AssignByBatch(splitContext, assignContext);
|
||||
// 2、按行分配
|
||||
AssignByRow(splitContext, assignContext);
|
||||
// 3、按块分配
|
||||
if (supportFd_){
|
||||
AssignByBlock(splitContext, assignContext);
|
||||
// 4、强制分配
|
||||
if (assignContext.coreCache.block == 0) {
|
||||
ForceAssign(splitContext, assignContext);
|
||||
}
|
||||
}
|
||||
result.bN2End[assignContext.curCoreIdx] = assignContext.curBN2Idx;
|
||||
result.gS1End[assignContext.curCoreIdx] = assignContext.curS1GIdx;
|
||||
result.s2End[assignContext.curCoreIdx] = assignContext.curS2Idx;
|
||||
result.maxCost = std::max(result.maxCost, assignContext.coreCache.cost);
|
||||
assignContext.unassignedCost -= assignContext.coreCache.cost;
|
||||
// 对之前的归约信息进行记录并清理
|
||||
if (supportFd_ && IsNeedRecordFDInfo(assignContext, result)) {
|
||||
RecordFDInfo(splitContext, assignContext, result);
|
||||
assignContext.preFdDataNum += assignContext.curKvSplitPart;
|
||||
assignContext.curKvSplitPart = 1U;
|
||||
}
|
||||
// 更新S2切分信息
|
||||
if (supportFd_ && assignContext.curS2Idx > assignContext.s1GCache.s2Start &&
|
||||
assignContext.curS2Idx <= assignContext.s1GCache.s2End) {
|
||||
assignContext.curKvSplitPart++;
|
||||
}
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::CalcSplitPlan(uint32_t coreNum,
|
||||
int64_t costLimit, const SplitContext &splitContext, SplitResult &result)
|
||||
{
|
||||
const CostInfo &costInfo = splitContext.costInfo;
|
||||
if (coreNum == 0U) {
|
||||
return;
|
||||
}
|
||||
result.maxCost = 0U;
|
||||
result.usedCoreNum = 0U;
|
||||
AssignContext assignContext {};
|
||||
assignContext.curBIdx = 0U;
|
||||
assignContext.curS1GIdx = 0U;
|
||||
assignContext.unassignedCost = costInfo.totalCost;
|
||||
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
|
||||
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
|
||||
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
|
||||
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
|
||||
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
|
||||
for (uint32_t i = 0; i < coreNum; ++i) {
|
||||
if (result.maxCost > costLimit) {
|
||||
return;
|
||||
}
|
||||
if (assignContext.isFinished || assignContext.unassignedCost <= 0) {
|
||||
break;
|
||||
}
|
||||
assignContext.curCoreIdx = i;
|
||||
AssignBlockToCore(coreNum, splitContext, assignContext, result);
|
||||
}
|
||||
result.usedCoreNum = assignContext.curCoreIdx + 1;
|
||||
}
|
||||
|
||||
void VllmQuantLightningIndexerMetadataCpuKernel::SplitFD(SplitResult &splitRes)
|
||||
{
|
||||
// 计算FD的总数据量
|
||||
uint64_t totalFDLoad = 0;
|
||||
for (uint32_t i = 0; i < splitRes.numOfFdHead; i++) {
|
||||
totalFDLoad += splitRes.fdRes.fdS2SplitNum[i] * splitRes.fdRes.fdMSize[i];
|
||||
}
|
||||
// 计算每个核处理的load
|
||||
uint64_t averageLoad = (totalFDLoad + aivCoreNum_ - 1U) / aivCoreNum_; //向上取整,避免核负载为0
|
||||
uint32_t curCoreIndex = 0;
|
||||
for (uint32_t i = 0; i < splitRes.numOfFdHead; i++) {
|
||||
uint32_t curFDVectorNum = splitRes.fdRes.fdS2SplitNum[i] * splitRes.fdRes.fdMSize[i] / averageLoad; // 计算当前归约任务所用核数,向下取整,避免使用核数超出总核数
|
||||
uint32_t curAveMSize = (splitRes.fdRes.fdMSize[i] + curFDVectorNum - 1U) / curFDVectorNum; // 计算当前归约任务每个核的行数,向上取整,避免行数为0
|
||||
curFDVectorNum = (splitRes.fdRes.fdMSize[i] + curAveMSize -1U)/ curAveMSize;
|
||||
for (uint32_t vid = 0; vid < curFDVectorNum; vid++) {
|
||||
splitRes.fdRes.fdIdx[curCoreIndex] = i;
|
||||
splitRes.fdRes.fdMStart[curCoreIndex] = vid * curAveMSize;
|
||||
splitRes.fdRes.fdMNum[curCoreIndex] =
|
||||
(vid < curFDVectorNum - 1) ? curAveMSize : (splitRes.fdRes.fdMSize[i] - vid * curAveMSize);
|
||||
curCoreIndex++;
|
||||
}
|
||||
}
|
||||
splitRes.fdRes.fdUsedVecNum = curCoreIndex;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::BalanceSchedule(SplitResult &splitRes)
|
||||
{
|
||||
SplitContext splitContext(batchSize_);
|
||||
// 1、划分基本块,统计信息
|
||||
CalcSplitInfo(splitContext);
|
||||
// 全空case
|
||||
if (splitContext.splitInfo.isKvSeqAllZero) {
|
||||
splitRes.usedCoreNum = 1U;
|
||||
splitRes.bN2End[0] = batchSize_ * numHeadsK_;
|
||||
splitRes.gS1End[0] = 0U;
|
||||
splitRes.s2End[0] = 0U;
|
||||
return true;
|
||||
}
|
||||
CalcCostInfo(splitContext);
|
||||
|
||||
splitRes.maxCost = INT64_MAX;
|
||||
splitRes.usedCoreNum = 1U;
|
||||
CalcSplitPlan(aicCoreNum_, splitRes.maxCost, splitContext, splitRes);
|
||||
// 3、存在FD任务,对FD进行负载均衡分配
|
||||
if (supportFd_ && splitRes.numOfFdHead > 0U) {
|
||||
SplitFD(splitRes);
|
||||
}
|
||||
splitRes.usedCoreNum = std::max(splitRes.usedCoreNum, 1U); // 至少使用1个core
|
||||
return true;
|
||||
}
|
||||
|
||||
bool VllmQuantLightningIndexerMetadataCpuKernel::GenMetaData(SplitResult &splitRes)
|
||||
{
|
||||
optiling::detail::QliMetaData* metaDataPtr = (optiling::detail::QliMetaData*)metaData_->GetData();
|
||||
// LI Metadata Generate
|
||||
for (size_t i = 0; i < aicCoreNum_; ++i) {
|
||||
if (i >= splitRes.usedCoreNum) {
|
||||
metaDataPtr->LIMetadata[i][LI_CORE_ENABLE_INDEX] = 0; // AIC disenable
|
||||
continue;
|
||||
}
|
||||
metaDataPtr->LIMetadata[i][LI_CORE_ENABLE_INDEX] = 1; // AIC enable
|
||||
// FA START
|
||||
metaDataPtr->LIMetadata[i][LI_BN2_START_INDEX] = i == 0 ? 0 : splitRes.bN2End[i-1];
|
||||
metaDataPtr->LIMetadata[i][LI_M_START_INDEX] = i == 0 ? 0 : splitRes.gS1End[i-1];
|
||||
metaDataPtr->LIMetadata[i][LI_S2_START_INDEX] = i == 0 ? 0 : splitRes.s2End[i-1];
|
||||
// FA END
|
||||
metaDataPtr->LIMetadata[i][LI_BN2_END_INDEX] = splitRes.bN2End[i];
|
||||
metaDataPtr->LIMetadata[i][LI_M_END_INDEX] = splitRes.gS1End[i];
|
||||
metaDataPtr->LIMetadata[i][LI_S2_END_INDEX] = splitRes.s2End[i];
|
||||
//
|
||||
metaDataPtr->LIMetadata[i][LI_FIRST_LD_DATA_WORKSPACE_IDX_INDEX] = splitRes.firstFdDataWorkspaceIdx[i];
|
||||
}
|
||||
|
||||
// LD Metadata Generate
|
||||
for (size_t i = 0; i < aivCoreNum_; ++i) {
|
||||
if (i >= splitRes.fdRes.fdUsedVecNum) {
|
||||
metaDataPtr->LDMetadata[i][LD_CORE_ENABLE_INDEX] = 0; // AIV disenable
|
||||
continue;
|
||||
}
|
||||
metaDataPtr->LDMetadata[i][LD_CORE_ENABLE_INDEX] = 1; // AIV enable
|
||||
uint32_t curFdIdx = splitRes.fdRes.fdIdx[i];
|
||||
metaDataPtr->LDMetadata[i][LD_BN2_IDX_INDEX] = splitRes.fdRes.fdBN2Idx[curFdIdx];
|
||||
metaDataPtr->LDMetadata[i][LD_M_IDX_INDEX] = splitRes.fdRes.fdMIdx[curFdIdx];
|
||||
metaDataPtr->LDMetadata[i][LD_WORKSPACE_IDX_INDEX] = splitRes.fdRes.fdWorkspaceIdx[curFdIdx];
|
||||
metaDataPtr->LDMetadata[i][LD_WORKSPACE_NUM_INDEX] = splitRes.fdRes.fdS2SplitNum[curFdIdx];
|
||||
metaDataPtr->LDMetadata[i][LD_M_START_INDEX] = splitRes.fdRes.fdMStart[i];
|
||||
metaDataPtr->LDMetadata[i][LD_M_NUM_INDEX] = splitRes.fdRes.fdMNum[i];
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
namespace {
|
||||
static const char *kernelType = "VllmQuantLightningIndexerMetadata";
|
||||
REGISTER_CPU_KERNEL(kernelType, VllmQuantLightningIndexerMetadataCpuKernel);
|
||||
}
|
||||
}; // namespace aicpu
|
||||
@@ -0,0 +1,310 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file vllm_quant_lightning_indexer_metadata_aicpu.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
#define QUANT_LIGHTNING_INDEXER_METADATA_AICPU_H
|
||||
|
||||
#include <string>
|
||||
#include <vector>
|
||||
#include <array>
|
||||
#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 {
|
||||
NORMAL_BLOCK = 0,
|
||||
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 {
|
||||
ASCEND910B = 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 };
|
||||
uint64_t maxS1GCost { 0 }; //新增
|
||||
|
||||
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 };
|
||||
int64_t s1GCost { 0 };
|
||||
int64_t s1GLastBlockCost { 0 };
|
||||
uint32_t s1GBlock { 0U };
|
||||
int64_t s1GNormalBlockCost { 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 curKvSplitPart { 1U };
|
||||
uint32_t preFdDataNum { 0U };
|
||||
|
||||
int64_t bN2Cost { 0 };
|
||||
uint32_t bN2Block { 0U };
|
||||
bool isFinished { false };
|
||||
BatchCache batchCache {};
|
||||
S1GCache s1GCache {};
|
||||
CoreCache coreCache {};
|
||||
};
|
||||
class VllmQuantLightningIndexerMetadataCpuKernel : public CpuKernel {
|
||||
public:
|
||||
VllmQuantLightningIndexerMetadataCpuKernel() = default;
|
||||
~VllmQuantLightningIndexerMetadataCpuKernel() = 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);
|
||||
uint32_t GetSparseSeqSize(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 CalcCost(uint32_t basicM, uint32_t basicS2);
|
||||
BlockCost<int64_t> CalcCostTable(uint32_t s1NormalSize, uint32_t s2NormalSize, uint32_t s1GTailSize,
|
||||
uint32_t s2TailSize);
|
||||
|
||||
// cache calculation
|
||||
void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache);
|
||||
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);
|
||||
void AssignByBlock(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void ForceAssign(const SplitContext &splitContext, AssignContext &assignContext);
|
||||
void AssignBlockToCore(uint32_t coreNum, 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(uint32_t coreNum, int64_t costLimit, const SplitContext &splitContext, SplitResult &result);
|
||||
|
||||
|
||||
private:
|
||||
CpuKernelContext* context_ = nullptr;
|
||||
// input
|
||||
Tensor *actSeqLenQ_ = nullptr;
|
||||
Tensor *actSeqLenKey_ = nullptr;
|
||||
// output
|
||||
Tensor *metaData_ = nullptr;
|
||||
// attributes
|
||||
std::string socVersion_ = "";
|
||||
bool supportFd_ = false;
|
||||
int32_t cmpRatio_ = 4;
|
||||
uint32_t aicCoreNum_ = 24U;
|
||||
uint32_t aivCoreNum_ = 48U;
|
||||
int32_t batchSize_ = 0;
|
||||
int32_t maxSeqlenQ_ = 0;
|
||||
int32_t maxSeqlenK_ = 0;
|
||||
int32_t numHeadsQ_ = 0;
|
||||
int32_t numHeadsK_ = 0;
|
||||
int32_t headDim_ = 0;
|
||||
int32_t queryQuantMode_ = 0;
|
||||
int32_t keyQuantMode_ = 0;
|
||||
int32_t sparseCount_ = 0;
|
||||
std::string layoutQuery_ = "BSND";
|
||||
std::string layoutKey_ = "BSND";
|
||||
int32_t sparseMode_ = 0;
|
||||
uint32_t attentionMode_ = 0;
|
||||
ValidSocVersion validSocVersion_ = ValidSocVersion::ASCEND910B;
|
||||
|
||||
// SplitParams
|
||||
int64_t preToken_ = INT64_MAX;
|
||||
int64_t nextToken_ = INT64_MAX;
|
||||
uint32_t groupSize_ = 0;
|
||||
uint32_t mBaseSize_ = 256;
|
||||
uint32_t s2BaseSize_ = 0;
|
||||
bool isS1G_ = true;
|
||||
bool isActQBatchPlus = false;
|
||||
|
||||
private:
|
||||
enum class ParamId : uint32_t {
|
||||
// input
|
||||
actSeqLenQ = 0,
|
||||
actSeqLenKV = 1,
|
||||
// output
|
||||
metaData = 0,
|
||||
};
|
||||
};
|
||||
} // namespace aicpu
|
||||
#endif
|
||||
@@ -0,0 +1,15 @@
|
||||
{
|
||||
"VllmQuantLightningIndexerMetadata":{
|
||||
"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