@@ -0,0 +1,41 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# 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
|
||||
kv_quant_sparse_attn_sharedkv_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME KvQuantSparseAttnSharedkv
|
||||
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 kv_quant_sparse_attn_sharedkv ACLNNTYPE aclnn)
|
||||
endif()
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_tiling_modules()
|
||||
target_sources(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
kv_quant_sparse_attn_sharedkv_check_consistancy.cpp
|
||||
kv_quant_sparse_attn_sharedkv_check_existance.cpp
|
||||
kv_quant_sparse_attn_sharedkv_check_feature.cpp
|
||||
kv_quant_sparse_attn_sharedkv_check_single_para.cpp
|
||||
kv_quant_sparse_attn_sharedkv_check.cpp
|
||||
kv_quant_sparse_attn_sharedkv_tiling.cpp
|
||||
)
|
||||
endif()
|
||||
@@ -0,0 +1,168 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
std::string SASDataTypeToSerialString(ge::DataType type)
|
||||
{
|
||||
const auto it = DATATYPE_TO_STRING_MAP.find(type);
|
||||
if (it != DATATYPE_TO_STRING_MAP.end()) {
|
||||
return it->second;
|
||||
} else {
|
||||
OP_LOGE("KvQuantSparseAttnSharedkv ", "datatype %d not support", type);
|
||||
return "UNDEFINED";
|
||||
}
|
||||
}
|
||||
|
||||
std::string KvQuantSASLayoutToSerialString(SASLayout layout)
|
||||
{
|
||||
switch (layout) {
|
||||
case SASLayout::BSND: return "BSND";
|
||||
case SASLayout::TND: return "TND";
|
||||
case SASLayout::PA_ND: return "PA_ND";
|
||||
default: return "UNKNOWN";
|
||||
}
|
||||
}
|
||||
|
||||
std::string GetShapeStr(gert::Shape shape)
|
||||
{
|
||||
std::ostringstream oss;
|
||||
oss << "[";
|
||||
if (shape.GetDimNum() > 0) {
|
||||
for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) {
|
||||
oss << shape.GetDim(i) << ", ";
|
||||
}
|
||||
oss << shape.GetDim(shape.GetDimNum() - 1);
|
||||
}
|
||||
oss << "]";
|
||||
return oss.str();
|
||||
}
|
||||
|
||||
bool KvQuantSASTilingCheck::HasAxis(const SASAxis &axis, const SASLayout &layout, const gert::Shape &shape) const
|
||||
{
|
||||
const auto& layoutIt = SAS_LAYOUT_AXIS_MAP.find(layout);
|
||||
if (layoutIt == SAS_LAYOUT_AXIS_MAP.end()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const std::vector<SASAxis>& axes = layoutIt->second;
|
||||
const auto& axisIt = std::find(axes.begin(), axes.end(), axis);
|
||||
if (axisIt == axes.end()) {
|
||||
return false;
|
||||
}
|
||||
const auto& dimIt = SAS_LAYOUT_DIM_MAP.find(layout);
|
||||
if (dimIt == SAS_LAYOUT_DIM_MAP.end() || dimIt->second != shape.GetDimNum()) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t KvQuantSASTilingCheck::GetAxisIdx(const SASAxis &axis, const SASLayout &layout) const
|
||||
{
|
||||
const std::vector<SASAxis>& axes = SAS_LAYOUT_AXIS_MAP.find(layout)->second;
|
||||
const auto& axisIt = std::find(axes.begin(), axes.end(), axis);
|
||||
return std::distance(axes.begin(), axisIt);
|
||||
}
|
||||
|
||||
uint32_t KvQuantSASTilingCheck::GetAxisNum(const gert::Shape &shape, const SASAxis &axis,const SASLayout &layout) const
|
||||
{
|
||||
return HasAxis(axis, layout, shape) ? shape.GetDim(GetAxisIdx(axis, layout)) : invalidDimValue_;
|
||||
}
|
||||
|
||||
void KvQuantSASTilingCheck::Init()
|
||||
{
|
||||
opName_ = sasInfo_.opName;
|
||||
platformInfo_ = sasInfo_.platformInfo;
|
||||
opParamInfo_ = sasInfo_.opParamInfo;
|
||||
socVersion_ = sasInfo_.socVersion;
|
||||
|
||||
bSize_ = sasInfo_.bSize;
|
||||
bSize_ = opParamInfo_.oriBlockTable.tensor->GetShape().GetStorageShape().GetDim(0);
|
||||
n1Size_ = sasInfo_.n1Size;
|
||||
n2Size_ = sasInfo_.n2Size;
|
||||
s1Size_ = sasInfo_.s1Size;
|
||||
s2Size_ = sasInfo_.s2Size;
|
||||
gSize_ = sasInfo_.gSize;
|
||||
qkHeadDim_ = sasInfo_.qkHeadDim;
|
||||
qTSize_ = sasInfo_.qTSize;
|
||||
|
||||
dSize_ = sasInfo_.dSize;
|
||||
dSizeV_ = sasInfo_.dSizeV;
|
||||
if (opParamInfo_.oriKv.tensor != nullptr) {
|
||||
dSizeOriKvInput_ = GetAxisNum(opParamInfo_.oriKv.tensor->GetStorageShape(), SASAxis::D, kvLayout_);
|
||||
}
|
||||
if (opParamInfo_.cmpKv.tensor != nullptr) {
|
||||
dSizeCmpKvInput_ = GetAxisNum(opParamInfo_.cmpKv.tensor->GetStorageShape(), SASAxis::D, kvLayout_);
|
||||
}
|
||||
|
||||
actualLenDimsQ_ = sasInfo_.actualLenDimsQ;
|
||||
maxActualseq_ = sasInfo_.maxActualseq;
|
||||
|
||||
ropeHeadDim_ = sasInfo_.ropeHeadDim;
|
||||
oriMaxBlockNumPerBatch_ = sasInfo_.oriMaxBlockNumPerBatch;
|
||||
cmpMaxBlockNumPerBatch_ = sasInfo_.cmpMaxBlockNumPerBatch;
|
||||
|
||||
oriBlockSize_ = sasInfo_.oriBlockSize;
|
||||
cmpBlockSize_ = sasInfo_.cmpBlockSize;
|
||||
|
||||
sparseBlockCount_ = sasInfo_.sparseBlockCount;
|
||||
sparseBlockSize_ = sasInfo_.sparseBlockSize;
|
||||
|
||||
tileSize_ = sasInfo_.tileSize;
|
||||
|
||||
cmpRatio_ = sasInfo_.cmpRatio;
|
||||
oriWinLeft_ = sasInfo_.oriWinLeft;
|
||||
oriWinRight_ = sasInfo_.oriWinRight;
|
||||
|
||||
oriMaskMode_ = sasInfo_.oriMaskMode;
|
||||
cmpMaskMode_ = sasInfo_.cmpMaskMode;
|
||||
|
||||
qType_ = sasInfo_.qType;
|
||||
oriKvType_ = sasInfo_.oriKvType;
|
||||
cmpKvType_ = sasInfo_.cmpKvType;
|
||||
outputType_ = sasInfo_.outputType;
|
||||
|
||||
qLayout_ = sasInfo_.qLayout;
|
||||
kvLayout_ = sasInfo_.kvLayout;
|
||||
outLayout_ = sasInfo_.outLayout;
|
||||
|
||||
if (opParamInfo_.cmpKv.tensor == nullptr) {
|
||||
perfMode_ = SASTemplateMode::SWA_TEMPLATE_MODE;
|
||||
} else if (opParamInfo_.cmpSparseIndices.tensor != nullptr) {
|
||||
perfMode_ = SASTemplateMode::SCFA_TEMPLATE_MODE;
|
||||
} else {
|
||||
perfMode_ = SASTemplateMode::CFA_TEMPLATE_MODE;
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::Process()
|
||||
{
|
||||
Init();
|
||||
if (CheckSinglePara() != ge::GRAPH_SUCCESS ||
|
||||
CheckParaExistence() != ge::GRAPH_SUCCESS ||
|
||||
CheckFeature() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,579 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_CHECK_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_CHECK_H
|
||||
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <exe_graph/runtime/tiling_context.h>
|
||||
#include <tiling/platform/platform_ascendc.h>
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "log/log.h"
|
||||
#include "log/error_code.h"
|
||||
#include "err/ops_err.h"
|
||||
#include "platform/platform_info.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
const std::string ORI_BLOCK_TABLE_NAME = "ori_block_table";
|
||||
const std::string CMP_BLOCK_TABLE_NAME = "cmp_block_table";
|
||||
const std::string SINKS_NAME = "sinks";
|
||||
|
||||
const std::string QUERY_NAME = "query";
|
||||
const std::string KEY_NAME = "key";
|
||||
const std::string VALUE_NAME = "value";
|
||||
|
||||
const std::string ORI_KV_NAME = "ori_kv";
|
||||
const std::string CMP_KV_NAME = "cmp_kv";
|
||||
const std::string ORI_SPARSE_INDICES_NAME = "ori_sparse_indices";
|
||||
const std::string CMP_SPARSE_INDICES_NAME = "cmp_sparse_indices";
|
||||
const std::string ATTEN_OUT_NAME = "attention_out";
|
||||
|
||||
const std::string CU_SEQLENS_Q_NAME = "cu_seqlens_q";
|
||||
const std::string SEQUSED_KV_NAME = "seqused_kv";
|
||||
|
||||
// // ------------------公共定义--------------------------
|
||||
struct SASTilingRequiredParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::StorageShape *shape;
|
||||
};
|
||||
|
||||
struct SASTilingOptionalParaInfo {
|
||||
const gert::CompileTimeTensorDesc *desc;
|
||||
const gert::Tensor *tensor;
|
||||
};
|
||||
|
||||
enum class SASLayout : uint32_t {
|
||||
BSND = 0,
|
||||
TND = 1,
|
||||
PA_ND = 2
|
||||
};
|
||||
|
||||
enum class SASAxis : uint32_t {
|
||||
B = 0,
|
||||
S = 1,
|
||||
N = 2,
|
||||
D = 3,
|
||||
K = 3, // sparse_indices的K和key的D枚举值相同,表达相同位置, 最后一维
|
||||
T = 5,
|
||||
Bn = 6, // block number
|
||||
Bs = 7 // block size
|
||||
};
|
||||
|
||||
enum class SASTemplateMode : uint32_t {
|
||||
SWA_TEMPLATE_MODE = 0,
|
||||
CFA_TEMPLATE_MODE = 1,
|
||||
SCFA_TEMPLATE_MODE = 2
|
||||
};
|
||||
|
||||
enum class KvStorageMode : uint32_t {
|
||||
BATCH_CONTINUOUS = 0,
|
||||
TENSOR_LIST = 1,
|
||||
PAGE_ATTENTION = 2
|
||||
};
|
||||
|
||||
struct KvQuantSASTilingShapeCompareParam {
|
||||
int64_t B = 1;
|
||||
int64_t S = 1;
|
||||
int64_t N = 1;
|
||||
int64_t D = 1;
|
||||
int64_t T = 1;
|
||||
// PA
|
||||
int64_t Bs = 1;
|
||||
int64_t Bn = 1;
|
||||
};
|
||||
|
||||
// ------------------算子原型索引常量定义----------------
|
||||
// Inputs Index
|
||||
constexpr uint32_t Q_INDEX = 0;
|
||||
constexpr uint32_t ORI_KV_INDEX = 1;
|
||||
constexpr uint32_t CMP_KV_INDEX = 2;
|
||||
constexpr uint32_t ORI_SPARSE_INDICES_INDEX = 3;
|
||||
constexpr uint32_t CMP_SPARSE_INDICES_INDEX = 4;
|
||||
constexpr uint32_t ORI_BLOCK_TABLE_INDEX = 5;
|
||||
constexpr uint32_t CMP_BLOCK_TABLE_INDEX = 6;
|
||||
constexpr uint32_t CU_SEQLENS_Q_INDEX = 7;
|
||||
constexpr uint32_t CU_SEQLENS_ORI_KV_INDEX = 8;
|
||||
constexpr uint32_t CU_SEQLENS_CMP_KV_INDEX = 9;
|
||||
constexpr uint32_t SEQUSED_Q_INDEX = 10;
|
||||
constexpr uint32_t SEQUSED_KV_INDEX = 11;
|
||||
constexpr uint32_t SINKS_INDEX = 12;
|
||||
constexpr uint32_t METADATA_INDEX = 13;
|
||||
// Outputs Index
|
||||
constexpr uint32_t ATTN_OUT_INDEX = 0;
|
||||
|
||||
// Attributes Index
|
||||
constexpr uint32_t ATTR_KV_QUANT_SCALE_INDEX = 0;
|
||||
constexpr uint32_t ATTR_TILE_SIZE_INDEX = 1;
|
||||
constexpr uint32_t ATTR_ROPE_HEAD_DIM_INDEX = 2;
|
||||
constexpr uint32_t ATTR_SOTFMAX_SCALE_INDEX = 3;
|
||||
constexpr uint32_t ATTR_CMP_RATIO_INDEX = 4;
|
||||
constexpr uint32_t ATTR_ORI_MASK_MODE_INDEX = 5;
|
||||
constexpr uint32_t ATTR_CMP_MASK_MODE_INDEX = 6;
|
||||
constexpr uint32_t ATTR_ORI_WIN_LEFT_INDEX = 7;
|
||||
constexpr uint32_t ATTR_ORI_WIN_RIGHT_INDEX = 8;
|
||||
constexpr uint32_t ATTR_LAYOUT_Q_INDEX = 9;
|
||||
constexpr uint32_t ATTR_LAYOUT_KV_INDEX = 10;
|
||||
constexpr uint32_t ATTR_ORIKV_STRIDE_INDEX = 11;
|
||||
constexpr uint32_t ATTR_CMPKV_STRIDE_INDEX = 12;
|
||||
|
||||
// Dim Index
|
||||
constexpr uint32_t DIM_IDX_ONE = 1;
|
||||
constexpr uint32_t DIM_IDX_TWO = 2;
|
||||
constexpr uint32_t DIM_IDX_THREE = 3;
|
||||
constexpr uint32_t DIM_IDX_FOUR = 4;
|
||||
|
||||
// Dim Num
|
||||
constexpr uint32_t DIM_NUM_ONE = 1;
|
||||
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 SPARSE_MODE_LOWER = 3;
|
||||
constexpr uint32_t MAX_BLOCK_SIZE = 1024;
|
||||
constexpr uint32_t COPYND2NZ_SRC_STRIDE_LIMITATION = 65535;
|
||||
constexpr uint32_t NUM_BYTES_FLOAT = 4;
|
||||
constexpr uint32_t NUM_BYTES_FLOAT16 = 2;
|
||||
constexpr uint32_t NUM_BYTES_BF16 = 2;
|
||||
constexpr uint32_t BYTE_BLOCK = 32;
|
||||
const uint32_t QSFA_MAX_AIC_CORE_NUM = 26; // 25 + 1 保证数组8字节对齐
|
||||
|
||||
const std::map<ge::DataType, std::string> DATATYPE_TO_STRING_MAP = {
|
||||
{ge::DT_UNDEFINED, "DT_UNDEFINED"}, // Used to indicate a DataType field has not been set.
|
||||
{ge::DT_FLOAT, "DT_FLOAT"}, // float type
|
||||
{ge::DT_FLOAT16, "DT_FLOAT16"}, // fp16 type
|
||||
{ge::DT_INT8, "DT_INT8"}, // int8 type
|
||||
{ge::DT_INT16, "DT_INT16"}, // int16 type
|
||||
{ge::DT_UINT16, "DT_UINT16"}, // uint16 type
|
||||
{ge::DT_UINT8, "DT_UINT8"}, // uint8 type
|
||||
{ge::DT_INT32, "DT_INT32"}, // uint32 type
|
||||
{ge::DT_INT64, "DT_INT64"}, // int64 type
|
||||
{ge::DT_UINT32, "DT_UINT32"}, // unsigned int32
|
||||
{ge::DT_UINT64, "DT_UINT64"}, // unsigned int64
|
||||
{ge::DT_BOOL, "DT_BOOL"}, // bool type
|
||||
{ge::DT_DOUBLE, "DT_DOUBLE"}, // double type
|
||||
{ge::DT_DUAL, "DT_DUAL"}, // dual output type
|
||||
{ge::DT_DUAL_SUB_INT8, "DT_DUAL_SUB_INT8"}, // dual output int8 type
|
||||
{ge::DT_DUAL_SUB_UINT8, "DT_DUAL_SUB_UINT8"}, // dual output uint8 type
|
||||
{ge::DT_COMPLEX32, "DT_COMPLEX32"}, // complex32 type
|
||||
{ge::DT_COMPLEX64, "DT_COMPLEX64"}, // complex64 type
|
||||
{ge::DT_COMPLEX128, "DT_COMPLEX128"}, // complex128 type
|
||||
{ge::DT_QINT8, "DT_QINT8"}, // qint8 type
|
||||
{ge::DT_QINT16, "DT_QINT16"}, // qint16 type
|
||||
{ge::DT_QINT32, "DT_QINT32"}, // qint32 type
|
||||
{ge::DT_QUINT8, "DT_QUINT8"}, // quint8 type
|
||||
{ge::DT_QUINT16, "DT_QUINT16"}, // quint16 type
|
||||
{ge::DT_RESOURCE, "DT_RESOURCE"}, // resource type
|
||||
{ge::DT_STRING_REF, "DT_STRING_REF"}, // string ref type
|
||||
{ge::DT_STRING, "DT_STRING"}, // string type
|
||||
{ge::DT_VARIANT, "DT_VARIANT"}, // dt_variant type
|
||||
{ge::DT_BF16, "DT_BFLOAT16"}, // dt_bfloat16 type
|
||||
{ge::DT_INT4, "DT_INT4"}, // dt_variant type
|
||||
{ge::DT_UINT1, "DT_UINT1"}, // dt_variant type
|
||||
{ge::DT_INT2, "DT_INT2"}, // dt_variant type
|
||||
{ge::DT_UINT2, "DT_UINT2"} // dt_variant type
|
||||
};
|
||||
|
||||
const std::map<SASLayout, std::vector<SASAxis>> SAS_LAYOUT_AXIS_MAP = {
|
||||
{SASLayout::BSND, {SASAxis::B, SASAxis::S, SASAxis::N, SASAxis::D}},
|
||||
{SASLayout::TND, {SASAxis::T, SASAxis::N, SASAxis::D}},
|
||||
{SASLayout::PA_ND, {SASAxis::Bn, SASAxis::Bs, SASAxis::N, SASAxis::D}},
|
||||
};
|
||||
|
||||
const std::map<SASLayout, size_t> SAS_LAYOUT_DIM_MAP = {
|
||||
{SASLayout::BSND, DIM_NUM_FOUR},
|
||||
{SASLayout::TND, DIM_NUM_THREE},
|
||||
{SASLayout::PA_ND, DIM_NUM_FOUR},
|
||||
};
|
||||
|
||||
std::string SASDataTypeToSerialString(ge::DataType type);
|
||||
std::string KvQuantSASLayoutToSerialString(SASLayout layout);
|
||||
std::string GetShapeStr(gert::Shape shape);
|
||||
|
||||
// -----------算子Tiling入参信息解析及Check类---------------
|
||||
|
||||
struct KvQuantSASParaInfo {
|
||||
SASTilingRequiredParaInfo q = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo oriKv = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cmpKv = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo oriSparseIndices = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cmpSparseIndices = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo oriBlockTable = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cmpBlockTable = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cuSeqLensQ = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cuSeqLensOriKv = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo cuSeqLensCmpKv = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo seqUsedQ = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo sequsedKv = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo sinks = {nullptr, nullptr};
|
||||
SASTilingOptionalParaInfo metadata = {nullptr, nullptr};
|
||||
SASTilingRequiredParaInfo attnOut = {nullptr, nullptr};
|
||||
|
||||
const int64_t *kvQuantMode = nullptr;
|
||||
const int64_t *tileSize = nullptr;
|
||||
const int64_t *ropeHeadDim = nullptr;
|
||||
const float *softmaxScale = nullptr;
|
||||
const int64_t *oriKvStride = nullptr;
|
||||
const int64_t *cmpKvStride = nullptr;
|
||||
const int64_t *cmpRatio = nullptr;
|
||||
const uint32_t *oriMaskMode = nullptr;
|
||||
const uint32_t *cmpMaskMode = nullptr;
|
||||
const int64_t *oriWinLeft = nullptr;
|
||||
const int64_t *oriWinRight = nullptr;
|
||||
const char *layoutQ = nullptr;
|
||||
const char *layoutKv = nullptr;
|
||||
};
|
||||
|
||||
// -----------算子Tiling入参信息类---------------
|
||||
class KvQuantSASTilingInfo {
|
||||
public:
|
||||
const char *opName = nullptr;
|
||||
fe::PlatFormInfos *platformInfo = nullptr;
|
||||
KvQuantSASParaInfo 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 gSize = 0;
|
||||
uint32_t qkHeadDim = 0;
|
||||
uint32_t qTSize = 0; // 仅TND时生效
|
||||
|
||||
uint32_t actualLenDimsQ = 0;
|
||||
uint32_t maxActualseq = 0;
|
||||
bool actualSeqLenFlag = false;
|
||||
bool isSameSeqAllKVTensor = true;
|
||||
bool isSameActualseq = true;
|
||||
uint32_t actualLenDimsKV = 0;
|
||||
|
||||
int64_t kvQuantMode = 0;
|
||||
int64_t tileSize = 0;
|
||||
int64_t ropeHeadDim = 0;
|
||||
uint32_t dSize = 0;
|
||||
uint32_t dSizeV = 0;
|
||||
uint32_t dSizeVInput = 0;
|
||||
float softmaxScale = 0;
|
||||
int64_t oriKvStride = 0;
|
||||
int64_t cmpKvStride = 0;
|
||||
int64_t cmpRatio = 0;
|
||||
uint64_t oriMaskMode = 0;
|
||||
uint64_t cmpMaskMode = 0;
|
||||
int64_t oriWinLeft = 0;
|
||||
int64_t oriWinRight = 0;
|
||||
int64_t sparseBlockSize = 0;
|
||||
int64_t sparseBlockCount = 0;
|
||||
// Mask
|
||||
int32_t sparseMode = 0;
|
||||
// Others Flag
|
||||
uint32_t sparseCount = 0;
|
||||
|
||||
// PageAttention
|
||||
uint32_t blockTypeSize = 0;
|
||||
uint32_t oriMaxBlockNumPerBatch = 0;
|
||||
int32_t oriBlockSize = 0;
|
||||
int32_t cmpBlockSize = 0;
|
||||
uint32_t cmpMaxBlockNumPerBatch = 0;
|
||||
uint32_t totalBlockNum = 0;
|
||||
|
||||
// DType
|
||||
ge::DataType qType = ge::DT_FLOAT16;
|
||||
ge::DataType oriKvType = ge::DT_FLOAT16;
|
||||
ge::DataType cmpKvType = ge::DT_FLOAT16;
|
||||
ge::DataType outputType = ge::DT_FLOAT16;
|
||||
|
||||
// Layout
|
||||
SASLayout qLayout = SASLayout::BSND;
|
||||
SASLayout kvLayout = SASLayout::PA_ND;
|
||||
SASLayout outLayout = SASLayout::BSND;
|
||||
};
|
||||
|
||||
class KvQuantSASInfoParser {
|
||||
public:
|
||||
explicit KvQuantSASInfoParser(gert::TilingContext *context) : context_(context) {}
|
||||
~KvQuantSASInfoParser() = default;
|
||||
|
||||
ge::graphStatus CheckRequiredInOutExistence() const;
|
||||
ge::graphStatus CheckRequiredAttrExistence() const;
|
||||
ge::graphStatus CheckRequiredParaExistence() const;
|
||||
|
||||
ge::graphStatus GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
|
||||
SASLayout &layout, const std::string &name) const;
|
||||
ge::graphStatus GetActualSeqLenQSize(uint32_t &size);
|
||||
ge::graphStatus GetOpName();
|
||||
ge::graphStatus GetNpuInfo();
|
||||
void GetOptionalInputParaInfo();
|
||||
void GetInputParaInfo();
|
||||
void GetOutputParaInfo();
|
||||
ge::graphStatus GetAttrParaInfo();
|
||||
|
||||
ge::graphStatus GetOpParaInfo();
|
||||
|
||||
ge::graphStatus GetInOutDataType();
|
||||
ge::graphStatus GetQueryAndOutLayout();
|
||||
ge::graphStatus GetKvLayout();
|
||||
void SetSASShape();
|
||||
ge::graphStatus GetN1Size();
|
||||
ge::graphStatus GetN2Size();
|
||||
ge::graphStatus GetGSize();
|
||||
ge::graphStatus GetBatchSize();
|
||||
ge::graphStatus GetQTSize();
|
||||
ge::graphStatus GetS1Size();
|
||||
ge::graphStatus GetS2SizeForPageAttention();
|
||||
ge::graphStatus GetS2Size();
|
||||
ge::graphStatus GetMaxBlockNumPerBatch();
|
||||
ge::graphStatus GetBlockSize();
|
||||
ge::graphStatus GetQkHeadDim();
|
||||
ge::graphStatus GetSparseBlockCount();
|
||||
ge::graphStatus GetActualseqInfo();
|
||||
ge::graphStatus GetDSizeQ();
|
||||
ge::graphStatus GetDSizeKV();
|
||||
ge::graphStatus GetSinks();
|
||||
void GenerateInfo(KvQuantSASTilingInfo &sasInfo);
|
||||
ge::graphStatus Parse(KvQuantSASTilingInfo &sasInfo);
|
||||
|
||||
public:
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
const char *opName_;
|
||||
fe::PlatFormInfos *platformInfo_;
|
||||
KvQuantSASParaInfo opParamInfo_;
|
||||
|
||||
bool HasAxis(const SASAxis &axis, const SASLayout &layout, const gert::Shape &shape) const;
|
||||
size_t GetAxisIdx(const SASAxis &axis, const SASLayout &layout) const;
|
||||
uint32_t GetAxisNum(const gert::Shape &shape, const SASAxis &axis,const SASLayout &layout) const;
|
||||
static constexpr int64_t invalidDimValue_ = std::numeric_limits<int64_t>::min();
|
||||
|
||||
// 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;
|
||||
uint32_t qTSize_ = 0;
|
||||
uint32_t qkHeadDim_ = 0;
|
||||
int64_t sparseBlockSize_ = 0;
|
||||
int64_t sparseBlockCount_ = 0;
|
||||
uint32_t maxActualseq_ = 0;
|
||||
bool isSameSeqAllKVTensor_ = true;
|
||||
uint32_t actualLenDimsKV_ = 0;
|
||||
uint32_t actualLenDimsQ_ = 0;
|
||||
uint32_t dSizeQ_ = 0;
|
||||
uint32_t dSizeKV_ = 0;
|
||||
// Layout
|
||||
SASLayout qLayout_ = SASLayout::BSND;
|
||||
SASLayout outLayout_ = SASLayout::BSND;
|
||||
SASLayout kvLayout_ = SASLayout::PA_ND;
|
||||
// PageAttention
|
||||
uint32_t oriMaxBlockNumPerBatch_ = 0;
|
||||
uint32_t cmpMaxBlockNumPerBatch_ = 0;
|
||||
int32_t oriBlockSize_ = 0;
|
||||
int32_t cmpBlockSize_ = 0;
|
||||
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
|
||||
ge::DataType qType_ = ge::DT_FLOAT16;
|
||||
ge::DataType oriKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType cmpKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType cmpSparseIndicesType_ = ge::DT_INT32;
|
||||
ge::DataType oriBlockTableType_ = ge::DT_INT32;
|
||||
ge::DataType cmpBlockTableType_ = ge::DT_INT32;
|
||||
ge::DataType cuSeqLensQType_ = ge::DT_INT32;
|
||||
ge::DataType seqsedKvType_ = ge::DT_INT32;
|
||||
ge::DataType sinksType_ = ge::DT_INT32;
|
||||
ge::DataType metadataType_ = ge::DT_INT32;
|
||||
ge::DataType outputType_ = ge::DT_FLOAT16;
|
||||
|
||||
gert::Shape qShape_{};
|
||||
gert::Shape oriKvShape_{};
|
||||
gert::Shape cmpKvShape_{};
|
||||
gert::Shape cmpSparseIndicesShape_{};
|
||||
};
|
||||
|
||||
class KvQuantSASTilingCheck {
|
||||
public:
|
||||
explicit KvQuantSASTilingCheck(const KvQuantSASTilingInfo &sasInfo) : sasInfo_(sasInfo) {};
|
||||
~KvQuantSASTilingCheck() = default;
|
||||
virtual ge::graphStatus Process();
|
||||
|
||||
private:
|
||||
void Init();
|
||||
bool HasAxis(const SASAxis &axis, const SASLayout &layout, const gert::Shape &shape) const;
|
||||
size_t GetAxisIdx(const SASAxis &axis, const SASLayout &layout) const;
|
||||
uint32_t GetAxisNum(const gert::Shape &shape, const SASAxis &axis,const SASLayout &layout) const;
|
||||
static constexpr int64_t invalidDimValue_ = std::numeric_limits<int64_t>::min();
|
||||
|
||||
void LogErrorDtypeSupport(const std::vector<ge::DataType> &expectDtypeList,
|
||||
const ge::DataType &actualDtype, const std::string &name) const;
|
||||
ge::graphStatus CheckDtypeSupport(const gert::CompileTimeTensorDesc *desc,
|
||||
const std::string &name) const;
|
||||
template <typename T> void LogErrorNumberSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name, const std::string subName) const;
|
||||
template <typename T> void LogErrorDimNumSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name) const;
|
||||
ge::graphStatus CheckDimNumSupport(const gert::StorageShape *shape,
|
||||
const std::vector<size_t> &expectDimNumList, const std::string &name) const;
|
||||
ge::graphStatus CheckShapeNumSupport(const gert::StorageShape *shape,
|
||||
const std::vector<int64_t> &expectShapeNumList, const std::string &name) const;
|
||||
ge::graphStatus CheckDimNumInLayoutSupport(const SASLayout &layout,
|
||||
const gert::StorageShape *shape, const std::string &name) const;
|
||||
void LogErrorLayoutSupport(const std::vector<SASLayout> &expectLayoutList,
|
||||
const SASLayout &actualLayout, const std::string &name) const;
|
||||
ge::graphStatus GetExpectedShape(gert::Shape &shapeExpected,
|
||||
const KvQuantSASTilingShapeCompareParam ¶m, const SASLayout &layout) const;
|
||||
ge::graphStatus CompareShape(KvQuantSASTilingShapeCompareParam ¶m,
|
||||
const gert::Shape &shape, const SASLayout &layout, const std::string &name) const;
|
||||
ge::graphStatus CheckLayoutSupport(const SASLayout &actualLayout, const std::string &name) const;
|
||||
ge::graphStatus CheckSingleParaQuery() const;
|
||||
ge::graphStatus CheckSingleParaKey() const;
|
||||
ge::graphStatus CheckSingleParaNumHeads() const;
|
||||
ge::graphStatus CheckSingleParaKvHeadNums() const;
|
||||
ge::graphStatus CheckSingleParaSparseMode() const;
|
||||
ge::graphStatus CheckSingleParaSparseBlockSize() const;
|
||||
ge::graphStatus CheckSingleParaCmpSparseIndices() const;
|
||||
ge::graphStatus CheckSingleParaBlockTable() const;
|
||||
ge::graphStatus CheckSingleParaCuSeqLensQ() const;
|
||||
ge::graphStatus CheckSingleParaSequsedKv() const;
|
||||
ge::graphStatus CheckSingleParaSinks() const;
|
||||
ge::graphStatus CheckSingleParaMetadata() const;
|
||||
|
||||
ge::graphStatus CheckSinglePara() const;
|
||||
ge::graphStatus CheckParaExistenceAntiquant() const;
|
||||
ge::graphStatus CheckCmpSparseIndicesExistence();
|
||||
ge::graphStatus CheckParaExistence();
|
||||
ge::graphStatus CheckCmpRatioExistence();
|
||||
ge::graphStatus GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
|
||||
const SASLayout &layout, const std::string &name) const;
|
||||
ge::graphStatus CheckSWAExistence();
|
||||
ge::graphStatus CheckCFAExistence();
|
||||
ge::graphStatus CheckSCFAExistence();
|
||||
ge::graphStatus CheckUnrequiredParaExistence() const;
|
||||
|
||||
ge::graphStatus CheckKVShapeForBatchContinuous();
|
||||
uint32_t GetTypeSize(ge::DataType dtype) const;
|
||||
ge::graphStatus CheckKVShapeForPageAttention();
|
||||
ge::graphStatus CheckKVShape();
|
||||
ge::graphStatus CheckKV();
|
||||
ge::graphStatus CheckTopK();
|
||||
ge::graphStatus CheckTopkShape();
|
||||
ge::graphStatus CheckBlockTable() const;
|
||||
ge::graphStatus CheckDTypeConsistency(const ge::DataType &actualDtype,
|
||||
const ge::DataType &expectDtype, const std::string &name) const;
|
||||
|
||||
ge::graphStatus CheckAttenOut();
|
||||
ge::graphStatus CheckAttenOutShape();
|
||||
ge::graphStatus CheckActualSeqLensQ();
|
||||
ge::graphStatus CheckActualSeqLensQShape();
|
||||
ge::graphStatus CheckActualSeqLensQDType();
|
||||
ge::graphStatus CheckActualSeqLens();
|
||||
ge::graphStatus CheckActualSeqLensDType();
|
||||
ge::graphStatus CheckActualSeqLensShape();
|
||||
ge::graphStatus CheckMultiParaConsistency();
|
||||
|
||||
ge::graphStatus CheckFeatureWinKV() const;
|
||||
ge::graphStatus CheckFeatureAntiquantShape() const;
|
||||
ge::graphStatus CheckFeatureAntiquantLayout() const;
|
||||
ge::graphStatus CheckFeatureAntiquantDtype() const;
|
||||
ge::graphStatus CheckFeatureAntiquantAttr() const;
|
||||
ge::graphStatus CheckFeatureAntiquantPa() const;
|
||||
ge::graphStatus CheckFeatureAntiquant() const;
|
||||
ge::graphStatus CheckFeature() const;
|
||||
|
||||
void SetSASShapeCompare();
|
||||
|
||||
private:
|
||||
const char *opName_;
|
||||
fe::PlatFormInfos *platformInfo_;
|
||||
KvQuantSASParaInfo opParamInfo_;
|
||||
const KvQuantSASTilingInfo &sasInfo_;
|
||||
|
||||
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 qkHeadDim_ = 0;
|
||||
uint32_t vHeadDim_ = 0;
|
||||
int64_t ropeHeadDim_ = 0;
|
||||
uint32_t qTSize_ = 0; // 仅TND时生效
|
||||
uint32_t kvTSize_ = 0; // 仅TND时生效
|
||||
KvStorageMode kvStorageMode_ = KvStorageMode::BATCH_CONTINUOUS;
|
||||
uint32_t sparseBlockCount_ = 0;
|
||||
uint32_t sparseBlockSize_ = 0;
|
||||
|
||||
uint32_t oriBlockNum_ = 0;
|
||||
uint32_t cmpBlockNum_ = 0;
|
||||
|
||||
uint32_t actualLenDimsQ_ = 0;
|
||||
|
||||
uint32_t oriBlockSize_ = 0;
|
||||
uint32_t cmpBlockSize_ = 0;
|
||||
uint32_t oriBlockTable_ = 0;
|
||||
uint32_t cmpBlockTable_ = 0;
|
||||
int64_t kv_quant_mode_ = 0;
|
||||
int64_t tileSize_ = 0;
|
||||
|
||||
int64_t oriWinLeft_ = 0;
|
||||
int64_t oriWinRight_ = 0;
|
||||
|
||||
int64_t cmpRatio_ = 0;
|
||||
|
||||
uint32_t dSize_ = sasInfo_.dSize;
|
||||
uint32_t dSizeV_ = sasInfo_.dSizeV;
|
||||
uint32_t dSizeVInput_ = sasInfo_.dSizeVInput;
|
||||
|
||||
uint32_t dSizeOriKvInput_ = 0;
|
||||
uint32_t dSizeCmpKvInput_ = 0;
|
||||
|
||||
uint32_t oriMaskMode_ = 0;
|
||||
uint32_t cmpMaskMode_ = 0;
|
||||
|
||||
SASLayout qLayout_ = SASLayout::BSND;
|
||||
SASLayout outLayout_ = SASLayout::BSND;
|
||||
SASLayout kvLayout_ = SASLayout::PA_ND;
|
||||
|
||||
uint32_t oriMaxBlockNumPerBatch_ = 0;
|
||||
uint32_t cmpMaxBlockNumPerBatch_ = 0;
|
||||
|
||||
uint32_t aicNum_ = 0;
|
||||
uint32_t aivNum_ = 0;
|
||||
platform_ascendc::SocVersion socVersion_ = platform_ascendc::SocVersion::ASCEND910B;
|
||||
uint64_t l2CacheSize_ = 0;
|
||||
|
||||
bool isSameSeqAllKVTensor_ = true;
|
||||
bool isSameActualseq_ = true;
|
||||
uint32_t maxActualseq_ = 0;
|
||||
|
||||
ge::DataType qType_ = ge::DT_FLOAT16;
|
||||
ge::DataType oriKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType cmpKvType_ = ge::DT_FLOAT16;
|
||||
ge::DataType outputType_ = ge::DT_FLOAT16;
|
||||
|
||||
SASTemplateMode perfMode_;
|
||||
|
||||
gert::Shape queryShapeCmp_{};
|
||||
gert::Shape keyShapeCmp_{};
|
||||
gert::Shape valueShapeCmp_{};
|
||||
gert::Shape topkShapeCmp_{};
|
||||
gert::Shape attenOutShapeCmp_{};
|
||||
};
|
||||
} // namespace optiling
|
||||
#endif // KVQUANT_SPARSE_ATTN_SHAREDKV_TILING_H
|
||||
@@ -0,0 +1,352 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check_consistancy.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckDTypeConsistency(const ge::DataType &actualDtype,
|
||||
const ge::DataType &expectDtype, const std::string &name) const
|
||||
{
|
||||
if (actualDtype != expectDtype) {
|
||||
OP_LOGE(opName_, "%s dtype should be %s, but it's %s.", name.c_str(),
|
||||
SASDataTypeToSerialString(expectDtype).c_str(),
|
||||
SASDataTypeToSerialString(actualDtype).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
|
||||
const SASLayout &layout, const std::string &name) const
|
||||
{
|
||||
if (tensor == nullptr) {
|
||||
OP_LOGE(opName_, "when layout of query is %s, %s must be provided.",
|
||||
KvQuantSASLayoutToSerialString(layout).c_str(), name.c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t shapeSize = tensor->GetShapeSize();
|
||||
if (shapeSize <= 0) {
|
||||
OP_LOGE(opName_, "the shape size of %s is %ld, it should be greater than 0.",
|
||||
name.c_str(), shapeSize);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
size = static_cast<uint32_t>(shapeSize);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::GetExpectedShape(gert::Shape &shapeExpected,
|
||||
const KvQuantSASTilingShapeCompareParam ¶m, const SASLayout &layout) const
|
||||
{
|
||||
if (layout == SASLayout::BSND) {
|
||||
shapeExpected = gert::Shape({param.B, param.S, param.N, param.D});
|
||||
} else if (layout == SASLayout::TND) {
|
||||
shapeExpected = gert::Shape({param.T, param.N, param.D});
|
||||
} else if (layout == SASLayout::PA_ND) {
|
||||
shapeExpected = gert::Shape({param.Bn, param.Bs, param.N, param.D});
|
||||
} else {
|
||||
OP_LOGE(opName_, "layout %s is unsupported", KvQuantSASLayoutToSerialString(layout).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CompareShape(KvQuantSASTilingShapeCompareParam ¶m,
|
||||
const gert::Shape &shape, const SASLayout &layout, const std::string &name) const
|
||||
{
|
||||
gert::Shape shapeExpected;
|
||||
if (GetExpectedShape(shapeExpected, param, layout) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (shape.GetDimNum() != shapeExpected.GetDimNum()) {
|
||||
OP_LOGE(opName_,
|
||||
"%s dimension is %zu, expected dimension is %zu.",
|
||||
name.c_str(), shape.GetDimNum(), shapeExpected.GetDimNum());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < shape.GetDimNum(); i++) {
|
||||
if (shape.GetDim(i) != shapeExpected.GetDim(i)) {
|
||||
OP_LOGE(opName_, "%s layout is %s, shape is %s, expected shape is %s.",
|
||||
name.c_str(), KvQuantSASLayoutToSerialString(layout).c_str(),
|
||||
GetShapeStr(shape).c_str(), GetShapeStr(shapeExpected).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void KvQuantSASTilingCheck::SetSASShapeCompare()
|
||||
{
|
||||
queryShapeCmp_ = opParamInfo_.q.shape->GetStorageShape();
|
||||
topkShapeCmp_ = opParamInfo_.cmpSparseIndices.tensor->GetShape().GetStorageShape();
|
||||
keyShapeCmp_ = opParamInfo_.oriKv.tensor->GetShape().GetStorageShape();
|
||||
valueShapeCmp_ = opParamInfo_.cmpKv.tensor->GetShape().GetStorageShape();
|
||||
attenOutShapeCmp_ = opParamInfo_.attnOut.shape->GetStorageShape();
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckBlockTable() const
|
||||
{
|
||||
if (kvStorageMode_ != KvStorageMode::PAGE_ATTENTION) {
|
||||
OP_CHECK_IF(opParamInfo_.oriBlockTable.tensor != nullptr,
|
||||
OP_LOGE(opName_, "when the layout_kv is %s, %s should be null",
|
||||
KvQuantSASLayoutToSerialString(kvLayout_).c_str(), ORI_BLOCK_TABLE_NAME.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
if (kvStorageMode_ != KvStorageMode::PAGE_ATTENTION) {
|
||||
OP_CHECK_IF(opParamInfo_.cmpBlockTable.tensor != nullptr,
|
||||
OP_LOGE(opName_, "when the layout_kv is %s, %s should be null",
|
||||
KvQuantSASLayoutToSerialString(kvLayout_).c_str(), CMP_BLOCK_TABLE_NAME.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint32_t oriBlockTableBatch = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(0);
|
||||
OP_CHECK_IF(oriBlockTableBatch != bSize_,
|
||||
OP_LOGE(opName_, "oriBlockTableBatch's first dimension(%u) should be equal to batch size(%u)",
|
||||
oriBlockTableBatch, bSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint32_t cmpBlockTableBatch = opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0);
|
||||
OP_CHECK_IF(cmpBlockTableBatch != bSize_,
|
||||
OP_LOGE(opName_, "cmpBlockTableBatch's first dimension(%u) should be equal to batch size(%u)",
|
||||
cmpBlockTableBatch, bSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckTopkShape()
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckAttenOutShape()
|
||||
{
|
||||
KvQuantSASTilingShapeCompareParam shapeParams;
|
||||
shapeParams.B = bSize_;
|
||||
shapeParams.N = n1Size_;
|
||||
shapeParams.S = s1Size_;
|
||||
shapeParams.D = 512; // 512:输出的head_dim
|
||||
shapeParams.T = qTSize_;
|
||||
if (CompareShape(shapeParams, attenOutShapeCmp_, outLayout_, ATTEN_OUT_NAME) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckAttenOut()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(opParamInfo_.attnOut.desc->GetDataType(),
|
||||
qType_, ATTEN_OUT_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckAttenOutShape()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckTopK()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckTopkShape()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckKVShapeForBatchContinuous()
|
||||
{
|
||||
KvQuantSASTilingShapeCompareParam shapeParams;
|
||||
shapeParams.B = bSize_;
|
||||
shapeParams.N = n2Size_;
|
||||
shapeParams.S = s2Size_;
|
||||
shapeParams.D = vHeadDim_;
|
||||
shapeParams.T = kvTSize_;
|
||||
if (CompareShape(shapeParams, valueShapeCmp_, kvLayout_, VALUE_NAME) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
uint32_t KvQuantSASTilingCheck::GetTypeSize(ge::DataType dtype) const
|
||||
{
|
||||
uint32_t typeSize = NUM_BYTES_FLOAT16;
|
||||
switch (dtype) {
|
||||
case ge::DT_FLOAT16:
|
||||
typeSize = NUM_BYTES_FLOAT16;
|
||||
break;
|
||||
case ge::DT_BF16:
|
||||
typeSize = NUM_BYTES_BF16;
|
||||
break;
|
||||
default:
|
||||
typeSize = NUM_BYTES_FLOAT16;
|
||||
}
|
||||
return typeSize;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckKVShapeForPageAttention()
|
||||
{
|
||||
int64_t blockNum = keyShapeCmp_.GetDim(0);
|
||||
KvQuantSASTilingShapeCompareParam shapeParams;
|
||||
shapeParams.Bn = blockNum;
|
||||
shapeParams.N = n2Size_;
|
||||
shapeParams.Bs = bSize_;
|
||||
shapeParams.T = kvTSize_;
|
||||
shapeParams.D = vHeadDim_;
|
||||
if (CompareShape(shapeParams, valueShapeCmp_, kvLayout_, VALUE_NAME) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckKVShape()
|
||||
{
|
||||
if (kvStorageMode_ == KvStorageMode::BATCH_CONTINUOUS) {
|
||||
return CheckKVShapeForBatchContinuous();
|
||||
}
|
||||
|
||||
if (kvStorageMode_ == KvStorageMode::PAGE_ATTENTION) {
|
||||
return CheckKVShapeForPageAttention();
|
||||
}
|
||||
|
||||
OP_LOGE(opName_, "storage mode of key and value is %u, it is incorrect.", static_cast<uint32_t>(kvStorageMode_));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckKV()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(cmpKvType_,
|
||||
oriKvType_, CMP_KV_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLensQ()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckActualSeqLensQDType() ||
|
||||
ge::GRAPH_SUCCESS != CheckActualSeqLensQShape()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLensQDType()
|
||||
{
|
||||
if (opParamInfo_.cuSeqLensQ.tensor == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
if (opParamInfo_.cuSeqLensQ.desc == nullptr) {
|
||||
OP_LOGE(opName_, "cuSeqLensQ is not empty,"
|
||||
"but cuSeqLensQ's dtype is nullptr.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.cuSeqLensQ.desc->GetDataType() != ge::DT_INT32) {
|
||||
OP_LOGE(opName_, "cuSeqLensQ's dtype is %s, it should be DT_INT32.",
|
||||
SASDataTypeToSerialString(opParamInfo_.cuSeqLensQ.desc->GetDataType()).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLensQShape()
|
||||
{
|
||||
if (opParamInfo_.cuSeqLensQ.tensor == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
uint32_t shapeSize = 0;
|
||||
if (GetActualSeqLenSize(shapeSize, opParamInfo_.cuSeqLensQ.tensor, qLayout_, "cuSeqLensQ") !=
|
||||
ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (shapeSize != bSize_ + 1) {
|
||||
OP_LOGE(opName_, "cuSeqLensQ shape size is %u, it should be equal to batch size[%u]",
|
||||
shapeSize, bSize_);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLens()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckActualSeqLensDType() ||
|
||||
ge::GRAPH_SUCCESS != CheckActualSeqLensShape()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLensDType()
|
||||
{
|
||||
if (opParamInfo_.sequsedKv.tensor == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
if (opParamInfo_.sequsedKv.desc == nullptr) {
|
||||
OP_LOGE(opName_, "sequsedKv is not empty,"
|
||||
"but sequsedKv's dtype is nullptr.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.sequsedKv.desc->GetDataType() != ge::DT_INT32) {
|
||||
OP_LOGE(opName_, "sequsedKv's dtype is %s, it should be DT_INT32.",
|
||||
SASDataTypeToSerialString(opParamInfo_.sequsedKv.desc->GetDataType()).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckActualSeqLensShape()
|
||||
{
|
||||
if (opParamInfo_.sequsedKv.tensor == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
uint32_t shapeSize = 0;
|
||||
if (GetActualSeqLenSize(shapeSize, opParamInfo_.sequsedKv.tensor, kvLayout_, "sequsedKv") !=
|
||||
ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (shapeSize != bSize_) {
|
||||
OP_LOGE(opName_, "sequsedKv shape size is %u, it should be equal to batch size[%u].",
|
||||
shapeSize, bSize_);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckMultiParaConsistency()
|
||||
{
|
||||
SetSASShapeCompare();
|
||||
if (ge::GRAPH_SUCCESS != CheckKV() ||
|
||||
ge::GRAPH_SUCCESS != CheckTopK() ||
|
||||
ge::GRAPH_SUCCESS != CheckAttenOut() ||
|
||||
ge::GRAPH_SUCCESS != CheckActualSeqLensQ() ||
|
||||
ge::GRAPH_SUCCESS != CheckActualSeqLens() ||
|
||||
ge::GRAPH_SUCCESS != CheckBlockTable()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,174 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check_existance.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
static constexpr uint32_t TopK_SIZE = 512;
|
||||
static constexpr uint32_t DIM_0 = 0;
|
||||
static constexpr uint32_t DIM_1 = 1;
|
||||
static constexpr uint32_t DIM_2 = 2;
|
||||
static constexpr uint32_t DIM_3 = 3;
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckParaExistenceAntiquant() const
|
||||
{
|
||||
if (kvLayout_ == SASLayout::BSND) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
} else if (kvLayout_ == SASLayout::PA_ND) {
|
||||
OP_CHECK_IF(opParamInfo_.sequsedKv.tensor == nullptr,
|
||||
OP_LOGE(opName_, "when layout_kv is PA_ND, actualSeqLengthsKv must not be null"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF((opParamInfo_.oriBlockTable.tensor == nullptr) && (opParamInfo_.cmpBlockTable.tensor == nullptr),
|
||||
OP_LOGE(opName_, "when layout_kv is PA_ND, oriBlockTable and cmpBlockTable must be one "),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckParaExistence()
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckCmpSparseIndicesExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckSWAExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckCFAExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckSCFAExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckCmpRatioExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckUnrequiredParaExistence() ||
|
||||
ge::GRAPH_SUCCESS != CheckParaExistenceAntiquant()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckUnrequiredParaExistence() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.oriSparseIndices.tensor != nullptr || opParamInfo_.oriSparseIndices.desc != nullptr,
|
||||
OP_LOGE(opName_, "oriSparseIndices is not supported now, it must be nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(opParamInfo_.cuSeqLensOriKv.tensor != nullptr || opParamInfo_.cuSeqLensOriKv.desc != nullptr,
|
||||
OP_LOGE(opName_, "cuSeqLensOriKv is not supported now, it must be nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(opParamInfo_.cuSeqLensCmpKv.tensor != nullptr || opParamInfo_.cuSeqLensCmpKv.desc != nullptr,
|
||||
OP_LOGE(opName_, "cuSeqLensCmpKv is not supported now, it must be nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckCmpSparseIndicesExistence()
|
||||
{
|
||||
if (opParamInfo_.cmpSparseIndices.tensor != nullptr) {
|
||||
if (qLayout_ == SASLayout::BSND) {
|
||||
if (opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_3) != 512 && opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_3) != 1024) {
|
||||
OP_LOGE(opName_, "When qLayout is BNSD, topK should be 512 or 1024, but got %ld", opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(3));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_1) != s1Size_) {
|
||||
OP_LOGE(opName_, "When qLayout is BNSD, cmpSparseIndices's S should be eaque to s1Size:%u, but got %ld", s1Size_, opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(1));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
} else {
|
||||
if (opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_2) != 512 && opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_2) != 1024) {
|
||||
OP_LOGE(opName_, "When qLayout is TND, topK should be 512 or 1024, but got %ld", opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(2));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(DIM_0) != qTSize_) {
|
||||
OP_LOGE(opName_, "When qLayout is TND, cmpSparseIndices's T should be eaque to qTSize:%u, but got %ld", qTSize_, opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetDim(0));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSWAExistence()
|
||||
{
|
||||
if (perfMode_ != SASTemplateMode::SWA_TEMPLATE_MODE) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor != nullptr && opParamInfo_.oriBlockTable.tensor == nullptr,
|
||||
OP_LOGE(opName_, "oriBlockTable must not be empty when cmpKv is not provided. "),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckCFAExistence()
|
||||
{
|
||||
if (perfMode_ != SASTemplateMode::CFA_TEMPLATE_MODE) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor == nullptr && opParamInfo_.cmpKv.tensor != nullptr,
|
||||
OP_LOGE(opName_, "oriKv must not be empty when cmpKv is provided and cmpSparseIndices is not provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor != nullptr && opParamInfo_.cmpKv.tensor == nullptr && opParamInfo_.cmpRatio != nullptr,
|
||||
OP_LOGE(opName_, "cmpKv must not be empty when cmpKv is provided and cmpSparseIndices is not provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor != nullptr && opParamInfo_.cmpKv.tensor != nullptr && opParamInfo_.cmpRatio == nullptr,
|
||||
OP_LOGE(opName_, "cmpRatio must not be empty when cmpKv is provided and cmpSparseIndices is not provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor != nullptr && opParamInfo_.cmpKv.tensor != nullptr && opParamInfo_.cmpBlockTable.tensor == nullptr,
|
||||
OP_LOGE(opName_, "cmpBlockTable must not be empty when cmpKv is provided and cmpSparseIndices is not provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSCFAExistence()
|
||||
{
|
||||
if (perfMode_ != SASTemplateMode::SCFA_TEMPLATE_MODE) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor != nullptr && opParamInfo_.cmpKv.tensor == nullptr && opParamInfo_.cmpSparseIndices.tensor != nullptr,
|
||||
OP_LOGE(opName_, "cmpKv must not be empty when cmpKv and cmpSparseIndices are provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor == nullptr && opParamInfo_.cmpKv.tensor != nullptr && opParamInfo_.cmpSparseIndices.tensor != nullptr,
|
||||
OP_LOGE(opName_, "oriKv must not be empty when cmpKv and cmpSparseIndices are provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor == nullptr && opParamInfo_.cmpKv.tensor == nullptr && opParamInfo_.cmpSparseIndices.tensor != nullptr,
|
||||
OP_LOGE(opName_, "oriKv and cmpKv must not be empty when cmpKv and cmpSparseIndices are provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckCmpRatioExistence()
|
||||
{
|
||||
if (perfMode_ == SASTemplateMode::SWA_TEMPLATE_MODE) {
|
||||
OP_CHECK_IF(*opParamInfo_.cmpRatio != 1 && *opParamInfo_.cmpRatio != 128 && *opParamInfo_.cmpRatio != 4,
|
||||
OP_LOGE(opName_, "when SWA mode, cmpRatio must be 1 or 4 or 128, but got %u", *opParamInfo_.cmpRatio),
|
||||
return ge::GRAPH_FAILED);
|
||||
} else if (perfMode_ == SASTemplateMode::CFA_TEMPLATE_MODE) {
|
||||
OP_CHECK_IF(*opParamInfo_.cmpRatio != 128 && *opParamInfo_.cmpRatio != 4,
|
||||
OP_LOGE(opName_, "when CFA mode, cmpRatio must be 4 or 128, but got %u", *opParamInfo_.cmpRatio),
|
||||
return ge::GRAPH_FAILED);
|
||||
} else {
|
||||
OP_CHECK_IF(*opParamInfo_.cmpRatio != 128 && *opParamInfo_.cmpRatio != 4,
|
||||
OP_LOGE(opName_, "when SCFA mode, cmpRatio must be 4 or 128, but got %u", *opParamInfo_.cmpRatio),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check_feature.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureWinKV() const
|
||||
{
|
||||
OP_CHECK_IF(oriWinLeft_ != 127, // 127:当前不泛化
|
||||
OP_LOGE(opName_, "oriWinLeft_ only support 127, but got %u", oriWinLeft_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(oriWinRight_ != 0, // 0:当前不泛化
|
||||
OP_LOGE(opName_, "oriWinRight_ only support 0, but got %u", oriWinRight_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquantShape() const
|
||||
{
|
||||
OP_CHECK_IF(bSize_ <= 0,
|
||||
OP_LOGE(opName_, "batch_size should be greater than 0, but got %u", bSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(qTSize_ <= 0 && (qLayout_ == SASLayout::TND),
|
||||
OP_LOGE(opName_, "T_size of query should be greater than 0, but got %u", qTSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(n1Size_ != 64 && n1Size_ != 128,
|
||||
OP_LOGE(opName_, "q_head_num only support 64 and 128, but got %u", n1Size_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(n2Size_ != 1,
|
||||
OP_LOGE(opName_, "kv_head_num only support 1, but got %u", n2Size_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(n1Size_ % n2Size_ != 0,
|
||||
OP_LOGE(opName_, "q_head_num(%u) must be divisible by kv_head_num(%u)", n1Size_, n2Size_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
std::vector<uint32_t> gSizeSupportList = {64, 128};
|
||||
OP_CHECK_IF(std::find(gSizeSupportList.begin(), gSizeSupportList.end(), gSize_) == gSizeSupportList.end(),
|
||||
OP_LOGE(opName_, "group num only support 64 and 128, but got %u", gSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(dSize_ != 512, // 512:当前不泛化
|
||||
OP_LOGE(opName_, "Head dim of input q only support 512, but got %u", dSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(dSizeV_ != 512, // 512:当前不泛化
|
||||
OP_LOGE(opName_, "dSizeV only support 512, but got %u", dSizeV_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(dSizeVInput_ != 640, // 640:当前不泛化
|
||||
OP_LOGE(opName_, "dSizeVInput only support 640, but got %u", dSizeVInput_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquantLayout() const
|
||||
{
|
||||
const std::vector<std::string> layoutSupportList = {
|
||||
"BSND",
|
||||
"TND"
|
||||
};
|
||||
std::string layoutQuery = opParamInfo_.layoutQ;
|
||||
OP_CHECK_IF(std::find(layoutSupportList.begin(), layoutSupportList.end(), layoutQuery) == layoutSupportList.end(),
|
||||
OP_LOGE(opName_, "layoutQuery only support BSND/TND, but got %s", layoutQuery.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquantDtype() const
|
||||
{
|
||||
OP_CHECK_IF(qType_ != ge::DT_BF16,
|
||||
OP_LOGE(opName_, "query dtype only support %s and %s, but got %s",
|
||||
SASDataTypeToSerialString(ge::DT_BF16).c_str(), SASDataTypeToSerialString(ge::DT_FLOAT16).c_str(),
|
||||
SASDataTypeToSerialString(qType_).c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquantAttr() const
|
||||
{
|
||||
OP_CHECK_IF(*opParamInfo_.kvQuantMode != 1,
|
||||
OP_LOGE(opName_, "kv_quant_mode_ only support 1, but got %ld",
|
||||
*opParamInfo_.kvQuantMode),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (*opParamInfo_.kvQuantMode == 1) {
|
||||
OP_CHECK_IF(*opParamInfo_.tileSize != 64, // 64:当前不泛化
|
||||
OP_LOGE(opName_, "tile_size only support 64, but got %ld",
|
||||
*opParamInfo_.tileSize),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
OP_CHECK_IF(*opParamInfo_.ropeHeadDim != 64, // 64:当前不泛化
|
||||
OP_LOGE(opName_, "rope_head_dim only support 64, but got %ld",
|
||||
*opParamInfo_.ropeHeadDim),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquantPa() const
|
||||
{
|
||||
OP_CHECK_IF(oriBlockSize_ <= 0 || oriBlockSize_ > static_cast<int32_t>(MAX_BLOCK_SIZE),
|
||||
OP_LOGE(opName_, "when page attention is enabled, oriBlockSize_(%ld) should be in range (0, %u].",
|
||||
oriBlockSize_, MAX_BLOCK_SIZE), return ge::GRAPH_FAILED);
|
||||
|
||||
if (cmpBlockSize_ != 0){
|
||||
OP_CHECK_IF(cmpBlockSize_ <= 0 || cmpBlockSize_ > static_cast<int32_t>(MAX_BLOCK_SIZE),
|
||||
OP_LOGE(opName_, "when page attention is enabled, cmpBlockSize_(%ld) should be in range (0, %u].",
|
||||
cmpBlockSize_, MAX_BLOCK_SIZE), return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeatureAntiquant() const
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckFeatureAntiquantAttr() ||
|
||||
ge::GRAPH_SUCCESS != CheckFeatureAntiquantShape() ||
|
||||
ge::GRAPH_SUCCESS != CheckFeatureAntiquantLayout() ||
|
||||
ge::GRAPH_SUCCESS != CheckFeatureAntiquantDtype() ||
|
||||
ge::GRAPH_SUCCESS != CheckFeatureWinKV() ||
|
||||
ge::GRAPH_SUCCESS != CheckFeatureAntiquantPa()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckFeature() const
|
||||
{
|
||||
return CheckFeatureAntiquant();
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,427 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file kv_quant_sparse_attn_sharedkv_check_single_para.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
#include "../op_kernel/kv_quant_sparse_attn_sharedkv_metadata.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
static constexpr uint32_t DIM_0 = 0;
|
||||
static constexpr uint32_t DIM_1 = 1;
|
||||
static constexpr uint32_t DIM_2 = 2;
|
||||
static constexpr uint32_t DIM_3 = 3;
|
||||
|
||||
const std::map<std::string, std::vector<ge::DataType>> DTYPE_SUPPORT_MAP = {
|
||||
{QUERY_NAME, {ge::DT_BF16}},
|
||||
{ORI_KV_NAME, {ge::DT_INT8, ge::DT_FLOAT8_E4M3FN}},
|
||||
{CMP_KV_NAME, {ge::DT_INT8, ge::DT_FLOAT8_E4M3FN}},
|
||||
{ATTEN_OUT_NAME, {ge::DT_FLOAT16, ge::DT_BF16}},
|
||||
{CMP_SPARSE_INDICES_NAME, {ge::DT_INT32}},
|
||||
{ORI_BLOCK_TABLE_NAME, {ge::DT_INT32}},
|
||||
{CMP_BLOCK_TABLE_NAME, {ge::DT_INT32}},
|
||||
{CU_SEQLENS_Q_NAME, {ge::DT_INT32}},
|
||||
{SEQUSED_KV_NAME, {ge::DT_INT32}},
|
||||
{SINKS_NAME, {ge::DT_FLOAT}},
|
||||
};
|
||||
|
||||
const std::map<std::string, std::vector<SASLayout>> LAYOUT_SUPPORT_MAP = {
|
||||
{QUERY_NAME, {SASLayout::BSND, SASLayout::TND}},
|
||||
{ORI_KV_NAME, {SASLayout::PA_ND}},
|
||||
{CMP_KV_NAME, {SASLayout::PA_ND}},
|
||||
{ATTEN_OUT_NAME, {SASLayout::BSND, SASLayout::TND}},
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
void KvQuantSASTilingCheck::LogErrorDimNumSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name) const
|
||||
{
|
||||
LogErrorNumberSupport(expectNumberList, actualValue, name, "dimension");
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckDimNumInLayoutSupport(const SASLayout &layout,
|
||||
const gert::StorageShape *shape, const std::string &name) const
|
||||
{
|
||||
const auto& dimIt = SAS_LAYOUT_DIM_MAP.find(layout);
|
||||
OP_CHECK_IF(shape->GetStorageShape().GetDimNum() != dimIt->second,
|
||||
OP_LOGE(opName_, "When layout is %s, %s dimension should be %zu, but it's %zu",
|
||||
KvQuantSASLayoutToSerialString(layout).c_str(), name.c_str(), dimIt->second,
|
||||
shape->GetStorageShape().GetDimNum()),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckDimNumSupport(const gert::StorageShape *shape,
|
||||
const std::vector<size_t> &expectDimNumList, const std::string &name) const
|
||||
{
|
||||
if (shape == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
if (std::find(expectDimNumList.begin(), expectDimNumList.end(),
|
||||
shape->GetStorageShape().GetDimNum()) == expectDimNumList.end()) {
|
||||
LogErrorDimNumSupport(expectDimNumList, shape->GetStorageShape().GetDimNum(), name);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckShapeNumSupport(const gert::StorageShape *shape,
|
||||
const std::vector<int64_t> &expectShapeNumList, const std::string &name) const
|
||||
{
|
||||
if (shape == nullptr) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
if (std::find(expectShapeNumList.begin(), expectShapeNumList.end(),
|
||||
shape->GetStorageShape().GetShapeSize()) == expectShapeNumList.end()) {
|
||||
LogErrorDimNumSupport(expectShapeNumList, shape->GetStorageShape().GetShapeSize(), name);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void KvQuantSASTilingCheck::LogErrorDtypeSupport(const std::vector<ge::DataType> &expectDtypeList,
|
||||
const ge::DataType &actualDtype, const std::string &name) const
|
||||
{
|
||||
std::ostringstream oss;
|
||||
for (size_t i = 0; i < expectDtypeList.size(); ++i) {
|
||||
oss << SASDataTypeToSerialString(expectDtypeList[i]);
|
||||
if (i < expectDtypeList.size() - 1) {
|
||||
oss << ", ";
|
||||
}
|
||||
}
|
||||
OP_LOGE(opName_, "Tensor %s only support dtype %s, but got %s",
|
||||
name.c_str(), oss.str().c_str(), SASDataTypeToSerialString(actualDtype).c_str());
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckDtypeSupport(const gert::CompileTimeTensorDesc *desc,
|
||||
const std::string &name) const
|
||||
{
|
||||
if (desc != nullptr) {
|
||||
const auto& it = DTYPE_SUPPORT_MAP.find(name);
|
||||
OP_CHECK_IF(it == DTYPE_SUPPORT_MAP.end(),
|
||||
OP_LOGE(opName_, "%s datatype support list should be specify in DTYPE_SUPPORT_MAP", name.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto &expectDtypeList = it->second;
|
||||
OP_CHECK_IF(std::find(
|
||||
expectDtypeList.begin(), expectDtypeList.end(), desc->GetDataType()) == expectDtypeList.end(),
|
||||
LogErrorDtypeSupport(expectDtypeList, desc->GetDataType(), name),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void KvQuantSASTilingCheck::LogErrorNumberSupport(const std::vector<T> &expectNumberList,
|
||||
const T &actualValue, const std::string &name, const std::string subName) const
|
||||
{
|
||||
std::ostringstream oss;
|
||||
for (size_t i = 0; i < expectNumberList.size(); ++i) {
|
||||
oss << std::to_string(expectNumberList[i]);
|
||||
if (i < expectNumberList.size() - 1) {
|
||||
oss << ", ";
|
||||
}
|
||||
}
|
||||
OP_LOGE(opName_, "%s %s only support %s, but got %s",
|
||||
name.c_str(), subName.c_str(), oss.str().c_str(), std::to_string(actualValue).c_str());
|
||||
}
|
||||
|
||||
|
||||
void KvQuantSASTilingCheck::LogErrorLayoutSupport(const std::vector<SASLayout> &expectLayoutList,
|
||||
const SASLayout &actualLayout, const std::string &name) const
|
||||
{
|
||||
std::ostringstream oss;
|
||||
for (size_t i = 0; i < expectLayoutList.size(); ++i) {
|
||||
oss << KvQuantSASLayoutToSerialString(expectLayoutList[i]);
|
||||
if (i < expectLayoutList.size() - 1) {
|
||||
oss << ", ";
|
||||
}
|
||||
}
|
||||
OP_LOGE(opName_, "Tensor %s only support layout %s, but got %s",
|
||||
name.c_str(), oss.str().c_str(), KvQuantSASLayoutToSerialString(actualLayout).c_str());
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaQuery() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.q.desc == nullptr,
|
||||
OP_LOGE(opName_, "Input q is required, but got nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.q.shape->GetStorageShape().GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input q cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
const std::vector<size_t> queryDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR};
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.q.desc, QUERY_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckLayoutSupport(qLayout_, QUERY_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(opParamInfo_.q.shape, queryDimNumList, QUERY_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(qLayout_, opParamInfo_.q.shape, QUERY_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaKey() const
|
||||
{
|
||||
const std::vector<size_t> keyDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR};
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor == nullptr,
|
||||
OP_LOGE(opName_, "input oriKv can not be nullptr, but it's empty"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.oriKv.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input oriKv cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriKv.desc, ORI_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckLayoutSupport(kvLayout_, ORI_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriKv.tensor->GetShape(), keyDimNumList, ORI_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.oriKv.tensor->GetShape(), ORI_KV_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
OP_CHECK_IF(oriBlockSize_ <= 0 || oriBlockSize_ > 1024,
|
||||
OP_LOGE(opName_, "when page attention is enabled, ori_block_size(%u) should be in range (0, %u].",
|
||||
oriBlockSize_, MAX_BLOCK_SIZE), return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(oriBlockSize_ % 16 > 0,
|
||||
OP_LOGE(opName_, "when page attention is enabled, ori_block_size(%u) should be 16-aligned.",
|
||||
oriBlockSize_), return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(dSizeOriKvInput_ != 640,
|
||||
OP_LOGE(opName_, "Dimension of OriKv only support 640, but got %u", dSizeOriKvInput_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (opParamInfo_.cmpKv.tensor != nullptr) {
|
||||
OP_CHECK_IF(opParamInfo_.cmpKv.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input cmpKv cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpKv.desc, CMP_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckLayoutSupport(kvLayout_, CMP_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpKv.tensor->GetShape(), keyDimNumList, CMP_KV_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.cmpKv.tensor->GetShape(), CMP_KV_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
OP_CHECK_IF(dSizeCmpKvInput_ != 640,
|
||||
OP_LOGE(opName_, "Dimension of CmpKv only support 640, but got %u", dSizeCmpKvInput_),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
uint32_t cmpKvN2Size_ = GetAxisNum(opParamInfo_.cmpKv.tensor->GetStorageShape(), SASAxis::N, kvLayout_);
|
||||
OP_CHECK_IF(cmpKvN2Size_ != n2Size_,
|
||||
OP_LOGE(opName_, "N2 size check failed! Expected cmpKvN2 == oriKvN2."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(cmpBlockSize_ <= 0 || cmpBlockSize_ > 1024,
|
||||
OP_LOGE(opName_, "when page attention is enabled, cmp_block_size(%ld) should be in range (0, %u].",
|
||||
cmpBlockSize_, MAX_BLOCK_SIZE), return ge::GRAPH_FAILED);
|
||||
|
||||
OP_CHECK_IF(cmpBlockSize_ % 16 > 0,
|
||||
OP_LOGE(opName_, "when page attention is enabled, cmp_block_size(%ld) should be 16-aligned.",
|
||||
cmpBlockSize_), return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckLayoutSupport(const SASLayout &actualLayout, const std::string &name) const
|
||||
{
|
||||
const auto& it = LAYOUT_SUPPORT_MAP.find(name);
|
||||
OP_CHECK_IF(it == LAYOUT_SUPPORT_MAP.end(),
|
||||
OP_LOGE(opName_, "%s layout support list should be specify in LAYOUT_SUPPORT_MAP", name.c_str()),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto &expectLayoutList = it->second;
|
||||
OP_CHECK_IF(std::find(
|
||||
expectLayoutList.begin(), expectLayoutList.end(), actualLayout) == expectLayoutList.end(),
|
||||
LogErrorLayoutSupport(expectLayoutList, actualLayout, name),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaNumHeads() const
|
||||
{
|
||||
OP_CHECK_IF(n1Size_ != 64 && n1Size_ != 128,
|
||||
OP_LOGE(opName_, "n1Size_ only support 64 and 128 now, but got %u.", n1Size_),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaKvHeadNums() const
|
||||
{
|
||||
OP_CHECK_IF(n2Size_ != 1,
|
||||
OP_LOGE(opName_, "n2Size_ only support 1 now, but got %u.", n2Size_),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaSparseMode() const
|
||||
{
|
||||
OP_CHECK_IF((*opParamInfo_.oriMaskMode != 4 || *opParamInfo_.cmpMaskMode != 3),
|
||||
OP_LOGE(opName_, "oriMaskMode only support 4 and cmpMaskMode only support 3, but got %u and %u.", *opParamInfo_.oriMaskMode, *opParamInfo_.cmpMaskMode),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaSparseBlockSize() const
|
||||
{
|
||||
OP_CHECK_IF(sparseBlockSize_ != 1,
|
||||
OP_LOGE(opName_, "sparseBlockSize_ only support 1, but got %u",
|
||||
sparseBlockSize_),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaCmpSparseIndices() const
|
||||
{
|
||||
const std::vector<size_t> cmpSparseIndicesDimNumList = {DIM_NUM_FOUR, DIM_NUM_THREE};
|
||||
|
||||
if (opParamInfo_.cmpSparseIndices.tensor != nullptr) {
|
||||
OP_CHECK_IF(opParamInfo_.cmpSparseIndices.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input cmpSparseIndices cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpSparseIndices.desc, CMP_SPARSE_INDICES_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpSparseIndices.tensor->GetShape(), cmpSparseIndicesDimNumList, CMP_SPARSE_INDICES_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaBlockTable() const
|
||||
{
|
||||
const std::vector<size_t> BlockTableDimNumList = {DIM_NUM_TWO};
|
||||
if (opParamInfo_.oriBlockTable.tensor != nullptr) {
|
||||
OP_CHECK_IF(opParamInfo_.oriBlockTable.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input oriBlockTable cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriBlockTable.desc, ORI_BLOCK_TABLE_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriBlockTable.tensor->GetShape(), BlockTableDimNumList, ORI_BLOCK_TABLE_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
|
||||
if (opParamInfo_.cmpBlockTable.tensor != nullptr) {
|
||||
OP_CHECK_IF(opParamInfo_.cmpBlockTable.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input cmpBlockTable cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpBlockTable.desc, CMP_BLOCK_TABLE_NAME) ||
|
||||
ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpBlockTable.tensor->GetShape(), BlockTableDimNumList, CMP_BLOCK_TABLE_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaCuSeqLensQ() const
|
||||
{
|
||||
if (qLayout_ == SASLayout::BSND) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
const std::vector<int64_t> cuSeqLensQDimNumList = {bSize_ + 1};
|
||||
OP_CHECK_IF((qLayout_ == SASLayout::TND && opParamInfo_.cuSeqLensQ.tensor == nullptr),
|
||||
OP_LOGE(opName_, "cuSeqLensQ can't be nullptr when layoutQ is TND"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (opParamInfo_.cuSeqLensQ.tensor != nullptr) {
|
||||
OP_CHECK_IF(opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() == 0,
|
||||
OP_LOGE(opName_, "Any dim of input cuSeqLensQ cannot be 0 "),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cuSeqLensQ.desc, CU_SEQLENS_Q_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
OP_CHECK_IF((opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() != bSize_ + 1),
|
||||
OP_LOGE(opName_, "cuSeqLensQ's shapeSize should be equal to bSize_+1:%u, but got %ld",
|
||||
bSize_ + 1, opParamInfo_.cuSeqLensQ.tensor->GetShapeSize()),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaSequsedKv() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.sequsedKv.tensor == nullptr,
|
||||
OP_LOGE(opName_, "input sequsedKv can not be nullptr, but it's empty"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.sequsedKv.desc, SEQUSED_KV_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.sequsedKv.tensor->GetShapeSize() != bSize_,
|
||||
OP_LOGE(opName_, "input sequsedKv's shapeSize is not equal to B: %u, it is %ld", bSize_, opParamInfo_.sequsedKv.tensor->GetShapeSize()),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaSinks() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.sinks.tensor == nullptr,
|
||||
OP_LOGE(opName_, "Input sinks is nullptr, which is not supported"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.sinks.desc, SINKS_NAME)) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
OP_CHECK_IF(opParamInfo_.sinks.tensor->GetShapeSize() != n1Size_,
|
||||
OP_LOGE(opName_, "Input sinks's shapeSize is not equal to n1: %u, it is %ld.", n1Size_, opParamInfo_.sinks.tensor->GetShapeSize()),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSingleParaMetadata() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.metadata.tensor == nullptr,
|
||||
OP_LOGE(opName_, "Input metadata is required, but got nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASTilingCheck::CheckSinglePara() const
|
||||
{
|
||||
if (ge::GRAPH_SUCCESS != CheckSingleParaQuery() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaKey() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaCmpSparseIndices() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaNumHeads() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaKvHeadNums() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaSparseMode() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaSparseBlockSize() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaBlockTable() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensQ() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaSequsedKv() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaMetadata() ||
|
||||
ge::GRAPH_SUCCESS != CheckSingleParaSinks() ) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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 kvquant_sparse_attn_sharedkv_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class KvQuantSparseAttnSharedkv : public OpDef {
|
||||
public:
|
||||
explicit KvQuantSparseAttnSharedkv(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("q")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("ori_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.IgnoreContiguous();
|
||||
this->Input("cmp_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT8_E4M3FN})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.IgnoreContiguous();
|
||||
this->Input("ori_sparse_indices")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cmp_sparse_indices")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("ori_block_table")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cmp_block_table")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cu_seqlens_q")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cu_seqlens_ori_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cu_seqlens_cmp_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("seqused_q")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("seqused_kv")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("sinks")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("metadata")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Output("attn_out")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16})
|
||||
.Format({ge::FORMAT_ND});
|
||||
this->Output("softmax_lse")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND});
|
||||
this->Attr("kv_quant_mode").AttrType(REQUIRED).Int(1);
|
||||
this->Attr("tile_size").AttrType(OPTIONAL).Int(64); // tile_size默认值64
|
||||
this->Attr("rope_head_dim").AttrType(OPTIONAL).Int(64); // rope_head_dim默认值64
|
||||
this->Attr("softmax_scale").AttrType(REQUIRED).Float(1.0);
|
||||
this->Attr("cmp_ratio").AttrType(REQUIRED).Int(1);
|
||||
this->Attr("ori_mask_mode").AttrType(REQUIRED).Int(4); // ori_mask_mode默认值4
|
||||
this->Attr("cmp_mask_mode").AttrType(REQUIRED).Int(3); // cmp_mask_mode默认值3
|
||||
this->Attr("ori_win_left").AttrType(OPTIONAL).Int(127); // ori_win_left默认值127
|
||||
this->Attr("ori_win_right").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("layout_q").AttrType(OPTIONAL).String("BSND");
|
||||
this->Attr("layout_kv").AttrType(OPTIONAL).String("PA_ND");
|
||||
this->Attr("ori_kv_stride0").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("cmp_kv_stride0").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("return_softmax_lse").AttrType(OPTIONAL).Bool(false);
|
||||
|
||||
OpAICoreConfig aicore_config;
|
||||
aicore_config.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("aclnnSupport.value", "support_aclnn");
|
||||
this->AICore().AddConfig("ascend950", aicore_config);
|
||||
}
|
||||
};
|
||||
OP_ADD(KvQuantSparseAttnSharedkv);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,62 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file sparse_attn_sharedkv_proto.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <register/op_impl_registry.h>
|
||||
#include "error/ops_error.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
constexpr uint32_t QUERY_INPUT_INDEX = 0;
|
||||
constexpr uint32_t RETURN_SOFTMAX_LSE_INDEX = 8;
|
||||
|
||||
ge::graphStatus InferShapeKvQuantSparseAttnSharedkv(gert::InferShapeContext *context)
|
||||
{
|
||||
OPS_ERR_IF(context == nullptr, OPS_LOG_E("KvQuantSparseAttnSharedkv", "InferShapeContext is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
const gert::Shape *queryShape = context->GetInputShape(QUERY_INPUT_INDEX);
|
||||
OPS_LOG_E_IF_NULL(context, queryShape, return ge::GRAPH_FAILED)
|
||||
gert::Shape *attentionOutShape = context->GetOutputShape(0);
|
||||
OPS_LOG_E_IF_NULL(context, attentionOutShape, return ge::GRAPH_FAILED)
|
||||
*attentionOutShape = *queryShape;
|
||||
|
||||
gert::Shape *softmaxLseShape = context->GetOutputShape(1);
|
||||
OPS_LOG_E_IF_NULL(context, attentionOutShape, return ge::GRAPH_FAILED)
|
||||
auto attr = context->GetAttrs();
|
||||
const bool *returnSoftmaxLsePtr = attr->GetAttrPointer<bool>(RETURN_SOFTMAX_LSE_INDEX);
|
||||
bool returnSoftmaxLse = (returnSoftmaxLsePtr != nullptr) ? *returnSoftmaxLsePtr : false;
|
||||
if (returnSoftmaxLse) {
|
||||
*softmaxLseShape = *queryShape;
|
||||
auto lastDimIdx = softmaxLseShape->GetDimNum() - 1;
|
||||
softmaxLseShape->SetDim(lastDimIdx, 1);
|
||||
} else {
|
||||
softmaxLseShape->SetDimNum(1);
|
||||
softmaxLseShape->SetDim(0, 0);
|
||||
}
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus InferDataTypeKvQuantSparseAttnSharedkv(gert::InferDataTypeContext *context)
|
||||
{
|
||||
OPS_ERR_IF(context == nullptr, OPS_LOG_E("KvQuantSparseAttnSharedkv", "InferShapeContext is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
const auto inputDataType = context->GetInputDataType(QUERY_INPUT_INDEX);
|
||||
context->SetOutputDataType(0, inputDataType);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP(KvQuantSparseAttnSharedkv).InferShape(InferShapeKvQuantSparseAttnSharedkv).InferDataType(InferDataTypeKvQuantSparseAttnSharedkv);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,669 @@
|
||||
/**
|
||||
* 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 kvquant_sparse_attn_sharedkv_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
#include "../op_kernel/kv_quant_sparse_attn_sharedkv_template_tiling_key.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_tiling.h"
|
||||
|
||||
using namespace ge;
|
||||
using namespace AscendC;
|
||||
using std::map;
|
||||
using std::string;
|
||||
using std::pair;
|
||||
namespace optiling {
|
||||
|
||||
struct SASCompileInfo {
|
||||
int64_t core_num;
|
||||
};
|
||||
|
||||
// --------------------------KvQuantSASInfoParser类成员函数定义-------------------------------------
|
||||
ge::graphStatus KvQuantSASInfoParser::CheckRequiredInOutExistence() const
|
||||
{
|
||||
OP_CHECK_IF(opParamInfo_.q.shape == nullptr, OP_LOGE(opName_, "Shape of tensor q is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::CheckRequiredAttrExistence() const
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::CheckRequiredParaExistence() const
|
||||
{
|
||||
if (CheckRequiredInOutExistence() != ge::GRAPH_SUCCESS ||
|
||||
CheckRequiredAttrExistence() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetOpName()
|
||||
{
|
||||
if (context_->GetNodeName() == nullptr) {
|
||||
OP_LOGE("KvQuantSparseAttnSharedkv", "opName got from TilingContext is nullptr");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
opName_ = context_->GetNodeName();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::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 ge::GRAPH_FAILED);
|
||||
|
||||
socVersion_ = ascendcPlatform.GetSocVersion();
|
||||
if (socVersion_ != platform_ascendc::SocVersion::ASCEND950) {
|
||||
OP_LOGE(opName_, "SOC Version[%d] is not support.", (int32_t)socVersion_);
|
||||
return GRAPH_FAILED;
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void KvQuantSASInfoParser::GetOptionalInputParaInfo()
|
||||
{
|
||||
opParamInfo_.oriKv.tensor = context_->GetOptionalInputTensor(ORI_KV_INDEX);
|
||||
opParamInfo_.oriKv.desc = context_->GetOptionalInputDesc(ORI_KV_INDEX);
|
||||
opParamInfo_.cmpKv.tensor = context_->GetOptionalInputTensor(CMP_KV_INDEX);
|
||||
opParamInfo_.cmpKv.desc = context_->GetOptionalInputDesc(CMP_KV_INDEX);
|
||||
opParamInfo_.oriSparseIndices.tensor = context_->GetOptionalInputTensor(ORI_SPARSE_INDICES_INDEX);
|
||||
opParamInfo_.oriSparseIndices.desc = context_->GetOptionalInputDesc(ORI_SPARSE_INDICES_INDEX);
|
||||
opParamInfo_.cmpSparseIndices.tensor = context_->GetOptionalInputTensor(CMP_SPARSE_INDICES_INDEX);
|
||||
opParamInfo_.cmpSparseIndices.desc = context_->GetOptionalInputDesc(CMP_SPARSE_INDICES_INDEX);
|
||||
opParamInfo_.oriBlockTable.tensor = context_->GetOptionalInputTensor(ORI_BLOCK_TABLE_INDEX);
|
||||
opParamInfo_.oriBlockTable.desc = context_->GetOptionalInputDesc(ORI_BLOCK_TABLE_INDEX);
|
||||
opParamInfo_.cmpBlockTable.tensor = context_->GetOptionalInputTensor(CMP_BLOCK_TABLE_INDEX);
|
||||
opParamInfo_.cmpBlockTable.desc = context_->GetOptionalInputDesc(CMP_BLOCK_TABLE_INDEX);
|
||||
opParamInfo_.sinks.tensor = context_->GetOptionalInputTensor(SINKS_INDEX);
|
||||
opParamInfo_.sinks.desc = context_->GetOptionalInputDesc(SINKS_INDEX);
|
||||
opParamInfo_.cuSeqLensQ.tensor = context_->GetOptionalInputTensor(CU_SEQLENS_Q_INDEX);
|
||||
opParamInfo_.cuSeqLensQ.desc = context_->GetOptionalInputDesc(CU_SEQLENS_Q_INDEX);
|
||||
opParamInfo_.cuSeqLensOriKv.tensor = context_->GetOptionalInputTensor(CU_SEQLENS_ORI_KV_INDEX);
|
||||
opParamInfo_.cuSeqLensOriKv.desc = context_->GetOptionalInputDesc(CU_SEQLENS_ORI_KV_INDEX);
|
||||
opParamInfo_.cuSeqLensCmpKv.tensor = context_->GetOptionalInputTensor(CU_SEQLENS_CMP_KV_INDEX);
|
||||
opParamInfo_.cuSeqLensCmpKv.desc = context_->GetOptionalInputDesc(CU_SEQLENS_CMP_KV_INDEX);
|
||||
opParamInfo_.seqUsedQ.tensor = context_->GetOptionalInputTensor(SEQUSED_Q_INDEX);
|
||||
opParamInfo_.seqUsedQ.desc = context_->GetOptionalInputDesc(SEQUSED_Q_INDEX);
|
||||
opParamInfo_.sequsedKv.tensor = context_->GetOptionalInputTensor(SEQUSED_KV_INDEX);
|
||||
opParamInfo_.sequsedKv.desc = context_->GetOptionalInputDesc(SEQUSED_KV_INDEX);
|
||||
opParamInfo_.metadata.desc = context_->GetOptionalInputDesc(METADATA_INDEX);
|
||||
opParamInfo_.metadata.tensor = context_->GetOptionalInputTensor(METADATA_INDEX);
|
||||
}
|
||||
|
||||
void KvQuantSASInfoParser::GetInputParaInfo()
|
||||
{
|
||||
opParamInfo_.q.desc = context_->GetInputDesc(Q_INDEX);
|
||||
opParamInfo_.q.shape = context_->GetInputShape(Q_INDEX);
|
||||
GetOptionalInputParaInfo();
|
||||
}
|
||||
|
||||
void KvQuantSASInfoParser::GetOutputParaInfo()
|
||||
{
|
||||
opParamInfo_.attnOut.desc = context_->GetOutputDesc(ATTN_OUT_INDEX);
|
||||
opParamInfo_.attnOut.shape = context_->GetOutputShape(ATTN_OUT_INDEX);
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetAttrParaInfo()
|
||||
{
|
||||
auto attrs = context_->GetAttrs();
|
||||
OP_CHECK_IF(attrs == nullptr, OPS_REPORT_VECTOR_INNER_ERR(context_->GetNodeName(), "attrs got from ge is nullptr"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo start");
|
||||
opParamInfo_.kvQuantMode = attrs->GetAttrPointer<int64_t>(ATTR_KV_QUANT_SCALE_INDEX);
|
||||
opParamInfo_.tileSize = attrs->GetAttrPointer<int64_t>(ATTR_TILE_SIZE_INDEX);
|
||||
opParamInfo_.ropeHeadDim = attrs->GetAttrPointer<int64_t>(ATTR_ROPE_HEAD_DIM_INDEX);
|
||||
opParamInfo_.softmaxScale = attrs->GetAttrPointer<float>(ATTR_SOTFMAX_SCALE_INDEX);
|
||||
opParamInfo_.oriKvStride = attrs->GetAttrPointer<int64_t>(ATTR_ORIKV_STRIDE_INDEX);
|
||||
opParamInfo_.cmpKvStride = attrs->GetAttrPointer<int64_t>(ATTR_CMPKV_STRIDE_INDEX);
|
||||
opParamInfo_.cmpRatio = attrs->GetAttrPointer<int64_t>(ATTR_CMP_RATIO_INDEX);
|
||||
opParamInfo_.oriMaskMode = attrs->GetAttrPointer<uint32_t>(ATTR_ORI_MASK_MODE_INDEX);
|
||||
opParamInfo_.cmpMaskMode = attrs->GetAttrPointer<uint32_t>(ATTR_CMP_MASK_MODE_INDEX);
|
||||
opParamInfo_.oriWinLeft = attrs->GetAttrPointer<int64_t>(ATTR_ORI_WIN_LEFT_INDEX);
|
||||
opParamInfo_.oriWinRight = attrs->GetAttrPointer<int64_t>(ATTR_ORI_WIN_RIGHT_INDEX);
|
||||
opParamInfo_.layoutQ = attrs->GetStr(ATTR_LAYOUT_Q_INDEX);
|
||||
opParamInfo_.layoutKv = attrs->GetStr(ATTR_LAYOUT_KV_INDEX);
|
||||
|
||||
OP_LOGI(context_->GetNodeName(), "GetAttrParaInfo end");
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetOpParaInfo()
|
||||
{
|
||||
GetInputParaInfo();
|
||||
GetOutputParaInfo();
|
||||
if (ge::GRAPH_SUCCESS != GetAttrParaInfo()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetInOutDataType()
|
||||
{
|
||||
qType_ = opParamInfo_.q.desc->GetDataType();
|
||||
outputType_ = opParamInfo_.attnOut.desc->GetDataType();
|
||||
if (opParamInfo_.oriKv.desc != nullptr) {
|
||||
oriKvType_ = opParamInfo_.oriKv.desc->GetDataType();
|
||||
}
|
||||
if (opParamInfo_.cmpKv.desc != nullptr) {
|
||||
cmpKvType_ = opParamInfo_.cmpKv.desc->GetDataType();
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetQueryAndOutLayout()
|
||||
{
|
||||
// 获取q和attnOut的Layout基准值
|
||||
// layoutQuery: {qLayout, outLayout}
|
||||
const map<string, pair<SASLayout, SASLayout>> layoutMap = {
|
||||
{"BSND", {SASLayout::BSND, SASLayout::BSND}},
|
||||
{"TND", {SASLayout::TND, SASLayout::TND }},
|
||||
};
|
||||
|
||||
std::string layout(opParamInfo_.layoutQ);
|
||||
auto it = layoutMap.find(layout);
|
||||
if (it != layoutMap.end()) {
|
||||
qLayout_ = it->second.first;
|
||||
outLayout_ = it->second.second;
|
||||
} else {
|
||||
OP_LOGE(opName_, "layout of Q is %s, it is unsupported.", layout.c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetKvLayout()
|
||||
{
|
||||
const map<string, SASLayout> layoutKVMap = {
|
||||
{"PA_ND", SASLayout::PA_ND},
|
||||
};
|
||||
|
||||
std::string layout(opParamInfo_.layoutKv);
|
||||
auto it = layoutKVMap.find(layout);
|
||||
if (it != layoutKVMap.end()) {
|
||||
kvLayout_ = it->second;
|
||||
} else {
|
||||
OP_LOGE(opName_, "layoutKV is %s, it is unsupported.", layout.c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
// =============Parser function====================
|
||||
|
||||
bool KvQuantSASInfoParser::HasAxis(const SASAxis &axis, const SASLayout &layout, const gert::Shape &shape) const
|
||||
{
|
||||
const auto& layoutIt = SAS_LAYOUT_AXIS_MAP.find(layout);
|
||||
if (layoutIt == SAS_LAYOUT_AXIS_MAP.end()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const std::vector<SASAxis>& axes = layoutIt->second;
|
||||
const auto& axisIt = std::find(axes.begin(), axes.end(), axis);
|
||||
if (axisIt == axes.end()) {
|
||||
return false;
|
||||
}
|
||||
const auto& dimIt = SAS_LAYOUT_DIM_MAP.find(layout);
|
||||
if (dimIt == SAS_LAYOUT_DIM_MAP.end() || dimIt->second != shape.GetDimNum()) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t KvQuantSASInfoParser::GetAxisIdx(const SASAxis &axis, const SASLayout &layout) const
|
||||
{
|
||||
const std::vector<SASAxis>& axes = SAS_LAYOUT_AXIS_MAP.find(layout)->second;
|
||||
const auto& axisIt = std::find(axes.begin(), axes.end(), axis);
|
||||
return std::distance(axes.begin(), axisIt);
|
||||
}
|
||||
|
||||
uint32_t KvQuantSASInfoParser::GetAxisNum(const gert::Shape &shape, const SASAxis &axis,const SASLayout &layout) const
|
||||
{
|
||||
return HasAxis(axis, layout, shape) ? shape.GetDim(GetAxisIdx(axis, layout)) : invalidDimValue_;
|
||||
}
|
||||
|
||||
void KvQuantSASInfoParser::SetSASShape()
|
||||
{
|
||||
qShape_ = opParamInfo_.q.shape->GetStorageShape();
|
||||
if (opParamInfo_.oriKv.tensor != nullptr) {
|
||||
oriKvShape_ = opParamInfo_.oriKv.tensor->GetStorageShape();
|
||||
}
|
||||
if (opParamInfo_.cmpKv.tensor != nullptr) {
|
||||
cmpKvShape_ = opParamInfo_.cmpKv.tensor->GetStorageShape();
|
||||
}
|
||||
if (opParamInfo_.cmpSparseIndices.tensor != nullptr) {
|
||||
cmpSparseIndicesShape_ = opParamInfo_.cmpSparseIndices.tensor->GetStorageShape();
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetN1Size()
|
||||
{
|
||||
n1Size_ = GetAxisNum(qShape_, SASAxis::N, qLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetN2Size()
|
||||
{
|
||||
if (opParamInfo_.oriKv.tensor != nullptr) {
|
||||
n2Size_ = GetAxisNum(oriKvShape_, SASAxis::N, kvLayout_);
|
||||
} else if (opParamInfo_.cmpKv.tensor != nullptr) {
|
||||
n2Size_ = GetAxisNum(cmpKvShape_, SASAxis::N, kvLayout_);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetGSize()
|
||||
{
|
||||
if (n2Size_ != 0) {
|
||||
gSize_ = n1Size_ / n2Size_;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor,
|
||||
SASLayout &layout, const std::string &name) const
|
||||
{
|
||||
if ((tensor == nullptr)) {
|
||||
OP_LOGE(opName_, "when layout of q is %s, %s must be provided.",
|
||||
KvQuantSASLayoutToSerialString(layout).c_str(), name.c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
int64_t shapeSize = tensor->GetShapeSize();
|
||||
if (shapeSize <= 0) {
|
||||
OP_LOGE(opName_, "the shape size of %s is %ld, it should be greater than 0.",
|
||||
name.c_str(), shapeSize);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
size = static_cast<uint32_t>(shapeSize);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetActualSeqLenQSize(uint32_t &size)
|
||||
{
|
||||
return GetActualSeqLenSize(size, opParamInfo_.sequsedKv.tensor, qLayout_, "cuSeqLensQ");
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetBatchSize()
|
||||
{
|
||||
// 获取B基准值
|
||||
// 1、非TND时, 以query的batch_size维度为基准;
|
||||
// 2、TND时, actual_seq_lens_q必须传入, 以actual_seq_lens_q数组的长度为B轴大小
|
||||
if (qLayout_ == SASLayout::TND) {
|
||||
return GetActualSeqLenQSize(bSize_);
|
||||
} else { // BSND
|
||||
bSize_ = GetAxisNum(qShape_, SASAxis::B, qLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetQTSize()
|
||||
{
|
||||
// 获取query的T基准值
|
||||
// 1、非TND时, 以query的batch_size维度为基准;
|
||||
// 2、TND时, actual_seq_lens_q必须传入, 以actual_seq_lens_q数组的长度为B轴大小
|
||||
qTSize_ = (qLayout_ == SASLayout::TND) ? GetAxisNum(qShape_, SASAxis::T, qLayout_) : 0;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetS1Size()
|
||||
{
|
||||
// 获取S1基准值
|
||||
// 1、非TND时, 以query的S维度为基准;
|
||||
// 2、TND时, actual_seq_lens_q必须传入, 以actual_seq_lens_q数组中的最大值为基准
|
||||
if (qLayout_ == SASLayout::TND) {
|
||||
s1Size_ = GetAxisNum(qShape_, SASAxis::T, qLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
} else { // BSND
|
||||
s1Size_ = GetAxisNum(qShape_, SASAxis::S, qLayout_);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetMaxBlockNumPerBatch()
|
||||
{
|
||||
if (opParamInfo_.oriBlockTable.tensor == nullptr) {
|
||||
OP_LOGE(opName_, "the layout_kv is %s, blockTable must be provided.", KvQuantSASLayoutToSerialString(kvLayout_).c_str());
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
uint32_t oriDimNum = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDimNum();
|
||||
if (oriDimNum != DIM_NUM_TWO) {
|
||||
OP_LOGE(opName_, "the dim num of ori_block_table is %u, it should be %u.", oriDimNum, DIM_NUM_TWO);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1) <= 0) {
|
||||
OP_LOGE(opName_, "%s's second dimension(%ld) should be greater than 0",
|
||||
ORI_BLOCK_TABLE_NAME.c_str(), opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
oriMaxBlockNumPerBatch_ = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1);
|
||||
|
||||
if (opParamInfo_.cmpBlockTable.tensor != nullptr) {
|
||||
uint32_t cmpDimNum = opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDimNum();
|
||||
if (cmpDimNum != DIM_NUM_TWO) {
|
||||
OP_LOGE(opName_, "the dim num of cmp_block_table is %u, it should be %u.", cmpDimNum, DIM_NUM_TWO);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
if (opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1) <= 0) {
|
||||
OP_LOGE(opName_, "%s's second dimension(%ld) should be greater than 0",
|
||||
CMP_BLOCK_TABLE_NAME.c_str(), opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1));
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
cmpMaxBlockNumPerBatch_ = opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetBlockSize()
|
||||
{
|
||||
if (opParamInfo_.oriKv.tensor != nullptr) {
|
||||
oriBlockSize_ = GetAxisNum(oriKvShape_, SASAxis::Bs, kvLayout_);
|
||||
}
|
||||
if (opParamInfo_.cmpKv.tensor != nullptr) {
|
||||
cmpBlockSize_ = GetAxisNum(cmpKvShape_, SASAxis::Bs, kvLayout_);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetS2SizeForPageAttention()
|
||||
{
|
||||
if (GetMaxBlockNumPerBatch() != ge::GRAPH_SUCCESS || GetBlockSize() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
s2Size_ = oriMaxBlockNumPerBatch_ * oriBlockSize_;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetS2Size()
|
||||
{
|
||||
// 获取S2基准值:PAGE_ATTENTION时, S2 = block_table.dim1 * block_size
|
||||
return GetS2SizeForPageAttention();
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetQkHeadDim()
|
||||
{
|
||||
// 获取qkHeadDim基准值
|
||||
// 以query的D维度为基准
|
||||
qkHeadDim_ = GetAxisNum(qShape_, SASAxis::D, qLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetSparseBlockCount()
|
||||
{
|
||||
if (opParamInfo_.cmpSparseIndices.tensor != nullptr) {
|
||||
sparseBlockCount_ = GetAxisNum(cmpSparseIndicesShape_, SASAxis::K, qLayout_);
|
||||
}
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetActualseqInfo()
|
||||
{
|
||||
maxActualseq_ = static_cast<uint32_t>(s2Size_);
|
||||
if (opParamInfo_.sequsedKv.tensor != nullptr) {
|
||||
actualLenDimsKV_ = opParamInfo_.sequsedKv.tensor->GetShapeSize();
|
||||
}
|
||||
if (opParamInfo_.cuSeqLensQ.tensor != nullptr) {
|
||||
actualLenDimsQ_ = opParamInfo_.cuSeqLensQ.tensor->GetShapeSize(); // cuSeqLensQ shape is B+1
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetDSizeQ() {
|
||||
dSizeQ_ = GetAxisNum(qShape_, SASAxis::D, qLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetDSizeKV() {
|
||||
dSizeKV_ = GetAxisNum(oriKvShape_, SASAxis::D, kvLayout_);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::GetSinks()
|
||||
{
|
||||
if(opParamInfo_.sequsedKv.tensor != nullptr){
|
||||
uint32_t oriDimNum = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDimNum();
|
||||
if(oriDimNum != DIM_NUM_ONE){
|
||||
OP_LOGE(opName_, "the dim num of sinks is %u, it should be %u.", oriDimNum, DIM_NUM_ONE);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
int64_t oriDimension = opParamInfo_.sequsedKv.tensor->GetStorageShape().GetDim(0);
|
||||
if(oriDimension != gSize_){
|
||||
OP_LOGE(opName_, "sinks's dimension(%ld) should be equal to query head num(%u).", oriDimension, gSize_);
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
void KvQuantSASInfoParser::GenerateInfo(KvQuantSASTilingInfo &sasInfo)
|
||||
{
|
||||
sasInfo.opName = opName_;
|
||||
sasInfo.platformInfo = platformInfo_;
|
||||
sasInfo.opParamInfo = opParamInfo_;
|
||||
sasInfo.socVersion = socVersion_;
|
||||
|
||||
sasInfo.bSize = bSize_;
|
||||
sasInfo.n1Size = n1Size_;
|
||||
sasInfo.n2Size = n2Size_;
|
||||
sasInfo.s1Size = s1Size_;
|
||||
sasInfo.s2Size = s2Size_;
|
||||
sasInfo.gSize = gSize_;
|
||||
sasInfo.qkHeadDim = qkHeadDim_;
|
||||
sasInfo.qTSize = qTSize_;
|
||||
sasInfo.sparseBlockCount = sparseBlockCount_;
|
||||
|
||||
sasInfo.qType = qType_;
|
||||
sasInfo.oriKvType = oriKvType_;
|
||||
sasInfo.cmpKvType = cmpKvType_;
|
||||
sasInfo.outputType = outputType_;
|
||||
sasInfo.dSize = dSizeQ_;
|
||||
sasInfo.dSizeV = 512;
|
||||
sasInfo.dSizeVInput = dSizeKV_;
|
||||
|
||||
sasInfo.totalBlockNum = (opParamInfo_.oriKv.tensor != nullptr) ?
|
||||
opParamInfo_.oriKv.tensor->GetStorageShape().GetDim(0) : 0;
|
||||
sasInfo.sparseBlockSize = 1; // 写死为1
|
||||
sasInfo.oriBlockSize = oriBlockSize_;
|
||||
sasInfo.cmpBlockSize = cmpBlockSize_;
|
||||
sasInfo.blockTypeSize = sizeof(float);
|
||||
sasInfo.oriMaxBlockNumPerBatch = oriMaxBlockNumPerBatch_;
|
||||
sasInfo.cmpMaxBlockNumPerBatch = cmpMaxBlockNumPerBatch_;
|
||||
|
||||
sasInfo.actualLenDimsQ = actualLenDimsQ_;
|
||||
sasInfo.actualLenDimsKV = actualLenDimsKV_;
|
||||
sasInfo.maxActualseq = maxActualseq_;
|
||||
sasInfo.actualSeqLenFlag = (opParamInfo_.sequsedKv.tensor != nullptr);
|
||||
sasInfo.isSameSeqAllKVTensor = isSameSeqAllKVTensor_;
|
||||
|
||||
sasInfo.kvQuantMode = *opParamInfo_.kvQuantMode;
|
||||
sasInfo.tileSize = *opParamInfo_.tileSize;
|
||||
sasInfo.ropeHeadDim = *opParamInfo_.ropeHeadDim;
|
||||
sasInfo.softmaxScale = *opParamInfo_.softmaxScale;
|
||||
sasInfo.oriKvStride = *opParamInfo_.oriKvStride;
|
||||
sasInfo.cmpKvStride = *opParamInfo_.cmpKvStride;
|
||||
sasInfo.cmpRatio = *opParamInfo_.cmpRatio;
|
||||
sasInfo.oriMaskMode = *opParamInfo_.oriMaskMode;
|
||||
sasInfo.cmpMaskMode = *opParamInfo_.cmpMaskMode;
|
||||
sasInfo.oriWinLeft = *opParamInfo_.oriWinLeft;
|
||||
sasInfo.oriWinRight = *opParamInfo_.oriWinRight;
|
||||
|
||||
sasInfo.qLayout = qLayout_;
|
||||
sasInfo.kvLayout = kvLayout_;
|
||||
sasInfo.outLayout = outLayout_;
|
||||
}
|
||||
|
||||
ge::graphStatus KvQuantSASInfoParser::Parse(KvQuantSASTilingInfo &sasInfo)
|
||||
{
|
||||
if (context_ == nullptr) {
|
||||
OP_LOGE("SparseFlashAttention", "tiling context is nullptr!");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
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 != GetInOutDataType() ||
|
||||
ge::GRAPH_SUCCESS != GetQueryAndOutLayout() ||
|
||||
ge::GRAPH_SUCCESS != GetKvLayout()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
SetSASShape();
|
||||
if (
|
||||
ge::GRAPH_SUCCESS != GetN1Size() ||
|
||||
ge::GRAPH_SUCCESS != GetN2Size() ||
|
||||
ge::GRAPH_SUCCESS != GetGSize() ||
|
||||
ge::GRAPH_SUCCESS != GetBatchSize() ||
|
||||
ge::GRAPH_SUCCESS != GetQTSize() ||
|
||||
ge::GRAPH_SUCCESS != GetS1Size() ||
|
||||
ge::GRAPH_SUCCESS != GetS2Size() ||
|
||||
ge::GRAPH_SUCCESS != GetQkHeadDim() ||
|
||||
ge::GRAPH_SUCCESS != GetSparseBlockCount() ||
|
||||
ge::GRAPH_SUCCESS != GetDSizeQ() ||
|
||||
ge::GRAPH_SUCCESS != GetDSizeKV()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
if (ge::GRAPH_SUCCESS != GetActualseqInfo()) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
GenerateInfo(sasInfo);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
// --------------------------TilingPrepare函数定义-------------------------------------
|
||||
static ge::graphStatus TilingPrepareForKvQuantSparseAttnSharedkv(gert::TilingParseContext * /* context */)
|
||||
{
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
// --------------------------SparseAttnSharedkvTiling类成员函数定义-----------------------
|
||||
ge::graphStatus KvQuantSparseAttnSharedkvTiling::DoOpTiling(KvQuantSASTilingInfo *tilingInfo)
|
||||
{
|
||||
if (tilingInfo->opParamInfo.cmpKv.tensor == nullptr) {
|
||||
OP_CHECK_IF(tilingInfo->opParamInfo.cmpSparseIndices.tensor != nullptr,
|
||||
OP_LOGE("KvQuantSparseAttnSharedkv", "cmpSparseIndices must be empty when cmpKv is not provided."),
|
||||
return ge::GRAPH_FAILED);
|
||||
perfMode_ = SASTemplateMode::SWA_TEMPLATE_MODE;
|
||||
} else if (tilingInfo->opParamInfo.cmpSparseIndices.tensor != nullptr) {
|
||||
perfMode_ = SASTemplateMode::SCFA_TEMPLATE_MODE;
|
||||
} else {
|
||||
perfMode_ = SASTemplateMode::CFA_TEMPLATE_MODE;
|
||||
}
|
||||
// -------------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);
|
||||
OP_LOGI(tilingInfo->opName, "SAS block dim: %u aiv Num: %u aic Num: %u.", blockDim, aivNum, aicNum);
|
||||
|
||||
// -------------set workspacesize-----------------
|
||||
constexpr uint32_t TRIPLE_BUFFER_NUM = 3;
|
||||
constexpr uint32_t M_BASE_SIZE = 64; // m轴基本块大小
|
||||
constexpr uint32_t S2_BASE_SIZE = 128; // S2轴基本块大小
|
||||
constexpr uint32_t D_SIZE = 512;
|
||||
constexpr uint32_t VEC_RES_ELEM_SIZE = 2; // 2: fp16/bf16
|
||||
constexpr uint32_t TOPK_MAX_SIZE = 2048; // TopK选取个数
|
||||
uint32_t workspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize();
|
||||
if (tilingInfo->gSize > 64) {
|
||||
workspaceSize += (S2_BASE_SIZE * D_SIZE * VEC_RES_ELEM_SIZE * TRIPLE_BUFFER_NUM * (aicNum >> 1));
|
||||
}
|
||||
size_t *workSpaces = context_->GetWorkspaceSizes(1);
|
||||
workSpaces[0] = workspaceSize;
|
||||
|
||||
// -------------set tilingdata-----------------
|
||||
tilingData_.baseParams.set_batchSize(tilingInfo->bSize);
|
||||
tilingData_.baseParams.set_kvSeqSize(tilingInfo->s2Size);
|
||||
tilingData_.baseParams.set_qSeqSize(tilingInfo->s1Size);
|
||||
tilingData_.baseParams.set_sparseBlockCount(tilingInfo->sparseBlockCount);
|
||||
tilingData_.baseParams.set_nNumOfQInOneGroup(tilingInfo->gSize);
|
||||
tilingData_.baseParams.set_paOriBlockSize(tilingInfo->oriBlockSize);
|
||||
tilingData_.baseParams.set_paCmpBlockSize(tilingInfo->cmpBlockSize);
|
||||
tilingData_.baseParams.set_oriMaxBlockNumPerBatch(tilingInfo->oriMaxBlockNumPerBatch);
|
||||
tilingData_.baseParams.set_cmpMaxBlockNumPerBatch(tilingInfo->cmpMaxBlockNumPerBatch);
|
||||
|
||||
tilingData_.baseParams.set_tileSize(tilingInfo->tileSize);
|
||||
tilingData_.baseParams.set_ropeHeadDim(tilingInfo->ropeHeadDim);
|
||||
tilingData_.baseParams.set_softmaxScale(tilingInfo->softmaxScale);
|
||||
tilingData_.baseParams.set_oriKvStride(tilingInfo->oriKvStride);
|
||||
tilingData_.baseParams.set_cmpKvStride(tilingInfo->cmpKvStride);
|
||||
tilingData_.baseParams.set_cmpRatio(tilingInfo->cmpRatio);
|
||||
tilingData_.baseParams.set_oriMaskMode(tilingInfo->oriMaskMode);
|
||||
tilingData_.baseParams.set_cmpMaskMode(tilingInfo->cmpMaskMode);
|
||||
tilingData_.baseParams.set_oriWinLeft(tilingInfo->oriWinLeft);
|
||||
tilingData_.baseParams.set_oriWinRight(tilingInfo->oriWinRight);
|
||||
tilingData_.baseParams.set_sparseBlockSize(tilingInfo->sparseBlockSize);
|
||||
tilingData_.baseParams.set_dSize(tilingInfo->dSize);
|
||||
tilingData_.baseParams.set_dSizeVInput(tilingInfo->dSizeVInput);
|
||||
|
||||
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 qType = static_cast<uint32_t>(tilingInfo->qType);
|
||||
uint32_t oriKvType = static_cast<uint32_t>(tilingInfo->oriKvType);
|
||||
uint32_t outputType = static_cast<uint32_t>(tilingInfo->outputType);
|
||||
uint32_t qLayout = static_cast<uint32_t>(tilingInfo->qLayout);
|
||||
uint32_t inputKvLayout = static_cast<uint32_t>(tilingInfo->kvLayout);
|
||||
uint32_t tilingKey =
|
||||
GET_TPL_TILING_KEY(0U, qLayout, inputKvLayout, static_cast<uint32_t>(perfMode_), static_cast<uint32_t>(tilingInfo->gSize > 64));
|
||||
context_->SetTilingKey(tilingKey);
|
||||
context_->SetScheduleMode(1);
|
||||
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
// --------------------------Tiling函数定义---------------------------
|
||||
ge::graphStatus TilingKvQuantSparseAttnSharedkv(gert::TilingContext *context)
|
||||
{
|
||||
OP_CHECK_IF(context == nullptr, OPS_REPORT_VECTOR_INNER_ERR("KvQuantSparseAttnSharedkv", "Tiling context is null."),
|
||||
return ge::GRAPH_FAILED);
|
||||
KvQuantSASTilingInfo sasInfo;
|
||||
KvQuantSASInfoParser sasInfoParser(context);
|
||||
if (sasInfoParser.Parse(sasInfo) != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
KvQuantSASTilingCheck sasTilingChecker(sasInfo);
|
||||
if (sasTilingChecker.Process() != ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
KvQuantSparseAttnSharedkvTiling tiling(context);
|
||||
return tiling.DoOpTiling(&sasInfo);
|
||||
}
|
||||
// --------------------------Tiling函数及TilingPrepare函数注册--------
|
||||
IMPL_OP_OPTILING(KvQuantSparseAttnSharedkv)
|
||||
.Tiling(TilingKvQuantSparseAttnSharedkv)
|
||||
.TilingParse<SASCompileInfo>(TilingPrepareForKvQuantSparseAttnSharedkv);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,85 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file sparse_attn_sharedkv_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef KV_QUANT_SPARSE_ATTN_SHAREDKV_TILING_H
|
||||
#define KV_QUANT_SPARSE_ATTN_SHAREDKV_TILING_H
|
||||
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include <exe_graph/runtime/tiling_context.h>
|
||||
#include <tiling/platform/platform_ascendc.h>
|
||||
#include "register/tilingdata_base.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "tiling/tiling_api.h"
|
||||
#include "log/log.h"
|
||||
#include "log/error_code.h"
|
||||
#include "err/ops_err.h"
|
||||
#include "platform/platform_info.h"
|
||||
#include "kv_quant_sparse_attn_sharedkv_check.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
std::string KvQuantSASLayoutToSerialString(SASLayout layout);
|
||||
|
||||
// -----------算子TilingData定义---------------
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseAttnSharedkvBaseParams)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, batchSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, qSeqSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, kvSeqSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, paOriBlockSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, paCmpBlockSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, oriMaxBlockNumPerBatch)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, cmpMaxBlockNumPerBatch)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, nNumOfQInOneGroup)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sparseBlockCount)
|
||||
TILING_DATA_FIELD_DEF(float, softmaxScale) // 即 scaleValue
|
||||
TILING_DATA_FIELD_DEF(int32_t, oriKvStride)
|
||||
TILING_DATA_FIELD_DEF(int32_t, cmpKvStride)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, tileSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, ropeHeadDim)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, cmpRatio)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, oriMaskMode)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, cmpMaskMode)
|
||||
TILING_DATA_FIELD_DEF(int32_t, oriWinLeft)
|
||||
TILING_DATA_FIELD_DEF(int32_t, oriWinRight)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, sparseBlockSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dSize)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dSizeVInput)
|
||||
END_TILING_DATA_DEF
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseAttnSharedkvBaseParamsOp, KvQuantSparseAttnSharedkvBaseParams)
|
||||
|
||||
BEGIN_TILING_DATA_DEF(KvQuantSparseAttnSharedkvTilingData)
|
||||
TILING_DATA_FIELD_DEF_STRUCT(KvQuantSparseAttnSharedkvBaseParams, baseParams);
|
||||
END_TILING_DATA_DEF
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(KvQuantSparseAttnSharedkv, KvQuantSparseAttnSharedkvTilingData)
|
||||
|
||||
// ---------------算子Tiling类---------------
|
||||
class KvQuantSparseAttnSharedkvTiling {
|
||||
public:
|
||||
explicit KvQuantSparseAttnSharedkvTiling(gert::TilingContext *context) : context_(context){};
|
||||
ge::graphStatus DoOpTiling(KvQuantSASTilingInfo *tilingInfo);
|
||||
|
||||
private:
|
||||
gert::TilingContext *context_ = nullptr;
|
||||
SASTemplateMode perfMode_ = SASTemplateMode::SWA_TEMPLATE_MODE;
|
||||
KvQuantSparseAttnSharedkvTilingData tilingData_;
|
||||
uint32_t blockDim_{0};
|
||||
uint64_t workspaceSize_{0};
|
||||
uint64_t tilingKey_{0};
|
||||
|
||||
KvQuantSASTilingInfo *sasInfo_ = nullptr;
|
||||
};
|
||||
|
||||
}
|
||||
#endif
|
||||
Reference in New Issue
Block a user