[INFO] Initializing environment for https://gitcode.com/pre-commit/pre-commit-hooks. [WARNING] repo `https://gitcode.com/pre-commit/pre-commit-hooks` uses deprecated stage names (commit, push) which will be removed in a future version. Hint: often `pre-commit autoupdate --repo https://gitcode.com/pre-commit/pre-commit-hooks` will fix this. if it does not -- consider reporting an issue to that repo. [INFO] Initializing environment for https://gitcode.com/pre-commit-clang/mirrors-clang-format. [INFO] Initializing environment for https://gitcode.com/gh_mirrors/ru/ruff-pre-commit. [INFO] Initializing environment for https://gitcode.com/gh_mirrors/co/codespell. [INFO] Installing environment for https://gitcode.com/pre-commit/pre-commit-hooks. [INFO] Once installed this environment will be reused. [INFO] This may take a few minutes... [INFO] Installing environment for https://gitcode.com/pre-commit-clang/mirrors-clang-format. [INFO] Once installed this environment will be reused. [INFO] This may take a few minutes... [INFO] Installing environment for https://gitcode.com/gh_mirrors/ru/ruff-pre-commit. [INFO] Once installed this environment will be reused. [INFO] This may take a few minutes... [INFO] Installing environment for https://gitcode.com/gh_mirrors/co/codespell. [INFO] Once installed this environment will be reused. [INFO] This may take a few minutes... trim trailing whitespace.................................................Failed - hook id: trailing-whitespace - exit code: 1 - files were modified by this hook Fixing torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py fix end of files.........................................................Passed check yaml...........................................(no files to check)Skipped check for added large files..............................................Passed check for merge conflicts................................................Passed detect private key.......................................................Passed check json...........................................(no files to check)Skipped clang-format.............................................................Failed - hook id: clang-format - files were modified by this hook Formatting [1/1] attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp ruff check...............................................................Failed - hook id: ruff-check - exit code: 1 - files were modified by this hook ::error title=Ruff (F841),file=/opt/cloud/agent_1766041207435_sAxVW/workspace/j_tuk1mAI2/torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py,line=95,col=13,endLine=95,endColumn=24::torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py:95:13: F841 Local variable `key_headnum` is assigned to but never used ruff format..............................................................Failed - hook id: ruff-format - files were modified by this hook 1 file reformatted codespell................................................................Passed All changes made by hooks: diff --git a/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp b/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp index 1b351de..b7fb735 100644 --- a/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp +++ b/attention/sparse_flash_mla/op_host/sparse_flash_mla_tiling.cpp @@ -19,8 +19,8 @@ using namespace ge; using namespace AscendC; using std::map; -using std::string; using std::pair; +using std::string; namespace optiling { @@ -33,56 +33,55 @@ static const std::string ORI_SPARSE_INDICES = "ori_sparse_indices"; static const std::string CMP_SPARSE_INDICES = "cmp_sparse_indices"; static const std::string ORI_BLOCK_TABLE_NAME = "ori_block_table"; static const std::string CMP_BLOCK_TABLE_NAME = "cmp_block_table"; -static const std::string SINKS_NAME = "sinks"; -static const std::string METADATA_NAME = "metadata"; -static const std::string ATTEN_OUT_NAME = "attn_out"; -static const std::string A2_A3_PLATFORM_LOG = "A2/A3"; -static const std::string A5_PLATFORM_LOG = "A5"; -static bool IsNonEmptyOptionalTensor(const gert::Tensor *tensor) -{ - return tensor != nullptr && tensor->GetShapeSize() > 0; -} - -static bool IsPowerOfTwoInRange(uint32_t value, uint32_t minValue, uint32_t maxValue) -{ - return value >= minValue && value <= maxValue && (value & (value - 1U)) == 0U; -} - -static bool IsA5Arch(NpuArch npuArch) -{ - return npuArch == NpuArch::DAV_3510; -} - -static bool IsPaBlockSizeSupport(NpuArch npuArch, int32_t blockSize) -{ - if (IsA5Arch(npuArch)) { - return blockSize >= 1 && blockSize <= static_cast(BLOCK_SIZE_LIMIT); - } - return blockSize >= 16 && blockSize <= static_cast(BLOCK_SIZE_LIMIT) && blockSize % 16 == 0; -} +static const std::string SINKS_NAME = "sinks"; +static const std::string METADATA_NAME = "metadata"; +static const std::string ATTEN_OUT_NAME = "attn_out"; +static const std::string A2_A3_PLATFORM_LOG = "A2/A3"; +static const std::string A5_PLATFORM_LOG = "A5"; +static bool IsNonEmptyOptionalTensor(const gert::Tensor *tensor) +{ + return tensor != nullptr && tensor->GetShapeSize() > 0; +} + +static bool IsPowerOfTwoInRange(uint32_t value, uint32_t minValue, uint32_t maxValue) +{ + return value >= minValue && value <= maxValue && (value & (value - 1U)) == 0U; +} + +static bool IsA5Arch(NpuArch npuArch) +{ + return npuArch == NpuArch::DAV_3510; +} + +static bool IsPaBlockSizeSupport(NpuArch npuArch, int32_t blockSize) +{ + if (IsA5Arch(npuArch)) { + return blockSize >= 1 && blockSize <= static_cast(BLOCK_SIZE_LIMIT); + } + return blockSize >= 16 && blockSize <= static_cast(BLOCK_SIZE_LIMIT) && blockSize % 16 == 0; +} static const std::map> DTYPE_SUPPORT_MAP = { - {QUERY_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, - {ORI_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, - {CMP_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, - {CU_SEQLENS_ORI_KV_NAME, {ge::DT_INT32}}, - {CU_SEQLENS_CMP_KV_NAME, {ge::DT_INT32}}, - {ORI_SPARSE_INDICES, {ge::DT_INT32}}, - {CMP_SPARSE_INDICES, {ge::DT_INT32}}, - {ATTEN_OUT_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, - {ORI_BLOCK_TABLE_NAME, {ge::DT_INT32}}, - {CMP_BLOCK_TABLE_NAME, {ge::DT_INT32}}, - {SINKS_NAME, {ge::DT_FLOAT}}, - {METADATA_NAME, {ge::DT_INT32}} -}; + {QUERY_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {ORI_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {CMP_KV_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {CU_SEQLENS_ORI_KV_NAME, {ge::DT_INT32}}, + {CU_SEQLENS_CMP_KV_NAME, {ge::DT_INT32}}, + {ORI_SPARSE_INDICES, {ge::DT_INT32}}, + {CMP_SPARSE_INDICES, {ge::DT_INT32}}, + {ATTEN_OUT_NAME, {ge::DT_FLOAT16, ge::DT_BF16}}, + {ORI_BLOCK_TABLE_NAME, {ge::DT_INT32}}, + {CMP_BLOCK_TABLE_NAME, {ge::DT_INT32}}, + {SINKS_NAME, {ge::DT_FLOAT}}, + {METADATA_NAME, {ge::DT_INT32}}}; static const std::map> LAYOUT_SUPPORT_MAP = { - {QUERY_NAME, {SMLALayout::BSND, SMLALayout::TND}}, - {ORI_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, - {CMP_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, - {ATTEN_OUT_NAME, {SMLALayout::BSND, SMLALayout::TND}}, - {ORI_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, - {CMP_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, + {QUERY_NAME, {SMLALayout::BSND, SMLALayout::TND}}, + {ORI_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, + {CMP_KV_NAME, {SMLALayout::PA_BBND, SMLALayout::TND, SMLALayout::BSND}}, + {ATTEN_OUT_NAME, {SMLALayout::BSND, SMLALayout::TND}}, + {ORI_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, + {CMP_SPARSE_INDICES, {SMLALayout::BSND, SMLALayout::TND}}, }; static const std::map DATATYPE_TO_STRING_MAP = { @@ -117,118 +116,123 @@ static const std::map DATATYPE_TO_STRING_MAP = { {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 -}; - -static uint64_t GetStorageShapeStride0(const gert::Shape &storageShape) -{ - if (storageShape.GetDimNum() <= DIM_NUM_ONE) { - return 0ULL; - } - - uint64_t stride0 = 1ULL; - for (size_t i = 1; i < storageShape.GetDimNum(); ++i) { - int64_t dim = storageShape.GetDim(i); - if (dim <= 0) { - return 0ULL; - } - stride0 *= static_cast(dim); - } - return stride0; -} - -template -static auto GetStride0FromStrideObject(const StrideT &stride, int) - -> decltype(stride.GetDimNum(), stride.GetStride(0), uint64_t()) -{ - if (stride.GetDimNum() <= 0) { - return 0ULL; - } - int64_t stride0 = stride.GetStride(0); - return stride0 > 0 ? static_cast(stride0) : 0ULL; -} - -template -static uint64_t GetStride0FromStrideObject(const StrideT &, ...) -{ - return 0ULL; -} - -template -static auto GetStride0FromStrideScalar(const StrideT &stride, int) - -> decltype(stride > 0, static_cast(stride)) -{ - return stride > 0 ? static_cast(stride) : 0ULL; -} - -template -static uint64_t GetStride0FromStrideScalar(const StrideT &, ...) -{ - return 0ULL; -} - -template -static uint64_t GetStride0FromStrideElement(const StrideT &stride) -{ - // CANN stride APIs return a dimension-wise stride array. In newer headers, stride[0] is scalar stride0. - // In compatibility headers it may be a stride object. Non-positive stride is treated as unavailable and - // falls back to the storage-shape contiguous calculation. - uint64_t stride0 = GetStride0FromStrideScalar(stride, 0); - if (stride0 > 0) { - return stride0; - } - return GetStride0FromStrideObject(stride, 0); -} - -template -static uint64_t GetStride0FromStrideArray(const StrideT *stride) -{ - if (stride == nullptr) { - return 0ULL; - } - return GetStride0FromStrideElement(stride[0]); -} - -template -static auto TryGetOptionalInputStride0(ContextT *context, uint32_t inputIndex, int) - -> decltype(context->GetOptionalInputStride(inputIndex), uint64_t()) -{ - return GetStride0FromStrideArray(context->GetOptionalInputStride(inputIndex)); -} - -template -static uint64_t TryGetOptionalInputStride0(ContextT *, uint32_t, ...) -{ - return 0ULL; -} - -// Compatibility path for CANN headers that do not expose GetOptionalInputStride. -// Some tiling contexts only provide real stride for view inputs through InputIsView/GetInputStride. -// Returning 0 means the stride is unavailable; the caller then falls back to storage-shape contiguous stride. -template -static auto TryGetInputViewStride0(ContextT *context, uint32_t inputIndex, int) - -> decltype(context->InputIsView(inputIndex), context->GetInputStride(inputIndex), uint64_t()) -{ - if (!context->InputIsView(inputIndex)) { - return 0ULL; - } - return GetStride0FromStrideArray(context->GetInputStride(inputIndex)); -} - -template -static uint64_t TryGetInputViewStride0(ContextT *, uint32_t, ...) -{ - return 0ULL; -} - -std::string SMLALayoutToSerialString(SMLALayout layout) -{ - switch (layout) { - case SMLALayout::BSND: return "BSND"; - case SMLALayout::TND: return "TND"; - case SMLALayout::PA_BBND: return "PA_BBND"; - default: return "UNKNOWN"; + {ge::DT_INT2, "DT_INT2"}, // dt_variant type + {ge::DT_UINT2, "DT_UINT2"} // dt_variant type +}; + +static uint64_t GetStorageShapeStride0(const gert::Shape &storageShape) +{ + if (storageShape.GetDimNum() <= DIM_NUM_ONE) { + return 0ULL; + } + + uint64_t stride0 = 1ULL; + for (size_t i = 1; i < storageShape.GetDimNum(); ++i) { + int64_t dim = storageShape.GetDim(i); + if (dim <= 0) { + return 0ULL; + } + stride0 *= static_cast(dim); + } + return stride0; +} + +template +static auto GetStride0FromStrideObject(const StrideT &stride, int) -> decltype(stride.GetDimNum(), stride.GetStride(0), + uint64_t()) +{ + if (stride.GetDimNum() <= 0) { + return 0ULL; + } + int64_t stride0 = stride.GetStride(0); + return stride0 > 0 ? static_cast(stride0) : 0ULL; +} + +template +static uint64_t GetStride0FromStrideObject(const StrideT &, ...) +{ + return 0ULL; +} + +template +static auto GetStride0FromStrideScalar(const StrideT &stride, int) -> decltype(stride > 0, + static_cast(stride)) +{ + return stride > 0 ? static_cast(stride) : 0ULL; +} + +template +static uint64_t GetStride0FromStrideScalar(const StrideT &, ...) +{ + return 0ULL; +} + +template +static uint64_t GetStride0FromStrideElement(const StrideT &stride) +{ + // CANN stride APIs return a dimension-wise stride array. In newer headers, stride[0] is scalar stride0. + // In compatibility headers it may be a stride object. Non-positive stride is treated as unavailable and + // falls back to the storage-shape contiguous calculation. + uint64_t stride0 = GetStride0FromStrideScalar(stride, 0); + if (stride0 > 0) { + return stride0; + } + return GetStride0FromStrideObject(stride, 0); +} + +template +static uint64_t GetStride0FromStrideArray(const StrideT *stride) +{ + if (stride == nullptr) { + return 0ULL; + } + return GetStride0FromStrideElement(stride[0]); +} + +template +static auto TryGetOptionalInputStride0(ContextT *context, uint32_t inputIndex, + int) -> decltype(context->GetOptionalInputStride(inputIndex), uint64_t()) +{ + return GetStride0FromStrideArray(context->GetOptionalInputStride(inputIndex)); +} + +template +static uint64_t TryGetOptionalInputStride0(ContextT *, uint32_t, ...) +{ + return 0ULL; +} + +// Compatibility path for CANN headers that do not expose GetOptionalInputStride. +// Some tiling contexts only provide real stride for view inputs through InputIsView/GetInputStride. +// Returning 0 means the stride is unavailable; the caller then falls back to storage-shape contiguous stride. +template +static auto TryGetInputViewStride0(ContextT *context, uint32_t inputIndex, + int) -> decltype(context->InputIsView(inputIndex), + context->GetInputStride(inputIndex), uint64_t()) +{ + if (!context->InputIsView(inputIndex)) { + return 0ULL; + } + return GetStride0FromStrideArray(context->GetInputStride(inputIndex)); +} + +template +static uint64_t TryGetInputViewStride0(ContextT *, uint32_t, ...) +{ + return 0ULL; +} + +std::string SMLALayoutToSerialString(SMLALayout layout) +{ + switch (layout) { + case SMLALayout::BSND: + return "BSND"; + case SMLALayout::TND: + return "TND"; + case SMLALayout::PA_BBND: + return "PA_BBND"; + default: + return "UNKNOWN"; } } @@ -293,8 +297,7 @@ ge::graphStatus SMLAInfoParser::CheckRequiredAttrExistence() const ge::graphStatus SMLAInfoParser::CheckRequiredParaExistence() const { - if (CheckRequiredInOutExistence() != ge::GRAPH_SUCCESS || - CheckRequiredAttrExistence() != ge::GRAPH_SUCCESS) { + if (CheckRequiredInOutExistence() != ge::GRAPH_SUCCESS || CheckRequiredAttrExistence() != ge::GRAPH_SUCCESS) { return ge::GRAPH_FAILED; } @@ -425,32 +428,32 @@ uint64_t SMLAInfoParser::GetOptionalInputStride0(uint32_t inputIndex) const } else if (inputIndex == CMP_KV_INDEX) { inputTensor = opParamInfo_.cmpKv.tensor; } - if (inputTensor == nullptr) { - return 0ULL; - } - - uint64_t stride0 = TryGetOptionalInputStride0(context_, inputIndex, 0); - if (stride0 > 0) { - return stride0; - } - - // Compatible with CANN packages that only expose view stride by normal input index. - stride0 = TryGetInputViewStride0(context_, inputIndex, 0); - if (stride0 > 0) { - return stride0; - } - - const gert::Shape &storageShape = inputTensor->GetStorageShape(); - stride0 = GetStorageShapeStride0(storageShape); - const char *inputName = inputIndex == ORI_KV_INDEX ? "ori_kv" : "cmp_kv"; - OP_LOGW(context_->GetNodeName(), - "Cannot get %s stride0 from tiling context stride APIs. Use storage shape to infer contiguous " - "stride0(%lu). Non-contiguous %s requires GetOptionalInputStride or GetInputStride support.", - inputName, stride0, inputName); - return stride0; -} -ge::graphStatus SMLAInfoParser::GetInOutDataType() -{ + if (inputTensor == nullptr) { + return 0ULL; + } + + uint64_t stride0 = TryGetOptionalInputStride0(context_, inputIndex, 0); + if (stride0 > 0) { + return stride0; + } + + // Compatible with CANN packages that only expose view stride by normal input index. + stride0 = TryGetInputViewStride0(context_, inputIndex, 0); + if (stride0 > 0) { + return stride0; + } + + const gert::Shape &storageShape = inputTensor->GetStorageShape(); + stride0 = GetStorageShapeStride0(storageShape); + const char *inputName = inputIndex == ORI_KV_INDEX ? "ori_kv" : "cmp_kv"; + OP_LOGW(context_->GetNodeName(), + "Cannot get %s stride0 from tiling context stride APIs. Use storage shape to infer contiguous " + "stride0(%lu). Non-contiguous %s requires GetOptionalInputStride or GetInputStride support.", + inputName, stride0, inputName); + return stride0; +} +ge::graphStatus SMLAInfoParser::GetInOutDataType() +{ qType_ = opParamInfo_.q.desc->GetDataType(); outputType_ = opParamInfo_.attnOut.desc->GetDataType(); if (opParamInfo_.oriKv.desc != nullptr) { @@ -492,8 +495,8 @@ ge::graphStatus SMLAInfoParser::GetSMLATemplateMode(SMLATilingInfo &smlaInfo) ge::graphStatus SMLAInfoParser::GetQueryAndOutLayout() { const map> layoutMap = { - {"BSND", {SMLALayout::BSND, SMLALayout::BSND}}, - {"TND", {SMLALayout::TND, SMLALayout::TND }}, + {"BSND", {SMLALayout::BSND, SMLALayout::BSND}}, + {"TND", {SMLALayout::TND, SMLALayout::TND}}, }; std::string layout(opParamInfo_.layoutQ); auto it = layoutMap.find(layout); @@ -517,41 +520,42 @@ ge::graphStatus SMLAInfoParser::GetQueryAndOutLayout() ge::graphStatus SMLAInfoParser::GetKvLayout() { const map layoutKVMap = { - {"PA_BBND", SMLALayout::PA_BBND}, - {"TND", SMLALayout::TND}, - {"BSND", SMLALayout::BSND}, + {"PA_BBND", SMLALayout::PA_BBND}, + {"TND", SMLALayout::TND}, + {"BSND", SMLALayout::BSND}, }; 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; - } - if (kvLayout_ != SMLALayout::PA_BBND && qLayout_ != kvLayout_) { - OP_LOGE(opName_, "layout_q and layout_kv only support BSND/BSND, TND/TND, BSND/PA_BBND " - "or TND/PA_BBND, but got %s/%s.", - SMLALayoutToSerialString(qLayout_).c_str(), SMLALayoutToSerialString(kvLayout_).c_str()); - return ge::GRAPH_FAILED; - } - return ge::GRAPH_SUCCESS; -} + if (it != layoutKVMap.end()) { + kvLayout_ = it->second; + } else { + OP_LOGE(opName_, "layoutKV is %s, it is unsupported.", layout.c_str()); + return ge::GRAPH_FAILED; + } + if (kvLayout_ != SMLALayout::PA_BBND && qLayout_ != kvLayout_) { + OP_LOGE(opName_, + "layout_q and layout_kv only support BSND/BSND, TND/TND, BSND/PA_BBND " + "or TND/PA_BBND, but got %s/%s.", + SMLALayoutToSerialString(qLayout_).c_str(), SMLALayoutToSerialString(kvLayout_).c_str()); + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} // =============Parser function==================== bool SMLAInfoParser::HasAxis(const SMLAAxis &axis, const SMLALayout &layout, const gert::Shape &shape) const { - const auto& layoutIt = SMLA_LAYOUT_AXIS_MAP.find(layout); + const auto &layoutIt = SMLA_LAYOUT_AXIS_MAP.find(layout); if (layoutIt == SMLA_LAYOUT_AXIS_MAP.end()) { return false; } - const std::vector& axes = layoutIt->second; - const auto& axisIt = std::find(axes.begin(), axes.end(), axis); + const std::vector &axes = layoutIt->second; + const auto &axisIt = std::find(axes.begin(), axes.end(), axis); if (axisIt == axes.end()) { return false; } - const auto& dimIt = SMLA_LAYOUT_DIM_MAP.find(layout); + const auto &dimIt = SMLA_LAYOUT_DIM_MAP.find(layout); if (dimIt == SMLA_LAYOUT_DIM_MAP.end() || dimIt->second != shape.GetDimNum()) { return false; } @@ -560,8 +564,8 @@ bool SMLAInfoParser::HasAxis(const SMLAAxis &axis, const SMLALayout &layout, con size_t SMLAInfoParser::GetAxisIdx(const SMLAAxis &axis, const SMLALayout &layout) const { - const std::vector& axes = SMLA_LAYOUT_AXIS_MAP.find(layout)->second; - const auto& axisIt = std::find(axes.begin(), axes.end(), axis); + const std::vector &axes = SMLA_LAYOUT_AXIS_MAP.find(layout)->second; + const auto &axisIt = std::find(axes.begin(), axes.end(), axis); return std::distance(axes.begin(), axisIt); } @@ -606,12 +610,11 @@ ge::graphStatus SMLAInfoParser::GetN2Size() if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { uint32_t cmpSparseIndicesN2Size_ = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::N, cmpSparseIndicesLayout_); OP_CHECK_IF(cmpKvN2Size_ != n2Size_ || n2Size_ != cmpSparseIndicesN2Size_, - OP_LOGE(opName_, "N2 size check failed! Expected oriKvN2 == cmpSparseIndicesN2."), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "N2 size check failed! Expected oriKvN2 == cmpSparseIndicesN2."), + return ge::GRAPH_FAILED); } - OP_CHECK_IF(cmpKvN2Size_ != n2Size_, - OP_LOGE(opName_, "N2 size check failed! Expected cmpKvN2 == oriKvN2."), - return ge::GRAPH_FAILED); + OP_CHECK_IF(cmpKvN2Size_ != n2Size_, OP_LOGE(opName_, "N2 size check failed! Expected cmpKvN2 == oriKvN2."), + return ge::GRAPH_FAILED); n2Size_ = cmpKvN2Size_; } return ge::GRAPH_SUCCESS; @@ -625,18 +628,17 @@ ge::graphStatus SMLAInfoParser::GetGSize() return ge::GRAPH_SUCCESS; } -ge::graphStatus SMLAInfoParser::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor, - SMLALayout &layout, const std::string &name) const +ge::graphStatus SMLAInfoParser::GetActualSeqLenSize(uint32_t &size, const gert::Tensor *tensor, SMLALayout &layout, + const std::string &name) const { if ((tensor == nullptr)) { - OP_LOGE(opName_, "when layout of q is %s, %s must be provided.", - SMLALayoutToSerialString(layout).c_str(), name.c_str()); + OP_LOGE(opName_, "when layout of q is %s, %s must be provided.", SMLALayoutToSerialString(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); + 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(shapeSize) - 1; @@ -680,13 +682,11 @@ ge::graphStatus SMLAInfoParser::GetS1Size() if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { if (cmpSparseIndicesLayout_ == SMLALayout::TND) { uint32_t cmpSparseIndicesT = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::T, cmpSparseIndicesLayout_); - OP_CHECK_IF(cmpSparseIndicesT != s1Size_, - OP_LOGE(opName_, "T size check failed !"), - return ge::GRAPH_FAILED); + OP_CHECK_IF(cmpSparseIndicesT != s1Size_, OP_LOGE(opName_, "T size check failed !"), + return ge::GRAPH_FAILED); } else { uint32_t cmpSparseIndicesS1 = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::S, cmpSparseIndicesLayout_); - OP_CHECK_IF(cmpSparseIndicesS1 != s1Size_, - OP_LOGE(opName_, "s1 size check failed !"), + OP_CHECK_IF(cmpSparseIndicesS1 != s1Size_, OP_LOGE(opName_, "s1 size check failed !"), return ge::GRAPH_FAILED); } } @@ -704,8 +704,8 @@ ge::graphStatus SMLAInfoParser::GetMaxBlockNumPerBatch() return ge::GRAPH_FAILED; } if (opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1) < 0) { - OP_LOGE(opName_, "%s's second dimension(%ld) should be non-negative number.", - ORI_BLOCK_TABLE_NAME.c_str(), opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1)); + OP_LOGE(opName_, "%s's second dimension(%ld) should be non-negative number.", ORI_BLOCK_TABLE_NAME.c_str(), + opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1)); return ge::GRAPH_FAILED; } oriMaxBlockNumPerBatch_ = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1); @@ -718,14 +718,14 @@ ge::graphStatus SMLAInfoParser::GetMaxBlockNumPerBatch() } if (qLayout_ == SMLALayout::TND || qLayout_ == SMLALayout::BSND) { if (opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0) != bSize_) { - OP_LOGE(opName_, "cmp_block_table's first dimension(%ld) should be equal to query's B(%u).", - opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0), bSize_); + OP_LOGE(opName_, "cmp_block_table's first dimension(%ld) should be equal to query's B(%u).", + opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0), bSize_); 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)); + 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); @@ -813,7 +813,7 @@ ge::graphStatus SMLAInfoParser::GetActualseqInfo() if (opParamInfo_.cuSeqLensQ.tensor != nullptr) { if (opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() != bSize_ + 1) { OP_LOGE(opName_, "cu_seqlens_q's dimension should be equal to %u, but now it's %ld.", bSize_ + 1, - opParamInfo_.cuSeqLensQ.tensor->GetShapeSize()); + opParamInfo_.cuSeqLensQ.tensor->GetShapeSize()); return ge::GRAPH_FAILED; } actualLenDimsQ_ = opParamInfo_.cuSeqLensQ.tensor->GetShapeSize() - 1; // cuSeqLensQ shape is B+1 @@ -836,42 +836,42 @@ ge::graphStatus SMLAInfoParser::GetActualseqInfo() if (opParamInfo_.sequsedCmpKv.tensor != nullptr) { actualLenDimsCmpKV_ = opParamInfo_.sequsedCmpKv.tensor->GetShapeSize(); if (opParamInfo_.sequsedCmpKv.tensor->GetShapeSize() != bSize_) { - OP_LOGE(opName_, "sequsedCmpKv's dimension should be equal to %u, but got %ld.", - bSize_, opParamInfo_.sequsedCmpKv.tensor->GetShapeSize()); + OP_LOGE(opName_, "sequsedCmpKv's dimension should be equal to %u, but got %ld.", bSize_, + opParamInfo_.sequsedCmpKv.tensor->GetShapeSize()); return ge::GRAPH_FAILED; } } if (opParamInfo_.cmpResidualKv.tensor != nullptr) { - cmpResidualKVSize_ = opParamInfo_.cmpResidualKv.tensor->GetShapeSize(); - if (opParamInfo_.cmpResidualKv.tensor->GetShapeSize() != bSize_) { - OP_LOGE(opName_, "cmpResidualKv's dimension should be equal to %u, but got %ld.", - bSize_, opParamInfo_.cmpResidualKv.tensor->GetShapeSize()); + cmpResidualKVSize_ = opParamInfo_.cmpResidualKv.tensor->GetShapeSize(); + if (opParamInfo_.cmpResidualKv.tensor->GetShapeSize() != bSize_) { + OP_LOGE(opName_, "cmpResidualKv's dimension should be equal to %u, but got %ld.", bSize_, + opParamInfo_.cmpResidualKv.tensor->GetShapeSize()); return ge::GRAPH_FAILED; } - } - if (!IsA5Arch(npuArch_) && IsNonEmptyOptionalTensor(opParamInfo_.oriTopkLength.tensor)) { - OP_LOGE(opName_, "ori_topk_length is reserved and does not support non-empty tensor on %s.", - A2_A3_PLATFORM_LOG.c_str()); - return ge::GRAPH_FAILED; - } - if (!IsA5Arch(npuArch_) && IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor)) { - OP_LOGE(opName_, "cmp_topk_length is reserved and does not support non-empty tensor on %s.", - A2_A3_PLATFORM_LOG.c_str()); - return ge::GRAPH_FAILED; - } + } + if (!IsA5Arch(npuArch_) && IsNonEmptyOptionalTensor(opParamInfo_.oriTopkLength.tensor)) { + OP_LOGE(opName_, "ori_topk_length is reserved and does not support non-empty tensor on %s.", + A2_A3_PLATFORM_LOG.c_str()); + return ge::GRAPH_FAILED; + } + if (!IsA5Arch(npuArch_) && IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor)) { + OP_LOGE(opName_, "cmp_topk_length is reserved and does not support non-empty tensor on %s.", + A2_A3_PLATFORM_LOG.c_str()); + return ge::GRAPH_FAILED; + } if (kvLayout_ == SMLALayout::PA_BBND) { if (opParamInfo_.sequsedOriKv.tensor != nullptr) { if (qLayout_ == SMLALayout::BSND) { if (opParamInfo_.sequsedOriKv.tensor->GetShapeSize() != bSize_) { - OP_LOGE(opName_, "sequsedOriKv's dimension should be equal to %u, but got %ld.", - bSize_, opParamInfo_.sequsedOriKv.tensor->GetShapeSize()); + OP_LOGE(opName_, "sequsedOriKv's dimension should be equal to %u, but got %ld.", bSize_, + opParamInfo_.sequsedOriKv.tensor->GetShapeSize()); return ge::GRAPH_FAILED; } } else { if (opParamInfo_.sequsedOriKv.tensor->GetShapeSize() != bSize_) { - OP_LOGE(opName_, "sequsedOriKv's dimension should be equal to bSize(%u), but got %ld.", - bSize_, opParamInfo_.sequsedOriKv.tensor->GetShapeSize()); + OP_LOGE(opName_, "sequsedOriKv's dimension should be equal to bSize(%u), but got %ld.", bSize_, + opParamInfo_.sequsedOriKv.tensor->GetShapeSize()); return ge::GRAPH_FAILED; } } @@ -884,8 +884,8 @@ ge::graphStatus SMLAInfoParser::GetActualseqInfo() } else if (kvLayout_ == SMLALayout::BSND) { actualLenDimsKV_ = actualLenDimsOriKV_; } else { - OP_LOGE(opName_, "oriKV and cmpKv only support PA_BBND, TND and BSND layout, but got %s.", - SMLALayoutToSerialString(kvLayout_).c_str()); + OP_LOGE(opName_, "oriKV and cmpKv only support PA_BBND, TND and BSND layout, but got %s.", + SMLALayoutToSerialString(kvLayout_).c_str()); return ge::GRAPH_FAILED; } return ge::GRAPH_SUCCESS; @@ -920,8 +920,8 @@ void SMLAInfoParser::GenerateInfo(SMLATilingInfo &smlaInfo) smlaInfo.outputType = outputType_; smlaInfo.perfMode = perfMode_; - smlaInfo.totalBlockNum = (opParamInfo_.oriKv.tensor != nullptr) ? - opParamInfo_.oriKv.tensor->GetStorageShape().GetDim(0) : 0; + smlaInfo.totalBlockNum = + (opParamInfo_.oriKv.tensor != nullptr) ? opParamInfo_.oriKv.tensor->GetStorageShape().GetDim(0) : 0; smlaInfo.sparseBlockSize = 1; smlaInfo.oriBlockSize = oriBlockSize_; smlaInfo.cmpBlockSize = cmpBlockSize_; @@ -965,35 +965,22 @@ ge::graphStatus SMLAInfoParser::Parse(SMLATilingInfo &smlaInfo) return ge::GRAPH_FAILED; } - if (ge::GRAPH_SUCCESS != GetOpName() || - ge::GRAPH_SUCCESS != GetNpuInfo() || - ge::GRAPH_SUCCESS != GetOpParaInfo() || - ge::GRAPH_SUCCESS != CheckRequiredParaExistence() || - ge::GRAPH_SUCCESS != CheckUnrequiredParaExistence()) { + if (ge::GRAPH_SUCCESS != GetOpName() || ge::GRAPH_SUCCESS != GetNpuInfo() || ge::GRAPH_SUCCESS != GetOpParaInfo() || + ge::GRAPH_SUCCESS != CheckRequiredParaExistence() || ge::GRAPH_SUCCESS != CheckUnrequiredParaExistence()) { return ge::GRAPH_FAILED; } - if (ge::GRAPH_SUCCESS != GetInOutDataType() || - ge::GRAPH_SUCCESS != GetQueryAndOutLayout() || - ge::GRAPH_SUCCESS != GetKvLayout() || - ge::GRAPH_SUCCESS != GetSMLATemplateMode(smlaInfo)) { + if (ge::GRAPH_SUCCESS != GetInOutDataType() || ge::GRAPH_SUCCESS != GetQueryAndOutLayout() || + ge::GRAPH_SUCCESS != GetKvLayout() || ge::GRAPH_SUCCESS != GetSMLATemplateMode(smlaInfo)) { return ge::GRAPH_FAILED; } SetSMLAShape(); - 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 != GetQHeadDim() || - ge::GRAPH_SUCCESS != GetValueHeadDim() || - ge::GRAPH_SUCCESS != GetSparseBlockCount() || - ge::GRAPH_SUCCESS != GetSinks() - ) { + 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 != GetQHeadDim() || + ge::GRAPH_SUCCESS != GetValueHeadDim() || ge::GRAPH_SUCCESS != GetSparseBlockCount() || + ge::GRAPH_SUCCESS != GetSinks()) { return ge::GRAPH_FAILED; } if (ge::GRAPH_SUCCESS != GetActualseqInfo()) { @@ -1038,7 +1025,7 @@ void SMLATilingCheck::Init() } void SMLATilingCheck::LogErrorDtypeSupport(const std::vector &expectDtypeList, - const ge::DataType &actualDtype, const std::string &name) const + const ge::DataType &actualDtype, const std::string &name) const { std::ostringstream oss; for (size_t i = 0; i < expectDtypeList.size(); ++i) { @@ -1047,29 +1034,28 @@ void SMLATilingCheck::LogErrorDtypeSupport(const std::vector &expe oss << ", "; } } - OP_LOGE(opName_, "Tensor %s only supports dtype %s, but got %s", - name.c_str(), oss.str().c_str(), SMLADataTypeToSerialString(actualDtype).c_str()); + OP_LOGE(opName_, "Tensor %s only supports dtype %s, but got %s", name.c_str(), oss.str().c_str(), + SMLADataTypeToSerialString(actualDtype).c_str()); } ge::graphStatus SMLATilingCheck::CheckDtypeSupport(const gert::CompileTimeTensorDesc *desc, - const std::string &name) const + const std::string &name) const { if (desc != nullptr) { - const auto& it = DTYPE_SUPPORT_MAP.find(name); + 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); + 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; } void SMLATilingCheck::LogErrorLayoutSupport(const std::vector &expectLayoutList, - const SMLALayout &actualLayout, const std::string &name) const + const SMLALayout &actualLayout, const std::string &name) const { std::ostringstream oss; for (size_t i = 0; i < expectLayoutList.size(); ++i) { @@ -1078,28 +1064,26 @@ void SMLATilingCheck::LogErrorLayoutSupport(const std::vector &expec oss << ", "; } } - OP_LOGE(opName_, "Tensor %s only supports layout %s, but got %s", - name.c_str(), oss.str().c_str(), SMLALayoutToSerialString(actualLayout).c_str()); + OP_LOGE(opName_, "Tensor %s only supports layout %s, but got %s", name.c_str(), oss.str().c_str(), + SMLALayoutToSerialString(actualLayout).c_str()); } ge::graphStatus SMLATilingCheck::CheckLayoutSupport(const SMLALayout &actualLayout, const std::string &name) const { - const auto& it = LAYOUT_SUPPORT_MAP.find(name); + 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); + OP_CHECK_IF(std::find(expectLayoutList.begin(), expectLayoutList.end(), actualLayout) == expectLayoutList.end(), + LogErrorLayoutSupport(expectLayoutList, actualLayout, name), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } template -void SMLATilingCheck::LogErrorNumberSupport(const std::vector &expectNumberList, - const T &actualValue, const std::string &name, const std::string subName) const +void SMLATilingCheck::LogErrorNumberSupport(const std::vector &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) { @@ -1108,26 +1092,27 @@ void SMLATilingCheck::LogErrorNumberSupport(const std::vector &expectNumberLi oss << ", "; } } - OP_LOGE(opName_, "%s %s only supports %s, but got %s", - name.c_str(), subName.c_str(), oss.str().c_str(), std::to_string(actualValue).c_str()); + OP_LOGE(opName_, "%s %s only supports %s, but got %s", name.c_str(), subName.c_str(), oss.str().c_str(), + std::to_string(actualValue).c_str()); } template -void SMLATilingCheck::LogErrorDimNumSupport(const std::vector &expectNumberList, - const T &actualValue, const std::string &name) const +void SMLATilingCheck::LogErrorDimNumSupport(const std::vector &expectNumberList, const T &actualValue, + const std::string &name) const { LogErrorNumberSupport(expectNumberList, actualValue, name, "dimension"); } ge::graphStatus SMLATilingCheck::CheckDimNumSupport(const gert::StorageShape *shape, - const std::vector &expectDimNumList, const std::string &name) const + const std::vector &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()) { + if (std::find(expectDimNumList.begin(), expectDimNumList.end(), shape->GetStorageShape().GetDimNum()) == + expectDimNumList.end()) { LogErrorDimNumSupport(expectDimNumList, shape->GetStorageShape().GetDimNum(), name); return ge::GRAPH_FAILED; } @@ -1135,14 +1120,14 @@ ge::graphStatus SMLATilingCheck::CheckDimNumSupport(const gert::StorageShape *sh return ge::GRAPH_SUCCESS; } -ge::graphStatus SMLATilingCheck::CheckDimNumInLayoutSupport(const SMLALayout &layout, - const gert::StorageShape *shape, const std::string &name) const +ge::graphStatus SMLATilingCheck::CheckDimNumInLayoutSupport(const SMLALayout &layout, const gert::StorageShape *shape, + const std::string &name) const { - const auto& dimIt = SMLA_LAYOUT_DIM_MAP.find(layout); + const auto &dimIt = SMLA_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 got %zu", - SMLALayoutToSerialString(layout).c_str(), name.c_str(), dimIt->second, - shape->GetStorageShape().GetDimNum()), + SMLALayoutToSerialString(layout).c_str(), name.c_str(), dimIt->second, + shape->GetStorageShape().GetDimNum()), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1154,8 +1139,7 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaQuery() const return ge::GRAPH_FAILED; } const std::vector queryDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.q.desc, QUERY_NAME) || + 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)) { @@ -1167,12 +1151,11 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaQuery() const ge::graphStatus SMLATilingCheck::CheckSingleParaOriKv() const { const std::vector oriKvDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriKv.desc, ORI_KV_NAME) || + 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(), oriKvDimNumList, ORI_KV_NAME) || - ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport( - kvLayout_, &opParamInfo_.oriKv.tensor->GetShape(), ORI_KV_NAME)) { + ge::GRAPH_SUCCESS != + CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.oriKv.tensor->GetShape(), ORI_KV_NAME)) { return ge::GRAPH_FAILED; } return ge::GRAPH_SUCCESS; @@ -1180,16 +1163,15 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriKv() const ge::graphStatus SMLATilingCheck::CheckSingleParaCmpKv() const { - if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE || \ + if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE || smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE) { const std::vector cmpKvDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpKv.desc, CMP_KV_NAME) || + 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(), cmpKvDimNumList, CMP_KV_NAME) || - ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport( - kvLayout_, &opParamInfo_.cmpKv.tensor->GetShape(), CMP_KV_NAME)) { + ge::GRAPH_SUCCESS != + CheckDimNumSupport(&opParamInfo_.cmpKv.tensor->GetShape(), cmpKvDimNumList, CMP_KV_NAME) || + ge::GRAPH_SUCCESS != + CheckDimNumInLayoutSupport(kvLayout_, &opParamInfo_.cmpKv.tensor->GetShape(), CMP_KV_NAME)) { return ge::GRAPH_FAILED; } } @@ -1205,9 +1187,9 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaCuSeqLensOriKv() const return ge::GRAPH_FAILED; } OP_CHECK_IF(opParamInfo_.cuSeqLensOriKv.tensor->GetShapeSize() != bSize_ + 1, - OP_LOGE(opName_, "Input cuSeqLensOriKv's shapeSize is not equal to B + 1: %u, it is %ld", bSize_ + 1, - opParamInfo_.cuSeqLensOriKv.tensor->GetShapeSize()), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "Input cuSeqLensOriKv's shapeSize is not equal to B + 1: %u, it is %ld", bSize_ + 1, + opParamInfo_.cuSeqLensOriKv.tensor->GetShapeSize()), + return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1220,9 +1202,9 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaCuSeqLensCmpKv() const return ge::GRAPH_FAILED; } OP_CHECK_IF(opParamInfo_.cuSeqLensCmpKv.tensor->GetShapeSize() != bSize_ + 1, - OP_LOGE(opName_, "Input cuSeqLensCmpKv's shapeSize is not equal to B + 1: %u, it is %ld", bSize_ + 1, - opParamInfo_.cuSeqLensCmpKv.tensor->GetShapeSize()), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "Input cuSeqLensCmpKv's shapeSize is not equal to B + 1: %u, it is %ld", bSize_ + 1, + opParamInfo_.cuSeqLensCmpKv.tensor->GetShapeSize()), + return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1236,45 +1218,44 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaKvHeadNums() const return ge::GRAPH_SUCCESS; } -ge::graphStatus SMLATilingCheck::CheckSingleParaOriSparseIndices() const -{ - if (opParamInfo_.oriSparseIndices.tensor == nullptr) { - return ge::GRAPH_SUCCESS; - } - OP_CHECK_IF(!IsA5Arch(npuArch_), - OP_LOGE(opName_, "ori_sparse_indices is only supported on %s.", A5_PLATFORM_LOG.c_str()), - return ge::GRAPH_FAILED); - OP_CHECK_IF(opParamInfo_.oriSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, - OP_LOGE(opName_, "ori_sparse_indices cannot be empty tensor."), - return ge::GRAPH_FAILED); - const std::vector oriSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriSparseIndices.desc, ORI_SPARSE_INDICES) || - ge::GRAPH_SUCCESS != CheckLayoutSupport(oriSparseIndicesLayout_, ORI_SPARSE_INDICES) || - ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriSparseIndices.tensor->GetShape(), - oriSparseIndicesDimNumList, ORI_SPARSE_INDICES) || - ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(oriSparseIndicesLayout_, - &opParamInfo_.oriSparseIndices.tensor->GetShape(), ORI_SPARSE_INDICES)) { - return ge::GRAPH_FAILED; - } - return ge::GRAPH_SUCCESS; -} +ge::graphStatus SMLATilingCheck::CheckSingleParaOriSparseIndices() const +{ + if (opParamInfo_.oriSparseIndices.tensor == nullptr) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(!IsA5Arch(npuArch_), + OP_LOGE(opName_, "ori_sparse_indices is only supported on %s.", A5_PLATFORM_LOG.c_str()), + return ge::GRAPH_FAILED); + OP_CHECK_IF(opParamInfo_.oriSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, + OP_LOGE(opName_, "ori_sparse_indices cannot be empty tensor."), return ge::GRAPH_FAILED); + const std::vector oriSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriSparseIndices.desc, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckLayoutSupport(oriSparseIndicesLayout_, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriSparseIndices.tensor->GetShape(), + oriSparseIndicesDimNumList, ORI_SPARSE_INDICES) || + ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(oriSparseIndicesLayout_, + &opParamInfo_.oriSparseIndices.tensor->GetShape(), + ORI_SPARSE_INDICES)) { + return ge::GRAPH_FAILED; + } + return ge::GRAPH_SUCCESS; +} ge::graphStatus SMLATilingCheck::CheckSingleParaCmpSparseIndices() const { if (smlaInfo_.perfMode == optiling::SMLATemplateMode::SCFA_TEMPLATE_MODE) { - OP_CHECK_IF(opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, - OP_LOGE(opName_, - "when cmp_sparse_indices is not nullptr(CSA), cmp_sparse_indices cannot be empty tensor."), - return ge::GRAPH_FAILED); + OP_CHECK_IF( + opParamInfo_.cmpSparseIndices.tensor->GetStorageShape().GetShapeSize() == 0, + OP_LOGE(opName_, "when cmp_sparse_indices is not nullptr(CSA), cmp_sparse_indices cannot be empty tensor."), + return ge::GRAPH_FAILED); const std::vector cmpSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpSparseIndices.desc, CMP_SPARSE_INDICES) || + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpSparseIndices.desc, CMP_SPARSE_INDICES) || ge::GRAPH_SUCCESS != CheckLayoutSupport(cmpSparseIndicesLayout_, CMP_SPARSE_INDICES) || ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpSparseIndices.tensor->GetShape(), - cmpSparseIndicesDimNumList, CMP_SPARSE_INDICES) || + cmpSparseIndicesDimNumList, CMP_SPARSE_INDICES) || ge::GRAPH_SUCCESS != CheckDimNumInLayoutSupport(cmpSparseIndicesLayout_, - &opParamInfo_.cmpSparseIndices.tensor->GetShape(), CMP_SPARSE_INDICES)) { + &opParamInfo_.cmpSparseIndices.tensor->GetShape(), + CMP_SPARSE_INDICES)) { return ge::GRAPH_FAILED; } } @@ -1287,17 +1268,16 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriBlockTable() const return ge::GRAPH_SUCCESS; } const std::vector oriBlockTableDimNumList = {DIM_NUM_TWO}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriBlockTable.desc, ORI_BLOCK_TABLE_NAME) || - ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriBlockTable.tensor->GetShape(), - oriBlockTableDimNumList, ORI_BLOCK_TABLE_NAME)) { + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.oriBlockTable.desc, ORI_BLOCK_TABLE_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.oriBlockTable.tensor->GetShape(), oriBlockTableDimNumList, + ORI_BLOCK_TABLE_NAME)) { return ge::GRAPH_FAILED; } - OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, oriBlockSize_), - OP_LOGE(opName_, - "oriBlockSize_ should be in [1, 1024] on %s or 16-aligned [16, 1024] on %s, but got: %d.", - A5_PLATFORM_LOG.c_str(), A2_A3_PLATFORM_LOG.c_str(), oriBlockSize_), - return ge::GRAPH_FAILED); + OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, oriBlockSize_), + OP_LOGE(opName_, + "oriBlockSize_ should be in [1, 1024] on %s or 16-aligned [16, 1024] on %s, but got: %d.", + A5_PLATFORM_LOG.c_str(), A2_A3_PLATFORM_LOG.c_str(), oriBlockSize_), + return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1308,41 +1288,38 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaCmpBlockTable() const } if (smlaInfo_.perfMode == optiling::SMLATemplateMode::SCFA_TEMPLATE_MODE || smlaInfo_.perfMode == optiling::SMLATemplateMode::CFA_TEMPLATE_MODE) { - const std::vector cmpBlockTableDimNumList = {DIM_NUM_TWO}; - if ( - ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpBlockTable.desc, CMP_BLOCK_TABLE_NAME) || - ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpBlockTable.tensor->GetShape(), - cmpBlockTableDimNumList, CMP_BLOCK_TABLE_NAME)) { - return ge::GRAPH_FAILED; - } - OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, cmpBlockSize_), - OP_LOGE(opName_, - "cmpBlockSize should be in [1, 1024] on %s or 16-aligned [16, 1024] on %s, " - "but got: %d.", - A5_PLATFORM_LOG.c_str(), A2_A3_PLATFORM_LOG.c_str(), cmpBlockSize_), - return ge::GRAPH_FAILED); - } - return ge::GRAPH_SUCCESS; -} + const std::vector cmpBlockTableDimNumList = {DIM_NUM_TWO}; + if (ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpBlockTable.desc, CMP_BLOCK_TABLE_NAME) || + ge::GRAPH_SUCCESS != CheckDimNumSupport(&opParamInfo_.cmpBlockTable.tensor->GetShape(), + cmpBlockTableDimNumList, CMP_BLOCK_TABLE_NAME)) { + return ge::GRAPH_FAILED; + } + OP_CHECK_IF(!IsPaBlockSizeSupport(npuArch_, cmpBlockSize_), + OP_LOGE(opName_, + "cmpBlockSize should be in [1, 1024] on %s or 16-aligned [16, 1024] on %s, " + "but got: %d.", + A5_PLATFORM_LOG.c_str(), A2_A3_PLATFORM_LOG.c_str(), cmpBlockSize_), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} ge::graphStatus SMLATilingCheck::CheckSingleParaSinks() const { OP_CHECK_IF(opParamInfo_.sinks.tensor->GetStorageShape().GetShapeSize() == 0, - OP_LOGE(opName_, "sinks cannot be empty tensor."), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "sinks cannot be empty tensor."), return ge::GRAPH_FAILED); if (opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum() != DIM_NUM_ONE) { - OP_LOGE(opName_, "the dim num of %s is %zu, it should be %u.", SINKS_NAME.c_str(), - opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum(), DIM_NUM_ONE); + OP_LOGE(opName_, "the dim num of %s is %zu, it should be %u.", SINKS_NAME.c_str(), + opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum(), DIM_NUM_ONE); return ge::GRAPH_FAILED; } if (opParamInfo_.sinks.tensor->GetStorageShape().GetDim(0) != n1Size_) { OP_LOGE(opName_, "%s's dimension(%ld) should be equal to query head num(%u).", SINKS_NAME.c_str(), - opParamInfo_.sinks.tensor->GetStorageShape().GetDim(0), n1Size_); + opParamInfo_.sinks.tensor->GetStorageShape().GetDim(0), n1Size_); return ge::GRAPH_FAILED; } OP_CHECK_IF(opParamInfo_.sinks.desc->GetDataType() != ge::DT_FLOAT, - OP_LOGE(opName_, "sinks's dtype must be DT_FLOAT."), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "sinks's dtype must be DT_FLOAT."), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1353,48 +1330,46 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaMetadata() const return ge::GRAPH_FAILED; } OP_CHECK_IF((opParamInfo_.metadata.tensor->GetShapeSize() != METADATA_LIMIT), - OP_LOGE(opName_, "input metadata dim 0 must be %u.", METADATA_LIMIT), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "input metadata dim 0 must be %u.", METADATA_LIMIT), return ge::GRAPH_FAILED); OP_CHECK_IF(opParamInfo_.metadata.desc->GetDataType() != ge::DT_INT32, - OP_LOGE(opName_, "metadata's dtype must be DT_INT32."), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "metadata's dtype must be DT_INT32."), return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaCmpRatio() const +{ + if (IsA5Arch(npuArch_)) { + if (opParamInfo_.cmpKv.tensor != nullptr) { + OP_CHECK_IF(cmpRatio_ < 1 || cmpRatio_ > 128, + OP_LOGE(opName_, "cmpRatio should be in range [1, 128] on %s, but got %ld.", + A5_PLATFORM_LOG.c_str(), cmpRatio_), + return ge::GRAPH_FAILED); + } + } else { + uint32_t expectedCmpRatio = 1; + const char *modeName = "SWA"; + const char *modeReason = "when cmp_kv is not provided"; + if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + expectedCmpRatio = 4; + modeName = "CSA"; + modeReason = "when cmp_sparse_indices is provided"; + } else if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE) { + expectedCmpRatio = 128; + modeName = "HCA"; + modeReason = "when cmp_sparse_indices is not provided"; + } + OP_CHECK_IF(cmpRatio_ != expectedCmpRatio, + OP_LOGE(opName_, "cmpRatio should be %u in %s on %s %s, but got %ld.", expectedCmpRatio, modeName, + A2_A3_PLATFORM_LOG.c_str(), modeReason, cmpRatio_), + return ge::GRAPH_FAILED); + } return ge::GRAPH_SUCCESS; } -ge::graphStatus SMLATilingCheck::CheckSingleParaCmpRatio() const -{ - if (IsA5Arch(npuArch_)) { - if (opParamInfo_.cmpKv.tensor != nullptr) { - OP_CHECK_IF(cmpRatio_ < 1 || cmpRatio_ > 128, - OP_LOGE(opName_, "cmpRatio should be in range [1, 128] on %s, but got %ld.", - A5_PLATFORM_LOG.c_str(), cmpRatio_), - return ge::GRAPH_FAILED); - } - } else { - uint32_t expectedCmpRatio = 1; - const char *modeName = "SWA"; - const char *modeReason = "when cmp_kv is not provided"; - if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { - expectedCmpRatio = 4; - modeName = "CSA"; - modeReason = "when cmp_sparse_indices is provided"; - } else if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE) { - expectedCmpRatio = 128; - modeName = "HCA"; - modeReason = "when cmp_sparse_indices is not provided"; - } - OP_CHECK_IF(cmpRatio_ != expectedCmpRatio, - OP_LOGE(opName_, "cmpRatio should be %u in %s on %s %s, but got %ld.", - expectedCmpRatio, modeName, A2_A3_PLATFORM_LOG.c_str(), modeReason, cmpRatio_), - return ge::GRAPH_FAILED); - } - return ge::GRAPH_SUCCESS; -} - -ge::graphStatus SMLATilingCheck::CheckSingleParaOriMaskMode() const -{ - return ge::GRAPH_SUCCESS; -} +ge::graphStatus SMLATilingCheck::CheckSingleParaOriMaskMode() const +{ + return ge::GRAPH_SUCCESS; +} ge::graphStatus SMLATilingCheck::CheckSingleParaCmpMaskMode() const { @@ -1414,53 +1389,41 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriWinRight() const ge::graphStatus SMLATilingCheck::CheckSingleParaCmpResidualKv() const { bool isCmpTemplate = smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || - smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE; - if (isCmpTemplate && *opParamInfo_.cmpMaskMode == 3 && cmpRatio_ != 1) { - OP_CHECK_IF(opParamInfo_.cmpResidualKv.tensor == nullptr, - OP_LOGE(opName_, "cmp_residual_kv is required when cmp_mask_mode=3 and cmp_ratio != 1"), - return ge::GRAPH_FAILED); - } - return ge::GRAPH_SUCCESS; -} - -ge::graphStatus SMLATilingCheck::CheckSingleParaTopkLength() const -{ - if (IsA5Arch(npuArch_)) { - return ge::GRAPH_SUCCESS; - } - OP_CHECK_IF(IsNonEmptyOptionalTensor(opParamInfo_.oriTopkLength.tensor), - OP_LOGE(opName_, "ori_topk_length is reserved and must be empty on %s.", - A2_A3_PLATFORM_LOG.c_str()), - return ge::GRAPH_FAILED); - OP_CHECK_IF(IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor), - OP_LOGE(opName_, "cmp_topk_length is reserved and must be empty on %s.", - A2_A3_PLATFORM_LOG.c_str()), - return ge::GRAPH_FAILED); - return ge::GRAPH_SUCCESS; -} - -ge::graphStatus SMLATilingCheck::CheckSinglePara() const -{ - if ( - ge::GRAPH_SUCCESS != CheckSingleParaQuery() || - ge::GRAPH_SUCCESS != CheckSingleParaOriKv() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpKv() || - ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensOriKv() || - ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensCmpKv() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpRatio() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpResidualKv() || - ge::GRAPH_SUCCESS != CheckSingleParaTopkLength() || - ge::GRAPH_SUCCESS != CheckSingleParaNumHeads() || - ge::GRAPH_SUCCESS != CheckSingleParaKvHeadNums() || - ge::GRAPH_SUCCESS != CheckSingleParaOriSparseIndices() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpSparseIndices() || - ge::GRAPH_SUCCESS != CheckSingleParaOriBlockTable() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpBlockTable() || - ge::GRAPH_SUCCESS != CheckSingleParaSinks() || - ge::GRAPH_SUCCESS != CheckSingleParaMetadata() || - ge::GRAPH_SUCCESS != CheckSingleParaOriMaskMode() || - ge::GRAPH_SUCCESS != CheckSingleParaCmpMaskMode() || - ge::GRAPH_SUCCESS != CheckSingleParaOriWinLeft() || + smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE; + if (isCmpTemplate && *opParamInfo_.cmpMaskMode == 3 && cmpRatio_ != 1) { + OP_CHECK_IF(opParamInfo_.cmpResidualKv.tensor == nullptr, + OP_LOGE(opName_, "cmp_residual_kv is required when cmp_mask_mode=3 and cmp_ratio != 1"), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSingleParaTopkLength() const +{ + if (IsA5Arch(npuArch_)) { + return ge::GRAPH_SUCCESS; + } + OP_CHECK_IF(IsNonEmptyOptionalTensor(opParamInfo_.oriTopkLength.tensor), + OP_LOGE(opName_, "ori_topk_length is reserved and must be empty on %s.", A2_A3_PLATFORM_LOG.c_str()), + return ge::GRAPH_FAILED); + OP_CHECK_IF(IsNonEmptyOptionalTensor(opParamInfo_.cmpTopkLength.tensor), + OP_LOGE(opName_, "cmp_topk_length is reserved and must be empty on %s.", A2_A3_PLATFORM_LOG.c_str()), + return ge::GRAPH_FAILED); + return ge::GRAPH_SUCCESS; +} + +ge::graphStatus SMLATilingCheck::CheckSinglePara() const +{ + if (ge::GRAPH_SUCCESS != CheckSingleParaQuery() || ge::GRAPH_SUCCESS != CheckSingleParaOriKv() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpKv() || ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensOriKv() || + ge::GRAPH_SUCCESS != CheckSingleParaCuSeqLensCmpKv() || ge::GRAPH_SUCCESS != CheckSingleParaCmpRatio() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpResidualKv() || ge::GRAPH_SUCCESS != CheckSingleParaTopkLength() || + ge::GRAPH_SUCCESS != CheckSingleParaNumHeads() || ge::GRAPH_SUCCESS != CheckSingleParaKvHeadNums() || + ge::GRAPH_SUCCESS != CheckSingleParaOriSparseIndices() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpSparseIndices() || ge::GRAPH_SUCCESS != CheckSingleParaOriBlockTable() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpBlockTable() || ge::GRAPH_SUCCESS != CheckSingleParaSinks() || + ge::GRAPH_SUCCESS != CheckSingleParaMetadata() || ge::GRAPH_SUCCESS != CheckSingleParaOriMaskMode() || + ge::GRAPH_SUCCESS != CheckSingleParaCmpMaskMode() || ge::GRAPH_SUCCESS != CheckSingleParaOriWinLeft() || ge::GRAPH_SUCCESS != CheckSingleParaOriWinRight()) { return ge::GRAPH_FAILED; } @@ -1469,23 +1432,19 @@ ge::graphStatus SMLATilingCheck::CheckSinglePara() const ge::graphStatus SMLATilingCheck::CheckExists(const void *pointer, const std::string &name) const { - OP_CHECK_IF(pointer == nullptr, - OP_LOGE(opName_, "%s should not be null", name.c_str()), - return ge::GRAPH_FAILED); + OP_CHECK_IF(pointer == nullptr, OP_LOGE(opName_, "%s should not be null", name.c_str()), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } ge::graphStatus SMLATilingCheck::CheckNotExists(const void *pointer, const std::string &name) const { - OP_CHECK_IF(pointer != nullptr, - OP_LOGE(opName_, "%s should be null", name.c_str()), - return ge::GRAPH_FAILED); + OP_CHECK_IF(pointer != nullptr, OP_LOGE(opName_, "%s should be null", name.c_str()), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } ge::graphStatus SMLATilingCheck::CheckExistsByMap(const std::map ¶mMap) const { - for (const auto& kv : paramMap) { + for (const auto &kv : paramMap) { if (CheckExists(kv.second, kv.first) != ge::GRAPH_SUCCESS) { return ge::GRAPH_FAILED; } @@ -1495,7 +1454,7 @@ ge::graphStatus SMLATilingCheck::CheckExistsByMap(const std::map ¶mMap) const { - for (const auto& kv : paramMap) { + for (const auto &kv : paramMap) { if (CheckNotExists(kv.second, kv.first) != ge::GRAPH_SUCCESS) { return ge::GRAPH_FAILED; } @@ -1504,7 +1463,7 @@ ge::graphStatus SMLATilingCheck::CheckNotExistsByMap(const std::map &existMap, - std::map ¬ExistMap) const + std::map ¬ExistMap) const { if (CheckExistsByMap(existMap) != ge::GRAPH_SUCCESS) { return ge::GRAPH_FAILED; @@ -1519,8 +1478,7 @@ ge::graphStatus SMLATilingCheck::CheckParaExistence() const { if (npuArch_ == NpuArch::DAV_3510) { OP_CHECK_IF((kvLayout_ == SMLALayout::TND && opParamInfo_.cuSeqLensOriKv.tensor == nullptr), - OP_LOGE(opName_, "cuSeqLensOriKv must be provided when kv layout is TND"), - return ge::GRAPH_FAILED); + OP_LOGE(opName_, "cuSeqLensOriKv must be provided when kv layout is TND"), return ge::GRAPH_FAILED); } else { if (kvLayout_ == SMLALayout::PA_BBND) { std::map ParamExistMap = { @@ -1538,43 +1496,41 @@ ge::graphStatus SMLATilingCheck::CheckParaExistence() const ge::graphStatus SMLATilingCheck::CheckFeatureShape() const { - OP_CHECK_IF(bSize_ <= 0, - OP_LOGE(opName_, "batch_size should be greater than 0, but got %u", bSize_), + 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_ == SMLALayout::TND), OP_LOGE(opName_, "T_size of query should be greater than 0, but got %u", qTSize_), return ge::GRAPH_FAILED); - if (IsA5Arch(npuArch_)) { - OP_CHECK_IF(n1Size_ < 1 || n1Size_ > 128, - OP_LOGE(opName_, "q_head_num should be in [1, 128] on %s, but got %u", - A5_PLATFORM_LOG.c_str(), n1Size_), - return ge::GRAPH_FAILED); - } else { - OP_CHECK_IF(!IsPowerOfTwoInRange(n1Size_, 1, 128), - OP_LOGE(opName_, "q_head_num should be power of two in [1, 128] on %s, but got %u", - A2_A3_PLATFORM_LOG.c_str(), n1Size_), - return ge::GRAPH_FAILED); - } - - OP_CHECK_IF(n2Size_ != 1, - OP_LOGE(opName_, "kv_head_num should be 1, but got %u", n2Size_), + if (IsA5Arch(npuArch_)) { + OP_CHECK_IF( + n1Size_ < 1 || n1Size_ > 128, + OP_LOGE(opName_, "q_head_num should be in [1, 128] on %s, but got %u", A5_PLATFORM_LOG.c_str(), n1Size_), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(!IsPowerOfTwoInRange(n1Size_, 1, 128), + OP_LOGE(opName_, "q_head_num should be power of two in [1, 128] on %s, but got %u", + A2_A3_PLATFORM_LOG.c_str(), n1Size_), + return ge::GRAPH_FAILED); + } + + OP_CHECK_IF(n2Size_ != 1, OP_LOGE(opName_, "kv_head_num should be 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); - if (IsA5Arch(npuArch_)) { - OP_CHECK_IF(gSize_ < 1 || gSize_ > 128, - OP_LOGE(opName_, "group num should be in [1, 128] on %s, but got %u", - A5_PLATFORM_LOG.c_str(), gSize_), - return ge::GRAPH_FAILED); - } else { - OP_CHECK_IF(!IsPowerOfTwoInRange(gSize_, 1, 128), - OP_LOGE(opName_, "group num should be power of two in [1, 128] on %s, but got %u", - A2_A3_PLATFORM_LOG.c_str(), gSize_), - return ge::GRAPH_FAILED); + if (IsA5Arch(npuArch_)) { + OP_CHECK_IF( + gSize_ < 1 || gSize_ > 128, + OP_LOGE(opName_, "group num should be in [1, 128] on %s, but got %u", A5_PLATFORM_LOG.c_str(), gSize_), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(!IsPowerOfTwoInRange(gSize_, 1, 128), + OP_LOGE(opName_, "group num should be power of two in [1, 128] on %s, but got %u", + A2_A3_PLATFORM_LOG.c_str(), gSize_), + return ge::GRAPH_FAILED); } OP_CHECK_IF(qHeadDim_ != DIM_LIMIT, @@ -1591,64 +1547,59 @@ ge::graphStatus SMLATilingCheck::CheckFeatureShape() const OP_CHECK_IF(!(qType_ == oriKvType_), OP_LOGE(opName_, - "Head dimension data type check failed! qType[%s] must be the same with oriKvType[%s].", - SMLADataTypeToSerialString(qType_).c_str(), - SMLADataTypeToSerialString(oriKvType_).c_str()), + "Head dimension data type check failed! qType[%s] must be the same with oriKvType[%s].", + SMLADataTypeToSerialString(qType_).c_str(), SMLADataTypeToSerialString(oriKvType_).c_str()), return ge::GRAPH_FAILED); - if (IsA5Arch(npuArch_)) { - OP_CHECK_IF(*opParamInfo_.oriMaskMode != 0 && *opParamInfo_.oriMaskMode != 3 && *opParamInfo_.oriMaskMode != 4, - OP_LOGE(opName_, "oriMaskMode should be {0, 3, 4} on %s, but got %u", - A5_PLATFORM_LOG.c_str(), *opParamInfo_.oriMaskMode), - return ge::GRAPH_FAILED); - OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 0 && *opParamInfo_.cmpMaskMode != 3, - OP_LOGE(opName_, "cmpMaskMode should be {0, 3} on %s, but got %u", - A5_PLATFORM_LOG.c_str(), *opParamInfo_.cmpMaskMode), - return ge::GRAPH_FAILED); - OP_CHECK_IF(topkValueMode_ != 1, - OP_LOGE(opName_, "topkValueMode should be 1, but got %ld", topkValueMode_), - return ge::GRAPH_FAILED); - OP_CHECK_IF(oriWinLeft_ < -1, - OP_LOGE(opName_, "oriWinLeft_ should be -1(unlimited) or non-negative on %s, but got %ld", - A5_PLATFORM_LOG.c_str(), oriWinLeft_), - return ge::GRAPH_FAILED); - OP_CHECK_IF(oriWinRight_ < -1, - OP_LOGE(opName_, "oriWinRight_ should be -1(unlimited) or non-negative on %s, but got %ld", - A5_PLATFORM_LOG.c_str(), oriWinRight_), - return ge::GRAPH_FAILED); - } else { - OP_CHECK_IF(*opParamInfo_.oriMaskMode != 4, - OP_LOGE(opName_, "oriMaskMode should be 4 on %s, but got %u", - A2_A3_PLATFORM_LOG.c_str(), *opParamInfo_.oriMaskMode), - return ge::GRAPH_FAILED); - if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || - smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { - OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 3, - OP_LOGE(opName_, "cmpMaskMode should be 3 on %s, but got %u", - A2_A3_PLATFORM_LOG.c_str(), *opParamInfo_.cmpMaskMode), - return ge::GRAPH_FAILED); - } - OP_CHECK_IF(oriWinLeft_ != 127, - OP_LOGE(opName_, "oriWinLeft_ should be 127 on %s, but got %ld", - A2_A3_PLATFORM_LOG.c_str(), oriWinLeft_), - return ge::GRAPH_FAILED); - OP_CHECK_IF(oriWinRight_ != 0, - OP_LOGE(opName_, "oriWinRight_ should be 0 on %s, but got %ld", - A2_A3_PLATFORM_LOG.c_str(), oriWinRight_), - return ge::GRAPH_FAILED); - } - return ge::GRAPH_SUCCESS; -} + if (IsA5Arch(npuArch_)) { + OP_CHECK_IF(*opParamInfo_.oriMaskMode != 0 && *opParamInfo_.oriMaskMode != 3 && *opParamInfo_.oriMaskMode != 4, + OP_LOGE(opName_, "oriMaskMode should be {0, 3, 4} on %s, but got %u", A5_PLATFORM_LOG.c_str(), + *opParamInfo_.oriMaskMode), + return ge::GRAPH_FAILED); + OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 0 && *opParamInfo_.cmpMaskMode != 3, + OP_LOGE(opName_, "cmpMaskMode should be {0, 3} on %s, but got %u", A5_PLATFORM_LOG.c_str(), + *opParamInfo_.cmpMaskMode), + return ge::GRAPH_FAILED); + OP_CHECK_IF(topkValueMode_ != 1, OP_LOGE(opName_, "topkValueMode should be 1, but got %ld", topkValueMode_), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinLeft_ < -1, + OP_LOGE(opName_, "oriWinLeft_ should be -1(unlimited) or non-negative on %s, but got %ld", + A5_PLATFORM_LOG.c_str(), oriWinLeft_), + return ge::GRAPH_FAILED); + OP_CHECK_IF(oriWinRight_ < -1, + OP_LOGE(opName_, "oriWinRight_ should be -1(unlimited) or non-negative on %s, but got %ld", + A5_PLATFORM_LOG.c_str(), oriWinRight_), + return ge::GRAPH_FAILED); + } else { + OP_CHECK_IF(*opParamInfo_.oriMaskMode != 4, + OP_LOGE(opName_, "oriMaskMode should be 4 on %s, but got %u", A2_A3_PLATFORM_LOG.c_str(), + *opParamInfo_.oriMaskMode), + return ge::GRAPH_FAILED); + if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 3, + OP_LOGE(opName_, "cmpMaskMode should be 3 on %s, but got %u", A2_A3_PLATFORM_LOG.c_str(), + *opParamInfo_.cmpMaskMode), + return ge::GRAPH_FAILED); + } + OP_CHECK_IF( + oriWinLeft_ != 127, + OP_LOGE(opName_, "oriWinLeft_ should be 127 on %s, but got %ld", A2_A3_PLATFORM_LOG.c_str(), oriWinLeft_), + return ge::GRAPH_FAILED); + OP_CHECK_IF( + oriWinRight_ != 0, + OP_LOGE(opName_, "oriWinRight_ should be 0 on %s, but got %ld", A2_A3_PLATFORM_LOG.c_str(), oriWinRight_), + return ge::GRAPH_FAILED); + } + return ge::GRAPH_SUCCESS; +} ge::graphStatus SMLATilingCheck::CheckFeatureLayout() const { - const std::vector layoutQuerySupportList = { - "BSND", - "TND" - }; + const std::vector layoutQuerySupportList = {"BSND", "TND"}; std::string layoutQuery = opParamInfo_.layoutQ; OP_CHECK_IF(std::find(layoutQuerySupportList.begin(), layoutQuerySupportList.end(), layoutQuery) == - layoutQuerySupportList.end(), + layoutQuerySupportList.end(), OP_LOGE(opName_, "layoutQuery only supports BSND/TND, but got %s", layoutQuery.c_str()), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; @@ -1658,8 +1609,8 @@ ge::graphStatus SMLATilingCheck::CheckFeatureDtype() const { OP_CHECK_IF(qType_ != ge::DT_BF16 && qType_ != ge::DT_FLOAT16, OP_LOGE(opName_, "query dtype only support %s and %s, but got %s", - SMLADataTypeToSerialString(ge::DT_BF16).c_str(), SMLADataTypeToSerialString(ge::DT_FLOAT16).c_str(), - SMLADataTypeToSerialString(qType_).c_str()), + SMLADataTypeToSerialString(ge::DT_BF16).c_str(), + SMLADataTypeToSerialString(ge::DT_FLOAT16).c_str(), SMLADataTypeToSerialString(qType_).c_str()), return ge::GRAPH_FAILED); return ge::GRAPH_SUCCESS; } @@ -1671,10 +1622,8 @@ ge::graphStatus SMLATilingCheck::CheckFeaturePa() const ge::graphStatus SMLATilingCheck::CheckFeature() const { - if (ge::GRAPH_SUCCESS != CheckFeatureShape() || - ge::GRAPH_SUCCESS != CheckFeatureLayout() || - ge::GRAPH_SUCCESS != CheckFeatureDtype() || - ge::GRAPH_SUCCESS != CheckFeaturePa()) { + if (ge::GRAPH_SUCCESS != CheckFeatureShape() || ge::GRAPH_SUCCESS != CheckFeatureLayout() || + ge::GRAPH_SUCCESS != CheckFeatureDtype() || ge::GRAPH_SUCCESS != CheckFeaturePa()) { return ge::GRAPH_FAILED; } return ge::GRAPH_SUCCESS; @@ -1683,24 +1632,23 @@ ge::graphStatus SMLATilingCheck::CheckFeature() const void SMLATilingCheck::SetSMLAShapeCompare() { queryShapeCmp_ = opParamInfo_.q.shape->GetStorageShape(); - oriKvShapeCmp_= opParamInfo_.oriKv.tensor->GetShape().GetStorageShape(); + oriKvShapeCmp_ = opParamInfo_.oriKv.tensor->GetShape().GetStorageShape(); attenOutShapeCmp_ = opParamInfo_.attnOut.shape->GetStorageShape(); if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { - cmpKvShapeCmp_= opParamInfo_.cmpKv.tensor->GetShape().GetStorageShape(); + cmpKvShapeCmp_ = opParamInfo_.cmpKv.tensor->GetShape().GetStorageShape(); } if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { cmpKvSparseIndicesCmp_ = opParamInfo_.cmpSparseIndices.tensor->GetShape().GetStorageShape(); } } -ge::graphStatus SMLATilingCheck::CheckDTypeConsistency(const ge::DataType &actualDtype, - const ge::DataType &expectDtype, const std::string &name) const +ge::graphStatus SMLATilingCheck::CheckDTypeConsistency(const ge::DataType &actualDtype, const ge::DataType &expectDtype, + const std::string &name) const { if (actualDtype != expectDtype) { OP_LOGE(opName_, "%s dtype should be the same to %s, but it's %s.", name.c_str(), - SMLADataTypeToSerialString(expectDtype).c_str(), - SMLADataTypeToSerialString(actualDtype).c_str()); + SMLADataTypeToSerialString(expectDtype).c_str(), SMLADataTypeToSerialString(actualDtype).c_str()); return ge::GRAPH_FAILED; } return ge::GRAPH_SUCCESS; @@ -1710,8 +1658,7 @@ ge::graphStatus SMLATilingCheck::CheckOriAndCmpKv() const { if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { - if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(cmpKvType_, - oriKvType_, CMP_KV_NAME)) { + if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(cmpKvType_, oriKvType_, CMP_KV_NAME)) { return ge::GRAPH_FAILED; } } @@ -1740,11 +1687,8 @@ ge::graphStatus SMLATilingCheck::CheckBlockTable() const ge::graphStatus SMLATilingCheck::CheckMultiParaConsistency() { SetSMLAShapeCompare(); - if ( - ge::GRAPH_SUCCESS != CheckOriAndCmpKv() || - ge::GRAPH_SUCCESS != CheckAttenOut() || - ge::GRAPH_SUCCESS != CheckActualSeqLensQ() || - ge::GRAPH_SUCCESS != CheckActualSeqLens() || + if (ge::GRAPH_SUCCESS != CheckOriAndCmpKv() || ge::GRAPH_SUCCESS != CheckAttenOut() || + ge::GRAPH_SUCCESS != CheckActualSeqLensQ() || ge::GRAPH_SUCCESS != CheckActualSeqLens() || ge::GRAPH_SUCCESS != CheckBlockTable()) { return ge::GRAPH_FAILED; } @@ -1754,12 +1698,8 @@ ge::graphStatus SMLATilingCheck::CheckMultiParaConsistency() ge::graphStatus SMLATilingCheck::Process() { Init(); - if ( - CheckSinglePara() != ge::GRAPH_SUCCESS || - CheckParaExistence() != ge::GRAPH_SUCCESS || - CheckFeature() != ge::GRAPH_SUCCESS || - CheckMultiParaConsistency() != ge::GRAPH_SUCCESS - ) { + if (CheckSinglePara() != ge::GRAPH_SUCCESS || CheckParaExistence() != ge::GRAPH_SUCCESS || + CheckFeature() != ge::GRAPH_SUCCESS || CheckMultiParaConsistency() != ge::GRAPH_SUCCESS) { return ge::GRAPH_FAILED; } return ge::GRAPH_SUCCESS; @@ -1781,7 +1721,8 @@ void SparseFlashMlaTiling::SplitBalanced(SMLATilingInfo *tilingInfo) sInnerSizeAlign_ = Align(sInnerSize_, BYTE_BLOCK); if (tilingInfo->npuArch == NpuArch::DAV_2201) { mBaseSize_ = tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE ? - tilingInfo->gSize : (256 / tilingInfo->gSize) * tilingInfo->gSize; + tilingInfo->gSize : + (256 / tilingInfo->gSize) * tilingInfo->gSize; } headDimAlign_ = Align(tilingInfo->qHeadDim, BYTE_BLOCK); CalcUbBmm(tilingInfo); @@ -1887,15 +1828,14 @@ ge::graphStatus SparseFlashMlaTiling::DoOpTiling(SMLATilingInfo *tilingInfo) uint32_t tilingKey; uint32_t splitG = 0U; - uint32_t headRatioOne = static_cast( - tilingInfo->npuArch == NpuArch::DAV_2201 && - tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE && - tilingInfo->gSize == 1U); + uint32_t headRatioOne = + static_cast(tilingInfo->npuArch == NpuArch::DAV_2201 && + tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE && tilingInfo->gSize == 1U); if (tilingInfo->npuArch == NpuArch::DAV_3510) { splitG = static_cast(tilingInfo->gSize > 64); } tilingKey = GET_TPL_TILING_KEY(0U, qLayout, inputKvLayout, static_cast(tilingInfo->perfMode), splitG, - headRatioOne); + headRatioOne); context_->SetScheduleMode(1); context_->SetTilingKey(tilingKey); diff --git a/torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py b/torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py index 9f99b84..1cea4e3 100644 --- a/torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py +++ b/torch_extension/cann_ops_transformer/ops/sparse_flash_mla.py @@ -10,7 +10,6 @@ from typing import Optional import torch -import torch_npu from torch.library import impl from cann_ops_transformer.op_builder.builder import OpBuilder from cann_ops_transformer.op_builder.builder import AS_LIBRARY @@ -25,7 +24,7 @@ class SparseFlashMlaOpBuilder(OpBuilder): def sources(self): """Path to C++ source code.""" - return ['ops/csrc/sparse_flash_mla.cpp'] + return ["ops/csrc/sparse_flash_mla.cpp"] def schema(self) -> str: """PyTorch operator signature.""" @@ -38,7 +37,6 @@ class SparseFlashMlaOpBuilder(OpBuilder): "int? ori_topk=None, int? cmp_topk=None, int? cmp_ratio=None, int? ori_mask_mode=None," "int? cmp_mask_mode=None, int? ori_win_left=None, int? ori_win_right=None, str? layout_q=None," "str? layout_kv=None, bool? has_ori_kv=None, bool? has_cmp_kv=None) -> Tensor", - "sparse_flash_mla(Tensor q, *," "Tensor? ori_kv=None, Tensor? cmp_kv=None, " "Tensor? ori_sparse_indices=None, Tensor? cmp_sparse_indices=None, " @@ -52,83 +50,145 @@ class SparseFlashMlaOpBuilder(OpBuilder): "float softmax_scale=1.0, int cmp_ratio=1, " "int ori_mask_mode=0, int cmp_mask_mode=0, " "int ori_win_left=-1, int ori_win_right=-1, " - "str layout_q=\"BSND\", str layout_kv=\"BSND\", " - "int topk_value_mode=1, bool return_softmax_lse=False) -> (Tensor, Tensor)" + 'str layout_q="BSND", str layout_kv="BSND", ' + "int topk_value_mode=1, bool return_softmax_lse=False) -> (Tensor, Tensor)", ] - + def register_meta(self): """ Registers the Meta implementation (Shape/Dtype inference). Essential for Autograd and FakeTensor support. """ + @torch.library.register_fake("cann_ops_transformer::" + SMLA_METADATA_OP_NAME) def sparse_flash_mla_metadata_meta( - num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None, - cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, - seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None, - seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None, - ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None, - batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, - max_seqlen_ori_kv: Optional[int] = None, max_seqlen_cmp_kv: Optional[int] = None, - ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None, cmp_ratio: Optional[int] = None, - ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None, - ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None, - layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None): + num_heads_q: int, + num_heads_kv: int, + head_dim: int, + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_ori_kv: Optional[torch.Tensor] = None, + cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, + seqused_q: Optional[torch.Tensor] = None, + seqused_ori_kv: Optional[torch.Tensor] = None, + seqused_cmp_kv: Optional[torch.Tensor] = None, + cmp_residual_kv: Optional[torch.Tensor] = None, + ori_topk_length: Optional[torch.Tensor] = None, + cmp_topk_length: Optional[torch.Tensor] = None, + batch_size: Optional[int] = None, + max_seqlen_q: Optional[int] = None, + max_seqlen_ori_kv: Optional[int] = None, + max_seqlen_cmp_kv: Optional[int] = None, + ori_topk: Optional[int] = None, + cmp_topk: Optional[int] = None, + cmp_ratio: Optional[int] = None, + ori_mask_mode: Optional[int] = None, + cmp_mask_mode: Optional[int] = None, + ori_win_left: Optional[int] = None, + ori_win_right: Optional[int] = None, + layout_q: Optional[str] = None, + layout_kv: Optional[str] = None, + has_ori_kv: Optional[bool] = None, + has_cmp_kv: Optional[bool] = None, + ): return torch.empty((SMLA_METADATA_SIZE), dtype=torch.int32, device="npu") - @impl(AS_LIBRARY, self.name, "Meta") - def sparse_flash_mla_meta(q, - ori_kv=None, cmp_kv=None, - ori_sparse_indices=None, cmp_sparse_indices=None, - ori_block_table=None, cmp_block_table=None, - cu_seqlens_q=None, cu_seqlens_ori_kv=None, - cu_seqlens_cmp_kv=None, seqused_q=None, - seqused_ori_kv=None, seqused_cmp_kv=None, - cmp_residual_kv=None, - ori_topk_length=None, cmp_topk_length=None, - sinks=None, metadata=None, - softmax_scale=1.0, cmp_ratio=1, - ori_mask_mode=0, cmp_mask_mode=0, - ori_win_left=-1, ori_win_right=-1, - layout_q='BSND', layout_kv='BSND', - topk_value_mode=1, return_softmax_lse=False): + def sparse_flash_mla_meta( + q, + ori_kv=None, + cmp_kv=None, + ori_sparse_indices=None, + cmp_sparse_indices=None, + ori_block_table=None, + cmp_block_table=None, + cu_seqlens_q=None, + cu_seqlens_ori_kv=None, + cu_seqlens_cmp_kv=None, + seqused_q=None, + seqused_ori_kv=None, + seqused_cmp_kv=None, + cmp_residual_kv=None, + ori_topk_length=None, + cmp_topk_length=None, + sinks=None, + metadata=None, + softmax_scale=1.0, + cmp_ratio=1, + ori_mask_mode=0, + cmp_mask_mode=0, + ori_win_left=-1, + ori_win_right=-1, + layout_q="BSND", + layout_kv="BSND", + topk_value_mode=1, + return_softmax_lse=False, + ): key_headnum = ori_kv.shape[1] if layout_kv == "TND" else ori_kv.shape[2] if layout_q == "BSND": ## 添加softmax_lse attn_out = torch.empty(q.shape, dtype=q.dtype, device="meta") if return_softmax_lse: - softmax_lse = torch.empty([q.shape[0], ori_kv.shape[2], q.shape[1], q.shape[2] / ori_kv.shape[2]], - dtype=torch.float32, device="meta") + softmax_lse = torch.empty( + [ + q.shape[0], + ori_kv.shape[2], + q.shape[1], + q.shape[2] / ori_kv.shape[2], + ], + dtype=torch.float32, + device="meta", + ) else: # 给一个空的合法张量,不能是 nullptr softmax_lse = torch.empty([], dtype=torch.float32, device="meta") else: attn_out = torch.empty(q.shape, dtype=q.dtype, device="meta") if return_softmax_lse: - softmax_lse = torch.empty([ori_kv.shape[1], q.shape[0], q.shape[1] / ori_kv.shape[1]], - dtype=torch.float32, device="meta") + softmax_lse = torch.empty( + [ori_kv.shape[1], q.shape[0], q.shape[1] / ori_kv.shape[1]], + dtype=torch.float32, + device="meta", + ) else: # 给一个空的合法张量,不能是 nullptr softmax_lse = torch.empty([], dtype=torch.float32, device="meta") return (attn_out, softmax_lse) + # Instantiate the builder sparse_flash_mla_op_builder = SparseFlashMlaOpBuilder() @impl(AS_LIBRARY, SMLA_METADATA_OP_NAME, "PrivateUse1") def sparse_flash_mla_metadata( - num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None, - cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, - seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None, - seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None, - ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None, - batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None, - max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None, - cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None, - ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None, - layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None): + num_heads_q: int, + num_heads_kv: int, + head_dim: int, + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_ori_kv: Optional[torch.Tensor] = None, + cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, + seqused_q: Optional[torch.Tensor] = None, + seqused_ori_kv: Optional[torch.Tensor] = None, + seqused_cmp_kv: Optional[torch.Tensor] = None, + cmp_residual_kv: Optional[torch.Tensor] = None, + ori_topk_length: Optional[torch.Tensor] = None, + cmp_topk_length: Optional[torch.Tensor] = None, + batch_size: Optional[int] = None, + max_seqlen_q: Optional[int] = None, + max_seqlen_ori_kv: Optional[int] = None, + max_seqlen_cmp_kv: Optional[int] = None, + ori_topk: Optional[int] = None, + cmp_topk: Optional[int] = None, + cmp_ratio: Optional[int] = None, + ori_mask_mode: Optional[int] = None, + cmp_mask_mode: Optional[int] = None, + ori_win_left: Optional[int] = None, + ori_win_right: Optional[int] = None, + layout_q: Optional[str] = None, + layout_kv: Optional[str] = None, + has_ori_kv: Optional[bool] = None, + has_cmp_kv: Optional[bool] = None, +): """ Dispatcher implementation: NPU. 'PrivateUse1' is dispatch key for custom NPU backends. @@ -151,68 +211,165 @@ def sparse_flash_mla_metadata( has_cmp_kv = True if has_cmp_kv is None else has_cmp_kv return op_module.sparse_flash_mla_metadata( - num_heads_q, num_heads_kv, head_dim, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q, - seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q, - max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left, - ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv) + num_heads_q, + num_heads_kv, + head_dim, + cu_seqlens_q, + cu_seqlens_ori_kv, + cu_seqlens_cmp_kv, + seqused_q, + seqused_ori_kv, + seqused_cmp_kv, + cmp_residual_kv, + ori_topk_length, + cmp_topk_length, + batch_size, + max_seqlen_q, + max_seqlen_ori_kv, + max_seqlen_cmp_kv, + ori_topk, + cmp_topk, + cmp_ratio, + ori_mask_mode, + cmp_mask_mode, + ori_win_left, + ori_win_right, + layout_q, + layout_kv, + has_ori_kv, + has_cmp_kv, + ) @torch.library.register_kernel("cann_ops_transformer::" + SMLA_METADATA_OP_NAME, None) def sparse_flash_mla_metadata_fallback( - num_heads_q: int, num_heads_kv: int, head_dim: int, cu_seqlens_q: Optional[torch.Tensor] = None, - cu_seqlens_ori_kv: Optional[torch.Tensor] = None, cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, - seqused_q: Optional[torch.Tensor] = None, seqused_ori_kv: Optional[torch.Tensor] = None, - seqused_cmp_kv: Optional[torch.Tensor] = None, cmp_residual_kv: Optional[torch.Tensor] = None, - ori_topk_length: Optional[torch.Tensor] = None, cmp_topk_length: Optional[torch.Tensor] = None, - batch_size: Optional[int] = None, max_seqlen_q: Optional[int] = None, max_seqlen_ori_kv: Optional[int] = None, - max_seqlen_cmp_kv: Optional[int] = None, ori_topk: Optional[int] = None, cmp_topk: Optional[int] = None, - cmp_ratio: Optional[int] = None, ori_mask_mode: Optional[int] = None, cmp_mask_mode: Optional[int] = None, - ori_win_left: Optional[int] = None, ori_win_right: Optional[int] = None, layout_q: Optional[str] = None, - layout_kv: Optional[str] = None, has_ori_kv: Optional[bool] = None, has_cmp_kv: Optional[bool] = None): + num_heads_q: int, + num_heads_kv: int, + head_dim: int, + cu_seqlens_q: Optional[torch.Tensor] = None, + cu_seqlens_ori_kv: Optional[torch.Tensor] = None, + cu_seqlens_cmp_kv: Optional[torch.Tensor] = None, + seqused_q: Optional[torch.Tensor] = None, + seqused_ori_kv: Optional[torch.Tensor] = None, + seqused_cmp_kv: Optional[torch.Tensor] = None, + cmp_residual_kv: Optional[torch.Tensor] = None, + ori_topk_length: Optional[torch.Tensor] = None, + cmp_topk_length: Optional[torch.Tensor] = None, + batch_size: Optional[int] = None, + max_seqlen_q: Optional[int] = None, + max_seqlen_ori_kv: Optional[int] = None, + max_seqlen_cmp_kv: Optional[int] = None, + ori_topk: Optional[int] = None, + cmp_topk: Optional[int] = None, + cmp_ratio: Optional[int] = None, + ori_mask_mode: Optional[int] = None, + cmp_mask_mode: Optional[int] = None, + ori_win_left: Optional[int] = None, + ori_win_right: Optional[int] = None, + layout_q: Optional[str] = None, + layout_kv: Optional[str] = None, + has_ori_kv: Optional[bool] = None, + has_cmp_kv: Optional[bool] = None, +): # 处理所有 tensor 都为 None 的情况 # 调用 NPU 实现 return sparse_flash_mla_metadata( - num_heads_q, num_heads_kv, head_dim, cu_seqlens_q, cu_seqlens_ori_kv, cu_seqlens_cmp_kv, seqused_q, - seqused_ori_kv, seqused_cmp_kv, cmp_residual_kv, ori_topk_length, cmp_topk_length, batch_size, max_seqlen_q, - max_seqlen_ori_kv, max_seqlen_cmp_kv, ori_topk, cmp_topk, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left, - ori_win_right, layout_q, layout_kv, has_ori_kv, has_cmp_kv) + num_heads_q, + num_heads_kv, + head_dim, + cu_seqlens_q, + cu_seqlens_ori_kv, + cu_seqlens_cmp_kv, + seqused_q, + seqused_ori_kv, + seqused_cmp_kv, + cmp_residual_kv, + ori_topk_length, + cmp_topk_length, + batch_size, + max_seqlen_q, + max_seqlen_ori_kv, + max_seqlen_cmp_kv, + ori_topk, + cmp_topk, + cmp_ratio, + ori_mask_mode, + cmp_mask_mode, + ori_win_left, + ori_win_right, + layout_q, + layout_kv, + has_ori_kv, + has_cmp_kv, + ) + torch.compiler.allow_in_graph(sparse_flash_mla_metadata) @impl(AS_LIBRARY, sparse_flash_mla_op_builder.name, "PrivateUse1") -def sparse_flash_mla(q, - ori_kv=None, cmp_kv=None, - ori_sparse_indices=None, cmp_sparse_indices=None, - ori_block_table=None, cmp_block_table=None, - cu_seqlens_q=None, cu_seqlens_ori_kv=None, - cu_seqlens_cmp_kv=None, seqused_q=None, - seqused_ori_kv=None, seqused_cmp_kv=None, - cmp_residual_kv=None, - ori_topk_length=None, cmp_topk_length=None, - sinks=None, metadata=None, - softmax_scale=1.0, cmp_ratio=1, - ori_mask_mode=0, cmp_mask_mode=0, - ori_win_left=-1, ori_win_right=-1, - layout_q='BSND', layout_kv='BSND', - topk_value_mode=1, return_softmax_lse=False): +def sparse_flash_mla( + q, + ori_kv=None, + cmp_kv=None, + ori_sparse_indices=None, + cmp_sparse_indices=None, + ori_block_table=None, + cmp_block_table=None, + cu_seqlens_q=None, + cu_seqlens_ori_kv=None, + cu_seqlens_cmp_kv=None, + seqused_q=None, + seqused_ori_kv=None, + seqused_cmp_kv=None, + cmp_residual_kv=None, + ori_topk_length=None, + cmp_topk_length=None, + sinks=None, + metadata=None, + softmax_scale=1.0, + cmp_ratio=1, + ori_mask_mode=0, + cmp_mask_mode=0, + ori_win_left=-1, + ori_win_right=-1, + layout_q="BSND", + layout_kv="BSND", + topk_value_mode=1, + return_softmax_lse=False, +): """ dispatcher implementation for NPU. 'PrivateUse1' is the combine key for custom NPU backends. """ op_module = sparse_flash_mla_op_builder.load() - return op_module.sparse_flash_mla(q, - ori_kv, cmp_kv, - ori_sparse_indices, cmp_sparse_indices, - ori_block_table, cmp_block_table, - cu_seqlens_q, cu_seqlens_ori_kv, - cu_seqlens_cmp_kv, seqused_q, - seqused_ori_kv, seqused_cmp_kv, - cmp_residual_kv, - ori_topk_length, cmp_topk_length, - sinks, metadata, - softmax_scale, cmp_ratio, - ori_mask_mode, cmp_mask_mode, - ori_win_left, ori_win_right, - layout_q, layout_kv, - topk_value_mode, return_softmax_lse) + return op_module.sparse_flash_mla( + q, + ori_kv, + cmp_kv, + ori_sparse_indices, + cmp_sparse_indices, + ori_block_table, + cmp_block_table, + cu_seqlens_q, + cu_seqlens_ori_kv, + cu_seqlens_cmp_kv, + seqused_q, + seqused_ori_kv, + seqused_cmp_kv, + cmp_residual_kv, + ori_topk_length, + cmp_topk_length, + sinks, + metadata, + softmax_scale, + cmp_ratio, + ori_mask_mode, + cmp_mask_mode, + ori_win_left, + ori_win_right, + layout_q, + layout_kv, + topk_value_mode, + return_softmax_lse, + )