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,29 @@
# -----------------------------------------------------------------------------------------------------------
# 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_op_to_compiled_list()
if (BUILD_OPEN_PROJECT)
target_sources(op_host_aclnn PRIVATE
vllm_quant_lightning_indexer_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME VllmQuantLightningIndexer
OPTIONS --cce-auto-sync=off
-Wno-deprecated-declarations
-mllvm -cce-vf-remove-membar=false
-mllvm -cce-aicore-hoist-movemask=false
)
if (NOT BUILD_OPS_RTY_KERNEL)
add_modules_sources(OPTYPE vllm_quant_lightning_indexer ACLNNTYPE aclnn)
endif()

View File

@@ -0,0 +1,153 @@
/**
 * 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_def.cpp
* \brief
*/
#include "register/op_def_registry.h"
namespace ops {
class VllmQuantLightningIndexer : public OpDef {
public:
explicit VllmQuantLightningIndexer(const char *name) : OpDef(name)
{
this->Input("query")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("key")
.ParamType(REQUIRED)
.DataType({ge::DT_INT8})
.Format({ge::FORMAT_ND})
.IgnoreContiguous();
this->Input("weights")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("query_dequant_scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("key_dequant_scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT16})
.Format({ge::FORMAT_ND})
.IgnoreContiguous();
this->Input("actual_seq_lengths_query")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("actual_seq_lengths_key")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("block_table")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Input("metadata")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
this->Output("sparse_indices").ParamType(REQUIRED).DataType({ge::DT_INT32}).Format({ge::FORMAT_ND});
this->Output("sparse_values").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND});
this->Attr("query_quant_mode").AttrType(REQUIRED).Int(0); // 0: 默认值,per-token-head
this->Attr("key_quant_mode").AttrType(REQUIRED).Int(0); // 0: 默认值,per-token-head
this->Attr("layout_query").AttrType(OPTIONAL).String("BSND");
this->Attr("layout_key").AttrType(OPTIONAL).String("BSND");
this->Attr("sparse_count").AttrType(OPTIONAL).Int(2048); // 2048: 默认值,筛选前2048
this->Attr("sparse_mode").AttrType(OPTIONAL).Int(3); // 3: 默认值,只计算下三角
this->Attr("pre_tokens").AttrType(OPTIONAL).Int(9223372036854775807); // 9223372036854775807: 默认值,int64的最大值
this->Attr("next_tokens").AttrType(OPTIONAL).Int(9223372036854775807); // 9223372036854775807: 默认值,int64的最大值
this->Attr("cmp_ratio").AttrType(OPTIONAL).Int(1); // 1: 压缩率
this->Attr("return_values").AttrType(OPTIONAL).Bool(false); // 是否返回sparse_values
this->Attr("stride").AttrType(OPTIONAL).Int(1); // stride参数
this->Attr("scale_stride").AttrType(OPTIONAL).Int(1); // scaleStride参数
OpAICoreConfig aicore_config;
aicore_config.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true);
this->AICore().AddConfig("ascend910b", aicore_config);
this->AICore().AddConfig("ascend910_93", aicore_config);
OpAICoreConfig aicore_config_950;
aicore_config_950.Input("query")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT8_E4M3FN})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("key")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT8_E4M3FN})
.Format({ge::FORMAT_ND})
.IgnoreContiguous();
aicore_config_950.Input("weights")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("query_dequant_scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("key_dequant_scale")
.ParamType(REQUIRED)
.DataType({ge::DT_FLOAT})
.Format({ge::FORMAT_ND})
.IgnoreContiguous();
aicore_config_950.Input("actual_seq_lengths_query")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("actual_seq_lengths_key")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("block_table")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Input("metadata")
.ParamType(OPTIONAL)
.DataType({ge::DT_INT32})
.Format({ge::FORMAT_ND})
.AutoContiguous();
aicore_config_950.Output("sparse_indices").ParamType(REQUIRED).DataType({ge::DT_INT32}).Format({ge::FORMAT_ND});
aicore_config_950.Output("sparse_values").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND});
aicore_config_950.DynamicCompileStaticFlag(true)
.DynamicFormatFlag(true)
.DynamicRankSupportFlag(true)
.DynamicShapeSupportFlag(true)
.NeedCheckSupportFlag(false)
.PrecisionReduceFlag(true)
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
.ExtendCfgInfo("opFile.value", "vllm_quant_lightning_indexer")
.ExtendCfgInfo("jitCompile.flag", "static_false,dynamic_false");
this->AICore().AddConfig("ascend950", aicore_config_950);
}
};
OP_ADD(VllmQuantLightningIndexer);
} // namespace ops

View File

@@ -0,0 +1,106 @@
/**
 * 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_infershape.cpp
* \brief
*/
#include <graph/utils/type_utils.h>
#include <register/op_impl_registry.h>
#include "err/ops_err.h"
#include "log/log.h"
using namespace ge;
namespace ops {
constexpr uint32_t QUERY_INDEX = 0;
constexpr uint32_t KEY_INDEX = 1;
constexpr uint32_t ATTR_QUERY_LAYOUT_INDEX = 2;
constexpr uint32_t ATTR_KV_LAYOUT_INDEX = 3;
constexpr uint32_t ATTR_SPARSE_COUNT_INDEX = 4;
constexpr uint32_t ATTR_RETURN_VALUE_INDEX = 9;
constexpr uint32_t DIM_NUM_3 = 3;
constexpr uint32_t DIM_NUM_4 = 4;
static ge::graphStatus InferShapeVllmQuantLightningIndexer(gert::InferShapeContext *context)
{
if (context == nullptr) {
OP_LOGE("VllmQuantLightningIndexer", "context is nullptr!");
return ge::GRAPH_FAILED;
}
const gert::Shape *queryShape = context->GetInputShape(QUERY_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context, queryShape);
const gert::Shape *keyShape = context->GetInputShape(KEY_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context, keyShape);
gert::Shape *sparseIndicesShape = context->GetOutputShape(0);
OP_CHECK_NULL_WITH_CONTEXT(context, sparseIndicesShape);
gert::Shape *sparseValuesShape = context->GetOutputShape(1);
auto attrs = context->GetAttrs();
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
const char *inputLayoutQueryPtr = attrs->GetAttrPointer<char>(ATTR_QUERY_LAYOUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context, inputLayoutQueryPtr);
const char *inputLayoutKeyPtr = attrs->GetAttrPointer<char>(ATTR_KV_LAYOUT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context, inputLayoutKeyPtr);
const int64_t *sparse_count = attrs->GetInt(ATTR_SPARSE_COUNT_INDEX);
OP_CHECK_NULL_WITH_CONTEXT(context, sparse_count);
std::string inputLayoutQueryPtrStr = std::string(inputLayoutQueryPtr);
std::string inputLayoutKeyPtrStr = std::string(inputLayoutKeyPtr);
if (inputLayoutQueryPtrStr != "TND" && inputLayoutQueryPtrStr != "BSND") {
OP_LOGE(context, "The input layout query should be TND or BSND, but got %s.", inputLayoutQueryPtrStr.c_str());
return GRAPH_FAILED;
}
int64_t keyHeadNum = (inputLayoutKeyPtrStr == "TND") ? keyShape->GetDim(1) : keyShape->GetDim(2);
if (inputLayoutQueryPtrStr == "BSND") {
sparseIndicesShape->SetDimNum(DIM_NUM_4);
sparseIndicesShape->SetDim(0, queryShape->GetDim(0)); // 0:Dim B
sparseIndicesShape->SetDim(1, queryShape->GetDim(1)); // 1:Dim S
sparseIndicesShape->SetDim(2, keyHeadNum); // 2:Dim N
sparseIndicesShape->SetDim(3, *sparse_count); // 3:Dim K
} else {
sparseIndicesShape->SetDimNum(DIM_NUM_3);
sparseIndicesShape->SetDim(0, queryShape->GetDim(0)); // 0:Dim T
sparseIndicesShape->SetDim(1, keyHeadNum); // 1:output shape's N Dim, 2: key shape's N Dim
sparseIndicesShape->SetDim(2, *sparse_count); // 2:Dim K
}
const bool *return_value = attrs->GetAttrPointer<bool>(ATTR_RETURN_VALUE_INDEX);
bool returnValueFlag = (return_value != nullptr) ? *return_value : false;
if (returnValueFlag) {
*sparseValuesShape = *sparseIndicesShape;
} else {
sparseValuesShape->SetDimNum(1);
sparseValuesShape->SetDim(0, 0);
}
OP_LOGD(context->GetNodeName(), "VllmQuantLightningIndexer InferShape end.");
return ge::GRAPH_SUCCESS;
}
static ge::graphStatus InferDataTypeVllmQuantLightningIndexer(gert::InferDataTypeContext *context)
{
if (context == nullptr) {
OP_LOGE("VllmQuantLightningIndexer", "InferDataTypeContext context is nullptr!");
return ge::GRAPH_FAILED;
}
OP_LOGD(context->GetNodeName(), "Enter VllmQuantLightningIndexer InferDataType impl.");
// default index data type is int32
ge::DataType outputType = ge::DT_INT32;
context->SetOutputDataType(0, outputType);
OP_LOGD(context->GetNodeName(), "VllmQuantLightningIndexer InferDataType end.");
return GRAPH_SUCCESS;
}
IMPL_OP_INFERSHAPE(VllmQuantLightningIndexer)
.InferShape(InferShapeVllmQuantLightningIndexer)
.InferDataType(InferDataTypeVllmQuantLightningIndexer);
} // namespace ops

View File

@@ -0,0 +1,948 @@
/**
 * 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_tiling.cpp
* \brief
*/
#include "vllm_quant_lightning_indexer_tiling.h"
#include "../op_kernel/vllm_quant_lightning_indexer_template_tiling_key.h"
using namespace ge;
using namespace AscendC;
using std::map;
using std::string;
namespace optiling {
// --------------------------QLIInfoParser类成员函数定义-------------------------------------
ge::graphStatus QLIInfoParser::CheckRequiredInOutExistence() const
{
OP_CHECK_IF(opParamInfo_.query.shape == nullptr, OP_LOGE(opName_, "Shape of tensor query is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.query.desc == nullptr, OP_LOGE(opName_, "Desc of tensor query is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.key.shape == nullptr, OP_LOGE(opName_, "Shape of tensor key is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.key.desc == nullptr, OP_LOGE(opName_, "Desc of tensor key is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.weights.shape == nullptr, OP_LOGE(opName_, "Shape of tensor weights is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.weights.desc == nullptr, OP_LOGE(opName_, "Desc of tensor weights is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.query_dequant_scale.shape == nullptr,
OP_LOGE(opName_, "Shape of tensor query_dequant_scale is nullptr"), return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.query_dequant_scale.desc == nullptr,
OP_LOGE(opName_, "Desc of tensor query_dequant_scale is nullptr"), return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.key_dequant_scale.shape == nullptr,
OP_LOGE(opName_, "Shape of tensor key_dequant_scale is nullptr"), return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.key_dequant_scale.desc == nullptr,
OP_LOGE(opName_, "Desc of tensor key_dequant_scale is nullptr"), return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.attenOut.shape == nullptr, OP_LOGE(opName_, "Shape of tensor output is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.attenOut.desc == nullptr, OP_LOGE(opName_, "Desc of tensor output is nullptr"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::CheckRequiredAttrExistence() const
{
OP_CHECK_IF(opParamInfo_.layOutQuery == nullptr, OP_LOGE(opName_, "attr layout_query is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.layOutKey == nullptr, OP_LOGE(opName_, "attr layout_key is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.sparseCount == nullptr, OP_LOGE(opName_, "attr sparse_count is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.sparseMode == nullptr, OP_LOGE(opName_, "attr sparse_mode is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.queryQuantMode == nullptr, OP_LOGE(opName_, "query_quant_mode is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.keyQuantMode == nullptr, OP_LOGE(opName_, "key_quant_mode is nullptr"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::CheckRequiredParaExistence() const
{
if (CheckRequiredInOutExistence() != ge::GRAPH_SUCCESS || CheckRequiredAttrExistence() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetOpName()
{
if (context_->GetNodeName() == nullptr) {
OP_LOGE("VllmQuantLightningIndexer", "opName got from TilingContext is nullptr");
return ge::GRAPH_FAILED;
}
opName_ = context_->GetNodeName();
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetNpuInfo()
{
platformInfo_ = context_->GetPlatformInfo();
OP_CHECK_IF(platformInfo_ == nullptr, OP_LOGE(opName_, "GetPlatformInfo is nullptr."), return ge::GRAPH_FAILED);
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo_);
uint32_t aivNum = ascendcPlatform.GetCoreNumAiv();
uint32_t aicNum = ascendcPlatform.GetCoreNumAic();
OP_CHECK_IF(aicNum == 0 || aivNum == 0, OP_LOGE(opName_, "num of core obtained is 0."), return GRAPH_FAILED);
socVersion_ = ascendcPlatform.GetSocVersion();
if ((socVersion_ != platform_ascendc::SocVersion::ASCEND910B) &&
(socVersion_ != platform_ascendc::SocVersion::ASCEND910_93) &&
(socVersion_ != platform_ascendc::SocVersion::ASCEND950)) {
OP_LOGE(opName_, "SOC Version[%d] is not support.", static_cast<int32_t>(socVersion_));
return GRAPH_FAILED;
}
OP_CHECK_IF(context_->GetWorkspaceSizes(1) == nullptr, OP_LOGE(opName_, "workSpaceSize got from ge is nullptr"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(context_->GetRawTilingData() == nullptr,
OP_LOGE(context_->GetNodeName(), "RawTilingData got from GE context is nullptr."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
void QLIInfoParser::GetOptionalInputParaInfo()
{
opParamInfo_.actualSeqLengthsQ.tensor = context_->GetOptionalInputTensor(ACTUAL_SEQ_Q_INDEX);
opParamInfo_.actualSeqLengthsQ.desc = context_->GetOptionalInputDesc(ACTUAL_SEQ_Q_INDEX);
opParamInfo_.actualSeqLengthsK.tensor = context_->GetOptionalInputTensor(ACTUAL_SEQ_K_INDEX);
opParamInfo_.actualSeqLengthsK.desc = context_->GetOptionalInputDesc(ACTUAL_SEQ_K_INDEX);
opParamInfo_.blockTable.tensor = context_->GetOptionalInputTensor(BLOCK_TABLE_INDEX);
opParamInfo_.blockTable.desc = context_->GetOptionalInputDesc(BLOCK_TABLE_INDEX);
opParamInfo_.metadata.tensor = context_->GetOptionalInputTensor(METADATA_INDEX);
opParamInfo_.metadata.desc = context_->GetOptionalInputDesc(METADATA_INDEX);
}
void QLIInfoParser::GetInputParaInfo()
{
opParamInfo_.query.desc = context_->GetInputDesc(QUERY_INDEX);
opParamInfo_.query.shape = context_->GetInputShape(QUERY_INDEX);
opParamInfo_.key.desc = context_->GetInputDesc(KEY_INDEX);
opParamInfo_.key.shape = context_->GetInputShape(KEY_INDEX);
opParamInfo_.weights.desc = context_->GetInputDesc(WEIGTHS_INDEX);
opParamInfo_.weights.shape = context_->GetInputShape(WEIGTHS_INDEX);
opParamInfo_.query_dequant_scale.desc = context_->GetInputDesc(QUERY_DEQUANT_SCALE_INDEX);
opParamInfo_.query_dequant_scale.shape = context_->GetInputShape(QUERY_DEQUANT_SCALE_INDEX);
opParamInfo_.key_dequant_scale.desc = context_->GetInputDesc(KEY_DEQUANT_SCALE_INDEX);
opParamInfo_.key_dequant_scale.shape = context_->GetInputShape(KEY_DEQUANT_SCALE_INDEX);
GetOptionalInputParaInfo();
}
void QLIInfoParser::GetOutputParaInfo()
{
opParamInfo_.attenOut.desc = context_->GetOutputDesc(vllm_quant_lightning_indexer);
opParamInfo_.attenOut.shape = context_->GetOutputShape(vllm_quant_lightning_indexer);
}
ge::graphStatus QLIInfoParser::GetAttrParaInfo()
{
auto attrs = context_->GetAttrs();
OP_CHECK_IF(attrs == nullptr, OP_LOGE(context_->GetNodeName(), "attrs got from ge is nullptr"),
return ge::GRAPH_FAILED);
OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo start");
opParamInfo_.layOutQuery = attrs->GetStr(ATTR_QUERY_LAYOUT_INDEX);
opParamInfo_.layOutKey = attrs->GetStr(ATTR_KEY_LAYOUT_INDEX);
opParamInfo_.queryQuantMode = attrs->GetAttrPointer<int64_t>(ATTR_QUERY_QUANT_MODE_INDEX);
opParamInfo_.keyQuantMode = attrs->GetAttrPointer<int64_t>(ATTR_KEY_QUANT_MODE_INDEX);
opParamInfo_.layOutQuery = attrs->GetStr(ATTR_QUERY_LAYOUT_INDEX);
opParamInfo_.layOutKey = attrs->GetStr(ATTR_KEY_LAYOUT_INDEX);
opParamInfo_.sparseCount = attrs->GetAttrPointer<int64_t>(ATTR_SPARSE_COUNT_INDEX);
opParamInfo_.sparseMode = attrs->GetAttrPointer<int64_t>(ATTR_SPARSE_MODE_INDEX);
opParamInfo_.preTokens = attrs->GetAttrPointer<int64_t>(ATTR_PRE_TOKENS_INDEX);
opParamInfo_.nextTokens = attrs->GetAttrPointer<int64_t>(ATTR_NEXT_TOKENS_INDEX);
opParamInfo_.cmpRatio = attrs->GetAttrPointer<int64_t>(ATTR_CMP_RATIO_INDEX);
opParamInfo_.returnValues = attrs->GetAttrPointer<bool>(ATTR_RETURN_VALUES_INDEX);
opParamInfo_.stride = attrs->GetAttrPointer<int64_t>(ATTR_STRIDE_INDEX);
opParamInfo_.scaleStride = attrs->GetAttrPointer<int64_t>(ATTR_SCALE_STRIDE_INDEX);
if (opParamInfo_.layOutQuery != nullptr) {
OP_LOGI(context_->GetNodeName(), "layout_query is:%s", opParamInfo_.layOutQuery);
}
if (opParamInfo_.layOutKey != nullptr) {
OP_LOGI(context_->GetNodeName(), "layout_key is:%s", opParamInfo_.layOutKey);
}
if (opParamInfo_.sparseCount != nullptr) {
OP_LOGI(context_->GetNodeName(), "selscted count is:%d", *opParamInfo_.sparseCount);
}
if (opParamInfo_.sparseMode != nullptr) {
OP_LOGI(context_->GetNodeName(), "sparse mode is:%d", *opParamInfo_.sparseMode);
}
if (opParamInfo_.preTokens != nullptr) {
OP_LOGI(context_->GetNodeName(), "preTokens is:%d", *opParamInfo_.preTokens);
}
if (opParamInfo_.nextTokens != nullptr) {
OP_LOGI(context_->GetNodeName(), "nextTokens is:%d", *opParamInfo_.nextTokens);
}
if (opParamInfo_.cmpRatio != nullptr) {
OP_LOGI(context_->GetNodeName(), "cmpRatio is:%d", *opParamInfo_.cmpRatio);
}
if (opParamInfo_.returnValues != nullptr) {
OP_LOGI(context_->GetNodeName(), "returnValues is:%s", *opParamInfo_.returnValues ? "true" : "false");
}
if (opParamInfo_.queryQuantMode != nullptr) {
OP_LOGI(context_->GetNodeName(), "query_quant_mode mode is:%d", *opParamInfo_.queryQuantMode);
}
if (opParamInfo_.keyQuantMode != nullptr) {
OP_LOGI(context_->GetNodeName(), "key_quant_mode mode is:%d", *opParamInfo_.keyQuantMode);
}
OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo end");
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::CheckAttrParaInfo()
{
std::string layout_key(opParamInfo_.layOutKey);
std::string layout_query(opParamInfo_.layOutQuery);
OP_CHECK_IF(
((std::string(opParamInfo_.layOutKey) != "PA_BSND")),
OP_LOGE(opName_, "input attr layout_key only supported PA_BSND,"
"but now layout_key is %s.", layout_key.c_str()),
return ge::GRAPH_FAILED);
if ((socVersion_ == platform_ascendc::SocVersion::ASCEND910B) ||
(socVersion_ == platform_ascendc::SocVersion::ASCEND910_93)) {
OP_CHECK_IF(!((*opParamInfo_.sparseCount > 0) && (*opParamInfo_.sparseCount <= SPARSE_LIMIT)),
OP_LOGE(opName_, "input attr sparse_count must > 0 and <= %d, but now sparse_count is %d",
SPARSE_LIMIT, *opParamInfo_.sparseCount),return ge::GRAPH_FAILED);
OP_CHECK_IF((*opParamInfo_.cmpRatio <= 0) || (*opParamInfo_.cmpRatio > 128) ||
((*opParamInfo_.cmpRatio & (*opParamInfo_.cmpRatio - 1)) != 0),
OP_LOGE(opName_, "input attr cmpRatio must > 0 and <= 128 and should be powers of 2, but now cmpRatio is %ld.",
*opParamInfo_.cmpRatio), return ge::GRAPH_FAILED);
} else if (socVersion_ == platform_ascendc::SocVersion::ASCEND950) {
OP_CHECK_IF(!((*opParamInfo_.sparseCount > 0) && (*opParamInfo_.sparseCount <= SPARSE_LIMIT)),
OP_LOGE(opName_, "input attr sparse_count must > 0 and <= %d, but now sparse_count is %d",
SPARSE_LIMIT, *opParamInfo_.sparseCount),return ge::GRAPH_FAILED);
OP_CHECK_IF((*opParamInfo_.cmpRatio != 1) && (*opParamInfo_.cmpRatio != 4) && (*opParamInfo_.cmpRatio != 128),
OP_LOGE(opName_, "input attr cmpRatio must be 1、4 or 128, but now cmpRatio is %ld.",
*opParamInfo_.cmpRatio), return ge::GRAPH_FAILED);
}
OP_CHECK_IF(((std::string(opParamInfo_.layOutQuery) != "BSND") && (std::string(opParamInfo_.layOutQuery) != "TND")),
OP_LOGE(opName_, "input attr layout_query only supported BSND or TND."), return ge::GRAPH_FAILED);
OP_CHECK_IF(
((std::string(opParamInfo_.layOutKey) != "PA_BSND") &&
(std::string(opParamInfo_.layOutQuery)) != (std::string(opParamInfo_.layOutKey))),
OP_LOGE(opName_, "outside of PA, input attr layout_query and input attr layout_key must be the same,"
"but now layout_key is %s, layout_query is %s.",
layout_key.c_str(), layout_query.c_str()), return ge::GRAPH_FAILED);
OP_CHECK_IF(!((*opParamInfo_.sparseMode == 0) || (*opParamInfo_.sparseMode == SPARSE_MODE_LOWER)),
OP_LOGE(opName_, "input attr sparse_mode only supported 0 or 3, but now sparseMode is %d.",
*opParamInfo_.sparseMode), return ge::GRAPH_FAILED);
OP_CHECK_IF(*opParamInfo_.preTokens != 9223372036854775807,
OP_LOGE(opName_, "input attr preTokens only supported 9223372036854775807, but now preTokens is %ld.",
*opParamInfo_.preTokens), return ge::GRAPH_FAILED);
OP_CHECK_IF(*opParamInfo_.nextTokens != 9223372036854775807,
OP_LOGE(opName_, "input attr nextTokens only supported 9223372036854775807, but now nextTokens is %ld.",
*opParamInfo_.nextTokens), return ge::GRAPH_FAILED);
OP_CHECK_IF(*opParamInfo_.queryQuantMode != 0, OP_LOGE(opName_, "input attr query_quant_mode only supported 0."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(*opParamInfo_.keyQuantMode != 0, OP_LOGE(opName_, "input attr key_quant_mode only supported 0."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(*opParamInfo_.returnValues, OP_LOGE(opName_, "input attr returnValues only supported False."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetOpParaInfo()
{
GetInputParaInfo();
GetOutputParaInfo();
if (ge::GRAPH_SUCCESS != GetAttrParaInfo()) {
return ge::GRAPH_FAILED;
}
if (ge::GRAPH_SUCCESS != CheckAttrParaInfo()) {
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetAndCheckInOutDataType()
{
inputQType_ = opParamInfo_.query.desc->GetDataType();
inputKType_ = opParamInfo_.key.desc->GetDataType();
weightsType_ = opParamInfo_.weights.desc->GetDataType();
inputQueryScaleType_ = opParamInfo_.query_dequant_scale.desc->GetDataType();
inputKeyScaleType_ = opParamInfo_.key_dequant_scale.desc->GetDataType();
outputType_ = opParamInfo_.attenOut.desc->GetDataType();
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo_);
socVersion_ = ascendcPlatform.GetSocVersion();
OP_CHECK_IF(!(inputQType_ == inputKType_),
OP_LOGE(opName_, "The data types of the input query and key must be the same."),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
!(inputQueryScaleType_ == inputKeyScaleType_),
OP_LOGE(opName_, "The data types of the input query_dequant_scale and key_dequant_scale must be the same."),
return ge::GRAPH_FAILED);
if ((socVersion_ == platform_ascendc::SocVersion::ASCEND910B) ||
(socVersion_ == platform_ascendc::SocVersion::ASCEND910_93)) {
OP_CHECK_IF(inputQType_ != ge::DT_INT8,
OP_LOGE(opName_, "The data types of the input query and key must be int8."), return ge::GRAPH_FAILED);
OP_CHECK_IF(
inputQueryScaleType_ != ge::DT_FLOAT16,
OP_LOGE(opName_, "The data types of the input query_dequant_scale and key_dequant_scale must be float16."),
return ge::GRAPH_FAILED);
} else if (socVersion_ == platform_ascendc::SocVersion::ASCEND950) {
OP_CHECK_IF(inputQType_ != ge::DT_FLOAT8_E4M3FN,
OP_LOGE(opName_, "The data types of the input query and key must be float8_e4m3."), return ge::GRAPH_FAILED);
OP_CHECK_IF(
inputQueryScaleType_ != ge::DT_FLOAT,
OP_LOGE(opName_, "The data types of the input query_dequant_scale and key_dequant_scale must be float."),
return ge::GRAPH_FAILED);
}
if ((socVersion_ == platform_ascendc::SocVersion::ASCEND910B) ||
(socVersion_ == platform_ascendc::SocVersion::ASCEND910_93)) {
OP_CHECK_IF(weightsType_ != ge::DT_FLOAT16,
OP_LOGE(opName_, "The data types of the input weights must be float16."), return ge::GRAPH_FAILED);
} else if (socVersion_ == platform_ascendc::SocVersion::ASCEND950) {
OP_CHECK_IF(weightsType_ != ge::DT_FLOAT,
OP_LOGE(opName_, "The data types of the input weights must be float."), return ge::GRAPH_FAILED);
}
OP_CHECK_IF(outputType_ != ge::DT_INT32,
OP_LOGE(opName_, "The data types of the output sparse_indices must be int32."),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetQueryKeyAndOutLayout()
{
// 获取query,key的Layout基准值
const map<string, DataLayout> layoutQueryMap = {{"BSND", DataLayout::BSND}, {"TND", DataLayout::TND}};
std::string layout_query(opParamInfo_.layOutQuery);
auto QLayout_ = layoutQueryMap.find(layout_query);
if (QLayout_ != layoutQueryMap.end()) {
qLayout_ = QLayout_->second;
}
const map<string, DataLayout> layoutKeyMap = {
{"BSND", DataLayout::BSND}, {"TND", DataLayout::TND},
{"PA_BSND", DataLayout::PA_BSND}, {"PA_BBND", DataLayout::PA_BSND}};
std::string layout_key(opParamInfo_.layOutKey);
auto KLayout = layoutKeyMap.find(layout_key);
if (KLayout != layoutKeyMap.end()) {
kLayout_ = KLayout->second;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetAndCheckOptionalInput()
{
if (kLayout_ == DataLayout::PA_BSND) {
OP_CHECK_IF(opParamInfo_.blockTable.tensor == nullptr,
OP_LOGE(opName_, "key layout only supported PA_BSND, input block_table must not be null"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
opParamInfo_.actualSeqLengthsK.tensor == nullptr,
OP_LOGE(opName_, "key layout only supported PA_BSND, input actual_seq_lengths_key must not be null"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.blockTable.desc->GetDataType() != ge::DT_INT32,
OP_LOGE(opName_, "input block_table data type only support int32"), return ge::GRAPH_FAILED);
} else {
OP_CHECK_IF(opParamInfo_.blockTable.tensor != nullptr,
OP_LOGE(opName_, "key layout is not PA_BSND, input block_table must be null"),
return ge::GRAPH_FAILED);
}
if (kLayout_ == DataLayout::TND) {
OP_CHECK_IF(opParamInfo_.actualSeqLengthsK.tensor == nullptr,
OP_LOGE(opName_, "when layout_key is TND, input actual_seq_lengths_key must not be null"),
return ge::GRAPH_FAILED);
}
OP_CHECK_IF(opParamInfo_.actualSeqLengthsK.tensor != nullptr &&
opParamInfo_.actualSeqLengthsK.desc->GetDataType() != ge::DT_INT32,
OP_LOGE(opName_, "input actual_seq_lengths_key data type only support int32"),
return ge::GRAPH_FAILED);
if (qLayout_ == DataLayout::TND) {
OP_CHECK_IF(opParamInfo_.actualSeqLengthsQ.tensor == nullptr,
OP_LOGE(opName_, "when layout_query is TND, input actual_seq_lengths_query must not be null"),
return ge::GRAPH_FAILED);
}
OP_CHECK_IF(opParamInfo_.actualSeqLengthsQ.tensor != nullptr &&
opParamInfo_.actualSeqLengthsQ.desc->GetDataType() != ge::DT_INT32,
OP_LOGE(opName_, "input actual_seq_lengths_query data type only support int32"),
return ge::GRAPH_FAILED);
OP_CHECK_IF(opParamInfo_.metadata.tensor == nullptr,
OP_LOGE(opName_, "input metadata must not be null"),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::CheckShapeDim()
{
OP_CHECK_IF((opParamInfo_.blockTable.tensor != nullptr) &&
(opParamInfo_.blockTable.tensor->GetStorageShape().GetDimNum() != DIM_NUM_TWO),
OP_LOGE(opName_, "the dim num of block_table's shape should be 2, but now is %u",
opParamInfo_.blockTable.tensor->GetStorageShape().GetDimNum()), return ge::GRAPH_FAILED);
OP_CHECK_IF(
((kLayout_ == DataLayout::PA_BSND)||(kLayout_ == DataLayout::BSND)) &&
(opParamInfo_.key.shape->GetStorageShape().GetDimNum() != DIM_NUM_FOUR),
OP_LOGE(opName_, "the dim num of key's shape should be 4, but now is %u",
opParamInfo_.key.shape->GetStorageShape().GetDimNum()), return ge::GRAPH_FAILED);
OP_CHECK_IF(
(kLayout_ == DataLayout::TND) && (opParamInfo_.key.shape->GetStorageShape().GetDimNum() != DIM_NUM_THREE),
OP_LOGE(opName_, "the dim num of key's shape should be 3, but now is %u",
opParamInfo_.key.shape->GetStorageShape().GetDimNum()), return ge::GRAPH_FAILED);
uint32_t qShapeDim = opParamInfo_.query.shape->GetStorageShape().GetDimNum();
uint32_t weightsShapeDim = opParamInfo_.weights.shape->GetStorageShape().GetDimNum();
uint32_t outShapeDim = opParamInfo_.attenOut.shape->GetStorageShape().GetDimNum();
uint32_t expectShapeDim = DIM_NUM_FOUR;
if (qLayout_ == DataLayout::TND) {
expectShapeDim = DIM_NUM_THREE;
}
OP_CHECK_IF(
qShapeDim != expectShapeDim,
OP_LOGE(opName_, "the dim num of query's shape should be %u, but now is %u", expectShapeDim, qShapeDim),
return ge::GRAPH_FAILED);
OP_CHECK_IF(outShapeDim != expectShapeDim,
OP_LOGE(opName_, "the dim num of sparse_indices's shape should be %u, but now is %u", expectShapeDim,
outShapeDim),
return ge::GRAPH_FAILED);
OP_CHECK_IF(!(weightsShapeDim == expectShapeDim - 1),
OP_LOGE(opName_, "the dim num of weights's shape should be %u, but now is %u", expectShapeDim - 1,
weightsShapeDim),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetN1Size()
{
if (qLayout_ == DataLayout::BSND) {
n1Size_ = static_cast<uint32_t>(opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_TWO));
} else {
// TND
n1Size_ = static_cast<uint32_t>(opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_ONE));
}
OP_LOGI(context_->GetNodeName(), "n1Size is %d", n1Size_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
const std::string &actualSeqLenName) const
{
size = static_cast<uint32_t>(tensor->GetShapeSize());
if (size <= 0) {
OP_LOGE(opName_, "%s's shape size is %u, it should be greater than 0.", actualSeqLenName.c_str(), size);
return ge::GRAPH_FAILED;
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetAndCheckN2Size()
{
// PA_BSND
if (kLayout_ == DataLayout::TND) {
n2Size_ = static_cast<uint32_t>(opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_ONE));
} else {
n2Size_ = static_cast<uint32_t>(opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_TWO));
}
OP_LOGI(context_->GetNodeName(), "N2 is %d", n2Size_);
OP_CHECK_IF(n2Size_ != 1, OP_LOGE(opName_, "key shape[2] is numhead, only support 1."), return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetGSize()
{
if (n1Size_ % n2Size_ != 0) {
OP_LOGE(opName_, "input query's head_num %u can not be a multiple of key's head_num %u.", n1Size_, n2Size_);
return ge::GRAPH_FAILED;
}
gSize_ = n1Size_ / n2Size_;
OP_CHECK_IF(gSize_ != G_SIZE_LIMIT,
OP_LOGE(opName_, "N1 is %u, N2 is %u, N1 divided by N2 must equal 64.", n1Size_, n2Size_),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetBatchSize()
{
// 获取B基准值
// 1、非TND时, 以query的batch_size维度为基准;
// 2、Q和K都为TND时, actual_seq_lens_q必须传入, 以actual_seq_lens_q数组的长度为B轴大小
// 3、Q为TND,K为PA_BSND时,以actual_seq_lens_k数组的长度为B轴大小
if (qLayout_ == DataLayout::BSND) {
bSize_ = opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_ZERO);
OP_LOGI(context_->GetNodeName(), "b: %d, s: %d, n: %d,d :%d",
opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_ZERO),
opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_ONE),
opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_TWO),
opParamInfo_.query.shape->GetStorageShape().GetDim(DIM_IDX_THREE));
return ge::GRAPH_SUCCESS;
} else { // TND
uint32_t bSizeQuery;
uint32_t bSizeKey;
GetActualSeqLenSize(bSizeQuery, opParamInfo_.actualSeqLengthsQ.tensor, "input actual_seq_lengths_query");
GetActualSeqLenSize(bSizeKey, opParamInfo_.actualSeqLengthsK.tensor, "input actual_seq_lengths_key");
if (kLayout_ == DataLayout::TND) {
OP_CHECK_IF(bSizeQuery != bSizeKey,
OP_LOGE(opName_, "the lengths of actual_seq_lengths_query and actual_seq_lengths_key is %u, %u respectively, they must be same.",
bSizeQuery, bSizeKey),
return ge::GRAPH_FAILED);
bSize_ = bSizeQuery;
} else {
if (bSizeQuery == bSizeKey + 1) {
batchSupperFlag_ = true;
}
OP_CHECK_IF((bSizeQuery != bSizeKey) && !batchSupperFlag_,
OP_LOGE(opName_, "the lengths of actual_seq_lengths_query and actual_seq_lengths_key is %u, %u respectively, they must be same.",
bSizeQuery, bSizeKey),
return ge::GRAPH_FAILED);
bSize_ = bSizeKey; // Q为TND,batch从Key中获取
}
return ge::GRAPH_SUCCESS;
}
}
ge::graphStatus QLIInfoParser::GetHeadDim()
{
// 以query的D维度为基准
uint32_t dIndex = DIM_IDX_TWO;
// 根据layout确定D维度在shape中的位置
switch (qLayout_) {
case DataLayout::TND:
// TND格式: [Total, N, D] -> D是第2维(索引2)
dIndex = DIM_IDX_TWO;
break;
case DataLayout::BSND:
// BSND格式: [Batch, SeqLen, N, D] -> D是第3维(索引3)
dIndex = DIM_IDX_THREE;
break;
default:
OP_LOGE(opName_, "unsupported layout for getting head dim.");
return ge::GRAPH_FAILED;
}
headDim_ = opParamInfo_.query.shape->GetStorageShape().GetDim(dIndex);
OP_CHECK_IF(headDim_ != HEAD_DIM_LIMIT, OP_LOGE(opName_, "input query's last dim head_dim only support 128, but now is %u.", headDim_),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetS1Size()
{
if (qLayout_ == DataLayout::BSND) {
s1Size_ = opParamInfo_.query.shape->GetStorageShape().GetDim(1);
}
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetAndCheckBlockSize()
{
blockSize_ = static_cast<uint32_t>(opParamInfo_.key.shape->GetStorageShape().GetDim(1));
OP_LOGI(context_->GetNodeName(), "blockSize_ is %d", blockSize_);
OP_CHECK_IF(
((blockSize_ % BLOCK_SIZE_FACTOR != 0) || (blockSize_ == 0) || (blockSize_ > BLOCK_SIZE_LIMIT)),
OP_LOGE(opName_, "input key's block_size must be a multiple of 16 and belong to (0, 1024], but now is %d.", blockSize_),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetS2SizeForPageAttention()
{
if (GetAndCheckBlockSize() != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
int32_t blockCount_ = static_cast<uint32_t>(opParamInfo_.key.shape->GetStorageShape().GetDim(0));
OP_CHECK_IF((blockCount_ == 0), OP_LOGE(opName_, "input key's block_count cannot be 0."), return ge::GRAPH_FAILED);
maxBlockNumPerBatch_ = opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(1);
s2Size_ = maxBlockNumPerBatch_ * blockSize_;
OP_LOGI(context_->GetNodeName(), "maxBlockNumPerBatch_ is %d, blockSize_ is %d, s2Size_ is %d",
maxBlockNumPerBatch_, blockSize_, s2Size_);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetS2SizeForBatchContinuous()
{
std::string layout_key(opParamInfo_.layOutKey);
if (kLayout_ == DataLayout::BSND) {
s2Size_ = opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_ONE);
} else if (kLayout_ == DataLayout::TND) {
s2Size_ = opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_ZERO);
}
OP_CHECK_IF((kLayout_ != DataLayout::BSND) && (kLayout_ != DataLayout::TND),
OP_LOGE(opName_, "the layout of key is %s, it is unsupported.", layout_key.c_str()),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::GetS2Size()
{
// 获取S2基准值
// 1、BATCH_CONTINUOUS时, 从key的S轴获取
// 3、PAGE_ATTENTION时, S2 = block_table.dim1 * block_size
if (kLayout_ == DataLayout::PA_BSND) {
return GetS2SizeForPageAttention();
}
return GetS2SizeForBatchContinuous();
}
ge::graphStatus QLIInfoParser::ValidateInputShapesMatch()
{
/*
TND:
query [T,N1,D],
key [BlockNum,BlockSize,N2,D],
weight [T,N1],
block_table [BatchSize, BatchMaxBlockNum],
act_seq_k [BatchSize]
act_seq_q [BatchSize],
out [T,N2,topk]
----------------------
BSND:
query [BatchSize,S1,N1,D],
key [BlockNum,BlockSize,N2,D],
weight [BatchSize,S1,N1],
block_table [BatchSize, BatchMaxBlockNum],
act_seq_k [BatchSize]
act_seq_q [BatchSize] 可选
out [BatchSize,S1,N2,topk]
*/
uint32_t queryWeightsN1Dim = 1;
uint32_t outN2Dim = 1;
if (qLayout_ == DataLayout::TND) {
// -----------------------check BatchSize-------------------
// bSize_ 来源于act_seq_q
OP_CHECK_IF((kLayout_ == DataLayout::PA_BSND) &&
((opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize() != bSize_) ||
(opParamInfo_.blockTable.tensor != nullptr &&
opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(0) != bSize_)),
OP_LOGE(
opName_,
"TND case input actual_seq_lengths_query, actual_seq_lengths_key, block_table dim 0 are %u, %u, %u "
"respectively, they must be same.",
bSize_, opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize(),
opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(0)),
return ge::GRAPH_FAILED);
OP_CHECK_IF((kLayout_ != DataLayout::PA_BSND) &&
(opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize() != bSize_),
OP_LOGE(
opName_,
"TND case input actual_seq_lengths_query, actual_seq_lengths_key, are %u, %u "
"respectively, they must be same.",
bSize_, opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize()),
return ge::GRAPH_FAILED);
// -----------------------check T-------------------
uint32_t qTsize = opParamInfo_.query.shape->GetStorageShape().GetDim(0);
OP_CHECK_IF((opParamInfo_.weights.shape->GetStorageShape().GetDim(0) != qTsize) ||
(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != qTsize),
OP_LOGE(opName_,
"TND case input query, weights, sparse_indices dim 0 are %u, %u, %u "
"respectively, they must be same.",
qTsize, opParamInfo_.weights.shape->GetStorageShape().GetDim(0),
opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0)),
return ge::GRAPH_FAILED);
} else {
// -----------------------check BatchSize-------------------
// bSize_ 来源于query
OP_CHECK_IF((kLayout_ == DataLayout::PA_BSND) &&
((opParamInfo_.weights.shape->GetStorageShape().GetDim(0) != bSize_) ||
(opParamInfo_.blockTable.tensor != nullptr &&
opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(0) != bSize_) ||
(opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize() != bSize_) ||
(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != bSize_)),
OP_LOGE(opName_,
"BSND case input query, weight, actual_seq_lengths_key, block_table, sparse_indices dim 0 are %u, %u, %u, %u, %u respectively, they must be same.",
bSize_, opParamInfo_.weights.shape->GetStorageShape().GetDim(0),
opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize(),
opParamInfo_.blockTable.tensor->GetStorageShape().GetDim(0),
opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0)),
return ge::GRAPH_FAILED);
OP_CHECK_IF((kLayout_ != DataLayout::PA_BSND) &&
((opParamInfo_.weights.shape->GetStorageShape().GetDim(0) != bSize_) ||
(opParamInfo_.actualSeqLengthsK.tensor != nullptr &&
opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize() != bSize_) ||
(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0) != bSize_)),
OP_LOGE(opName_,
"BSND case input query, weight, actual_seq_lengths_key, sparse_indices dim 0 are %u, %u, %u, %u respectively, they must be same.",
bSize_, opParamInfo_.weights.shape->GetStorageShape().GetDim(0),
opParamInfo_.actualSeqLengthsK.tensor->GetShapeSize(),
opParamInfo_.attenOut.shape->GetStorageShape().GetDim(0)),
return ge::GRAPH_FAILED);
OP_CHECK_IF(
(opParamInfo_.actualSeqLengthsQ.tensor != nullptr) &&
(opParamInfo_.actualSeqLengthsQ.tensor->GetShapeSize() != bSize_),
OP_LOGE(
opName_,
"BSND case input query, actual_seq_lengths_query dim 0 are %u, %ld respectively, they must be same",
bSize_, opParamInfo_.actualSeqLengthsQ.tensor->GetShapeSize()),
return ge::GRAPH_FAILED);
// -----------------------check S1-------------------
OP_CHECK_IF(
(opParamInfo_.weights.shape->GetStorageShape().GetDim(1) != s1Size_) ||
(opParamInfo_.attenOut.shape->GetStorageShape().GetDim(1) != s1Size_),
OP_LOGE(opName_, "BSND case input query, weight, sparse_indices dim 1 are %u, %u, %u, they must be same.",
s1Size_, opParamInfo_.weights.shape->GetStorageShape().GetDim(1),
opParamInfo_.attenOut.shape->GetStorageShape().GetDim(1)),
return ge::GRAPH_FAILED);
queryWeightsN1Dim = DIM_IDX_TWO;
outN2Dim = DIM_IDX_TWO;
}
// -----------------------check N1-------------------
OP_CHECK_IF((opParamInfo_.weights.shape->GetStorageShape().GetDim(queryWeightsN1Dim) != n1Size_),
OP_LOGE(opName_, "input query, weight shape dim N1 must be same, but now are %u, %u respectively, they must be same.",
opParamInfo_.weights.shape->GetStorageShape().GetDim(queryWeightsN1Dim), n1Size_),
return ge::GRAPH_FAILED);
// -----------------------check D-------------------
OP_CHECK_IF(
((kLayout_ != DataLayout::TND && opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_THREE) != headDim_)
|| (kLayout_ == DataLayout::TND && opParamInfo_.key.shape->GetStorageShape().GetDim(DIM_IDX_TWO) != headDim_)),
OP_LOGE(opName_, "input query, key shape last dim must be same."), return ge::GRAPH_FAILED);
// -----------------------check N2-------------------
OP_CHECK_IF((opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim) != n2Size_),
OP_LOGE(opName_, "input query and output sparse_indices shape n2 dim must be same."),
return ge::GRAPH_FAILED);
// -----------------------check sparse_count-------------------
OP_CHECK_IF((opParamInfo_.attenOut.shape->GetStorageShape().GetDim(outN2Dim + 1) != *opParamInfo_.sparseCount),
OP_LOGE(opName_, "output sparse_indices shape last dim must be same as attr sparse_count."),
return ge::GRAPH_FAILED);
// -----------------------check metadata-------------------
OP_CHECK_IF((opParamInfo_.metadata.tensor->GetShapeSize() != METADATA_LIMIT),
OP_LOGE(opName_, "input metadata dim 0 must be %u.", METADATA_LIMIT),
return ge::GRAPH_FAILED);
return ge::GRAPH_SUCCESS;
}
ge::graphStatus QLIInfoParser::CheckScaleShape()
{
uint32_t qShapeDim = opParamInfo_.query.shape->GetStorageShape().GetDimNum();
uint32_t kShapeDim = opParamInfo_.key.shape->GetStorageShape().GetDimNum();
uint32_t qDequantScaleShapeDim = opParamInfo_.query_dequant_scale.shape->GetStorageShape().GetDimNum();
uint32_t kDequantScaleShapeDim = opParamInfo_.key_dequant_scale.shape->GetStorageShape().GetDimNum();
OP_CHECK_IF(qDequantScaleShapeDim != (qShapeDim - 1),
OP_LOGE(opName_, "the dim num of query_dequant_scale's shape should be %u, but now is %u",
qShapeDim - 1, qDequantScaleShapeDim),
return ge::GRAPH_FAILED);
OP_CHECK_IF(kDequantScaleShapeDim != (kShapeDim - 1),
OP_LOGE(opName_, "the dim num of key_dequant_scale's shape should be %u, but now is %u", kShapeDim - 1,
kDequantScaleShapeDim),
return ge::GRAPH_FAILED);
// check q scale
for (uint32_t i = 0; i < (qShapeDim - 1); i++) {
uint32_t dimValueQueryScale = opParamInfo_.query_dequant_scale.shape->GetStorageShape().GetDim(i);
uint32_t dimValueQuery = opParamInfo_.query.shape->GetStorageShape().GetDim(i);
OP_CHECK_IF(dimValueQueryScale != dimValueQuery,
OP_LOGE(opName_, "query_dequant_scale's shape[%u] %u and query's shape[%u] %u is not same", i,
dimValueQueryScale, i, dimValueQuery),
return ge::GRAPH_FAILED);
}
// check k scale
for (uint32_t i = 0; i < (kShapeDim - 1); i++) {
uint32_t dimValueKeyScale = opParamInfo_.key_dequant_scale.shape->GetStorageShape().GetDim(i);
uint32_t dimValueKey = opParamInfo_.key.shape->GetStorageShape().GetDim(i);
OP_CHECK_IF(dimValueKeyScale != dimValueKey,
OP_LOGE(opName_, "key_dequant_scale's shape[%u] %u and key's shape[%u] %u is not same", i,
dimValueKeyScale, i, dimValueKey),
return ge::GRAPH_FAILED);
}
return ge::GRAPH_SUCCESS;
}
void QLIInfoParser::GenerateInfo(QLITilingInfo &QLIInfo)
{
QLIInfo.opName = opName_;
QLIInfo.platformInfo = platformInfo_;
QLIInfo.opParamInfo = opParamInfo_;
QLIInfo.socVersion = socVersion_;
QLIInfo.bSize = bSize_;
QLIInfo.n1Size = n1Size_;
QLIInfo.n2Size = n2Size_;
QLIInfo.s1Size = s1Size_;
QLIInfo.s2Size = s2Size_;
QLIInfo.gSize = gSize_;
QLIInfo.inputQType = inputQType_;
QLIInfo.inputKType = inputKType_;
QLIInfo.outputType = outputType_;
QLIInfo.blockSize = blockSize_;
QLIInfo.maxBlockNumPerBatch = maxBlockNumPerBatch_;
QLIInfo.pageAttentionFlag = (kLayout_ == DataLayout::PA_BSND);
QLIInfo.batchSupperFlag = batchSupperFlag_;
QLIInfo.sparseMode = *opParamInfo_.sparseMode;
QLIInfo.sparseCount = *opParamInfo_.sparseCount;
QLIInfo.preTokens = *opParamInfo_.preTokens;
QLIInfo.nextTokens = *opParamInfo_.nextTokens;
QLIInfo.cmpRatio = *opParamInfo_.cmpRatio;
QLIInfo.returnValues = *opParamInfo_.returnValues;
QLIInfo.stride = *opParamInfo_.stride;
QLIInfo.scaleStride = *opParamInfo_.scaleStride;
QLIInfo.inputQLayout = qLayout_;
QLIInfo.inputKLayout = kLayout_;
}
ge::graphStatus QLIInfoParser::ParseAndCheck(QLITilingInfo &QLIInfo)
{
if (ge::GRAPH_SUCCESS != GetOpName() || ge::GRAPH_SUCCESS != GetNpuInfo() || ge::GRAPH_SUCCESS != GetOpParaInfo() ||
ge::GRAPH_SUCCESS != CheckRequiredParaExistence()) {
return ge::GRAPH_FAILED;
}
if (ge::GRAPH_SUCCESS != GetAndCheckInOutDataType() || ge::GRAPH_SUCCESS != GetQueryKeyAndOutLayout() ||
ge::GRAPH_SUCCESS != GetAndCheckOptionalInput()) {
return ge::GRAPH_FAILED;
}
if (ge::GRAPH_SUCCESS != CheckShapeDim() || ge::GRAPH_SUCCESS != GetN1Size() ||
ge::GRAPH_SUCCESS != GetAndCheckN2Size() || ge::GRAPH_SUCCESS != GetGSize()) {
return ge::GRAPH_FAILED;
}
if (ge::GRAPH_SUCCESS != GetBatchSize() || ge::GRAPH_SUCCESS != GetS1Size() || ge::GRAPH_SUCCESS != GetHeadDim() ||
ge::GRAPH_SUCCESS != GetS2Size()) {
return ge::GRAPH_FAILED;
}
if (ge::GRAPH_SUCCESS != ValidateInputShapesMatch() || ge::GRAPH_SUCCESS != CheckScaleShape()) {
return ge::GRAPH_FAILED;
}
GenerateInfo(QLIInfo);
return ge::GRAPH_SUCCESS;
}
// --------------------------TilingPrepare函数定义-------------------------------------
static ge::graphStatus TilingPrepareForVllmQuantLightningIndexer(gert::TilingParseContext * /* context */)
{
return ge::GRAPH_SUCCESS;
}
// --------------------------VllmQuantLightningIndexerTiling类成员函数定义-----------------------
ge::graphStatus VllmQuantLightningIndexerTiling::DoTiling(QLITilingInfo *tilingInfo)
{
// -------------set blockdim-----------------
auto ascendcPlatform = platform_ascendc::PlatformAscendC(tilingInfo->platformInfo);
uint32_t aivNum = ascendcPlatform.GetCoreNumAiv();
uint32_t aicNum = ascendcPlatform.GetCoreNumAic();
uint32_t blockDim = ascendcPlatform.CalcTschBlockDim(aivNum, aicNum, aivNum);
context_->SetBlockDim(blockDim);
// -------------set workspacesize-----------------
constexpr uint32_t MM1_RES_ELEM_SIZE = 4; // 4: fp32
constexpr uint32_t DOUBLE_BUFFER = 2; // 双Buffer
constexpr uint32_t M_BASE_SIZE = 512; // m轴基本块大小
constexpr uint32_t S2_BASE_SIZE = 512; // S2轴基本块大小
constexpr uint32_t V1_RES_ELEM_SIZE = 4; // 4: int32
constexpr uint32_t V1_RES_ELEM_TYPE = 2; // 保留Index和Value 2种数据
constexpr uint32_t V1_DECODE_PARAM_ELEM_SIZE = 8; // 8: int64
constexpr uint32_t V1_DECODE_PARAM_NUM = 16; // Decode参数个数
constexpr uint32_t V1_DECODE_DATA_NUM = 2; // Decode每个核需要存储头和尾部两块数据
constexpr uint32_t S1_BASE_SIZE = 8; // S1轴基本块的大小
constexpr uint32_t TOPK_MAX_SIZE = 2048; // TopK选取个数
constexpr uint32_t ASCEND950_S1_BASE_SIZE = 4; // Ascend 950 S1轴基本块的大小
constexpr uint32_t ASCEND950_S2_BASE_SIZE = 128; // Ascend 950 S2轴基本块的大小
uint32_t workspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize();
// 主流程需Workspace大小
platform_ascendc::SocVersion socVersion_ = ascendcPlatform.GetSocVersion();
if (socVersion_ == platform_ascendc::SocVersion::ASCEND950) {
constexpr uint32_t s1Base = ASCEND950_S1_BASE_SIZE;
constexpr uint32_t s2Base = ASCEND950_S2_BASE_SIZE;
workspaceSize += s1Base * ((tilingInfo->s2Size + s2Base - 1) / s2Base) * s2Base * sizeof(uint32_t) * aicNum;
} else {
uint32_t mm1ResSize = M_BASE_SIZE * S2_BASE_SIZE;
workspaceSize += mm1ResSize * MM1_RES_ELEM_SIZE * DOUBLE_BUFFER * aicNum;
// Decode流程(LD)需要Workspace大小
// 临时存储Decode中间结果大小: 2(头/尾)*8(s1Base)*2(idx/value)*2048(K)*sizeof(int32)*24=6M
workspaceSize += V1_DECODE_DATA_NUM * S1_BASE_SIZE * V1_RES_ELEM_TYPE * TOPK_MAX_SIZE * V1_RES_ELEM_SIZE * aicNum;
// 临时存储Decode中间参数信息大小: 2(头/尾)*8(s1Base)*16(paramNum)*sizeof(int64_t)*24=48k
workspaceSize += V1_DECODE_DATA_NUM * S1_BASE_SIZE * V1_DECODE_PARAM_NUM * V1_DECODE_PARAM_ELEM_SIZE * aicNum;
}
size_t *workSpaces = context_->GetWorkspaceSizes(1);
workSpaces[0] = workspaceSize;
// -------------set tilingdata-----------------
tilingData_.set_bSize(tilingInfo->bSize);
tilingData_.set_s2Size(tilingInfo->s2Size);
tilingData_.set_s1Size(tilingInfo->s1Size);
tilingData_.set_sparseCount(tilingInfo->sparseCount);
tilingData_.set_gSize(tilingInfo->gSize);
tilingData_.set_blockSize(tilingInfo->blockSize);
tilingData_.set_maxBlockNumPerBatch(tilingInfo->maxBlockNumPerBatch);
tilingData_.set_sparseMode(tilingInfo->sparseMode);
tilingData_.set_cmpRatio(tilingInfo->cmpRatio);
tilingData_.set_returnValues(tilingInfo->returnValues);
tilingData_.set_usedCoreNum(blockDim);
tilingData_.set_batchSupperFlag(tilingInfo->batchSupperFlag);
tilingData_.set_stride(tilingInfo->stride);
tilingData_.set_scaleStride(tilingInfo->scaleStride);
tilingData_.SaveToBuffer(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity());
context_->GetRawTilingData()->SetDataSize(tilingData_.GetDataSize());
// -------------set tilingkey-----------------
// DT_Q, DT_KV, DT_OUT, PAGE_ATTENTION, FLASH_DECODE, LAYOUT_T, KV_LAYOUT_T
uint32_t inputQType = static_cast<uint32_t>(tilingInfo->inputQType);
uint32_t inputKType = static_cast<uint32_t>(tilingInfo->inputKType);
uint32_t outputType = static_cast<uint32_t>(tilingInfo->outputType);
uint32_t pageAttentionFlag = static_cast<uint32_t>(tilingInfo->pageAttentionFlag);
uint32_t inputQLayout = static_cast<uint32_t>(tilingInfo->inputQLayout);
uint32_t inputKLayout = static_cast<uint32_t>(tilingInfo->inputKLayout);
uint32_t tilingKey =
GET_TPL_TILING_KEY(inputQType, inputKType, outputType, pageAttentionFlag, inputQLayout, inputKLayout);
context_->SetTilingKey(tilingKey);
context_->SetScheduleMode(1);
return ge::GRAPH_SUCCESS;
}
// --------------------------Tiling函数定义---------------------------
ge::graphStatus TilingForVllmQuantLightningIndexer(gert::TilingContext *context)
{
OP_CHECK_IF(context == nullptr, OP_LOGE("VllmQuantLightningIndexer", "Tiling context is null."),
return ge::GRAPH_FAILED);
QLITilingInfo QLIInfo;
QLIInfoParser QLIInfoParser(context);
if (QLIInfoParser.ParseAndCheck(QLIInfo) != ge::GRAPH_SUCCESS) {
return ge::GRAPH_FAILED;
}
VllmQuantLightningIndexerTiling QLITiling(context);
return QLITiling.DoTiling(&QLIInfo);
}
// --------------------------Tiling及函数TilingPrepare函数注册--------
IMPL_OP_OPTILING(VllmQuantLightningIndexer)
.Tiling(TilingForVllmQuantLightningIndexer)
.TilingParse<QLICompileInfo>(TilingPrepareForVllmQuantLightningIndexer);
} // namespace optiling

View File

@@ -0,0 +1,263 @@
/**
 * 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_tiling.h
* \brief
*/
#ifndef QUANT_LIGHTNING_INDEXER_TILING_H
#define QUANT_LIGHTNING_INDEXER_TILING_H
#include "err/ops_err.h"
#include "exe_graph/runtime/tiling_context.h"
#include "platform/platform_info.h"
#include "register/op_def_registry.h"
#include "register/tilingdata_base.h"
#include "tiling/platform/platform_ascendc.h"
#include "tiling/tiling_api.h"
namespace optiling {
// ------------------公共定义--------------------------
struct TilingRequiredParaInfo {
const gert::CompileTimeTensorDesc *desc;
const gert::StorageShape *shape;
};
struct TilingOptionalParaInfo {
const gert::CompileTimeTensorDesc *desc;
const gert::Tensor *tensor;
};
enum class DataLayout : uint32_t {
BSND = 0,
TND = 1,
PA_BSND = 2
};
// ------------------算子原型索引常量定义----------------
// Inputs Index
constexpr uint32_t QUERY_INDEX = 0;
constexpr uint32_t KEY_INDEX = 1;
constexpr uint32_t WEIGTHS_INDEX = 2;
constexpr uint32_t QUERY_DEQUANT_SCALE_INDEX = 3;
constexpr uint32_t KEY_DEQUANT_SCALE_INDEX = 4;
constexpr uint32_t ACTUAL_SEQ_Q_INDEX = 5;
constexpr uint32_t ACTUAL_SEQ_K_INDEX = 6;
constexpr uint32_t BLOCK_TABLE_INDEX = 7;
constexpr uint32_t METADATA_INDEX = 8;
constexpr uint32_t vllm_quant_lightning_indexer = 0;
// Attributes Index
constexpr uint32_t ATTR_QUERY_QUANT_MODE_INDEX = 0;
constexpr uint32_t ATTR_KEY_QUANT_MODE_INDEX = 1;
constexpr uint32_t ATTR_QUERY_LAYOUT_INDEX = 2;
constexpr uint32_t ATTR_KEY_LAYOUT_INDEX = 3;
constexpr uint32_t ATTR_SPARSE_COUNT_INDEX = 4;
constexpr uint32_t ATTR_SPARSE_MODE_INDEX = 5;
constexpr uint32_t ATTR_PRE_TOKENS_INDEX = 6;
constexpr uint32_t ATTR_NEXT_TOKENS_INDEX = 7;
constexpr uint32_t ATTR_CMP_RATIO_INDEX = 8;
constexpr uint32_t ATTR_RETURN_VALUES_INDEX = 9;
constexpr uint32_t ATTR_STRIDE_INDEX = 10;
constexpr uint32_t ATTR_SCALE_STRIDE_INDEX = 11;
// Dim Index
constexpr uint32_t DIM_IDX_ZERO = 0;
constexpr uint32_t DIM_IDX_ONE = 1;
constexpr uint32_t DIM_IDX_TWO = 2;
constexpr uint32_t DIM_IDX_THREE = 3;
// Dim Num
constexpr uint32_t DIM_NUM_TWO = 2;
constexpr uint32_t DIM_NUM_THREE = 3;
constexpr uint32_t DIM_NUM_FOUR = 4;
// 入参限制常量
constexpr uint32_t HEAD_DIM_LIMIT = 128;
constexpr uint32_t SPARSE_LIMIT = 2048;
constexpr uint32_t G_SIZE_LIMIT = 64;
constexpr uint32_t BLOCK_SIZE_LIMIT = 1024;
constexpr uint32_t BLOCK_SIZE_FACTOR = 16;
constexpr uint32_t SPARSE_MODE_LOWER = 3;
constexpr uint32_t METADATA_LIMIT = 1024;
// -----------算子TilingData定义---------------
BEGIN_TILING_DATA_DEF(QLITilingData)
TILING_DATA_FIELD_DEF(uint32_t, bSize)
TILING_DATA_FIELD_DEF(uint32_t, n2Size)
TILING_DATA_FIELD_DEF(uint32_t, gSize)
TILING_DATA_FIELD_DEF(uint32_t, s1Size)
TILING_DATA_FIELD_DEF(uint32_t, s2Size)
TILING_DATA_FIELD_DEF(uint32_t, sparseCount)
TILING_DATA_FIELD_DEF(uint32_t, usedCoreNum)
TILING_DATA_FIELD_DEF(uint32_t, blockSize)
TILING_DATA_FIELD_DEF(uint32_t, maxBlockNumPerBatch)
TILING_DATA_FIELD_DEF(uint32_t, sparseMode)
TILING_DATA_FIELD_DEF(uint32_t, cmpRatio)
TILING_DATA_FIELD_DEF(uint32_t, batchSupperFlag)
TILING_DATA_FIELD_DEF(uint32_t, returnValues)
TILING_DATA_FIELD_DEF(int64_t, stride)
TILING_DATA_FIELD_DEF(int64_t, scaleStride)
END_TILING_DATA_DEF
REGISTER_TILING_DATA_CLASS(VllmQuantLightningIndexer, QLITilingData)
// -----------算子CompileInfo定义-------------------
struct QLICompileInfo {};
// -----------算子Tiling入参结构体定义---------------
struct QLIParaInfo {
TilingRequiredParaInfo query = {nullptr, nullptr};
TilingRequiredParaInfo key = {nullptr, nullptr};
TilingRequiredParaInfo weights = {nullptr, nullptr};
TilingRequiredParaInfo query_dequant_scale = {nullptr, nullptr};
TilingRequiredParaInfo key_dequant_scale = {nullptr, nullptr};
TilingOptionalParaInfo actualSeqLengthsQ = {nullptr, nullptr};
TilingOptionalParaInfo actualSeqLengthsK = {nullptr, nullptr};
TilingOptionalParaInfo blockTable = {nullptr, nullptr};
TilingOptionalParaInfo metadata = {nullptr, nullptr};
TilingRequiredParaInfo attenOut = {nullptr, nullptr};
const int64_t *queryQuantMode = nullptr;
const int64_t *keyQuantMode = nullptr;
const char *layOutQuery = nullptr;
const char *layOutKey = nullptr;
const int64_t *blockSize = nullptr;
const int64_t *sparseMode = nullptr;
const int64_t *sparseCount = nullptr;
const int64_t *preTokens = nullptr;
const int64_t *nextTokens = nullptr;
const int64_t *cmpRatio = nullptr;
const int32_t *batchSupperFlag = nullptr;
const bool *returnValues = nullptr;
const int64_t *stride = nullptr;
const int64_t *scaleStride = nullptr;
};
// -----------算子Tiling入参信息类---------------
class QLITilingInfo {
public:
const char *opName = nullptr;
fe::PlatFormInfos *platformInfo = nullptr;
QLIParaInfo opParamInfo;
// Base Param
platform_ascendc::SocVersion socVersion = platform_ascendc::SocVersion::ASCEND910B;
uint32_t bSize = 0;
uint32_t n1Size = 0;
uint32_t n2Size = 0;
uint32_t s1Size = 0;
int64_t s2Size = 0;
uint32_t qkHeadDim = 0;
uint32_t gSize = 0;
// PageAttention
bool pageAttentionFlag = false;
int32_t blockSize = 0;
uint32_t maxBlockNumPerBatch = 0;
// Mask
int32_t sparseMode = 0;
// Others Flag
uint32_t sparseCount = 0;
int64_t preTokens = 0;
int64_t nextTokens = 0;
uint32_t cmpRatio = 1;
bool batchSupperFlag = false;
bool returnValues = false;
int64_t stride = 1;
int64_t scaleStride = 1;
// DType
ge::DataType inputQType = ge::DT_FLOAT16;
ge::DataType inputKType = ge::DT_FLOAT16;
ge::DataType outputType = ge::DT_INT32;
// Layout
DataLayout inputQLayout = DataLayout::BSND;
DataLayout inputKLayout = DataLayout::PA_BSND;
};
// -----------算子Tiling入参信息解析及Check类---------------
class QLIInfoParser {
public:
explicit QLIInfoParser(gert::TilingContext *context) : context_(context) {}
~QLIInfoParser() = default;
ge::graphStatus CheckRequiredInOutExistence() const;
ge::graphStatus CheckRequiredAttrExistence() const;
ge::graphStatus CheckRequiredParaExistence() const;
ge::graphStatus GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
const std::string &actualSeqLenName) const;
ge::graphStatus GetOpName();
ge::graphStatus GetNpuInfo();
void GetOptionalInputParaInfo();
void GetInputParaInfo();
void GetOutputParaInfo();
ge::graphStatus GetAttrParaInfo();
ge::graphStatus CheckAttrParaInfo();
ge::graphStatus GetOpParaInfo();
ge::graphStatus ValidateInputShapesMatch();
ge::graphStatus CheckScaleShape();
ge::graphStatus GetAndCheckInOutDataType();
ge::graphStatus GetBatchSize();
ge::graphStatus GetHeadDim();
ge::graphStatus GetS1Size();
ge::graphStatus GetAndCheckOptionalInput();
ge::graphStatus CheckShapeDim();
ge::graphStatus GetAndCheckBlockSize();
ge::graphStatus GetS2SizeForPageAttention();
ge::graphStatus GetS2SizeForBatchContinuous();
ge::graphStatus GetS2Size();
ge::graphStatus GetQueryKeyAndOutLayout();
ge::graphStatus GetN1Size();
ge::graphStatus GetAndCheckN2Size();
ge::graphStatus GetGSize();
ge::graphStatus GetAttenMaskInfo();
ge::graphStatus GetActualSeqInfo();
void GenerateInfo(QLITilingInfo &QLIInfo);
ge::graphStatus ParseAndCheck(QLITilingInfo &QLIInfo);
public:
gert::TilingContext *context_ = nullptr;
const char *opName_;
fe::PlatFormInfos *platformInfo_;
QLIParaInfo opParamInfo_;
// BaseParams
uint32_t bSize_ = 0;
uint32_t n1Size_ = 0;
uint32_t n2Size_ = 0;
uint32_t gSize_ = 0;
uint32_t s1Size_ = 0;
int64_t s2Size_ = 0;
uint32_t headDim_ = 0;
bool batchSupperFlag_ = false;
// Layout
DataLayout qLayout_ = DataLayout::BSND;
DataLayout kLayout_ = DataLayout::PA_BSND;
// PageAttention
uint32_t maxBlockNumPerBatch_ = 0;
int32_t blockSize_ = 0;
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
ge::DataType inputQType_ = ge::DT_FLOAT16;
ge::DataType inputKType_ = ge::DT_FLOAT16;
ge::DataType weightsType_ = ge::DT_FLOAT16;
ge::DataType inputQueryScaleType_ = ge::DT_FLOAT16;
ge::DataType inputKeyScaleType_ = ge::DT_FLOAT16;
ge::DataType blockTableType_ = ge::DT_FLOAT16;
ge::DataType inputKRopeType_ = ge::DT_FLOAT16;
ge::DataType outputType_ = ge::DT_FLOAT16;
};
// ---------------算子Tiling类---------------
class VllmQuantLightningIndexerTiling {
public:
explicit VllmQuantLightningIndexerTiling(gert::TilingContext *context) : context_(context) {};
ge::graphStatus DoTiling(QLITilingInfo *tilingInfo);
private:
gert::TilingContext *context_ = nullptr;
QLITilingData tilingData_;
};
} // namespace optiling
#endif // QUANT_LIGHTNING_INDEXER_TILING_H