init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -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)

View 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模式。

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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"
}
}
}