init v0.23.0

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

View File

@@ -0,0 +1,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()

View File

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

View File

@@ -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 &param, const SASLayout &layout) const;
ge::graphStatus CompareShape(KvQuantSASTilingShapeCompareParam &param,
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

View File

@@ -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 &param, 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 &param,
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;
}
}

View File

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

View File

@@ -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();
}
}

View File

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

View File

@@ -0,0 +1,128 @@
/**
 * Copyright (c) 2026 Huawei Technologies Co., Ltd.
 * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
 * CANN Open Software License Agreement Version 2.0 (the "License").
 * Please refer to the License for details. You may not use this file except in compliance with the License.
 * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
 * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
 * See LICENSE in the root of the software repository for the full text of the License.
 */
/*!
* \file 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

View File

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

View File

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

View File

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