From 42ac585a04d1f9f18e86a69cba9e10b4826dcc7f Mon Sep 17 00:00:00 2001 From: cann-robot Date: Tue, 30 Jun 2026 15:54:35 +0800 Subject: [PATCH 1/2] =?UTF-8?q?bsnd=20padding=E5=8A=9F=E8=83=BD+fdworkspac?= =?UTF-8?q?eID=E8=AE=A1=E7=AE=97=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: INeedAnID # message auto-generated for no-merge-commit merge: !7879 merge bsnd_padding into master bsnd padding功能+fdworkspaceID计算修复 Created-by: INeedAnID Commit-by: INeedAnID Merged-by: cann-robot Description: ## 描述 bsnd padding功能+fdworkspaceID计算修复 ## 关联的Issue ## 测试 ## 文档更新 ## 类型标签 - [x] 🐛 Bug修复 - [ ] ✨ 新特性 - [ ] ⚡ 性能优化 - [ ] ♻️ 重构 - [ ] 🧪 测试 - [ ] 📦 构建/CI - [ ] 🔧 配置变更 - [ ] 📝 文档更新 - [ ] ⬆️ 依赖升级 - [ ] 🔒 安全修复 - [ ] 🧹 代码清理 - [ ] ❓ 其他,请描述: See merge request: cann/ops-transformer!7879 --- .../arch35/mixed_quant_sparse_flash_mla_scfa_kernel.h | 5 ++--- .../op_kernel/arch35/sparse_flash_mla_scfa_kernel.h | 9 ++++++++- .../op_kernel/arch35/sparse_flash_mla_swa_kernel.h | 9 ++++++++- 3 files changed, 18 insertions(+), 5 deletions(-) diff --git a/attention/mixed_quant_sparse_flash_mla/op_kernel/arch35/mixed_quant_sparse_flash_mla_scfa_kernel.h b/attention/mixed_quant_sparse_flash_mla/op_kernel/arch35/mixed_quant_sparse_flash_mla_scfa_kernel.h index 057f0eca23..da51229124 100644 --- a/attention/mixed_quant_sparse_flash_mla/op_kernel/arch35/mixed_quant_sparse_flash_mla_scfa_kernel.h +++ b/attention/mixed_quant_sparse_flash_mla/op_kernel/arch35/mixed_quant_sparse_flash_mla_scfa_kernel.h @@ -489,9 +489,8 @@ __aicore__ inline void MixedQuantSparseFlashMlaScfa if (s1NoNeedCalc || s2NoNeedCalc) { continue; } - runParam.s2SplitIdx = s2SplitIdxCounter; - if (runParam.isS2Split && gS1Index == runParam.gs1LoopEndIdx - 1) { - s2SplitIdxCounter++; + if (runParam.isS2Split) { + runParam.s2SplitIdx = s2SplitIdxCounter++; } if constexpr (IS_SPLIT_G) { maxS2LoopCnt -= runParam.s2LoopEndIdx; diff --git a/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_scfa_kernel.h b/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_scfa_kernel.h index 5cdb7598ed..708c32b9b8 100644 --- a/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_scfa_kernel.h +++ b/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_scfa_kernel.h @@ -299,7 +299,14 @@ __aicore__ inline void SparseFlashMlaScfaKernel::Pa } int64_t s1Size = GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); - if (s1Size > s2Size || (LAYOUT_T == SMLA_LAYOUT::TND && hasActualSeqQlen)) { + int64_t expectQs; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + expectQs = GetSeqLen(bIdx, false, hasCuSeqlensQ, + actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + } else { + expectQs = constInfo.s1Size; + } + if (s1Size > s2Size || s1Size < expectQs) { constInfo.needInit = 1; break; } diff --git a/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel.h b/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel.h index f9f823428b..aa697a0133 100644 --- a/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel.h +++ b/attention/sparse_flash_mla/op_kernel/arch35/sparse_flash_mla_swa_kernel.h @@ -280,7 +280,14 @@ __aicore__ inline void SparseFlashMlaSwaKernel::Par } int64_t s1Size = GetSeqLen(bIdx, hasActualSeqQlen, hasCuSeqlensQ, actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); - if (s1Size > s2Size || (LAYOUT_T == SMLA_LAYOUT::TND && hasActualSeqQlen)) { + int64_t expectQs; + if constexpr (LAYOUT_T == SMLA_LAYOUT::TND) { + expectQs = GetSeqLen(bIdx, false, hasCuSeqlensQ, + actualSeqQlenGm, cuSeqlensQGm, constInfo.s1Size); + } else { + expectQs = constInfo.s1Size; + } + if (s1Size > s2Size || s1Size < expectQs) { constInfo.needInit = 1; break; } From 0bf0706974e7e26dd52d2db933265ee79c9118d3 Mon Sep 17 00:00:00 2001 From: Wei_NaChuan Date: Wed, 1 Jul 2026 14:41:52 +0800 Subject: [PATCH 2/2] =?UTF-8?q?SMLA=E6=8C=89stride=E6=95=B0=E7=BB=84?= =?UTF-8?q?=E6=8C=87=E9=92=88=E8=8E=B7=E5=8F=96KV=20stride0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: weinachuan !7403 merge codex/smla-optional-stride into master SMLA按stride数组指针获取KV stride0 Created-by: Wei_NaChuan Commit-by: weinachuan Merged-by: cann-robot Description: ## 描述 SparseFlashMla tiling 中获取 optional input stride 时,按 `GetOptionalInputStride` / `GetInputStride` 返回的 stride 数组指针读取第 0 维 stride,判空后使用 `stride[0]`。同时保留对象型 stride 兼容分支,以及无法获取 stride 时回退连续内存推导的行为。 新增 SMLA stride 调测脚本和一键使用README,构造 BSND + SCFA 场景下 0 轴非连续的 `ori_kv` 和 `cmp_kv`,打印输入 stride 信息,并对比非连续 KV view 与 contiguous KV clone 的输出,用于确认内部 CANN 包中新 stride 接口路径生效。 [#3186](https://gitcode.com/cann/ops-transformer/issues/3186) - 执行 `git diff --check`,通过。 - 执行调测脚本语法检查,通过。 - 编译 `sparse_flash_mla` 与 `sparse_flash_mla_metadata` custom 包,通过。 新增 SMLA stride 调测脚本使用说明。 - [ ] 🐛 Bug修复 - [ ] ✨ 新特性 - [ ] ⚡ 性能优化 - [ ] ♻️ 重构 - [ ] 🧪 测试 - [ ] 📦 构建/CI - [ ] 🔧 配置变更 - [x] 📝 文档更新 - [ ] ⬆️ 依赖升级 - [ ] 🔒 安全修复 - [x] 🧹 代码清理 - [ ] ❓ 其他,请描述: See merge request: cann/ops-transformer!7403 --- attention/sparse_flash_mla/README.md | 38 +- .../docs/aclnnSparseFlashMla.md | 75 ++- .../examples/test_aclnn_sparse_flash_mla.cpp | 26 +- .../op_host/sparse_flash_mla_tiling.cpp | 516 ++++++++------- .../arch22/sparse_flash_mla_swa_block_cube.h | 27 +- .../sparse_flash_mla_swa_block_vector.h | 6 +- .../sparse_flash_mla_metadata/CMakeLists.txt | 6 +- attention/sparse_flash_mla_metadata/README.md | 37 +- .../docs/aclnnSparseFlashMlaMetadata.md | 586 +++++++++++------- .../op_host/sparse_flash_mla_metadata_check.h | 13 +- .../tests/CMakeLists.txt | 16 + .../tests/ut/CMakeLists.txt | 16 + .../tests/ut/op_host/CMakeLists.txt | 16 + .../tests/ut/op_host/op_api/CMakeLists.txt | 13 + .../test_aclnn_sparse_flash_mla_metadata.cpp | 80 +++ docs/zh/op_api_list.md | 5 +- docs/zh/op_list.md | 2 +- docs/zh/torch_api_list.md | 8 +- .../docs/zh/sparse_flash_mla.md | 30 +- 19 files changed, 931 insertions(+), 585 deletions(-) create mode 100644 attention/sparse_flash_mla_metadata/tests/CMakeLists.txt create mode 100644 attention/sparse_flash_mla_metadata/tests/ut/CMakeLists.txt create mode 100644 attention/sparse_flash_mla_metadata/tests/ut/op_host/CMakeLists.txt create mode 100644 attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/CMakeLists.txt create mode 100644 attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/test_aclnn_sparse_flash_mla_metadata.cpp diff --git a/attention/sparse_flash_mla/README.md b/attention/sparse_flash_mla/README.md index 0ca90b96b0..74a0b773cf 100644 --- a/attention/sparse_flash_mla/README.md +++ b/attention/sparse_flash_mla/README.md @@ -2,7 +2,7 @@ ## 产品支持情况 | 产品 | 是否支持 | -| ------------------------------------------------------------ | :------: | +| :------------------------------------------------------------ | :------: | |Ascend 950PR/Ascend 950DT | √ | |Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ | |Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ | @@ -11,7 +11,7 @@ |Atlas 训练系列产品 | × | ## 功能说明 -- 算子功能:`SparseFlashMla`算子旨在完成以下公式描述的Attention计算,支持C1A(Sliding Window Attention,SWA)、C4A(Compressed Sparse Attention,CSA)、C128A(Heavily Compressed Attention,HCA)三类Attention计算场景。 +- 算子功能:`SparseFlashMla`算子旨在完成以下公式描述的Attention计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。 - 计算公式: @@ -163,77 +163,77 @@ metadata 可选输入 - `SparseFlashMlaMetadata`生成的任务切分结果。 + 配套metadata前置接口生成的任务切分结果。 INT32 ND softmax_scale 可选属性 - 对应公式中的softmax_scale。默认值为1.0。 + 对应公式中的softmax_scale。 FLOAT - cmp_ratio 可选属性 - 表示`cmp_kv`相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入`ori_kv`时不参与压缩KV计算。支持1、4、128。默认值为1。 + 表示`cmp_kv`相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入`ori_kv`时不参与压缩KV计算。支持1、4、128。 INT - ori_mask_mode 可选属性 - 表示`q`和`ori_kv`计算的mask模式。
0: No Mask。
3: RightDownCausal模式。
4: Band模式。默认值为0。 + 表示`q`和`ori_kv`计算的mask模式。
0: No Mask。
3: RightDownCausal模式。
4: Band模式。 INT - cmp_mask_mode 可选属性 - 表示`q`和`cmp_kv`计算的mask模式。
0: No Mask。
3: RightDownCausal模式。默认值为0。 + 表示`q`和`cmp_kv`计算的mask模式。
0: No Mask。
3: RightDownCausal模式。 INT - ori_win_left 可选属性 - 表示`q`和`ori_kv`计算中`q`对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限。默认值为-1。 + 表示`q`和`ori_kv`计算中`q`对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限。 INT - ori_win_right 可选属性 - 表示`q`和`ori_kv`计算中`q`对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限。默认值为-1。 + 表示`q`和`ori_kv`计算中`q`对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限。 INT - layout_q 可选属性 - 表示输入`q`的数据排布格式,支持"BSND"和"TND"。默认值为"BSND"。 + 表示输入`q`的数据排布格式,支持"BSND"和"TND"。 STRING - layout_kv 可选属性 - 表示输入`ori_kv`和`cmp_kv`的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。默认值为"BSND"。 + 表示输入`ori_kv`和`cmp_kv`的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。 STRING - topk_value_mode 可选属性 - 表示TopK索引取值模式。默认值为1。 + 表示TopK索引取值模式。 INT - return_softmax_lse 可选属性 - 表示是否返回`softmax_lse`。默认值为False。 + 表示是否返回`softmax_lse`。 BOOL - @@ -257,10 +257,10 @@ ## 约束说明 - 该接口支持训练、推理场景下使用。 - 该接口支持aclgraph模式。 -- 该接口当前支持三种计算场景:C1A(Sliding Window Attention,SWA)场景仅传入`ori_kv`;C4A(Compressed Sparse Attention,CSA)场景传入`ori_kv`、`cmp_kv`及`cmp_sparse_indices`;C128A(Heavily Compressed Attention,HCA)场景传入`ori_kv`及`cmp_kv`。 +- 该接口当前支持三种计算场景:SWA(Sliding Window Attention)场景仅传入`ori_kv`;CSA(Compressed Sparse Attention)场景传入`ori_kv`、`cmp_kv`及`cmp_sparse_indices`;HCA(Heavily Compressed Attention)场景传入`ori_kv`及`cmp_kv`。 - 通用规格约束如下: - - N2仅支持1,D仅支持512。 - - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时不参与压缩KV计算。C4A场景传4,C128A场景传128。 + - N2仅支持1,D仅支持512。其中,`ori_kv`和`cmp_kv`的D_kv由nope(448)和rope(64)拼接而成。 + - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时不参与压缩KV计算。CSA场景传4,HCA场景传128。 - `ori_mask_mode`仅支持4,`cmp_mask_mode`仅支持3,`ori_win_left`仅支持127,`ori_win_right`仅支持0。 - `cmp_sparse_indices`的TopK长度支持512或1024。 - PageAttention的block_size支持16的倍数,且不超过1024。 @@ -285,7 +285,7 @@ - `ori_kv`和`cmp_kv`的shape分别为[ori\_block\_num, ori\_block\_size, KV\_N, D]和[cmp\_block\_num, cmp\_block\_size, KV\_N, D],其中ori\_block\_num和cmp\_block\_num为PageAttention时block总数,ori\_block\_size和cmp\_block\_size为一个block的token数,ori\_block\_size和cmp\_block\_size取值为16的倍数,最大支持1024,KV_N仅支持1。 - `ori_block_table`和`cmp_block_table`的shape为2维,其中第一维长度为B,第二维长度不小于所有batch中最大的S2和S3对应的block数量,即S2\_max / block\_size和S3\_max / block\_size向上取整。 - `metadata`为算子实际需要使用的分核结果,目前该参数必传,shape大小固定为[1024]。 -- `layout_kv`支持输入"BSND"、"TND"和"PA_BBND",默认值为"BSND",需满足上述`layout_q`和`layout_kv`组合约束。 +- `layout_kv`支持输入"BSND"、"TND"和"PA_BBND",需满足上述`layout_q`和`layout_kv`组合约束。 - 当输入为PA_BBND时,`seqused_ori_kv`和`ori_block_table`必须传入;当输入为BSND时,`seqused_ori_kv`可用于表达每个batch的`ori_kv`有效长度;当输入为TND时,`ori_kv`有效长度由`cu_seqlens_ori_kv`表达。 - 当输入为BSND时,`ori_kv`和`cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, S2, N2,D],cmp_kv的shape为[B, S3, N2,D]。 - 当输入为TND时,`cu_seqlens_ori_kv`必须传入;若存在`cmp_kv`,`cu_seqlens_cmp_kv`也必须传入。 @@ -294,7 +294,7 @@ - 目前暂不支持对`ori_kv`进行稀疏计算,因此设置`ori_sparse_indices`无效。 - 除`ori_topk_length`和`cmp_topk_length`等预留输入可不传或传入空Tensor外,其余已传入Tensor不支持为空。 - `seqused_cmp_kv`为所有`layout_kv`下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由`cmp_kv` shape、`cu_seqlens_cmp_kv`或PA block table相关语义推导。 -- `cmp_residual_kv`为主算子和metadata算子的可选入参;传入后用于按`cmp_len * cmp_ratio + residual`恢复cmp侧mask使用的压缩前KV长度,其中`cmp_len`优先来自显式传入的`seqused_cmp_kv`。 +- `cmp_residual_kv`为主接口和metadata前置接口的可选入参;传入后用于按`cmp_len * cmp_ratio + residual`恢复cmp侧mask使用的压缩前KV长度,其中`cmp_len`优先来自显式传入的`seqused_cmp_kv`。 - `q`、`ori_kv`、`cmp_kv`数据排布格式支持从多种维度解读,B(Batch)表示输入样本批量大小、S(Seq-Length)表示输入样本序列长度、H(Hidden-Size)表示隐藏层的大小、N(Head-Num)表示多头数、D(Head-Dim)表示hidden层最小的单元尺寸,且满足D=H/N、T表示所有Batch输入样本序列长度的累加和。 - Q\_S和S1表示q shape中的S,S2表示ori_kv shape中的S,S3表示cmp_kv shape中的S;Q\_N和N1表示num\_q\_heads,KV\_N和N2表示num\_ori_kv\_heads和num\_cmp_kv\_heads;Q\_T和T1表示q shape中的输入样本序列长度的累加和。 @@ -303,4 +303,4 @@ | 调用方式 | 样例代码 | 说明 | | --------- | ------------------------------------------------------------ | ------------------------------------------------------------ | | aclnn API | [test_aclnnSparseFlashMla](./examples/test_aclnn_sparse_flash_mla.cpp) | 通过[aclnnSparseFlashMla](./docs/aclnnSparseFlashMla.md)调用SparseFlashMla算子 | -| PyTorch API | [sparse_flash_mla](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md) | 通过`cann_ops_transformer.ops.sparse_flash_mla`调用SparseFlashMla算子 | +| PyTorch API | [sparse_flash_mla](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md) | 通过`cann_ops_transformer.sparse_flash_mla`调用SparseFlashMla算子 | diff --git a/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md b/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md index 0ae7b4a366..28094ad2e1 100644 --- a/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md +++ b/attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md @@ -13,7 +13,7 @@ ## 功能说明 -- 接口功能:`SparseFlashMla`算子实现基于共享KV(Key=Value)的稀疏注意力计算,支持C1A(Sliding Window Attention,SWA)、C4A(Compressed Sparse Attention,CSA)、C128A(Heavily Compressed Attention,HCA)三类Attention计算场景。该算子适用于大语言模型训练、推理场景,通过滑动窗口和KV压缩机制大幅降低长序列注意力计算的开销。 +- 接口功能:`SparseFlashMla`算子实现基于共享KV(Key=Value)的稀疏注意力计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。该算子适用于大语言模型训练、推理场景,通过滑动窗口和KV压缩机制大幅降低长序列注意力计算的开销。 - 计算公式: @@ -25,7 +25,7 @@ - 滑动窗口部分(oriKv):对第$i_{S1}$个Query token,其因果对角线位置为$\text{ori\_threshold} = S2_{act} - S1_{act} + i_{S1} + 1$,窗口范围为$[\max(\text{ori\_threshold} - \text{ori\_win\_left} - 1, 0), \text{ori\_threshold} + \text{ori\_win\_right})$。 - - 压缩KV部分(cmpKv):因果边界阈值为$\text{cmp\_threshold} = \lfloor \frac{\text{ori\_threshold}}{\text{cmp\_ratio}} \rfloor$。C128A场景取$[0, \text{cmp\_threshold})$内的连续压缩KV;C4A场景通过TopK索引从压缩KV中按需收集,仅保留$\text{begin\_idx} < \text{cmp\_threshold}$的块。 + - 压缩KV部分(cmpKv):因果边界阈值为$\text{cmp\_threshold} = \lfloor \frac{\text{ori\_threshold}}{\text{cmp\_ratio}} \rfloor$。HCA场景取$[0, \text{cmp\_threshold})$内的连续压缩KV;CSA场景通过TopK索引从压缩KV中按需收集,仅保留$\text{begin\_idx} < \text{cmp\_threshold}$的块。 注意力计算采用Online Softmax(Flash Attention V2),S2方向按512分块循环,sinks作为每行softmax的初始最大值: @@ -166,7 +166,7 @@ aclnnStatus aclnnSparseFlashMla( oriKvOptional(aclTensor*) 输入 原始KV输入张量,Key与Value共享同一份数据。 - C1A/C4A/C128A场景必须传入。 + SWA、CSA、HCA场景必须传入。 BFLOAT16、FLOAT16 ND @@ -175,7 +175,7 @@ aclnnStatus aclnnSparseFlashMla(
  • layoutKv为BSND时:(B, S2, N2, D)
  • layoutKv为TND时:(T2, N2, D)
  • - N2仅支持1,D仅支持512。 + N2仅支持1,D仅支持512,由nope(448)和rope(64)拼接而成。 √ @@ -183,7 +183,7 @@ aclnnStatus aclnnSparseFlashMla( cmpKvOptional(aclTensor*) 输入 压缩KV输入张量,Key与Value共享同一份数据。 - C4A/C128A场景必须传入,C1A场景不传入。 + CSA、HCA场景必须传入,SWA场景不传入。 BFLOAT16、FLOAT16 ND @@ -192,7 +192,7 @@ aclnnStatus aclnnSparseFlashMla(
  • layoutKv为BSND时:(B, S3, N2, D)
  • layoutKv为TND时:(T3, N2, D)
  • - N2仅支持1,D仅支持512。 + N2仅支持1,D仅支持512,由nope(448)和rope(64)拼接而成。 √ @@ -210,7 +210,7 @@ aclnnStatus aclnnSparseFlashMla( cmpSparseIndicesOptional(aclTensor*) 输入 代表离散取cmpKvCache的TopK索引。 - C4A场景必须传入,C1A/C128A场景不传入。 + CSA场景必须传入,SWA、HCA场景不传入。 INT32 ND @@ -236,7 +236,7 @@ aclnnStatus aclnnSparseFlashMla( cmpBlockTableOptional(aclTensor*) 输入 PageAttention中cmpKvCache存储使用的block映射表。 - C4A/C128A场景且layoutKv为PA_BBND时必须传入。 + CSA、HCA场景且layoutKv为PA_BBND时必须传入。 INT32 ND (B, cmp_max_block_num_per_batch) @@ -306,7 +306,7 @@ aclnnStatus aclnnSparseFlashMla( cmpResidualKvOptional(aclTensor*) 输入 压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度。 - 可选输入。传入时shape必须为(B,),第b个batch按cmp_len * cmpRatio + cmpResidualKvOptional[b]恢复压缩前KV长度;在C4A/C128A、cmpRatio不等于1且cmpMaskMode为3场景必传。该参数是主算子和SparseFlashMlaMetadata的可选入参,layoutKvOptional为BSND、TND、PA_BBND时均可使用。 + 可选输入。传入时shape必须为(B,),第b个batch按cmp_len * cmpRatio + cmpResidualKvOptional[b]恢复压缩前KV长度;在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传。该参数是主算子和SparseFlashMlaMetadata的可选入参,layoutKvOptional为BSND、TND、PA_BBND时均可使用。 INT32 ND (B,) @@ -366,7 +366,7 @@ aclnnStatus aclnnSparseFlashMla( cmpRatio(int64_t) 输入 cmpKv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。 - 支持1、4、128;仅传入oriKv时不参与压缩KV计算,C4A场景传4,C128A场景传128。 + 支持1、4、128;仅传入oriKv时不参与压缩KV计算,CSA场景传4,HCA场景传128。 - - - @@ -386,7 +386,7 @@ aclnnStatus aclnnSparseFlashMla( cmpMaskMode(int64_t) 输入 q和cmpKv计算的mask模式。 - 0: No Mask。
    3: RightDownCausal模式。C1A场景下该参数不生效。 + 0: No Mask。
    3: RightDownCausal模式。SWA场景下该参数不生效。 - - - @@ -416,7 +416,7 @@ aclnnStatus aclnnSparseFlashMla( layoutQOptional(char*) 输入 标识输入q的数据排布格式。 - 支持"BSND"和"TND",默认值为"BSND"。 + 支持"BSND"和"TND"。 - - - @@ -426,7 +426,7 @@ aclnnStatus aclnnSparseFlashMla( layoutKvOptional(char*) 输入 标识输入oriKvOptional和cmpKvOptional的数据排布格式。 - 支持"PA_BBND"、"BSND"和"TND",默认值为"BSND"。 + 支持"PA_BBND"、"BSND"和"TND"。 - - - @@ -446,7 +446,7 @@ aclnnStatus aclnnSparseFlashMla( returnSoftmaxLse(bool) 输入 是否返回softmaxLse。 - 支持true或false,默认值为false。 + 支持true或false。 - - - @@ -545,7 +545,7 @@ aclnnStatus aclnnSparseFlashMla( oriMaskMode不为4,或cmpMaskMode不为3。 - C1A场景cmpRatio不为1,或cmpRatio与C4A/C128A场景不匹配。 + SWA场景cmpRatio不为1,或cmpRatio与CSA、HCA场景不匹配。 oriWinLeft不为127,或oriWinRight不为0。 @@ -662,9 +662,9 @@ aclnnStatus aclnnSparseFlashMla( | 场景 | oriKvOptional | cmpKvOptional | cmpSparseIndicesOptional | 说明 | | :--- | :----- | :----- | :----------------- | :--- | - | C1A | 必须传入 | 不传入 | 不传入 | 仅滑动窗口注意力 | - | C4A | 必须传入 | 必须传入 | 必须传入 | 滑动窗口 + TopK稀疏压缩KV | - | C128A | 必须传入 | 必须传入 | 不传入 | 滑动窗口 + 稠密压缩KV | + | SWA | 必须传入 | 不传入 | 不传入 | 仅滑动窗口注意力 | + | CSA | 必须传入 | 必须传入 | 必须传入 | 滑动窗口 + TopK稀疏压缩KV | + | HCA | 必须传入 | 必须传入 | 不传入 | 滑动窗口 + 稠密压缩KV | - Layout约束 @@ -680,6 +680,11 @@ aclnnStatus aclnnSparseFlashMla( 调用示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/编译与运行样例.md)。 ```c++ +/*! + * \file test_aclnn_sparse_flash_mla.cpp + * \brief SparseFlashMla + SparseFlashMlaMetadata 算子调用示例(CSA) + */ + #include #include #include @@ -812,7 +817,7 @@ int main() { // 1. (固定写法)device/stream初始化,参考acl API手册 // 根据自己的实际device填写deviceId - int32_t deviceId = 5; + int32_t deviceId = 0; aclrtContext context = nullptr; aclrtStream stream = nullptr; auto ret = Init(deviceId, &context, &stream); @@ -849,12 +854,13 @@ int main() std::vector cmpBlockTableShape = {B, (cmpKvLen + cmpBlockSize - 1) / cmpBlockSize}; std::vector cuSeqLensQShape = {B + 1}; std::vector seqUsedOriKvShape = {B}; + std::vector seqUsedCmpKvShape = {B}; std::vector cmpResidualKvShape = {B}; std::vector sinksShape = {N1}; std::vector metadataShape = {1024}; std::vector attnOutShape = {T1, N1, D}; std::vector softmaxLseShape = {T1, N1, 1}; - // 对全部 optional 输入调用 Contiguous,optional 输入传 shape 为 {0} 的空 tensor。 + // 对全部 5 个输入调用 Contiguous,optional 输入传 shape 为 {0} 的空 tensor。 std::vector emptyShape = {0}; void* qDeviceAddr = nullptr; @@ -870,8 +876,6 @@ int main() void* seqUsedOriKvDeviceAddr = nullptr; void* seqUsedCmpKvDeviceAddr = nullptr; void* cmpResidualKvDeviceAddr = nullptr; - void* oriTopkLengthDeviceAddr = nullptr; - void* cmpTopkLengthDeviceAddr = nullptr; void* sinksDeviceAddr = nullptr; void* metadataDeviceAddr = nullptr; void* attnOutDeviceAddr = nullptr; @@ -890,8 +894,6 @@ int main() aclTensor* seqUsedOriKv = nullptr; aclTensor* seqUsedCmpKv = nullptr; aclTensor* cmpResidualKv = nullptr; - aclTensor* oriTopkLength = nullptr; - aclTensor* cmpTopkLength = nullptr; aclTensor* sinks = nullptr; aclTensor* metadata = nullptr; aclTensor* attnOut = nullptr; @@ -920,7 +922,8 @@ int main() } std::vector emptyHostData; std::vector seqUsedOriKvHostData(B, static_cast(s2Act)); - std::vector cmpResidualKvHostData(B, 0); + std::vector seqUsedCmpKvHostData(B, static_cast(cmpKvLen)); + std::vector cmpResidualKvHostData(B, static_cast(s2Act % cmpRatio)); std::vector sinksHostData(N1, 1.0f); std::vector metadataHostData(1024, 0); std::vector attnOutHostData = MakeFp16Data(attnOutSize, 0.0f); @@ -960,14 +963,10 @@ int main() CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv); CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); + ret = CreateAclTensor(seqUsedCmpKvHostData, seqUsedCmpKvShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv); CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &oriTopkLengthDeviceAddr, aclDataType::ACL_INT32, &oriTopkLength); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &cmpTopkLengthDeviceAddr, aclDataType::ACL_INT32, &cmpTopkLength); - CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(sinksHostData, sinksShape, &sinksDeviceAddr, aclDataType::ACL_FLOAT, &sinks); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata); @@ -986,8 +985,8 @@ int main() // 3. 调用CANN算子库API,需要修改为具体的Api名称 ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, - seqUsedQ, seqUsedOriKv, seqUsedCmpKv, - cmpResidualKv, oriTopkLength, cmpTopkLength, + seqUsedQ, seqUsedOriKv, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, N1, N2, D, B, S1, S2, cmpKvLen, 0, K, cmpRatio, oriMaskMode, cmpMaskMode, @@ -1021,13 +1020,13 @@ int main() oriBlockTable, cmpBlockTable, cuSeqLensQ, nullptr, nullptr, nullptr, seqUsedOriKv, - nullptr, nullptr, nullptr, nullptr, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, sinks, metadata, softmaxScale, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft, oriWinRight, layoutQ, layoutKv, - 1, 0, 0, + 1, false, attnOut, softmaxLse, &workspaceSize, &executor); @@ -1060,6 +1059,8 @@ int main() aclDestroyTensor(cuSeqLensCmpKv); aclDestroyTensor(seqUsedQ); aclDestroyTensor(seqUsedOriKv); + aclDestroyTensor(seqUsedCmpKv); + aclDestroyTensor(cmpResidualKv); aclDestroyTensor(sinks); aclDestroyTensor(metadata); aclDestroyTensor(attnOut); @@ -1078,6 +1079,12 @@ int main() if (seqUsedOriKvDeviceAddr != nullptr) { aclrtFree(seqUsedOriKvDeviceAddr); } + if (seqUsedCmpKvDeviceAddr != nullptr) { + aclrtFree(seqUsedCmpKvDeviceAddr); + } + if (cmpResidualKvDeviceAddr != nullptr) { + aclrtFree(cmpResidualKvDeviceAddr); + } aclrtFree(sinksDeviceAddr); aclrtFree(metadataDeviceAddr); aclrtFree(attnOutDeviceAddr); diff --git a/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp b/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp index dfa868e667..838cb71db0 100644 --- a/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp +++ b/attention/sparse_flash_mla/examples/test_aclnn_sparse_flash_mla.cpp @@ -145,7 +145,7 @@ int main() { // 1. (固定写法)device/stream初始化,参考acl API手册 // 根据自己的实际device填写deviceId - int32_t deviceId = 5; + int32_t deviceId = 0; aclrtContext context = nullptr; aclrtStream stream = nullptr; auto ret = Init(deviceId, &context, &stream); @@ -182,6 +182,8 @@ int main() std::vector cmpBlockTableShape = {B, (cmpKvLen + cmpBlockSize - 1) / cmpBlockSize}; std::vector cuSeqLensQShape = {B + 1}; std::vector seqUsedOriKvShape = {B}; + std::vector seqUsedCmpKvShape = {B}; + std::vector cmpResidualKvShape = {B}; std::vector sinksShape = {N1}; std::vector metadataShape = {1024}; std::vector attnOutShape = {T1, N1, D}; @@ -200,6 +202,8 @@ int main() void* cuSeqLensCmpKvDeviceAddr = nullptr; void* seqUsedQDeviceAddr = nullptr; void* seqUsedOriKvDeviceAddr = nullptr; + void* seqUsedCmpKvDeviceAddr = nullptr; + void* cmpResidualKvDeviceAddr = nullptr; void* sinksDeviceAddr = nullptr; void* metadataDeviceAddr = nullptr; void* attnOutDeviceAddr = nullptr; @@ -216,6 +220,8 @@ int main() aclTensor* cuSeqLensCmpKv = nullptr; aclTensor* seqUsedQ = nullptr; aclTensor* seqUsedOriKv = nullptr; + aclTensor* seqUsedCmpKv = nullptr; + aclTensor* cmpResidualKv = nullptr; aclTensor* sinks = nullptr; aclTensor* metadata = nullptr; aclTensor* attnOut = nullptr; @@ -244,6 +250,8 @@ int main() } std::vector emptyHostData; std::vector seqUsedOriKvHostData(B, static_cast(s2Act)); + std::vector seqUsedCmpKvHostData(B, static_cast(cmpKvLen)); + std::vector cmpResidualKvHostData(B, static_cast(s2Act % cmpRatio)); std::vector sinksHostData(N1, 1.0f); std::vector metadataHostData(1024, 0); std::vector attnOutHostData = MakeFp16Data(attnOutSize, 0.0f); @@ -283,6 +291,10 @@ int main() CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv); CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(seqUsedCmpKvHostData, seqUsedCmpKvShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); + ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv); + CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(sinksHostData, sinksShape, &sinksDeviceAddr, aclDataType::ACL_FLOAT, &sinks); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata); @@ -302,7 +314,7 @@ int main() ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, seqUsedQ, seqUsedOriKv, - nullptr, nullptr, nullptr, nullptr, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, N1, N2, D, B, S1, S2, cmpKvLen, 0, K, cmpRatio, oriMaskMode, cmpMaskMode, @@ -336,7 +348,7 @@ int main() oriBlockTable, cmpBlockTable, cuSeqLensQ, nullptr, nullptr, nullptr, seqUsedOriKv, - nullptr, nullptr, nullptr, nullptr, + seqUsedCmpKv, cmpResidualKv, nullptr, nullptr, sinks, metadata, softmaxScale, cmpRatio, oriMaskMode, cmpMaskMode, @@ -375,6 +387,8 @@ int main() aclDestroyTensor(cuSeqLensCmpKv); aclDestroyTensor(seqUsedQ); aclDestroyTensor(seqUsedOriKv); + aclDestroyTensor(seqUsedCmpKv); + aclDestroyTensor(cmpResidualKv); aclDestroyTensor(sinks); aclDestroyTensor(metadata); aclDestroyTensor(attnOut); @@ -393,6 +407,12 @@ int main() if (seqUsedOriKvDeviceAddr != nullptr) { aclrtFree(seqUsedOriKvDeviceAddr); } + if (seqUsedCmpKvDeviceAddr != nullptr) { + aclrtFree(seqUsedCmpKvDeviceAddr); + } + if (cmpResidualKvDeviceAddr != nullptr) { + aclrtFree(cmpResidualKvDeviceAddr); + } aclrtFree(sinksDeviceAddr); aclrtFree(metadataDeviceAddr); aclrtFree(attnOutDeviceAddr); 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 d0eb0f3386..2b330cf117 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 @@ -117,91 +117,114 @@ 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 GetStride0FromStrideImpl(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 GetStride0FromStrideImpl(const StrideT &, ...) -{ - return 0ULL; -} - -template -static uint64_t GetStride0FromStride(const StrideT &stride) -{ - return GetStride0FromStrideImpl(stride, 0); -} - -template -static uint64_t GetStride0FromStride(const StrideT *stride) -{ - if (stride == nullptr) { - return 0ULL; - } - return GetStride0FromStride(*stride); -} - -template -static auto TryGetOptionalInputStride0(ContextT *context, uint32_t inputIndex, int) - -> decltype(context->GetOptionalInputStride(inputIndex), uint64_t()) -{ - return GetStride0FromStride(context->GetOptionalInputStride(inputIndex)); -} - -template -static uint64_t TryGetOptionalInputStride0(ContextT *, uint32_t, ...) -{ - return 0ULL; -} - -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 GetStride0FromStride(context->GetInputStride(inputIndex)); -} - -template -static uint64_t TryGetInputViewStride0(ContextT *, uint32_t, ...) -{ - return 0ULL; -} - -std::string SMLALayoutToSerialString(SMLALayout layout) -{ - switch (layout) { + {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"; @@ -250,11 +273,11 @@ ge::graphStatus SMLAInfoParser::CheckRequiredInOutExistence() const OP_LOGE(opName_, "tensor of oriBlockTable is nullptr when layoutKv is PA_BBND"), return ge::GRAPH_FAILED); } - if (perfMode_ == SMLATemplateMode::CFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::HCA_TEMPLATE_MODE) { OP_CHECK_IF(opParamInfo_.cmpKv.tensor == nullptr, OP_LOGE(opName_, "tensor of cmpKv is nullptr"), return ge::GRAPH_FAILED); } - if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { OP_CHECK_IF(opParamInfo_.cmpKv.tensor == nullptr, OP_LOGE(opName_, "tensor of cmpKv is nullptr"), return ge::GRAPH_FAILED); OP_CHECK_IF(opParamInfo_.cmpSparseIndices.tensor == nullptr, OP_LOGE(opName_, "cmpSparseIndices is nullptr"), @@ -430,11 +453,17 @@ uint64_t SMLAInfoParser::GetOptionalInputStride0(uint32_t inputIndex) const return stride0; } - const gert::Shape &storageShape = inputTensor->GetStorageShape(); - return GetStorageShapeStride0(storageShape); -} -ge::graphStatus SMLAInfoParser::GetInOutDataType() -{ + 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) { @@ -450,16 +479,16 @@ ge::graphStatus SMLAInfoParser::GetSMLATemplateMode(SMLATilingInfo &smlaInfo) { if (opParamInfo_.oriKv.desc != nullptr) { if (opParamInfo_.cmpKv.desc != nullptr && opParamInfo_.cmpSparseIndices.tensor != nullptr) { - perfMode_ = SMLATemplateMode::SCFA_TEMPLATE_MODE; + perfMode_ = SMLATemplateMode::CSA_TEMPLATE_MODE; } else if (opParamInfo_.cmpKv.desc != nullptr && opParamInfo_.cmpSparseIndices.tensor == nullptr) { - perfMode_ = SMLATemplateMode::CFA_TEMPLATE_MODE; + perfMode_ = SMLATemplateMode::HCA_TEMPLATE_MODE; } else if (opParamInfo_.cmpKv.desc == nullptr && opParamInfo_.cmpSparseIndices.tensor == nullptr) { perfMode_ = SMLATemplateMode::SWA_TEMPLATE_MODE; } else { OP_LOGE(opName_, "When cmpSparseIndices is not nullptr, cmpKv cannot be nullptr."); return ge::GRAPH_FAILED; } - if (perfMode_ == SMLATemplateMode::CFA_TEMPLATE_MODE || perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::HCA_TEMPLATE_MODE || perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { if (kvLayout_ == SMLALayout::TND && opParamInfo_.cuSeqLensCmpKv.tensor == nullptr) { OP_LOGE(opName_, "the layout_kv is %s, seqlens_cmp_kv must be provided.", SMLALayoutToSerialString(kvLayout_).c_str()); @@ -565,7 +594,7 @@ void SMLAInfoParser::SetSMLAShape() if (opParamInfo_.cmpKv.tensor != nullptr) { cmpKvShape_ = opParamInfo_.cmpKv.tensor->GetStorageShape(); } - if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { if (opParamInfo_.cmpSparseIndices.tensor != nullptr) { cmpSparseIndicesShape_ = opParamInfo_.cmpSparseIndices.tensor->GetStorageShape(); } else { @@ -581,35 +610,45 @@ ge::graphStatus SMLAInfoParser::CheckContiguous() const { bool oriKeyNonContiguous = false; bool cmpKeyNonContiguous = false; - size_t checkStartIdx = (kvLayout_ == SMLALayout::PA_BBND) ? 1 : 0; - if (opParamInfo_.oriKv.tensor != nullptr && !oriKeyStridesVec_.empty()) { - std::vector oriExpectedStrides; - if (kvLayout_ == SMLALayout::BSND || kvLayout_ == SMLALayout::PA_BBND) { - oriExpectedStrides = {oriKvShape_.GetDim(1) * oriKvShape_.GetDim(2) * oriKvShape_.GetDim(3), - oriKvShape_.GetDim(2) * oriKvShape_.GetDim(3), oriKvShape_.GetDim(3), 1}; - } else if (kvLayout_ == SMLALayout::TND) { - oriExpectedStrides = {oriKvShape_.GetDim(1) * oriKvShape_.GetDim(2), oriKvShape_.GetDim(2), 1}; - } + size_t checkStartIdx = (kvLayout_ == SMLALayout::PA_BBND) ? 1 : 0; + if (opParamInfo_.oriKv.tensor != nullptr && !oriKeyStridesVec_.empty()) { + std::vector oriExpectedStrides; + if (kvLayout_ == SMLALayout::BSND || kvLayout_ == SMLALayout::PA_BBND) { + uint64_t dim1 = static_cast(oriKvShape_.GetDim(1)); + uint64_t dim2 = static_cast(oriKvShape_.GetDim(2)); + uint64_t dim3 = static_cast(oriKvShape_.GetDim(3)); + oriExpectedStrides = {dim1 * dim2 * dim3, dim2 * dim3, dim3, 1}; + } else if (kvLayout_ == SMLALayout::TND) { + uint64_t dim1 = static_cast(oriKvShape_.GetDim(1)); + uint64_t dim2 = static_cast(oriKvShape_.GetDim(2)); + oriExpectedStrides = {dim1 * dim2, dim2, 1}; + } OP_CHECK_IF(oriKeyStridesVec_.size() != oriExpectedStrides.size(), OP_LOGE(opName_, "oriKey strideVec size[%zu] not match kvLayout expect len[%zu].", - oriKeyStridesVec_.size(), oriExpectedStrides.size()), - return ge::GRAPH_FAILED); - oriKeyNonContiguous = oriKeyStridesVec_[checkStartIdx] != oriExpectedStrides[checkStartIdx]; - } - if (opParamInfo_.cmpKv.tensor != nullptr && !cmpKeyStridesVec_.empty()) { - std::vector cmpExpectedStrides; - if (kvLayout_ == SMLALayout::BSND || kvLayout_ == SMLALayout::PA_BBND) { - cmpExpectedStrides = {cmpKvShape_.GetDim(1) * cmpKvShape_.GetDim(2) * cmpKvShape_.GetDim(3), - cmpKvShape_.GetDim(2) * cmpKvShape_.GetDim(3), cmpKvShape_.GetDim(3), 1}; - } else if (kvLayout_ == SMLALayout::TND) { - cmpExpectedStrides = {cmpKvShape_.GetDim(1) * cmpKvShape_.GetDim(2), cmpKvShape_.GetDim(2), 1}; - } + oriKeyStridesVec_.size(), oriExpectedStrides.size()), + return ge::GRAPH_FAILED); + oriKeyNonContiguous = static_cast(oriKeyStridesVec_[checkStartIdx]) != + oriExpectedStrides[checkStartIdx]; + } + if (opParamInfo_.cmpKv.tensor != nullptr && !cmpKeyStridesVec_.empty()) { + std::vector cmpExpectedStrides; + if (kvLayout_ == SMLALayout::BSND || kvLayout_ == SMLALayout::PA_BBND) { + uint64_t dim1 = static_cast(cmpKvShape_.GetDim(1)); + uint64_t dim2 = static_cast(cmpKvShape_.GetDim(2)); + uint64_t dim3 = static_cast(cmpKvShape_.GetDim(3)); + cmpExpectedStrides = {dim1 * dim2 * dim3, dim2 * dim3, dim3, 1}; + } else if (kvLayout_ == SMLALayout::TND) { + uint64_t dim1 = static_cast(cmpKvShape_.GetDim(1)); + uint64_t dim2 = static_cast(cmpKvShape_.GetDim(2)); + cmpExpectedStrides = {dim1 * dim2, dim2, 1}; + } OP_CHECK_IF(cmpKeyStridesVec_.size() != cmpExpectedStrides.size(), OP_LOGE(opName_, "cmpKey strideVec size[%zu] not match kvLayout expect len[%zu].", - cmpKeyStridesVec_.size(), cmpExpectedStrides.size()), - return ge::GRAPH_FAILED); - cmpKeyNonContiguous = cmpKeyStridesVec_[checkStartIdx] != cmpExpectedStrides[checkStartIdx]; - } + cmpKeyStridesVec_.size(), cmpExpectedStrides.size()), + return ge::GRAPH_FAILED); + cmpKeyNonContiguous = static_cast(cmpKeyStridesVec_[checkStartIdx]) != + cmpExpectedStrides[checkStartIdx]; + } OP_CHECK_IF(oriKeyNonContiguous, OP_LOGE(opName_, "oriKey only support non-continuous keying on the 0-axis."), @@ -634,7 +673,7 @@ ge::graphStatus SMLAInfoParser::GetN2Size() } if (opParamInfo_.cmpKv.tensor != nullptr) { uint32_t cmpKvN2Size_ = GetAxisNum(cmpKvShape_, SMLAAxis::N, kvLayout_); - if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::CSA_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."), @@ -708,7 +747,7 @@ ge::graphStatus SMLAInfoParser::GetS1Size() } else { // BSND s1Size_ = GetAxisNum(qShape_, SMLAAxis::S, qLayout_); } - if (perfMode_ == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (perfMode_ == SMLATemplateMode::CSA_TEMPLATE_MODE) { if (cmpSparseIndicesLayout_ == SMLALayout::TND) { uint32_t cmpSparseIndicesT = GetAxisNum(cmpSparseIndicesShape_, SMLAAxis::T, cmpSparseIndicesLayout_); OP_CHECK_IF(cmpSparseIndicesT != s1Size_, @@ -734,11 +773,11 @@ ge::graphStatus SMLAInfoParser::GetMaxBlockNumPerBatch() 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(%lld) should be non-negative number.", - ORI_BLOCK_TABLE_NAME.c_str(), opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1)); - 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)); + return ge::GRAPH_FAILED; + } oriMaxBlockNumPerBatch_ = opParamInfo_.oriBlockTable.tensor->GetStorageShape().GetDim(1); if (opParamInfo_.cmpBlockTable.tensor != nullptr) { @@ -747,18 +786,18 @@ ge::graphStatus SMLAInfoParser::GetMaxBlockNumPerBatch() 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 (qLayout_ == SMLALayout::TND || qLayout_ == SMLALayout::BSND) { - if (opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(0) != bSize_) { - OP_LOGE(opName_, "cmp_block_table's first dimension(%lld) 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(%lld) should be greater than 0", - CMP_BLOCK_TABLE_NAME.c_str(), opParamInfo_.cmpBlockTable.tensor->GetStorageShape().GetDim(1)); - return ge::GRAPH_FAILED; - } + 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_); + 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; @@ -912,12 +951,13 @@ ge::graphStatus SMLAInfoParser::GetActualseqInfo() return ge::GRAPH_FAILED; } } else if (kvLayout_ == SMLALayout::TND) { - } else if (kvLayout_ == SMLALayout::BSND) { - actualLenDimsKV_ = actualLenDimsOriKV_; - } else { - OP_LOGE(opName_, "oriKV and cmpKv only support PA_BBND, TND and BSND layout, but got %d.", kvLayout_); - return ge::GRAPH_FAILED; - } + } 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()); + return ge::GRAPH_FAILED; + } return ge::GRAPH_SUCCESS; } @@ -1218,8 +1258,8 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriKv() const ge::graphStatus SMLATilingCheck::CheckSingleParaCmpKv() const { - if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE || \ - smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE) { + if (smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE || \ + smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE) { const std::vector cmpKvDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; if ( ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpKv.desc, CMP_KV_NAME) || @@ -1300,11 +1340,11 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriSparseIndices() const 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(SCFA), cmp_sparse_indices cannot be empty tensor."), - return ge::GRAPH_FAILED); + if (smlaInfo_.perfMode == optiling::SMLATemplateMode::CSA_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); const std::vector cmpSparseIndicesDimNumList = {DIM_NUM_THREE, DIM_NUM_FOUR}; if ( ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpSparseIndices.desc, CMP_SPARSE_INDICES) || @@ -1344,8 +1384,8 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaCmpBlockTable() const if (kvLayout_ != SMLALayout::PA_BBND) { return ge::GRAPH_SUCCESS; } - if (smlaInfo_.perfMode == optiling::SMLATemplateMode::SCFA_TEMPLATE_MODE || - smlaInfo_.perfMode == optiling::SMLATemplateMode::CFA_TEMPLATE_MODE) { + if (smlaInfo_.perfMode == optiling::SMLATemplateMode::CSA_TEMPLATE_MODE || + smlaInfo_.perfMode == optiling::SMLATemplateMode::HCA_TEMPLATE_MODE) { const std::vector cmpBlockTableDimNumList = {DIM_NUM_TWO}; if ( ge::GRAPH_SUCCESS != CheckDtypeSupport(opParamInfo_.cmpBlockTable.desc, CMP_BLOCK_TABLE_NAME) || @@ -1365,14 +1405,14 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaCmpBlockTable() const 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); - if (opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum() != DIM_NUM_ONE) { - OP_LOGE(opName_, "the dim num of %s is %u, it should be %u.", SINKS_NAME.c_str(), - opParamInfo_.sinks.tensor->GetStorageShape().GetDimNum(), DIM_NUM_ONE); - return ge::GRAPH_FAILED; - } + OP_CHECK_IF(opParamInfo_.sinks.tensor->GetStorageShape().GetShapeSize() == 0, + 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); + 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_); @@ -1399,34 +1439,40 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaMetadata() const return ge::GRAPH_SUCCESS; } -ge::graphStatus SMLATilingCheck::CheckSingleParaCmpRatio() const -{ - if (IsA5Arch(npuArch_)) { - if (opParamInfo_.cmpKv.tensor != nullptr) { +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; - if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { - expectedCmpRatio = 4; - } else if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE) { - expectedCmpRatio = 128; - } - OP_CHECK_IF(cmpRatio_ != expectedCmpRatio, - OP_LOGE(opName_, "cmpRatio should be %u in current template mode on %s, but got %ld.", - expectedCmpRatio, A2_A3_PLATFORM_LOG.c_str(), cmpRatio_), - return ge::GRAPH_FAILED); - } - return ge::GRAPH_SUCCESS; -} - -ge::graphStatus SMLATilingCheck::CheckSingleParaOriMaskMode() const -{ - return ge::GRAPH_SUCCESS; -} + 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::CSA_TEMPLATE_MODE) { + expectedCmpRatio = 4; + modeName = "CSA"; + modeReason = "when cmp_sparse_indices is provided"; + } else if (smlaInfo_.perfMode == SMLATemplateMode::HCA_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::CheckSingleParaCmpMaskMode() const { @@ -1445,8 +1491,8 @@ ge::graphStatus SMLATilingCheck::CheckSingleParaOriWinRight() const ge::graphStatus SMLATilingCheck::CheckSingleParaCmpResidualKv() const { - bool isCmpTemplate = smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || - smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE; + bool isCmpTemplate = smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::CSA_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"), @@ -1629,42 +1675,42 @@ ge::graphStatus SMLATilingCheck::CheckFeatureShape() const 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 %d", - 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 %d", - 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 %d", 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 %lld", - 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 %lld", - 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 %d", - A2_A3_PLATFORM_LOG.c_str(), *opParamInfo_.oriMaskMode), - return ge::GRAPH_FAILED); - OP_CHECK_IF(*opParamInfo_.cmpMaskMode != 3, - OP_LOGE(opName_, "cmpMaskMode should be 3 on %s, but got %d", - 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 %lld", - 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 %lld", - A2_A3_PLATFORM_LOG.c_str(), oriWinRight_), - return ge::GRAPH_FAILED); + 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); + 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; } @@ -1714,11 +1760,11 @@ void SMLATilingCheck::SetSMLAShapeCompare() queryShapeCmp_ = opParamInfo_.q.shape->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) { + if (smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE) { cmpKvShapeCmp_= opParamInfo_.cmpKv.tensor->GetShape().GetStorageShape(); } - if (smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE) { cmpKvSparseIndicesCmp_ = opParamInfo_.cmpSparseIndices.tensor->GetShape().GetStorageShape(); } } @@ -1737,8 +1783,8 @@ ge::graphStatus SMLATilingCheck::CheckDTypeConsistency(const ge::DataType &actua ge::graphStatus SMLATilingCheck::CheckOriAndCmpKv() const { - if (smlaInfo_.perfMode == SMLATemplateMode::CFA_TEMPLATE_MODE || - smlaInfo_.perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (smlaInfo_.perfMode == SMLATemplateMode::HCA_TEMPLATE_MODE || + smlaInfo_.perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE) { if (ge::GRAPH_SUCCESS != CheckDTypeConsistency(cmpKvType_, oriKvType_, CMP_KV_NAME)) { return ge::GRAPH_FAILED; @@ -1809,7 +1855,7 @@ void SparseFlashMlaTiling::SplitBalanced(SMLATilingInfo *tilingInfo) { sInnerSizeAlign_ = Align(sInnerSize_, BYTE_BLOCK); if (tilingInfo->npuArch == NpuArch::DAV_2201) { - mBaseSize_ = tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE ? + mBaseSize_ = tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE ? tilingInfo->gSize : (256 / tilingInfo->gSize) * tilingInfo->gSize; } headDimAlign_ = Align(tilingInfo->qHeadDim, BYTE_BLOCK); @@ -1841,7 +1887,7 @@ ge::graphStatus SparseFlashMlaTiling::DoOpTiling(SMLATilingInfo *tilingInfo) uint32_t workspaceSize = ascendcPlatform.GetLibApiWorkSpaceSize(); if (tilingInfo->npuArch == NpuArch::DAV_3510) { - if (tilingInfo->gSize > 64 || tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (tilingInfo->gSize > 64 || tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE) { constexpr uint32_t TRIPLE_BUFFER_NUM = 3; constexpr uint32_t S2_BASE_SIZE = 128; constexpr uint32_t D_SIZE = 512; @@ -1861,7 +1907,7 @@ ge::graphStatus SparseFlashMlaTiling::DoOpTiling(SMLATilingInfo *tilingInfo) workspaceSize += PRELOAD_NUM * mmResUbSize_ * VEC1_RES_ELEM_SIZE * aicNum; workspaceSize += PRELOAD_NUM * bmm2ResUbSize_ * MM2_RES_ELEM_SIZE * aicNum; workspaceSize += PRELOAD_NUM * bmm2ResUbSize_ * VEC2_RES_ELEM_SIZE * aicNum; - if (tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE) { + if (tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE) { constexpr uint32_t MERGE_CACHE_GM_BUF_NUM = 3; workspaceSize += MERGE_CACHE_GM_BUF_NUM * 512 * 512 * 2 * aicNum; } @@ -1920,7 +1966,7 @@ ge::graphStatus SparseFlashMlaTiling::DoOpTiling(SMLATilingInfo *tilingInfo) uint32_t splitG = 0U; uint32_t headRatioOne = static_cast( tilingInfo->npuArch == NpuArch::DAV_2201 && - tilingInfo->perfMode == SMLATemplateMode::SCFA_TEMPLATE_MODE && + tilingInfo->perfMode == SMLATemplateMode::CSA_TEMPLATE_MODE && tilingInfo->gSize == 1U); if (tilingInfo->npuArch == NpuArch::DAV_3510) { splitG = static_cast(tilingInfo->gSize > 64); diff --git a/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h b/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h index e6e9b06568..4eb0ed250c 100644 --- a/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h +++ b/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_cube.h @@ -416,8 +416,8 @@ __aicore__ inline void SWACubeBlock::ComputeMm1(const RunInfo &info, cons uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? static_cast(constInfo.kvSeqSize) * seqStride : constInfo.oriKvStride0; - uint64_t curS2 = (uint64_t)info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint; - uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + \ + uint64_t curS2 = info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara); } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { @@ -449,8 +449,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm1(const RunInfo &info, cons if (oriSizeCur > 0) { uint32_t copyFinishRowCnt = 0; - uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + \ - info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t curS2Offset = info.s2StartPoint + nL1 * N_SPLIT_SIZE; while (copyFinishRowCnt < oriSizeCur) { // 由于ori_left的存在, 即使第一块搬运也可能并非是pa_block的零点位 copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; @@ -531,8 +530,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm1(const RunInfo &info, cons nd2nzPara.nValue = oriSizeCur; uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? static_cast(constInfo.kvSeqSize) * seqStride : constInfo.oriKvStride0; - uint64_t curS2 = (uint64_t)info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + \ - nL1 * N_SPLIT_SIZE; + uint64_t curS2 = info.s2StartPoint + nL1 * N_SPLIT_SIZE; uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + \ (uint64_t)info.n2Idx * headStride + kL1 * D_SPLIT_SIZE; DataCopy(bL1Tensor, oriKvGm[offset], nd2nzPara); @@ -562,8 +560,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm1(const RunInfo &info, cons if (oriSizeCur > 0) { nd2nzPara.nValue = oriSizeCur; - uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + \ - info.s2StartPoint + nL1 * N_SPLIT_SIZE; + uint64_t curS2Offset = info.s2StartPoint + nL1 * N_SPLIT_SIZE; DataCopy(bL1Tensor, oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + kL1 * D_SPLIT_SIZE], nd2nzPara); @@ -849,9 +846,8 @@ __aicore__ inline void SWACubeBlock::ComputeMm2(const RunInfo &info, cons uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? static_cast(constInfo.kvSeqSize) * seqStride : constInfo.oriKvStride0; - uint64_t curS2 = (uint64_t)info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + \ - kL1 * K_L0_SPLIT_SIZE; - uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + \ + uint64_t curS2 = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; DataCopy(subvTensor, oriKvGm[offset], nd2nzPara); } else if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::TND) { @@ -881,8 +877,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm2(const RunInfo &info, cons if constexpr (KV_LAYOUT_T == SMLA_LAYOUT::PA_BBND) { if (oriSizeCur > 0) { copyFinishRowCnt = 0; - uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + \ - info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + uint64_t curS2Offset = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; while (copyFinishRowCnt < oriSizeCur) { copyRowCnt = constInfo.paOriBlockSize - curS2Offset % constInfo.paOriBlockSize; if (copyFinishRowCnt + copyRowCnt > oriSizeCur) { @@ -967,8 +962,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm2(const RunInfo &info, cons nd2nzPara.nValue = oriSizeCur; uint64_t batchStride = (constInfo.oriKvStride0 == 0) ? static_cast(constInfo.kvSeqSize) * seqStride : constInfo.oriKvStride0; - uint64_t curS2 = (uint64_t)info.s2Idx * constInfo.s2BaseSize + info.s2StartPoint + - kL1 * K_L0_SPLIT_SIZE; + uint64_t curS2 = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; uint64_t offset = (uint64_t)info.bIdx * batchStride + (uint64_t)curS2 * seqStride + (uint64_t)info.n2Idx * headStride + nL1 * N_SPLIT_SIZE; DataCopy(subvTensor, oriKvGm[offset], nd2nzPara); @@ -1001,8 +995,7 @@ __aicore__ inline void SWACubeBlock::ComputeMm2(const RunInfo &info, cons if (oriSizeCur > 0) { nd2nzPara.nValue = oriSizeCur; - uint64_t curS2Offset = static_cast(info.s2Idx) * constInfo.s2BaseSize + \ - info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; + uint64_t curS2Offset = info.s2StartPoint + kL1 * K_L0_SPLIT_SIZE; DataCopy(bL1Tensor[(kL1 - kOffset) * K_L0_SPLIT_SIZE * N_SPLIT_SIZE], oriKvGm[info.tensorBOffset + curS2Offset * constInfo.headDim + nL1 * N_SPLIT_SIZE], nd2nzPara); diff --git a/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h b/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h index 74e5519eac..0a1a842b45 100644 --- a/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h +++ b/attention/sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_swa_block_vector.h @@ -390,12 +390,16 @@ __aicore__ inline void SWAVectorBlock::ElewiseCompute(const RunInfo &info int32_t noMaskCmpSize = info.cmpMaskRight + s1StartIdx + 1; int32_t actNoMaskCmpSize = 0; int32_t right = 0; + // relativeS2Idx is the cmp tile index relative to the current batch. + // s2Idx may include the global task index. + int64_t cmpTileStart = static_cast(info.relativeS2Idx) * constInfo.s2BaseSize + + static_cast(info.s2StartPoint); for (uint32_t i = s1StartIdx; i <= s1EndIdx; i++) { actNoMaskCmpSize = noMaskCmpSize / constInfo.cmpRatio; if (actNoMaskCmpSize <= 0) { right = 0; } else { - right = actNoMaskCmpSize - info.s2StartPoint; + right = actNoMaskCmpSize - cmpTileStart; } dealTempSize = constInfo.gSize - gStartIdx; if (i == s1EndIdx) { diff --git a/attention/sparse_flash_mla_metadata/CMakeLists.txt b/attention/sparse_flash_mla_metadata/CMakeLists.txt index fb93e0db32..e99a153f31 100644 --- a/attention/sparse_flash_mla_metadata/CMakeLists.txt +++ b/attention/sparse_flash_mla_metadata/CMakeLists.txt @@ -9,9 +9,11 @@ # ----------------------------------------------------------------------------------------------------------- file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) -list(REMOVE_ITEM CURRENT_DIRS tests) +if(NOT ENABLE_TEST AND NOT BENCHMARK) + list(REMOVE_ITEM CURRENT_DIRS tests) +endif() foreach(SUB_DIR ${CURRENT_DIRS}) if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") add_subdirectory(${SUB_DIR}) endif() -endforeach() \ No newline at end of file +endforeach() diff --git a/attention/sparse_flash_mla_metadata/README.md b/attention/sparse_flash_mla_metadata/README.md index e01b2e969e..fe1eff6f78 100644 --- a/attention/sparse_flash_mla_metadata/README.md +++ b/attention/sparse_flash_mla_metadata/README.md @@ -12,6 +12,7 @@ ## 功能说明 - 算子功能:`SparseFlashMlaMetadata`算子完成`SparseFlashMla`算子的tiling计算,包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及Q和K的分块的索引,供后续`SparseFlashMla`算子使用。 +- 场景简称:SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)。 @@ -122,105 +123,105 @@ batch_size 可选属性 - 表示输入样本批量大小。默认值为0。 + 表示输入样本批量大小;传入0时表示由接口推导。 INT - max_seqlen_q 可选属性 - 表示所有batch中`q`的最大有效token数。默认值为0。 + 表示所有batch中`q`的最大有效token数;传入0时表示由接口推导。 INT - max_seqlen_ori_kv 可选属性 - 表示所有batch中`ori_kv`的最大有效token数。默认值为0。 + 表示所有batch中`ori_kv`的最大有效token数;传入0时表示由接口推导。 INT - max_seqlen_cmp_kv 可选属性 - 表示所有batch中`cmp_kv`的最大有效token数。默认值为0。 + 表示所有batch中`cmp_kv`的最大有效token数;传入0时表示由接口推导。 INT - ori_topk 可选属性 - 预留参数,表示从`ori_kv`中筛选出的关键稀疏token个数。默认值为0,当前仅支持0。 + 预留参数,表示从`ori_kv`中筛选出的关键稀疏token个数;当前仅支持0。 INT - cmp_topk 可选属性 - 表示从`cmp_kv`中筛选出的关键稀疏token个数,支持0、512、1024。默认值为0。 + 表示从`cmp_kv`中筛选出的关键稀疏token个数,支持0、512、1024。 INT - cmp_ratio 可选属性 - 表示`cmp_kv`相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入`ori_kv`时不参与压缩KV计算。支持1、4、128。默认值为1。 + 表示`cmp_kv`相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入`ori_kv`时不参与压缩KV计算。支持1、4、128。 INT - ori_mask_mode 可选属性 - 表示`q`和`ori_kv`计算的mask模式。
    0: No Mask。
    3: RightDownCausal模式。
    4: Band模式。默认值为0。 + 表示`q`和`ori_kv`计算的mask模式。
    0: No Mask。
    3: RightDownCausal模式。
    4: Band模式。 INT - cmp_mask_mode 可选属性 - 表示`q`和`cmp_kv`计算的mask模式。
    0: No Mask。
    3: RightDownCausal模式。默认值为0。 + 表示`q`和`cmp_kv`计算的mask模式。
    0: No Mask。
    3: RightDownCausal模式。 INT - ori_win_left 可选属性 - 表示`q`和`ori_kv`计算中`q`对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限。默认值为-1。 + 表示`q`和`ori_kv`计算中`q`对过去token计算的数量,支持-1或非负数,其中-1表示窗口不受限。 INT - ori_win_right 可选属性 - 表示`q`和`ori_kv`计算中`q`对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限。默认值为-1。 + 表示`q`和`ori_kv`计算中`q`对未来token计算的数量,支持-1或非负数,其中-1表示窗口不受限。 INT - layout_q 可选属性 - 表示输入`q`的数据排布格式,支持"BSND"和"TND"。默认值为"BSND"。 + 表示输入`q`的数据排布格式,支持"BSND"和"TND"。 STRING - layout_kv 可选属性 - 表示输入`ori_kv`和`cmp_kv`的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。默认值为"BSND"。 + 表示输入`ori_kv`和`cmp_kv`的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。 STRING - has_ori_kv 可选属性 - 表示`SparseFlashMla`主算子是否传入`ori_kv`。默认值为true。 + 表示`SparseFlashMla`主算子是否传入`ori_kv`。 BOOL - has_cmp_kv 可选属性 - 表示`SparseFlashMla`主算子是否传入`cmp_kv`。默认值为true。 + 表示`SparseFlashMla`主算子是否传入`cmp_kv`。 BOOL - @@ -240,9 +241,9 @@ - 该接口支持aclgraph模式。 - 通用规格约束如下: - `num_heads_kv`仅支持1,`head_dim`仅支持512。 - - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时不参与压缩KV计算。C4A场景传4,C128A场景传128。 + - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时不参与压缩KV计算。CSA场景传4,HCA场景传128。 - `ori_mask_mode`仅支持4,`cmp_mask_mode`仅支持3,`ori_win_left`仅支持127,`ori_win_right`仅支持0。 - - `cmp_topk`在C4A场景支持512或1024,C1A/C128A场景传0。 + - `cmp_topk`在CSA场景支持512或1024,SWA、HCA场景传0。 - `layout_q`和`layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致。 - `ori_topk_length`和`cmp_topk_length`为预留输入,全平台均不支持传入非空Tensor。 - 产品型号约束如下: @@ -254,4 +255,4 @@ | 调用方式 | 样例代码 | 说明 | | --------- | ------------------------------------------------------------ | ------------------------------------------------------------ | | aclnn API | [test_aclnnSparseFlashMlaMetadata](./examples/test_aclnn_sparse_flash_mla_metadata.cpp) | 通过[aclnnSparseFlashMlaMetadata](./docs/aclnnSparseFlashMlaMetadata.md)调用SparseFlashMlaMetadata算子 | -| PyTorch API | [sparse_flash_mla_metadata](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md) | 通过`cann_ops_transformer.ops.sparse_flash_mla_metadata`生成SparseFlashMla主算子使用的metadata | +| PyTorch API | [sparse_flash_mla_metadata](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md) | 通过`cann_ops_transformer.sparse_flash_mla_metadata`生成SparseFlashMla主算子使用的metadata | diff --git a/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md b/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md index 8c1e729391..65e5a015b0 100644 --- a/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md +++ b/attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md @@ -16,6 +16,7 @@ - 接口功能:该算子为AICPU算子,`SparseFlashMlaMetadata`算子为`SparseFlashMla`算子的前序算子,负责根据输入的序列长度信息和注意力配置参数,生成负载均衡的分核元数据(metadata)。该元数据包含每个AICore上FlashAttention计算任务的Batch、Head、Query分块和KV分块的索引,以及每个VectorCore上FlashDecode归约任务的索引信息。 **该算子不建议单独使用,建议与aclnnSparseFlashMla算子配合使用,形成完整的工作流。** +- 场景简称:SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)。 - 计算公式: 该算子为AICPU调度算子,不涉及数值计算。核心流程为:解析各Batch的Q/KV序列长度 → 根据mask模式计算每个S1G块的有效S2范围 → 基于开销模型进行负载均衡分核 → 输出分核元数据。 @@ -196,7 +197,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( cmpResidualKvOptional(aclTensor*) 输入 压缩KV余数,用于按cmp_len * cmpRatio + residual恢复cmp侧mask使用的压缩前KV长度。 - 在C4A/C128A、cmpRatio不等于1且cmpMaskMode为3场景必传,layoutKvOptional为BSND、TND、PA_BBND时均可使用。 + 在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,layoutKvOptional为BSND、TND、PA_BBND时均可使用。 INT32 ND (B,) @@ -256,7 +257,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( batchSize(int64_t) 输入 输入样本批量大小。 - 默认值为0。layoutQ为TND时从cuSeqLensQ推断,无需手动指定。 + 传入0时表示从cuSeqLensQ推断;layoutQ为TND时无需手动指定。 - - - @@ -266,7 +267,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( maxSeqlenQ(int64_t) 输入 所有Batch中q的最大有效token数。 - 默认值为0。 + 传入0时表示由接口推导。 - - - @@ -276,7 +277,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( maxSeqlenOriKv(int64_t) 输入 所有Batch中oriKv的最大有效token数。 - 默认值为0。 + 传入0时表示由接口推导。 - - - @@ -286,7 +287,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( maxSeqlenCmpKv(int64_t) 输入 所有Batch中cmpKv的最大有效token数。 - 默认值为0。 + 传入0时表示由接口推导。 - - - @@ -296,7 +297,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( oriTopk(int64_t) 输入 从oriKv中筛选的稀疏token个数。 - 当前暂不支持,默认值为0,当前仅支持0。 + 当前暂不支持传入非0值,仅支持0。 - - - @@ -306,7 +307,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( cmpTopk(int64_t) 输入 从cmpKv中筛选的稀疏token个数。 - C4A场景下仅支持512或1024,C1A/C128A场景下为0。 + CSA场景下仅支持512或1024,SWA、HCA场景下为0。 - - - @@ -316,7 +317,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( cmpRatio(int64_t) 输入 cmpKv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。 - 支持1、4、128;仅传入oriKv时不参与压缩KV计算,C4A场景传4,C128A场景传128。 + 支持1、4、128;仅传入oriKv时不参与压缩KV计算,CSA场景传4,HCA场景传128。 - - - @@ -326,7 +327,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( oriMaskMode(int64_t) 输入 q和oriKv计算的mask模式。 - 0: No Mask。
    3: RightDownCausal模式。
    4: Band模式。默认值为4。 + 0: No Mask。
    3: RightDownCausal模式。
    4: Band模式。 - - - @@ -336,7 +337,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( cmpMaskMode(int64_t) 输入 q和cmpKv计算的mask模式。 - 0: No Mask。
    3: RightDownCausal模式。默认值为3。 + 0: No Mask。
    3: RightDownCausal模式。 - - - @@ -346,7 +347,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( oriWinLeft(int64_t) 输入 滑动窗口向左扩展的token数。 - 支持-1或非负数,其中-1表示窗口不受限。默认值为127。 + 支持-1或非负数,其中-1表示窗口不受限。 - - - @@ -356,7 +357,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( oriWinRight(int64_t) 输入 滑动窗口向右扩展的token数。 - 支持-1或非负数,其中-1表示窗口不受限。默认值为0。 + 支持-1或非负数,其中-1表示窗口不受限。 - - - @@ -366,7 +367,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( layoutQOptional(char*) 输入 标识输入q的数据排布格式。 - 支持"BSND"和"TND",默认值为"BSND"。 + 支持"BSND"和"TND"。 - - - @@ -376,7 +377,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( layoutKvOptional(char*) 输入 标识输入KV的数据排布格式。 - 支持"PA_BBND"、"BSND"和"TND",默认值为"BSND"。 + 支持"PA_BBND"、"BSND"和"TND"。 - - - @@ -386,7 +387,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( hasOriKv(bool) 输入 是否传入oriKv。 - 默认值为true。 + 根据是否传入oriKv设置。 - - - @@ -396,7 +397,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( hasCmpKv(bool) 输入 是否传入cmpKv。 - C1A场景为false,C4A/C128A场景为true。默认值为true。 + SWA场景为false,CSA、HCA场景为true。根据是否传入cmpKv设置。 - - - @@ -487,7 +488,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( oriWinLeft不为127,或oriWinRight不为0。 - C1A场景cmpRatio不为1,或cmpRatio与C4A/C128A场景不匹配。 + SWA场景cmpRatio不为1,或cmpRatio与CSA、HCA场景不匹配。 cmpTopk不为0、512或1024。 @@ -610,7 +611,7 @@ aclnnStatus aclnnSparseFlashMlaMetadata( - layoutKvOptional为PA_BBND时,`sequsedOriKvOptional`必须传入。BSND场景可选传入`sequsedOriKvOptional`覆盖每个batch的oriKv有效长度;TND场景使用`cuSeqlensOriKvOptional`表达oriKv序列边界。 - layoutKvOptional为TND时,`cuSeqlensOriKvOptional`必须传入;若hasCmpKv为true,`cuSeqlensCmpKvOptional`也必须传入。 - `sequsedCmpKvOptional`为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。 - - `cmpResidualKvOptional`为`aclnnSparseFlashMlaMetadata`和`aclnnSparseFlashMla`的可选输入,在C4A/C128A、cmpRatio不等于1且cmpMaskMode为3场景必传,用于恢复cmp侧mask使用的压缩前长度。 + - `cmpResidualKvOptional`为`aclnnSparseFlashMlaMetadata`和`aclnnSparseFlashMla`的可选输入,在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,用于恢复cmp侧mask使用的压缩前长度。 - 该算子为AICPU算子,在Host侧CPU上执行,不占用NPU计算资源。 @@ -619,242 +620,357 @@ aclnnStatus aclnnSparseFlashMlaMetadata( 调用示例代码如下,仅供参考,具体编译和执行过程请参考[编译与运行样例](../../../docs/zh/context/编译与运行样例.md)。 ```c++ +/** + * @file test_aclnn_sparse_flash_mla_metadata.cpp + */ #include #include +#include +#include +#include +#include +#include #include "acl/acl.h" #include "aclnnop/aclnn_sparse_flash_mla_metadata.h" -#include "../../sparse_flash_mla/op_kernel/arch22/sparse_flash_mla_metadata.h" - -#define CHECK_RET(cond, return_expr) \ - do { \ - if (!(cond)) { \ - return_expr; \ - } \ - } while (0) +#define CHECK_LOG_RET(cond, ret_val, fmt, ...) \ + do { \ + if (!(cond)) { \ + printf(fmt "\n", ##__VA_ARGS__); \ + return (ret_val); \ + } \ + } while (0) + +// 参考 sparse_flash_mla_metadata.h +constexpr uint32_t AIC_CORE_MAX_NUM = 36; +constexpr uint32_t AIV_CORE_MAX_NUM = 72; +constexpr uint32_t SMLA_METADATA_TOTAL_SIZE = 1024; +constexpr uint32_t FA_METADATA_SIZE = 9; +constexpr uint32_t FD_METADATA_SIZE = 8; + +// FA Metadata Index Definitions +constexpr uint32_t FA_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FA_BN2_START_INDEX = 1; +constexpr uint32_t FA_M_START_INDEX = 2; +constexpr uint32_t FA_S2_START_INDEX = 3; +constexpr uint32_t FA_BN2_END_INDEX = 4; +constexpr uint32_t FA_M_END_INDEX = 5; +constexpr uint32_t FA_S2_END_INDEX = 6; +constexpr uint32_t FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX = 7; +constexpr uint32_t FA_S2_MAX_NUM = 8; + +// FD Metadata Index Definitions +constexpr uint32_t FD_CORE_ENABLE_INDEX = 0; +constexpr uint32_t FD_BN2_IDX_INDEX = 1; +constexpr uint32_t FD_M_IDX_INDEX = 2; +constexpr uint32_t FD_WORKSPACE_IDX_INDEX = 3; +constexpr uint32_t FD_WORKSPACE_NUM_INDEX = 4; +constexpr uint32_t FD_M_START_INDEX = 5; +constexpr uint32_t FD_M_NUM_INDEX = 6; + +struct SmlaMetadata { + uint32_t faMetadata[AIC_CORE_MAX_NUM][FA_METADATA_SIZE]; + uint32_t fdMetadata[AIV_CORE_MAX_NUM][FD_METADATA_SIZE]; +}; + +struct ScopeGuard +{ + explicit ScopeGuard(std::function onExitScope) : m_exitFunc(std::move(onExitScope)), + m_isDismissed(false) {} + // 禁止拷贝 + ScopeGuard(const ScopeGuard&) = delete; + ScopeGuard& operator=(const ScopeGuard&) = delete; + + ~ScopeGuard() + { + if (!m_isDismissed) { + m_exitFunc(); + } + } + + void Dismiss() + { + m_isDismissed = true; + } + + std::function m_exitFunc; + bool m_isDismissed; +}; + +struct Tensor { + void *hostAddr { nullptr }; + void *deviceAddr { nullptr }; + aclTensor *data { nullptr }; +}; + +struct ArgScenario { + bool hasCuSeq { false }; + bool hasSeqused { false }; +}; + +struct ArgContext { + // required input + int64_t numHeadsQ { 0 }; + int64_t numHeadsKv { 0 }; + int64_t headDim { 0 }; + // optional input + Tensor cuSeqlensQOptional {}; + Tensor cuSeqlensOriKvOptional {}; + Tensor cuSeqlensCmpKvOptional {}; + Tensor sequsedQOptional {}; + Tensor sequsedOriKvOptional {}; + Tensor sequsedCmpKvOptional {}; + Tensor cmpResidualKvOptional {}; + Tensor oriTopkLengthOptional {}; + Tensor cmpTopkLengthOptional {}; + int64_t batchSize { 0 }; + int64_t maxSeqlenQ { 0 }; + int64_t maxSeqlenOriKv { 0 }; + int64_t maxSeqlenCmpKv { 0 }; + int64_t oriTopk { 0 }; + int64_t cmpTopk { 0 }; + int64_t cmpRatio { 0 }; + int64_t oriMaskMode { 0 }; + int64_t cmpMaskMode { 0 }; + int64_t oriWinLeft { -1 }; + int64_t oriWinRight { -1 }; + char *layoutQOptional { nullptr }; + char *layoutKvOptional { nullptr }; + bool hasOriKv { true }; + bool hasCmpKv { true }; + // output + Tensor metadata {}; +}; + +int64_t GetShapeSize(const std::vector& shape) +{ + int64_t shapeSize = 1; + for (auto i : shape) { + shapeSize *= i; + } + return shapeSize; +} -#define LOG_PRINT(message, ...) \ - do { \ - printf(message, ##__VA_ARGS__); \ - } while (0) +aclnnStatus Init(int32_t deviceId, aclrtStream* stream) +{ + // 固定写法,初始化 + auto ret = aclInit(nullptr); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclInit failed. ERROR: %d", ret); + ret = aclrtSetDevice(deviceId); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSetDevice failed. ERROR: %d", ret); + ret = aclrtCreateStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtCreateStream failed. ERROR: %d", ret); + return ACL_SUCCESS; +} -int64_t GetShapeSize(const std::vector& shape) +void Finalize(int32_t deviceId, aclrtStream stream) { - int64_t shapeSize = 1; - for (auto i : shape) { - shapeSize *= i; - } - return shapeSize; + aclrtDestroyStream(stream); + aclrtResetDevice(deviceId); + aclFinalize(); } -int Init(int32_t deviceId, aclrtContext* context, aclrtStream* stream) +aclnnStatus CreateTensor(aclDataType dataType, const std::vector &shape, Tensor &tensor) { - auto ret = aclInit(nullptr); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); - ret = aclrtSetDevice(deviceId); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret); - ret = aclrtCreateContext(context, deviceId); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateContext failed. ERROR: %d\n", ret); return ret); - ret = aclrtSetCurrentContext(*context); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetCurrentContext failed. ERROR: %d\n", ret); return ret); - ret = aclrtCreateStream(stream); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret); - return 0; + auto size = GetShapeSize(shape) * aclDataTypeSize(dataType); + // 调用aclrtMallocHost申请host侧内存 + auto ret = aclrtMallocHost(&(tensor.hostAddr), size); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMallocHost failed. ERROR: %d", ret); + memset(tensor.hostAddr, 0, size); + // 调用aclrtMalloc申请device侧内存 + ret = aclrtMalloc(&(tensor.deviceAddr), size, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMalloc failed. ERROR: %d", ret); + // 调用aclCreateTensor接口创建aclTensor + tensor.data = aclCreateTensor(shape.data(), shape.size(), dataType, nullptr, 0, aclFormat::ACL_FORMAT_ND, + shape.data(), shape.size(), tensor.deviceAddr); + CHECK_LOG_RET(tensor.data != nullptr, ACL_ERROR_FAILURE, "aclCreateTensor failed"); + // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 + ret = aclrtMemcpy(tensor.deviceAddr, size, tensor.hostAddr, size, ACL_MEMCPY_HOST_TO_DEVICE); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d", ret); + return ACL_SUCCESS; } -template -int CreateAclTensor(const std::vector& hostData, const std::vector& shape, void** deviceAddr, - aclDataType dataType, aclTensor** tensor) +void DestroyTensor(Tensor &tensor) { - auto size = GetShapeSize(shape) * sizeof(T); - if (size > 0) { - auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); - ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); - } else { - *deviceAddr = nullptr; - } - - std::vector strides(shape.size(), 1); - for (int64_t i = static_cast(shape.size()) - 2; i >= 0; i--) { - strides[i] = shape[i + 1] * strides[i + 1]; - } - - *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, - shape.data(), shape.size(), *deviceAddr); - return 0; + if (tensor.data != nullptr) { + aclDestroyTensor(tensor.data); + tensor.data = nullptr; + } + if (tensor.deviceAddr != nullptr) { + aclrtFree(tensor.deviceAddr); + tensor.deviceAddr = nullptr; + } + if (tensor.hostAddr != nullptr) { + aclrtFreeHost(tensor.hostAddr); + tensor.hostAddr = nullptr; + } } -void PrintMetadataSummary(const optiling::detail::SasMetadata& meta) +void DestroyArgs(ArgContext &context) { - printf("AIC core0 enable=%u, bn2_end=%u, m_end=%u, s2_end=%u\n", - meta.faMetadata[0][optiling::FA_CORE_ENABLE_INDEX], - meta.faMetadata[0][optiling::FA_BN2_END_INDEX], - meta.faMetadata[0][optiling::FA_M_END_INDEX], - meta.faMetadata[0][optiling::FA_S2_END_INDEX]); + DestroyTensor(context.metadata); + DestroyTensor(context.cuSeqlensQOptional); + DestroyTensor(context.cuSeqlensOriKvOptional); + DestroyTensor(context.cuSeqlensCmpKvOptional); + DestroyTensor(context.sequsedQOptional); + DestroyTensor(context.sequsedOriKvOptional); + DestroyTensor(context.sequsedCmpKvOptional); + DestroyTensor(context.cmpResidualKvOptional); + DestroyTensor(context.oriTopkLengthOptional); + DestroyTensor(context.cmpTopkLengthOptional); + + if (context.layoutQOptional != nullptr) { + free(context.layoutQOptional); + context.layoutQOptional = nullptr; + } + if (context.layoutKvOptional != nullptr) { + free(context.layoutKvOptional); + context.layoutKvOptional = nullptr; + } } -int main() +aclnnStatus CreateArgs(const ArgScenario &scenario, ArgContext &context) { - // 1. (固定写法)device/stream初始化,参考acl API手册 - // 根据自己的实际device填写deviceId - int32_t deviceId = 5; - aclrtContext context = nullptr; - aclrtStream stream = nullptr; - auto ret = Init(deviceId, &context, &stream); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); - - int64_t B = 4; - int64_t S1 = 128; - int64_t S2 = 8192; - int64_t N1 = 64; - int64_t N2 = 1; - int64_t D = 512; - int64_t K = 512; - int64_t s2Act = 4096; - int64_t cmpRatio = 4; - int64_t cmpKvLen = s2Act / cmpRatio; - int64_t oriMaskMode = 4; - int64_t cmpMaskMode = 3; - int64_t oriWinLeft = 127; - int64_t oriWinRight = 0; - - // 2. 构造输入与输出,需要根据API的接口自定义构造 - std::vector cuSeqLensQShape = {B + 1}; - std::vector seqUsedOriKvShape = {B}; - std::vector cmpResidualKvShape = {B}; - std::vector metadataShape = {optiling::SMLA_META_SIZE}; - // 对全部 optional 输入调用 Contiguous,optional 输入传 shape 为 {0} 的空 tensor。 - std::vector emptyShape = {0}; - std::vector seqUsedQShape = emptyShape; - - void* cuSeqLensQDeviceAddr = nullptr; - void* cuSeqLensOriKvDeviceAddr = nullptr; - void* cuSeqLensCmpKvDeviceAddr = nullptr; - void* seqUsedQDeviceAddr = nullptr; - void* seqUsedOriKvDeviceAddr = nullptr; - void* seqUsedCmpKvDeviceAddr = nullptr; - void* cmpResidualKvDeviceAddr = nullptr; - void* oriTopkLengthDeviceAddr = nullptr; - void* cmpTopkLengthDeviceAddr = nullptr; - void* metadataDeviceAddr = nullptr; - - aclTensor* cuSeqLensQ = nullptr; - aclTensor* cuSeqLensOriKv = nullptr; - aclTensor* cuSeqLensCmpKv = nullptr; - aclTensor* seqUsedQ = nullptr; - aclTensor* seqUsedOriKv = nullptr; - aclTensor* seqUsedCmpKv = nullptr; - aclTensor* cmpResidualKv = nullptr; - aclTensor* oriTopkLength = nullptr; - aclTensor* cmpTopkLength = nullptr; - aclTensor* metadata = nullptr; - - std::vector cuSeqLensQHostData(B + 1); - for (int64_t i = 0; i <= B; i++) { - cuSeqLensQHostData[i] = static_cast(i * S1); - } - std::vector emptyHostData; - std::vector seqUsedOriKvHostData(B, static_cast(s2Act)); - std::vector cmpResidualKvHostData(B, 0); - std::vector metadataHostData(optiling::SMLA_META_SIZE, 0); - - ret = CreateAclTensor(cuSeqLensQHostData, cuSeqLensQShape, &cuSeqLensQDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensQ); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensOriKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensOriKv); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &cuSeqLensCmpKvDeviceAddr, aclDataType::ACL_INT32, &cuSeqLensCmpKv); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, seqUsedQShape, &seqUsedQDeviceAddr, aclDataType::ACL_INT32, &seqUsedQ); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(seqUsedOriKvHostData, seqUsedOriKvShape, &seqUsedOriKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedOriKv); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &seqUsedCmpKvDeviceAddr, aclDataType::ACL_INT32, &seqUsedCmpKv); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(cmpResidualKvHostData, cmpResidualKvShape, &cmpResidualKvDeviceAddr, aclDataType::ACL_INT32, &cmpResidualKv); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &oriTopkLengthDeviceAddr, aclDataType::ACL_INT32, &oriTopkLength); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(emptyHostData, emptyShape, &cmpTopkLengthDeviceAddr, aclDataType::ACL_INT32, &cmpTopkLength); - CHECK_RET(ret == ACL_SUCCESS, return ret); - ret = CreateAclTensor(metadataHostData, metadataShape, &metadataDeviceAddr, aclDataType::ACL_INT32, &metadata); - CHECK_RET(ret == ACL_SUCCESS, return ret); - - char layoutQ[] = "TND"; - char layoutKv[] = "PA_BBND"; - - uint64_t workspaceSize = 0; - aclOpExecutor* executor = nullptr; - - // 3. 调用CANN算子库API,需要修改为具体的Api名称 - ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( - cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, - seqUsedQ, seqUsedOriKv, seqUsedCmpKv, - cmpResidualKv, oriTopkLength, cmpTopkLength, - N1, N2, D, B, S1, S2, cmpKvLen, - 0, K, cmpRatio, - oriMaskMode, cmpMaskMode, - oriWinLeft, oriWinRight, - layoutQ, layoutKv, - true, true, - metadata, - &workspaceSize, &executor); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); - return ret); - - void* workspaceAddr = nullptr; - if (workspaceSize > 0) { - ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); - } - - ret = aclnnSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); return ret); - - // 4. (固定写法)同步等待任务执行结束 - ret = aclrtSynchronizeStream(stream); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); - - optiling::detail::SasMetadata result {}; - ret = aclrtMemcpy(&result, sizeof(result), metadataDeviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); - CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("copy metadata result failed. ERROR: %d\n", ret); return ret); - - // 5.获取输出的值,将device侧内存上的结果拷贝至host侧,需要根据具体API的接口定义修改 - PrintMetadataSummary(result); - CHECK_RET(result.faMetadata[0][optiling::FA_CORE_ENABLE_INDEX] == 1U, - LOG_PRINT("metadata validation failed: core0 is not enabled\n"); return 1); - // 分核可能在 batch 内按行切分,此时 bn2_end 仍为 0,m_end 已推进。 - CHECK_RET(result.faMetadata[0][optiling::FA_BN2_END_INDEX] > 0U || - result.faMetadata[0][optiling::FA_M_END_INDEX] > 0U, - LOG_PRINT("metadata validation failed: core0 has no assigned work\n"); return 1); - - // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 - aclDestroyTensor(cuSeqLensQ); - aclDestroyTensor(cuSeqLensOriKv); - aclDestroyTensor(cuSeqLensCmpKv); - aclDestroyTensor(seqUsedQ); - aclDestroyTensor(seqUsedOriKv); - aclDestroyTensor(metadata); - - // 7. 释放device资源 - if (cuSeqLensQDeviceAddr != nullptr) { - aclrtFree(cuSeqLensQDeviceAddr); - } - if (seqUsedOriKvDeviceAddr != nullptr) { - aclrtFree(seqUsedOriKvDeviceAddr); - } - if (metadataDeviceAddr != nullptr) { - aclrtFree(metadataDeviceAddr); - } - if (workspaceSize > 0) { - aclrtFree(workspaceAddr); - } - aclrtDestroyStream(stream); - aclrtDestroyContext(context); - aclrtResetDevice(deviceId); - aclFinalize(); - - return 0; + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + aclnnStatus ret; + + context.numHeadsQ = 64; + context.numHeadsKv = 1; + context.headDim = 512; + ret = CreateTensor(aclDataType::ACL_INT32, { SMLA_METADATA_TOTAL_SIZE }, context.metadata); // 1024: Fix size + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create metadata failed. Error: %d", ret); + context.oriTopk = 0; + context.cmpTopk = 0; + context.cmpRatio = 128; + context.oriMaskMode = 4; + context.cmpMaskMode = 3; + context.oriWinLeft = 127; + context.oriWinRight = 0; + context.layoutQOptional = (char *)malloc(sizeof(char) * 16); + context.layoutKvOptional = (char *)malloc(sizeof(char) * 16); + CHECK_LOG_RET(context.layoutQOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutQOptional failed"); + CHECK_LOG_RET(context.layoutKvOptional != nullptr, ACL_ERROR_FAILURE, "Create layoutKvOptional failed"); + strcpy(context.layoutQOptional, "BSND"); // BSND,TND + strcpy(context.layoutKvOptional, "BSND"); // BSND,TND,PA_BBND + context.hasOriKv = true; + context.hasCmpKv = true; + + context.batchSize = 4; + context.maxSeqlenOriKv = 1024; + context.maxSeqlenCmpKv = 1024; + context.maxSeqlenQ = 1024; + + if (scenario.hasCuSeq) { + // (B+1,), first element is always 0 + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize + 1 }, context.cuSeqlensCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cuSeqlensCmpKvOptional failed. Error: %d", ret); + } + + if (scenario.hasSeqused) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedQOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedQOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedOriKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedOriKvOptional failed. Error: %d", ret); + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.sequsedCmpKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create sequsedCmpKvOptional failed. Error: %d", ret); + } + + if (context.hasCmpKv && context.cmpRatio != 1 && context.cmpMaskMode == 3) { + // (B,) + ret = CreateTensor(aclDataType::ACL_INT32, { context.batchSize }, context.cmpResidualKvOptional); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create cmpResidualKvOptional failed. Error: %d", ret); + } + + argsGuard.Dismiss(); + return ACL_SUCCESS; +} + +int main() { + // 1. (固定写法)device/stream初始化,参考对外接口列表 + // 根据自己的实际device填写deviceId + int32_t deviceId = 0; + aclrtStream stream; + auto ret = Init(deviceId, &stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Init acl failed. ERROR: %d", ret); + ScopeGuard sysGuard([&] { Finalize(deviceId, stream); }); + + // 2. 构造输入与输出,需要根据API的接口定义构造 + ArgScenario scenario {}; + scenario.hasCuSeq = false; + scenario.hasSeqused = false; + ArgContext context {}; + ret = CreateArgs(scenario, context); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "Create input arguments failed. ERROR: %d", ret); + ScopeGuard argsGuard([&] { DestroyArgs(context); }); + + // 3. 调用CANN算子库API,需要修改为具体的API + // 调用aclnnSparseFlashMlaMetadata第一段接口 + uint64_t workspaceSize = 0; + aclOpExecutor *executor = nullptr; + void *workspaceAddr = nullptr; + ret = aclnnSparseFlashMlaMetadataGetWorkspaceSize( + context.cuSeqlensQOptional.data, context.cuSeqlensOriKvOptional.data, context.cuSeqlensCmpKvOptional.data, + context.sequsedQOptional.data, context.sequsedOriKvOptional.data, context.sequsedCmpKvOptional.data, + context.cmpResidualKvOptional.data, context.oriTopkLengthOptional.data, context.cmpTopkLengthOptional.data, + context.numHeadsQ, context.numHeadsKv, context.headDim, context.batchSize, context.maxSeqlenQ, + context.maxSeqlenOriKv, context.maxSeqlenCmpKv, context.oriTopk, context.cmpTopk, context.cmpRatio, + context.oriMaskMode, context.cmpMaskMode, context.oriWinLeft, context.oriWinRight, context.layoutQOptional, + context.layoutKvOptional, context.hasOriKv, context.hasCmpKv, context.metadata.data, &workspaceSize, &executor); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, + "aclnnSparseFlashMlaMetadataGetWorkspaceSize failed. ERROR: %d\n", ret); + + if (workspaceSize > static_cast(0)) { + ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "allocate workspace failed. ERROR: %d\n", ret); + } + ScopeGuard workspaceGuard([&] { + if (workspaceAddr != nullptr) { + aclrtFree(workspaceAddr); + workspaceAddr = nullptr; + } + }); + + // 调用aclnnSparseFlashMlaMetadata第二段接口 + ret = aclnnSparseFlashMlaMetadata(workspaceAddr, workspaceSize, executor, stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclnnSparseFlashMlaMetadata failed. ERROR: %d\n", ret); + + // 4. (固定写法)同步等待任务执行结束 + ret = aclrtSynchronizeStream(stream); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtSynchronizeStream failed. ERROR: %d\n", ret); + + // 5. 打印输出 + SmlaMetadata result {}; + ret = aclrtMemcpy(&result, sizeof(result), context.metadata.deviceAddr, sizeof(result), ACL_MEMCPY_DEVICE_TO_HOST); + CHECK_LOG_RET(ret == ACL_SUCCESS, ret, "aclrtMemcpy failed. ERROR: %d\n", ret); + + for (uint32_t i = 0; i < AIC_CORE_MAX_NUM; ++i) { + printf("AIC Core%u\n", i); + printf(" Core Enable : %u\n", result.faMetadata[i][FA_CORE_ENABLE_INDEX]); + printf(" Start BN2 : %u\n", result.faMetadata[i][FA_BN2_START_INDEX]); + printf(" Start M : %u\n", result.faMetadata[i][FA_M_START_INDEX]); + printf(" Start S2 : %u\n", result.faMetadata[i][FA_S2_START_INDEX]); + printf(" End BN2 : %u\n", result.faMetadata[i][FA_BN2_END_INDEX]); + printf(" End M : %u\n", result.faMetadata[i][FA_M_END_INDEX]); + printf(" End S2 : %u\n", result.faMetadata[i][FA_S2_END_INDEX]); + printf(" First Worksapce Index : %u\n", result.faMetadata[i][FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX]); + printf(" Max S2 Block Num : %u\n", result.faMetadata[i][FA_S2_MAX_NUM]); + } + for (uint32_t i = 0; i < AIV_CORE_MAX_NUM; ++i) { + printf("AIV Core%u\n", i); + printf(" Core Enable : %u\n", result.fdMetadata[i][FD_CORE_ENABLE_INDEX]); + printf(" FD Task BN2 Idx : %u\n", result.fdMetadata[i][FD_BN2_IDX_INDEX]); + printf(" FD Task M Idx : %u\n", result.fdMetadata[i][FD_M_IDX_INDEX]); + printf(" FD Task S2 Idx : %u\n", result.fdMetadata[i][FD_WORKSPACE_IDX_INDEX]); + printf(" FD Task Workspace Num : %u\n", result.fdMetadata[i][FD_WORKSPACE_NUM_INDEX]); + printf(" FD Subtask M Start : %u\n", result.fdMetadata[i][FD_M_START_INDEX]); + printf(" FD Subtask M Num : %u\n", result.fdMetadata[i][FD_M_NUM_INDEX]); + } + + return 0; } ``` diff --git a/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h b/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h index 0b1f68fb63..a3426d01a7 100644 --- a/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h +++ b/attention/sparse_flash_mla_metadata/op_host/sparse_flash_mla_metadata_check.h @@ -241,8 +241,17 @@ aclnnStatus CheckSingleParamSmla(int64_t batchSize, int64_t maxSeqlenQ, int64_t SMLA_A5_PLATFORM_LOG.c_str(), cmpRatio); } else { int64_t expectedCmpRatio = (cmpTopk > 0) ? 4 : 128; - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "cmp_ratio should be %lld on %s, but got %lld", - expectedCmpRatio, SMLA_A2_A3_PLATFORM_LOG.c_str(), cmpRatio); + if (cmpTopk > 0) { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "cmp_ratio should be %lld on %s when cmp_topk is non-zero(CSA with " + "cmp_sparse_indices), but got %lld", + expectedCmpRatio, SMLA_A2_A3_PLATFORM_LOG.c_str(), cmpRatio); + } else { + OP_LOGE(ACLNN_ERR_PARAM_INVALID, + "cmp_ratio should be %lld on %s when cmp_topk is 0(HCA without " + "cmp_sparse_indices), but got %lld", + expectedCmpRatio, SMLA_A2_A3_PLATFORM_LOG.c_str(), cmpRatio); + } } return ACLNN_ERR_PARAM_INVALID; } diff --git a/attention/sparse_flash_mla_metadata/tests/CMakeLists.txt b/attention/sparse_flash_mla_metadata/tests/CMakeLists.txt new file mode 100644 index 0000000000..d5a84231c2 --- /dev/null +++ b/attention/sparse_flash_mla_metadata/tests/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# 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(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/attention/sparse_flash_mla_metadata/tests/ut/CMakeLists.txt b/attention/sparse_flash_mla_metadata/tests/ut/CMakeLists.txt new file mode 100644 index 0000000000..d5a84231c2 --- /dev/null +++ b/attention/sparse_flash_mla_metadata/tests/ut/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# 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(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/attention/sparse_flash_mla_metadata/tests/ut/op_host/CMakeLists.txt b/attention/sparse_flash_mla_metadata/tests/ut/op_host/CMakeLists.txt new file mode 100644 index 0000000000..d5a84231c2 --- /dev/null +++ b/attention/sparse_flash_mla_metadata/tests/ut/op_host/CMakeLists.txt @@ -0,0 +1,16 @@ +# ----------------------------------------------------------------------------------------------------------- +# 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(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/CMakeLists.txt b/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/CMakeLists.txt new file mode 100644 index 0000000000..77e0ca1492 --- /dev/null +++ b/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/CMakeLists.txt @@ -0,0 +1,13 @@ +# ----------------------------------------------------------------------------------------------------------- +# 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. +# ----------------------------------------------------------------------------------------------------------- + +if(UT_TEST_ALL OR OP_API_UT) + add_modules_ut_sources(UT_NAME ${OP_API_MODULE_NAME} MODE PRIVATE DIR ${CMAKE_CURRENT_SOURCE_DIR}) +endif() diff --git a/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/test_aclnn_sparse_flash_mla_metadata.cpp b/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/test_aclnn_sparse_flash_mla_metadata.cpp new file mode 100644 index 0000000000..90cea5ecf2 --- /dev/null +++ b/attention/sparse_flash_mla_metadata/tests/ut/op_host/op_api/test_aclnn_sparse_flash_mla_metadata.cpp @@ -0,0 +1,80 @@ +/** + * Copyright (c) 2026 Huawei Technologies Co., Ltd. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + */ + +#include + +#include "gtest/gtest.h" +#include "../../../../op_host/sparse_flash_mla_metadata_check.h" +#include "op_api_ut_common/tensor_desc.h" + +namespace { +constexpr int64_t NUM_HEADS_Q = 64; +constexpr int64_t NUM_HEADS_KV = 1; +constexpr int64_t HEAD_DIM = 512; +constexpr int64_t BATCH_SIZE = 2; +constexpr int64_t MAX_SEQLEN_Q = 16; +constexpr int64_t MAX_SEQLEN_ORI_KV = 16; +constexpr int64_t MAX_SEQLEN_CMP_KV = 4; +constexpr int64_t ORI_TOPK = 0; +constexpr int64_t CMP_TOPK = 512; +constexpr int64_t CMP_RATIO = 4; +constexpr int64_t ORI_MASK_MODE = 4; +constexpr int64_t CMP_MASK_MODE = 3; +constexpr int64_t ORI_WIN_LEFT = 127; +constexpr int64_t ORI_WIN_RIGHT = 0; +constexpr uint32_t AIC_CORE_NUM = 36; +constexpr uint32_t AIV_CORE_NUM = 72; + +struct TndC4aParams { + std::unique_ptr cuSeqlensQ = + TensorDesc({3}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); + std::unique_ptr cuSeqlensOriKv = + TensorDesc({3}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); + std::unique_ptr cuSeqlensCmpKv = + TensorDesc({3}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); + std::unique_ptr cmpResidualKv = + TensorDesc({2}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); + std::unique_ptr metadata = + TensorDesc({optiling::SMLA_METADATA_TOTAL_SIZE}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); +}; + +aclnnStatus RunParamsCheck(const TndC4aParams ¶ms, const aclTensor *cmpResidualKv) +{ + return ParamsCheck(params.cuSeqlensQ.get(), params.cuSeqlensOriKv.get(), params.cuSeqlensCmpKv.get(), + nullptr, nullptr, nullptr, cmpResidualKv, nullptr, nullptr, NUM_HEADS_Q, NUM_HEADS_KV, + HEAD_DIM, 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, "TND", + "TND", true, true, AIC_CORE_NUM, AIV_CORE_NUM, "Ascend910B", params.metadata.get()); +} +} // namespace + +class SparseFlashMlaMetadataApiTest : public testing::Test {}; + +TEST_F(SparseFlashMlaMetadataApiTest, tnd_c4a_requires_cmp_residual_kv) +{ + TndC4aParams params; + + EXPECT_EQ(RunParamsCheck(params, nullptr), ACLNN_ERR_PARAM_INVALID); +} + +TEST_F(SparseFlashMlaMetadataApiTest, tnd_c4a_accepts_cmp_residual_kv) +{ + TndC4aParams params; + + EXPECT_EQ(RunParamsCheck(params, params.cmpResidualKv.get()), ACLNN_SUCCESS); +} + +TEST_F(SparseFlashMlaMetadataApiTest, cmp_residual_kv_batch_must_match) +{ + TndC4aParams params; + auto wrongBatchResidual = TensorDesc({1}, ACL_INT32, ACL_FORMAT_ND).ToAclType(); + + EXPECT_EQ(RunParamsCheck(params, wrongBatchResidual.get()), ACLNN_ERR_PARAM_INVALID); +} diff --git a/docs/zh/op_api_list.md b/docs/zh/op_api_list.md index b35b34435b..2606463ad7 100644 --- a/docs/zh/op_api_list.md +++ b/docs/zh/op_api_list.md @@ -121,6 +121,7 @@ |[aclnnMatmulAllReduceV2](../../mc2/matmul_all_reduce/docs/aclnnMatmulAllReduceV2.md)|完成MatMul计算与AllReduce通信融合。|默认非确定性实现,支持配置开启| 默认确定性实现 | |[aclnnMatmulReduceScatter](../../mc2/matmul_reduce_scatter/docs/aclnnMatmulReduceScatter.md)|完成mm + reduce_scatter_base计算。|默认非确定性实现,支持配置开启| 默认确定性实现 | |[aclnnMatmulReduceScatterV2](../../mc2/matmul_reduce_scatter_v2/docs/aclnnMatmulReduceScatterV2.md)|aclnnMatmulReduceScatterV2接口是对[aclnnMatmulReduceScatter](../../mc2/matmul_reduce_scatter/docs/aclnnMatmulReduceScatter.md)接口的功能扩展。|默认确定性实现| 默认确定性实现 | +|[aclnnMixedQuantSparseFlashMla](../../attention/mixed_quant_sparse_flash_mla/docs/aclnnMixedQuantSparseFlashMla.md)|支持量化场景下SWA、CSA、HCA三类Attention计算场景。|默认确定性实现| 默认确定性实现 | |[aclnnMixedQuantSparseFlashMlaMetadata](../../attention/mixed_quant_sparse_flash_mla_metadata/docs/aclnnMixedQuantSparseFlashMlaMetadata.md)| aclnnMixedQuantSparseFlashMla接口的前置接口,用于计算aclnnMixedQuantSparseFlashMla的负载均衡。| - | 默认确定性实现 | |[aclnnMlaPreprocess](../../attention/mla_preprocess/docs/aclnnMlaPreprocess.md)|Multi-Head Latent Attention前处理的计算。|默认确定性实现| - | |[aclnnMlaPreprocessV2](../../attention/mla_preprocess_v2/docs/aclnnMlaPreprocessV2.md)|推理场景,Multi-Head Latent Attention前处理的计算。主要计算过程如下:|默认确定性实现| - | @@ -218,9 +219,9 @@ |[aclnnScatterPaKvCache](../../attention/scatter_pa_kv_cache/docs/aclnnScatterPaKvCache.md)|更新KvCache中指定位置的key和value。|默认确定性实现| 默认确定性实现 | |[aclnnSparseFlashAttention](../../attention/sparse_flash_attention/docs/aclnnSparseFlashAttention.md)|根据sparse_indices选取重要性较高的key和value进行attention运算,得到attention_out输出。|默认确定性实现| 默认确定性实现 | |[aclnnSparseFlashAttentionGrad](../../attention/sparse_flash_attention_grad/docs/aclnnSparseFlashAttentionGrad.md)|根据topkIndices对key和value选取大小为selectedBlockSize的数据重排,接着进行训练场景下计算注意力的反向输出。|默认非确定性实现,支持配置开启| 默认确定性实现 | -|[aclnnSparseFlashMla](../../attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md)|支持C1A、C4A、C128A三类Attention计算场景。|默认确定性实现| 默认确定性实现 | +|[aclnnSparseFlashMla](../../attention/sparse_flash_mla/docs/aclnnSparseFlashMla.md)|支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。|默认确定性实现| 默认确定性实现 | |[aclnnSparseFlashMlaMetadata](../../attention/sparse_flash_mla_metadata/docs/aclnnSparseFlashMlaMetadata.md)|生成aclnnSparseFlashMla主算子使用的任务切分metadata。|默认确定性实现| 默认确定性实现 | -|[aclnnSparseFlashMlaGrad](../../attention/sparse_flash_mla_grad/docs/aclnnSparseFlashMlaGrad.md)|计算SparseFlashMla训练场景下注意力的反向输出,支持Sliding Window Attention、Compressed Attention以及Sparse Compressed Attention。|默认非确定性实现,不支持配置开启| 默认非确定性实现,支持配置开启 | +|[aclnnSparseFlashMlaGrad](../../attention/sparse_flash_mla_grad/docs/aclnnSparseFlashMlaGrad.md)|计算SparseFlashMla训练场景下注意力的反向输出,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)以及HCA(Heavily Compressed Attention)。|默认非确定性实现,不支持配置开启| 默认非确定性实现,支持配置开启 | |[aclnnSparseLightningIndexerGradKLLoss](../../attention/sparse_lightning_indexer_grad_kl_loss/docs/aclnnSparseLightningIndexerGradKLLoss.md)|LightningIndexer的反向算子,再额外融合了Loss计算功能。|默认非确定性实现,不支持配置开启| 默认确定性实现 | |[aclnnSparseLightningIndexerKLLossGrad](../../attention/sparse_lightning_indexer_kl_loss_grad/docs/aclnnSparseLightningIndexerKLLossGrad.md)|LightningIndexer的反向算子,支持输出Loss计算所需Index部分的分数。|默认非确定性实现,不支持配置开启| 默认非确定性实现,支持配置开启 | |[aclnnSwigluGatedMlp](../../experimental/ffn/swiglu_gated_mlp/docs/aclnnSwigluGatedMlp.md)|完成融合SwiGLU门控MLP计算,包括首个MatMul、SwiGLU激活和第二个MatMul。|默认确定性实现| - | diff --git a/docs/zh/op_list.md b/docs/zh/op_list.md index 4c9d6b36e9..4d4e53b90a 100644 --- a/docs/zh/op_list.md +++ b/docs/zh/op_list.md @@ -484,7 +484,7 @@ ✗ ✗ AI Core - 支持Sliding Window Attention(SWA)、Compressed Flash Attention(CFA)和Sparse Compressed Flash Attention(SCFA)。 + 支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)和HCA(Heavily Compressed Attention)。 attention diff --git a/docs/zh/torch_api_list.md b/docs/zh/torch_api_list.md index 78d01fd35c..7f371239fa 100644 --- a/docs/zh/torch_api_list.md +++ b/docs/zh/torch_api_list.md @@ -23,9 +23,13 @@ |[lightning_indexer_metadata](../../torch_extension/cann_ops_transformer/docs/zh/lightning_indexer.md)|lightning_indexer接口的前置接口,用于计算lightning_indexer的负载均衡。|默认确定性实现|默认确定性实现| |[mhc_post](../../torch_extension/cann_ops_transformer/docs/zh/mhc_post.md)|实现MHC Post组件的前向计算,用于Transformer模型中多层残差连接的后处理阶段。该算子将残差矩阵变换与输出状态投影融合为单次计算,避免多次独立算子调用带来的额外开销。|默认确定性实现|-| |[mhc_pre_sinkhorn](../../torch_extension/cann_ops_transformer/docs/zh/mhc_pre_sinkhorn.md)|基于一系列计算得到MHC架构中hidden层的$\mathbf{H}'_{\text{res}}$和$\mathbf{H}_{\text{post}}$投影矩阵以及Attention或MLP层的输入矩阵$\mathbf{h}_{\text{in}}$。对$\mathbf{H}'_{\text{res}}$矩阵执行Sinkhorn迭代归一化变换,最终得到双随机矩阵$\mathbf{H}_{\text{res}}$;支持输出中间计算结果,用于反向梯度计算。|默认确定性实现|-| +|[mixed_quant_sparse_flash_mla](../../torch_extension/cann_ops_transformer/docs/zh/mixed_quant_sparse_flash_mla.md)|量化场景下基于共享KV完成MixedQuantSparseFlashMla稀疏注意力计算。|默认确定性实现|默认确定性实现| |[mixed_quant_sparse_flash_mla_metadata](../../torch_extension/cann_ops_transformer/docs/zh/mixed_quant_sparse_flash_mla.md)|mixed_quant_sparse_flash_mla接口的前置接口,用于计算mixed_quant_sparse_flash_mla的负载均衡。|-|默认确定性实现| |[quant_lightning_indexer_metadata](../../torch_extension/cann_ops_transformer/docs/zh/quant_lightning_indexer.md)|quant_lightning_indexer接口的前置接口,用于计算quant_lightning_indexer的负载均衡。|默认确定性实现|默认确定性实现| |[sparse_flash_mla](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md)|基于共享KV完成SparseFlashMla稀疏注意力计算。|默认确定性实现|默认确定性实现| -|[sparse_flash_mla_metadata](../../torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md)|生成SparseFlashMla主算子使用的任务切分metadata。|默认支持确定性计算;默认支持batch invariance。| |[inplace_partial_rotary_mul](../../torch_extension/cann_ops_transformer/docs/zh/inplace_partial_rotary_mul.md)|执行单路旋转位置编码的Inplace计算,直接修改输入张量,不产生新的输出张量。|默认确定性实现|默认确定性实现| -|[compressor](../../torch_extension/cann_ops_transformer/docs/zh/compressor.md)|将每4或128个token的KV cache压缩成一个,然后每个token与这些压缩的KV cache进行DSA计算。|默认支持确定性计算。| \ No newline at end of file +|[compressor](../../torch_extension/cann_ops_transformer/docs/zh/compressor.md)|将每4或128个token的KV cache压缩成一个,然后每个token与这些压缩的KV cache进行DSA计算。|默认支持确定性计算。| +|[get_low_latency_ccl_buffer_size](../../torch_extension/cann_ops_transformer/docs/zh/get_low_latency_ccl_buffer_size.md)|计算low_latency_dispatch/low_latency_combine所需的HCCL通信buffer_size(单位MB),为MoeDistributeBuffer的静态方法,可在初始化前调用。|默认支持确定性计算|-| +|[low_latency_dispatch](../../torch_extension/cann_ops_transformer/docs/zh/low_latency_dispatch.md)|完成MoE并行部署下token的低时延dispatch分发,支持动态量化与EP域alltoallv通信,需与low_latency_combine配套使用。|默认支持确定性计算|-| +|[low_latency_combine](../../torch_extension/cann_ops_transformer/docs/zh/low_latency_combine.md)|与low_latency_dispatch配套,按dispatch原路返回完成token的低时延combine反向聚合(乘路由权重再相加)。|默认支持确定性计算|-| +|[mega_moe](../../torch_extension/cann_ops_transformer/docs/zh/mega_moe.md)|MoE端到端通算融合算子,将Dispatch+GroupMatmul1+SwiGLUQuant+GroupMatmul2+Combine融合为单算子;配套get_mega_moe_ccl_buffer_size、get_symm_buffer_for_mega_moe使用。|-|默认支持确定性计算| diff --git a/torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md b/torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md index 5cf2ca3db1..d1ff0f8f58 100644 --- a/torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md +++ b/torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md @@ -1,4 +1,4 @@ -# sparse_flash_mla / sparse_flash_mla_metadata +# sparse_flash_mla ## 产品支持情况 @@ -19,7 +19,7 @@ - **C4A(Compressed Sparse Attention,CSA)**:同时使用`ori_kv`、`cmp_kv`和`cmp_sparse_indices`,对原始KV窗口和TopK选择出的压缩KV共同做注意力。 - **C128A(Heavily Compressed Attention,HCA)**:同时使用`ori_kv`和`cmp_kv`,对原始KV窗口和连续压缩KV段共同做注意力。 - `sparse_flash_mla_metadata`是`SparseFlashMlaMetadata`的torch扩展接口,用于在主算子执行前生成metadata。metadata记录AICore/AIVCore的任务切分结果,主算子必须传入该metadata。典型调用流程如下: + `sparse_flash_mla_metadata`是`sparse_flash_mla`的metadata前置接口,用于在主接口执行前生成metadata。metadata记录AICore/AIVCore的任务切分结果,主接口必须传入该metadata。典型调用流程如下: 1. 准备`q`、`ori_kv`、`cmp_kv`、序列长度、block table、sinks等输入。 2. 调用`sparse_flash_mla_metadata`生成`metadata`。 @@ -40,7 +40,7 @@ ## 函数原型 ```python -cann_ops_transformer.ops.sparse_flash_mla( +cann_ops_transformer.sparse_flash_mla( q, *, ori_kv=None, @@ -74,7 +74,7 @@ cann_ops_transformer.ops.sparse_flash_mla( ``` ```python -cann_ops_transformer.ops.sparse_flash_mla_metadata( +cann_ops_transformer.sparse_flash_mla_metadata( num_heads_q, num_heads_kv, head_dim, @@ -253,7 +253,7 @@ q = torch.randn(B, S1, N1, D, dtype=dtype, device="npu") ori_kv = torch.randn(B, S2, N2, D, dtype=dtype, device="npu") sinks = torch.zeros(N1, dtype=torch.float32, device="npu") -metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( +metadata = cann_ops_transformer.sparse_flash_mla_metadata( N1, N2, D, @@ -264,7 +264,7 @@ metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( cmp_topk=0, cmp_ratio=cmp_ratio, ori_mask_mode=4, - cmp_mask_mode=0, + cmp_mask_mode=3, ori_win_left=127, ori_win_right=0, layout_q="BSND", @@ -273,7 +273,7 @@ metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( has_cmp_kv=False, ) -attn_out, softmax_lse = cann_ops_transformer.ops.sparse_flash_mla( +attn_out, softmax_lse = cann_ops_transformer.sparse_flash_mla( q, ori_kv=ori_kv, sinks=sinks, @@ -281,7 +281,7 @@ attn_out, softmax_lse = cann_ops_transformer.ops.sparse_flash_mla( softmax_scale=1.0 / math.sqrt(D), cmp_ratio=cmp_ratio, ori_mask_mode=4, - cmp_mask_mode=0, + cmp_mask_mode=3, ori_win_left=127, ori_win_right=0, layout_q="BSND", @@ -321,7 +321,7 @@ cmp_kv = torch.randn(B, S3, N2, D, dtype=dtype, device="npu") cmp_residual_kv = torch.zeros(B, dtype=torch.int32, device="npu") sinks = torch.zeros(N1, dtype=torch.float32, device="npu") -metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( +metadata = cann_ops_transformer.sparse_flash_mla_metadata( N1, N2, D, @@ -343,7 +343,7 @@ metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( has_cmp_kv=True, ) -attn_out, softmax_lse = cann_ops_transformer.ops.sparse_flash_mla( +attn_out, softmax_lse = cann_ops_transformer.sparse_flash_mla( q, ori_kv=ori_kv, cmp_kv=cmp_kv, @@ -401,7 +401,7 @@ sinks = torch.zeros(N1, dtype=torch.float32, device="npu") cmp_sparse_indices = torch.full((sum(q_lens), N2, K), -1, dtype=torch.int32, device="npu") cmp_sparse_indices[:, :, :1] = torch.arange(1, dtype=torch.int32, device="npu").view(1, 1, 1) -metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( +metadata = cann_ops_transformer.sparse_flash_mla_metadata( N1, N2, D, @@ -416,6 +416,7 @@ metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( cmp_ratio=cmp_ratio, ori_mask_mode=4, cmp_mask_mode=3, + cmp_residual_kv=cmp_residual_kv, ori_win_left=127, ori_win_right=0, layout_q="TND", @@ -424,7 +425,7 @@ metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( has_cmp_kv=True, ) -attn_out, softmax_lse = cann_ops_transformer.ops.sparse_flash_mla( +attn_out, softmax_lse = cann_ops_transformer.sparse_flash_mla( q, ori_kv=ori_kv, cmp_kv=cmp_kv, @@ -522,7 +523,7 @@ def run_one_rank(name, q, ori_kv, cmp_kv, q_lens, ori_prefix_lens, cmp_lens, res cmp_residual_kv = torch.tensor(residuals, dtype=torch.int32, device="npu") cmp_sparse_indices = make_cmp_sparse_indices(q_lens, ori_prefix_lens, cmp_lens) - metadata = cann_ops_transformer.ops.sparse_flash_mla_metadata( + metadata = cann_ops_transformer.sparse_flash_mla_metadata( N1, N2, D, @@ -537,6 +538,7 @@ def run_one_rank(name, q, ori_kv, cmp_kv, q_lens, ori_prefix_lens, cmp_lens, res cmp_ratio=cmp_ratio, ori_mask_mode=4, cmp_mask_mode=3, + cmp_residual_kv=cmp_residual_kv, ori_win_left=127, ori_win_right=0, layout_q="TND", @@ -545,7 +547,7 @@ def run_one_rank(name, q, ori_kv, cmp_kv, q_lens, ori_prefix_lens, cmp_lens, res has_cmp_kv=True, ) - attn_out, softmax_lse = cann_ops_transformer.ops.sparse_flash_mla( + attn_out, softmax_lse = cann_ops_transformer.sparse_flash_mla( q, ori_kv=ori_kv, cmp_kv=cmp_kv,