已合并
smla/qsmla/mqsmla maskmode泛化资料修改 #10901
smla/qsmla/mqsmla maskmode泛化资料修改 #10901
已合并
jinying创建于 9 天前
9 个文件变更+8-424
@@ -318,7 +318,7 @@
318- 通用规格约束如下:318- 通用规格约束如下:
319 - kv_n仅支持1,q_d仅支持512。其中,`ori_kv``cmp_kv`的kv_d由nope、rope、scale、padding拼接而成,详见`quant_mode`。319 - kv_n仅支持1,q_d仅支持512。其中,`ori_kv``cmp_kv`的kv_d由nope、rope、scale、padding拼接而成,详见`quant_mode`。
320 - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时,`cmp_ratio`不参与压缩KV计算,需保持默认值1;支持1到128。320 - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时,`cmp_ratio`不参与压缩KV计算,需保持默认值1;支持1到128。
321- - `ori_mask_mode`支持0、3和4,`cmp_mask_mode`支持0和3,`ori_win_left`和`ori_win_right`支持-1或非负数,-1表示对应方向不受限。321+ - `ori_mask_mode`支持0、3和4,`cmp_mask_mode`支持0和3,`ori_win_left`和`ori_win_right`支持-1或非负数,-1表示对应方向不受限,只有`ori_mask_mode`为4时,`ori_win_left`和`ori_win_right`可以>=0
322 - `rope_head_dim`仅支持64。322 - `rope_head_dim`仅支持64。
323 - `layout_q``layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致;PA_BBND场景下`block_size`支持1到1024。323 - `layout_q``layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致;PA_BBND场景下`block_size`支持1到1024。
324-`layout_q`为TND时,功能使用限制如下:324-`layout_q`为TND时,功能使用限制如下:
@@ -920,7 +920,6 @@ metadataOptional校验
920 </td>920 </td>
921 <td>921 <td>
922 <ul>922 <ul>
923- <li>oriKvOptional稀疏场景下,cmpMaskMode为0和oriMaskMode必须为0</li>
924 <li>oriMaskMode支持0、3、4</li>923 <li>oriMaskMode支持0、3、4</li>
925 </ul>924 </ul>
926 </td>925 </td>
@@ -942,10 +941,7 @@ metadataOptional校验
942 </ul>941 </ul>
943 </td>942 </td>
944 <td>943 <td>
945- <li>当oriKvOptional/cmpKvOptional/cmpSparseIndicesOptional/oriSparseIndicesOptional传入时,cmpMaskMode为0和oriMaskMode必须为0</li>944+ <li>cmpKvOptional传入时,cmpMaskMode必须为0 </li>
946- <li>当cmpKvOptional不传时,oriMaskMode为3、4</li>
947- <li>当oriMaskMode为3时,cmpMaskMode必须为3</li>
948- <li>当oriMaskMode为4时,cmpMaskMode必须为3</li>
949 </td>945 </td>
950 </tr>946 </tr>
951 <tr>947 <tr>
@@ -304,7 +304,7 @@
304- 通用规格约束如下:304- 通用规格约束如下:
305 - N2仅支持1,D仅支持512。305 - N2仅支持1,D仅支持512。
306 - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时,`cmp_ratio`不参与压缩KV计算,需保持默认值1;支持1到128。306 - `cmp_ratio`表示`cmp_kv`相对于压缩前KV长度的压缩倍率;仅传入`ori_kv`时,`cmp_ratio`不参与压缩KV计算,需保持默认值1;支持1到128。
307- - `ori_mask_mode`支持0/3/4,`cmp_mask_mode`支持0/3,`ori_win_left`支持-1或非负数,`ori_win_right`支持-1或非负数。307+ - `ori_mask_mode`支持0/3/4,`cmp_mask_mode`支持0/3,`ori_win_left`支持-1或非负数,`ori_win_right`支持-1或非负数,只有`ori_mask_mode`为4时,`ori_win_left`和`ori_win_right`可以>=0
308 - `layout_q``layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致;PA_BBND场景下`block_size`支持1到1024。308 - `layout_q``layout_kv`组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下`layout_q`和`layout_kv`必须一致;PA_BBND场景下`block_size`支持1到1024。
309 - 全平台均不支持传入非空Tensor。309 - 全平台均不支持传入非空Tensor。
310 310 
@@ -979,7 +979,6 @@ metadataOptional校验
979 </td>979 </td>
980 <td>980 <td>
981 <ul>981 <ul>
982- <li>oriKvOptional稀疏场景下,cmpMaskMode为0和oriMaskMode必须为0</li>
983 <li>oriMaskMode支持0、3、4</li>982 <li>oriMaskMode支持0、3、4</li>
984 </ul>983 </ul>
985 </td>984 </td>
@@ -1000,7 +999,6 @@ metadataOptional校验
1000 </td>999 </td>
1001 <td>1000 <td>
1002 <ul>1001 <ul>
1003- <li>当oriKvOptional/cmpKvOptional/cmpSparseIndicesOptional/oriSparseIndicesOptional传入时,cmpMaskMode为0和oriMaskMode必须为0</li>
1004 <li>cmpKv未传入时,cmpMaskMode必须为0 </li>1002 <li>cmpKv未传入时,cmpMaskMode必须为0 </li>
1005 </ul>1003 </ul>
1006 </td>1004 </td>
@@ -282,7 +282,7 @@
282 - SWA稀疏ori_kv场景下,`ori_topk_length`必须传入,配套Metadata接口的`ori_topk`为`ori_sparse_indices`最后一维K,且`ori_topk_length`的元素取值应在[0, K]范围内;其他场景`ori_topk_length`传入nullptr或空Tensor。282 - SWA稀疏ori_kv场景下,`ori_topk_length`必须传入,配套Metadata接口的`ori_topk`为`ori_sparse_indices`最后一维K,且`ori_topk_length`的元素取值应在[0, K]范围内;其他场景`ori_topk_length`传入nullptr或空Tensor。
283- 产品型号约束如下:283- 产品型号约束如下:
284 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term><term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:Q\_N支持1、2、4、8、16、32、64、128,KV\_N只支持1;cmp_ratio在SWA场景保持默认值1,CSA支持传入4,HCA支持传入128;block_size取值为16的倍数,最大支持1024;SWA稀疏ori_kv场景支持`ori_sparse_indices`和`ori_topk_length`,`ori_mask_mode`为0,`ori_win_left`和`ori_win_right`为非负数;非SWA稀疏ori_kv场景的`ori_mask_mode`为4、`ori_win_left`为127、`ori_win_right`为0,`cmp_sparse_indices`的最后一维K2当前支持512或1024,`cmp_mask_mode`仅支持3。284 - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term><term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:Q\_N支持1、2、4、8、16、32、64、128,KV\_N只支持1;cmp_ratio在SWA场景保持默认值1,CSA支持传入4,HCA支持传入128;block_size取值为16的倍数,最大支持1024;SWA稀疏ori_kv场景支持`ori_sparse_indices`和`ori_topk_length`,`ori_mask_mode`为0,`ori_win_left`和`ori_win_right`为非负数;非SWA稀疏ori_kv场景的`ori_mask_mode`为4、`ori_win_left`为127、`ori_win_right`为0,`cmp_sparse_indices`的最后一维K2当前支持512或1024,`cmp_mask_mode`仅支持3。
285- - <term>Ascend 950PR/Ascend 950DT</term>:Q\_N支持2、4、8、16、32、64、128,不支持1,KV\_N只支持1。`ori_mask_mode`支持0、3、4,`cmp_mask_mode`支持0、3;`ori_mask_mode`为4时,`ori_win_left`和`ori_win_right`支持-1或非负数,-1表示对应方向不受限。285+ - <term>Ascend 950PR/Ascend 950DT</term>:Q\_N支持1-128,KV\_N只支持1。`ori_mask_mode`支持0、3、4,`cmp_mask_mode`支持0、3;`ori_win_left`和`ori_win_right`支持-1或非负数,-1表示对应方向不受限。只有`ori_mask_mode`为4时,`ori_win_left`和`ori_win_right`可以>=0。
286 286 
287-`layout_q`为TND时,功能使用限制如下:287-`layout_q`为TND时,功能使用限制如下:
288 - `q`的shape需要为[Q\_T, Q\_N, D]。288 - `q`的shape需要为[Q\_T, Q\_N, D]。
@@ -304,7 +304,7 @@
304 - 当输入为BSND时,`ori_kv``cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, ORI\_KV\_S, KV\_N, D],cmp_kv的shape为[B, CMP\_KV\_S, KV\_N, D]。304 - 当输入为BSND时,`ori_kv``cmp_kv`的layout都必须为BSND,ori_kv的shape为[B, ORI\_KV\_S, KV\_N, D],cmp_kv的shape为[B, CMP\_KV\_S, KV\_N, D]。
305 - 当输入为TND时,`cu_seqlens_ori_kv`必须传入;若存在`cmp_kv``cu_seqlens_cmp_kv`也必须传入。305 - 当输入为TND时,`cu_seqlens_ori_kv`必须传入;若存在`cmp_kv``cu_seqlens_cmp_kv`也必须传入。
306- `return_softmax_lse`为False时返回占位Tensor;为True时返回softmax的log-sum-exp结果。306- `return_softmax_lse`为False时返回占位Tensor;为True时返回softmax的log-sum-exp结果。
307-- SWA稀疏ori_kv场景仅支持SWA模板,仅传入`ori_kv`,必须同时传入`ori_sparse_indices`和`ori_topk_length`,并设置`ori_mask_mode`为0、`ori_win_left`和`ori_win_right`为非负数;配套Metadata接口的`ori_topk`为K,该场景不传入`cmp_kv`。307+- SWA稀疏ori_kv场景仅支持SWA模板,仅传入`ori_kv`,必须同时传入`ori_sparse_indices`和`ori_topk_length`,`ori_win_left`和`ori_win_right`仅在`ori_mask_mode`4时取非负数;配套Metadata接口的`ori_topk`为K,该场景不传入`cmp_kv`。
308- `ori_topk_length`表示每个q token和KV head的实际有效索引条目数,取值应在[0, K]范围内。`ori_sparse_indices`的[0, ori_topk_length)区间为左对齐的有效索引条目,[ori_topk_length, K)区间为无效或填充条目,建议填-1。308- `ori_topk_length`表示每个q token和KV head的实际有效索引条目数,取值应在[0, K]范围内。`ori_sparse_indices`的[0, ori_topk_length)区间为左对齐的有效索引条目,[ori_topk_length, K)区间为无效或填充条目,建议填-1。
309-`cmp_topk_length`等预留输入可不传或传入空Tensor外,其余已传入Tensor不支持为空。309-`cmp_topk_length`等预留输入可不传或传入空Tensor外,其余已传入Tensor不支持为空。
310- `seqused_cmp_kv`为所有`layout_kv`下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由`cmp_kv` shape、`cu_seqlens_cmp_kv`或PA block table相关语义推导。310- `seqused_cmp_kv`为所有`layout_kv`下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由`cmp_kv` shape、`cu_seqlens_cmp_kv`或PA block table相关语义推导。
@@ -412,7 +412,7 @@ aclnnStatus aclnnSparseFlashMla(
412 <td>cmpMaskMode(int64_t)</td>412 <td>cmpMaskMode(int64_t)</td>
413 <td>输入</td>413 <td>输入</td>
414 <td>q和cmpKv计算的mask模式。</td>414 <td>q和cmpKv计算的mask模式。</td>
415- <td>0: No Mask。<br/>3: RightDownCausal模式。cmpKv未传入时该参数取默认值1。</td>415+ <td>0: No Mask。<br/>3: RightDownCausal模式。cmpKv未传入时该参数取默认值0。</td>
416 <td>-</td>416 <td>-</td>
417 <td>-</td>417 <td>-</td>
418 <td>-</td>418 <td>-</td>
@@ -689,7 +689,7 @@ aclnnStatus aclnnSparseFlashMla(
689 689 
690- 使用约束690- 使用约束
691 691 
692- - SWA稀疏ori_kv场景仅支持SWA模板,仅传入`oriKvOptional`,同时传入`oriSparseIndicesOptional`和`oriTopkLengthOptional`;配套Metadata接口的`oriTopk`为oriSparseIndicesOptional最后一维K,`oriMaskMode`必须为0,`oriWinLeft`和`oriWinRight`为非负数,且`cmpKvOptional`不传入。692+ - SWA稀疏ori_kv场景仅支持SWA模板,仅传入`oriKvOptional`,同时传入`oriSparseIndicesOptional`和`oriTopkLengthOptional`;配套Metadata接口的`oriTopk`为oriSparseIndicesOptional最后一维K,`oriMaskMode`为0/3/4,`oriWinLeft`和`oriWinRight`仅在`oriMaskMode`4时取非负数,且`cmpKvOptional`不传入。
693 - SWA稀疏ori_kv场景下,`oriTopkLengthOptional`的元素表示实际有效索引条目数,取值应在[0, K]范围内。对每个q token和KV head,`oriSparseIndicesOptional`的[0, oriTopkLengthOptional)区间为左对齐的有效索引条目,[oriTopkLengthOptional, K)区间为无效或填充条目,建议填-1;其他场景`oriTopkLengthOptional`传入nullptr或空Tensor。693 - SWA稀疏ori_kv场景下,`oriTopkLengthOptional`的元素表示实际有效索引条目数,取值应在[0, K]范围内。对每个q token和KV head,`oriSparseIndicesOptional`的[0, oriTopkLengthOptional)区间为左对齐的有效索引条目,[oriTopkLengthOptional, K)区间为无效或填充条目,建议填-1;其他场景`oriTopkLengthOptional`传入nullptr或空Tensor。
694 -`cmpTopkLengthOptional`等预留输入可传入nullptr或空Tensor外,其余已传入Tensor不支持为空。694 -`cmpTopkLengthOptional`等预留输入可传入nullptr或空Tensor外,其余已传入Tensor不支持为空。
695 - `metadataOptional`参数必须传入,由`aclnnSparseFlashMlaMetadata`算子生成,shape固定为(1024,)。695 - `metadataOptional`参数必须传入,由`aclnnSparseFlashMlaMetadata`算子生成,shape固定为(1024,)。
@@ -1,404 +0,0 @@
1-# Sparse Flash MLA 系列算子拦截说明
2- 
3-## 1. 文档范围
4- 
5-本文档说明以下三个算子新增的 Host 侧参数拦截代码及当前实际拦截链路:
6- 
7-- `sparse_flash_mla`
8-- `mixed_quant_sparse_flash_mla`
9-- `quant_sparse_flash_mla`
10- 
11-拦截规则参考:
12- 
13-- `torch_extension/cann_ops_transformer/docs/zh/sparse_flash_mla.md`
14-- `torch_extension/cann_ops_transformer/docs/zh/mixed_quant_sparse_flash_mla.md`
15-- `torch_extension/cann_ops_transformer/docs/zh/quant_sparse_flash_mla.md`
16- 
17-代码采用与 `flash_attn/op_host/checkers` 相同的分层形式,将检查过程划分为:
18- 
19-1. 单参数检查(`CheckSinglePara`
20-2. 参数存在性检查(`CheckParaExistence`
21-3. 特性交叉检查(`CheckFeature`
22-4. 多参数一致性检查(`CheckMultiPara`
23- 
24-三个算子的公共检查代码统一放在本目录;混合量化和全量化算子的目录只保存各自的适配入口和差异规则。
25-当前新增 Checker 代码均予以保留并参与编译;`mixed_quant_sparse_flash_mla`
26-`quant_sparse_flash_mla`已使用新增 Checker;`sparse_flash_mla`按架构分流,DAV_2201
27-(Atlas A2/A3)使用原有旧 Checker,DAV_3510(Atlas A5)使用新增 Checker。
28- 
29-## 2. 代码结构
30- 
31-```text
32-attention/
33-├── sparse_flash_mla/op_host/checkers/
34-│ ├── base_checker.{h,cpp} # Checker 基类及公共检查工具
35-│ ├── checker_context.h # 三算子统一检查上下文
36-│ ├── checker_adapter.h # TilingInfo 到统一上下文的适配
37-│ ├── checker_runner.{h,cpp} # 四阶段检查编排
38-│ ├── common_checker.{h,cpp} # Q/KV/输出、布局和公共属性
39-│ ├── seq_len_checker.{h,cpp} # 序列长度类 Tensor
40-│ ├── sparse_compression_checker.{h,cpp} # 稀疏索引、TopK 和压缩参数
41-│ ├── mask_checker.{h,cpp} # Mask 与窗口联动
42-│ ├── paged_attention_checker.{h,cpp} # Paged Attention
43-│ ├── sinks_checker.{h,cpp} # Sinks
44-│ ├── metadata_checker.{h,cpp} # Metadata
45-│ ├── softmax_lse_checker.{h,cpp} # Softmax LSE 输出
46-│ ├── sparse_flash_mla_checker.{h,cpp} # sparse_flash_mla 入口
47-│ └── checker_sources.cmake # 公共源码一次性注册
48-├── mixed_quant_sparse_flash_mla/op_host/checkers/
49-│ ├── mixed_quant_variant_checker.{h,cpp} # 混合量化特有规则
50-│ └── mixed_quant_sparse_flash_mla_checker.{h,cpp}
51-└── quant_sparse_flash_mla/op_host/checkers/
52- ├── quant_variant_checker.{h,cpp} # 全量化及 descale 规则
53- └── quant_sparse_flash_mla_checker.{h,cpp}
54-```
55- 
56-`RegisterCommonCheckers` 按以下顺序注册公共 Checker:
57- 
58-```text
59-Common
60- → SeqLen
61- → SparseCompression
62- → Mask
63- → PagedAttention
64- → Sinks
65- → Metadata
66- → SoftmaxLse
67-```
68- 
69-混合量化和全量化入口会在公共 Checker 后追加各自的差异 Checker;当前 M/Q 均已启用该执行链。
70-`sparse_flash_mla`在 DAV_3510 上执行上述公共 Checker,DAV_2201 继续执行原有旧 Checker。
71- 
72-任一阶段、任一 Checker 返回失败后立即停止,Tiling 返回 `GRAPH_FAILED`
73- 
74-## 3. 新增 Checker 公共拦截
75- 
76-### 3.1 Tensor 通用要求
77- 
78-| 检查项 | 拦截规则 |
79-| --- | --- |
80-| 输入/输出存在性 | `q``ori_kv``attention_out``sinks``metadata`必须存在;`cmp_kv``softmax_lse`可选 |
81-| 数据格式 | 所有被检查的 Tensor 仅支持 `ND` |
82-| 空 Tensor | Q、KV、输出、索引、长度、Block Table、Sinks 等任一维度小于等于0时拦截 |
83-| Q 布局 | 仅支持 `BSND``TND` |
84-| KV 布局 | 仅支持 `BSND``TND``PA_BBND` |
85-| 布局组合 | 非 `PA_BBND` 时,`layout_q``layout_kv`必须相同;PA 时 Q 可为 `BSND``TND` |
86-| Q/输出一致性 | `attention_out`的 rank 和 shape 必须与`q`完全相同 |
87-| Q 轴范围 | `q_n``[1, 128]`内,Q head dim 固定为512,batch/sequence/token维均大于0 |
88-| KV 轴范围 | `kv_n`固定为1,序列、token、block 数和 block size 均必须有效 |
89-| Softmax scale | `softmax_scale`必须为有限值,拒绝 NaN 和 Inf |
90- 
91-### 3.2 三算子数据类型和 Head Dim 差异
92- 
93-| 算子 | q | ori_kv/cmp_kv | attention_out | KV head dim |
94-| --- | --- | --- | --- | --- |
95-| `sparse_flash_mla` | FP16/BF16 | FP16/BF16 | FP16/BF16 | 512 |
96-| `mixed_quant_sparse_flash_mla` | BF16 | FP8_E4M3FN | BF16 | `quant_mode=1`时608;`quant_mode=2`时584 |
97-| `quant_sparse_flash_mla` | HIFLOAT8 | HIFLOAT8 | BF16 | 512 |
98- 
99-`sparse_flash_mla`还要求`q`、所有非空 KV 和`attention_out`的数据类型完全一致。
100- 
101-在 Atlas A2/A3 上,`sparse_flash_mla``q_n`仅支持:
102- 
103-```text
104-1, 2, 4, 8, 16, 32, 64, 128
105-```
106- 
107-Ascend 950 上只检查`1 <= q_n <= 128`
108- 
109-### 3.3 序列长度 Tensor
110- 
111-以下 Tensor 均要求 `int32``ND`、一维且非空:
112- 
113-- `cu_seqlens_q`
114-- `cu_seqlens_ori_kv`
115-- `cu_seqlens_cmp_kv`
116-- `seqused_q`
117-- `seqused_ori_kv`
118-- `seqused_cmp_kv`
119-- `cmp_residual_kv`
120- 
121-存在性和 shape 规则如下:
122- 
123-| 参数 | 存在性 | shape |
124-| --- | --- | --- |
125-| `cu_seqlens_q` | Q 为 TND 时必传;其他布局禁止传入 | `(b+1,)` |
126-| `cu_seqlens_ori_kv` | KV 为 TND 时必传;其他布局禁止传入 | `(b+1,)` |
127-| `cu_seqlens_cmp_kv` | KV 为 TND且`cmp_kv`存在时必传;无`cmp_kv`时禁止传入 | `(b+1,)` |
128-| `seqused_q` | 可选 | `(b,)` |
129-| `seqused_ori_kv` | PA 场景通常必传;`ORI_SPARSE`中ori侧`mask_mode=0`,或`ORI_CMP_SPARSE`中两侧`mask_mode=0`,且传入`ori_topk_length`时可不传 | `(b,)` |
130-| `seqused_cmp_kv` | PA 且`cmp_kv`存在时通常必传;仅`ORI_CMP_SPARSE`中两侧`mask_mode=0`且传入`cmp_topk_length`时可不传 | `(b,)` |
131-| `cmp_residual_kv` | 由压缩模式和 Mask 联动决定 | `(b,)` |
132- 
133-Tensor 内部数值(例如累积长度单调性、首尾值及 residual 范围)在 Tiling 阶段不可读取,由调用方保证。
134- 
135-### 3.4 稀疏索引、TopK Length 与压缩属性
136- 
137-公共规则:
138- 
139-- `cmp_ratio > 0`
140-- SWA(未传`cmp_kv`)要求`cmp_ratio=1`
141-- `topk_value_mode`当前只支持1。
142-- `cmp_sparse_indices``cmp_topk_length``cmp_residual_kv`以及 cmp 侧长度/Block Table 均依赖`cmp_kv`
143-- 稀疏索引仅支持 `int32``ND`
144- - BSND:`(b, q_s, kv_n, topk)`
145- - TND:`(q_t, kv_n, topk)`
146-- TopK Length 仅支持 `int32``ND`
147- - BSND:`(b, q_s, kv_n)`
148- - TND:`(q_t, kv_n)`
149- 
150-三个算子统一采用以下成对规则:
151- 
152-- `ori_mask_mode=0`且传入`ori_sparse_indices`时,必须传`ori_topk_length`;其他情况禁止传`ori_topk_length`
153-- `cmp_mask_mode=0`且传入`cmp_sparse_indices`时,必须传`cmp_topk_length`;其他情况禁止传`cmp_topk_length`
154- 
155-全稀疏分为两种模式:
156- 
157-- `ORI_SPARSE`:存在`ori_sparse_indices`且不存在`cmp_kv`,只检查ori侧稀疏参数;ori侧`mask_mode=0`时,
158- `ori_topk_length`必须传入,并可替代`seqused_ori_kv`
159-- `ORI_CMP_SPARSE``ori_sparse_indices``cmp_kv``cmp_sparse_indices`同时存在,检查ori、cmp两侧
160- 稀疏参数;两侧`mask_mode=0`时,`ori_topk_length``cmp_topk_length`必须传入,并可分别替代
161- `seqused_ori_kv``seqused_cmp_kv`
162- 
163-因此,`sparse_flash_mla`与另外两个算子一样,允许在`ori_mask_mode=0`时传入`ori_sparse_indices`及配套的`ori_topk_length`
164- 
165-### 3.5 Mask 和窗口
166- 
167-| 参数 | 支持范围 |
168-| --- | --- |
169-| `ori_mask_mode` | 0、3、4 |
170-| `cmp_mask_mode` | 0、3 |
171-| `ori_win_left``ori_win_right` | -1或非负数 |
172- 
173-交叉规则:
174- 
175-- 非 Sliding Window(`ori_mask_mode != 4`)要求`ori_win_left=ori_win_right=-1`
176-- `ori_mask_mode=4`时才允许使用非负窗口值。
177-- SWA 要求`cmp_mask_mode=0`
178-- `sparse_flash_mla`的 HCA/CSA 要求`cmp_mask_mode=3`
179-- 混合量化算子的 Causal/Sliding Window 模式(`ori_mask_mode=3/4`)在存在`cmp_kv`时要求`cmp_mask_mode=3`
180-- 三个算子传入`ori_sparse_indices`时,均要求`ori_mask_mode=cmp_mask_mode=0`
181-- Atlas A2/A3 上的`sparse_flash_mla`固定要求`ori_win_left=127``ori_win_right=0`
182- 
183-### 3.6 Paged Attention
184- 
185-| 检查项 | 拦截规则 |
186-| --- | --- |
187-| Block Table dtype/format | `int32``ND` |
188-| Block Table rank | 2 |
189-| Block Table shape | 第一维必须等于 batch,第二维必须大于0 |
190-| 非 PA 布局 | 禁止传入`ori_block_table``cmp_block_table` |
191-| PA 布局 | 必须传`ori_block_table``cmp_block_table``cmp_kv`同时存在或同时不存在 |
192-| Seqused 联动 | Block Table 通常要求对应`seqused_*_kv``ORI_SPARSE`可用`ori_topk_length`替代ori侧,`ORI_CMP_SPARSE`可用对应`topk_length`分别替代两侧,且需满足各模式的`mask_mode=0`约束 |
193-| KV block size | 范围 `[1, 1024]` |
194-| A2/A3 sparse block size | 额外要求16对齐 |
195- 
196-Block Table 内的具体 block id 值由调用方保证。
197- 
198-### 3.7 Sinks、Metadata 和 Softmax LSE
199- 
200-| 参数 | 拦截规则 |
201-| --- | --- |
202-| `sinks` | 当前版本必传;`float32``ND`、shape为`(q_n,)` |
203-| `metadata` | 当前版本必传;`int32``ND`、shape为`(1024,)` |
204-| `softmax_lse` | 可存在或不存在;仅在`return_softmax_lse=true`且存在时要求`float32``ND` |
205- 
206-`return_softmax_lse=false``softmax_lse`不存在时,跳过该参数的所有检查。仅在
207-`return_softmax_lse=true``softmax_lse`存在时检查以下 shape:
208- 
209-| 条件 | shape |
210-| --- | --- |
211-| `return_softmax_lse=true`且 Q 为 BSND | `(b, kv_n, q_s, q_n/kv_n)` |
212-| `return_softmax_lse=true`且 Q 为 TND | `(kv_n, q_t, q_n/kv_n)` |
213- 
214-启用并传入`softmax_lse`时,同时检查`q_n`可被`kv_n`整除。
215- 
216-## 4. 逐参数校验明细
217- 
218-下表中的适用范围使用以下缩写:
219- 
220-- S:`sparse_flash_mla`
221-- M:`mixed_quant_sparse_flash_mla`
222-- Q:`quant_sparse_flash_mla`
223- 
224-“检查内容”描述新增 Checker 已实现的 Host 侧检查规则;当前 M/Q 和 DAV_3510 上的 S 已实际执行
225-这些规则,DAV_2201 上的 S 仍执行旧 Checker。Tensor 内部元素值无法在 Tiling 阶段读取的,会
226-单独标明“值由用户保证”。
227- 
228-### 4.1 核心输入和量化参数
229- 
230-| 参数 | 适用 | 存在性 | dtype、format、rank/shape | 一致性及交叉检查 |
231-| --- | --- | --- | --- | --- |
232-| `q` | S/M/Q | 必传;desc 和 shape 均不可为空 | S:FP16/BF16;M:BF16;Q:HIFLOAT8。仅支持ND;BSND为4维,TND为3维;任一维度必须大于0 | `q_d=512``1 <= q_n <= 128`;S在A2/A3上`q_n`仅支持1、2、4、8、16、32、64、128;shape必须与`attn_out`相同;布局与`layout_q`一致 |
233-| `ori_kv` | S/M/Q | 当前版本必传 | S:FP16/BF16;M:FP8_E4M3FN;Q:HIFLOAT8。仅支持ND;TND为3维,BSND/PA_BBND为4维;任一维度必须大于0 | `kv_n=1`;S/Q的head dim为512;M在`quant_mode=1/2`时分别为608/584;BSND的batch必须等于Q的batch;PA block size在`[1,1024]`内,S在A2/A3上还要求16对齐 |
234-| `cmp_kv` | S/M/Q | 可选 | 存在时执行与`ori_kv`相同的dtype、ND、rank、非空和轴检查 | 不存在时禁止传入所有cmp侧从属Tensor,并要求`cmp_ratio=1``cmp_mask_mode=0`;存在时BSND batch必须与Q一致;S要求其dtype与`q``ori_kv`一致 |
235-| `q_descale` | Q | 当前版本必传 | `float32`、ND、shape为`(1,)` | 仅在`quant_mode=1`下支持;当前Q算子只支持该模式 |
236-| `ori_kv_descale` | Q | 当前版本必传 | `float32`、ND、shape为`(1,)` | 与`ori_kv`配套使用 |
237-| `cmp_kv_descale` | Q | `cmp_kv`存在时必传;`cmp_kv`不存在时禁止传入 | `float32`、ND、shape为`(1,)` | 存在状态必须与`cmp_kv`完全一致 |
238-| `sinks` | S/M/Q | 当前版本必传 | `float32`、ND、1维、非空 | shape必须为`(q_n,)`,长度等于Q头数 |
239-| `metadata` | S/M/Q | 当前版本必传 | `int32`、ND、shape为`(1024,)` | Checker只能检查描述信息,无法验证其内容是否由本次调用的同一组参数生成 |
240- 
241-### 4.2 稀疏压缩参数
242- 
243-| 参数 | 适用 | 单参数及shape检查 | 存在性和交叉检查 | 无法检查的内容 |
244-| --- | --- | --- | --- | --- |
245-| `ori_sparse_indices` | S/M/Q | `int32`、ND、非空;BSND为`(b,q_s,kv_n,topk)`,TND为`(q_t,kv_n,topk)` | 三个算子均可选;传入时要求`ori_mask_mode=cmp_mask_mode=0`,且必须按规则传入`ori_topk_length` | 每个元素是否为-1或合法ori token索引 |
246-| `cmp_sparse_indices` | S/M/Q | `int32`、ND、非空;BSND为`(b,q_s,kv_n,topk)`,TND为`(q_t,kv_n,topk)`;最后一维必须大于0 | 依赖`cmp_kv`,无`cmp_kv`时禁止传入;S中传入后识别为CSA,A2/A3上TopK仅支持512或1024;M/Q在`cmp_mask_mode=0`时要求`cmp_topk_length` | 每个元素是否为-1或合法cmp token索引 |
247-| `ori_topk_length` | S/M/Q | `int32`、ND、非空;BSND为`(b,q_s,kv_n)`,TND为`(q_t,kv_n)` | 三个算子规则相同:仅当`ori_sparse_indices`存在且`ori_mask_mode=0`时必传,其他情况禁止传入 | 每个位置的TopK长度是否小于等于索引最后一维 |
248-| `cmp_topk_length` | S/M/Q | `int32`、ND、非空;BSND为`(b,q_s,kv_n)`,TND为`(q_t,kv_n)` | 三个算子规则相同:仅当`cmp_kv``cmp_sparse_indices`存在且`cmp_mask_mode=0`时必传,其他情况禁止传入 | 每个位置的TopK长度是否有效 |
249-| `cmp_residual_kv` | S/M/Q | `int32`、ND、1维、非空,shape为`(b,)` | 无`cmp_kv`时禁止传入;S的HCA/CSA必传;M/Q在`cmp_mask_mode=3``cmp_ratio!=1`时必传 | 每个元素是否位于`[0,cmp_ratio)`,以及压缩前后长度恢复关系 |
250- 
251-### 4.3 序列长度参数
252- 
253-| 参数 | 适用 | 单参数及shape检查 | 存在性和布局联动 | 无法检查的内容 |
254-| --- | --- | --- | --- | --- |
255-| `cu_seqlens_q` | S/M/Q | `int32`、ND、1维、非空,shape为`(b+1,)` | `layout_q=TND`时必传;其他Q布局禁止传入 | 首元素是否为0、是否单调非递减、末元素是否等于`q_t` |
256-| `cu_seqlens_ori_kv` | S/M/Q | `int32`、ND、1维、非空,shape为`(b+1,)` | `layout_kv=TND`时必传;其他KV布局禁止传入 | 首元素、单调性、末元素和ori KV实际长度 |
257-| `cu_seqlens_cmp_kv` | S/M/Q | `int32`、ND、1维、非空,shape为`(b+1,)` | `layout_kv=TND``cmp_kv`存在时必传;其他情况禁止传入 | 首元素、单调性、末元素和cmp KV实际长度 |
258-| `seqused_q` | S/M/Q | 存在时要求`int32`、ND、1维、非空,shape为`(b,)` | 可选,不作为PA必选参数 | 每个元素是否非负且不超过对应Q长度 |
259-| `seqused_ori_kv` | S/M/Q | 存在时要求`int32`、ND、1维、非空,shape为`(b,)` | BSND等非PA场景可选;PA场景通常必传;`ORI_SPARSE``ORI_CMP_SPARSE`满足对应`mask_mode=0`约束并传入`ori_topk_length`时可不传 | 每个元素是否非负且不超过对应ori KV长度 |
260-| `seqused_cmp_kv` | S/M/Q | 存在时要求`int32`、ND、1维、非空,shape为`(b,)` | 无`cmp_kv`时禁止传入;PA且`cmp_kv`存在时通常必传;仅`ORI_CMP_SPARSE`中两侧`mask_mode=0`且传入`cmp_topk_length`时可不传 | 每个元素是否非负且不超过对应cmp KV长度 |
261- 
262-### 4.4 Paged Attention 参数
263- 
264-| 参数 | 适用 | 单参数及shape检查 | 存在性和交叉检查 | 无法检查的内容 |
265-| --- | --- | --- | --- | --- |
266-| `ori_block_table` | S/M/Q | `int32`、ND、2维、非空;第一维必须等于batch | `layout_kv=PA_BBND`时必传,非PA禁止传入;通常要求`seqused_ori_kv``ORI_SPARSE``ORI_CMP_SPARSE`满足对应`mask_mode=0`约束时可用`ori_topk_length`替代 | block id是否为正整数、是否越界,以及第二维是否覆盖全部有效KV |
267-| `cmp_block_table` | S/M/Q | `int32`、ND、2维、非空;第一维必须等于batch | 仅在`layout_kv=PA_BBND``cmp_kv`存在时必传;其存在状态必须与`cmp_kv`一致;通常要求`seqused_cmp_kv`,仅`ORI_CMP_SPARSE`中两侧`mask_mode=0`时可用`cmp_topk_length`替代 | block id是否合法,以及第二维是否覆盖全部有效cmp KV |
268- 
269-`ori_kv``cmp_kv`自身的PA block size检查包含:范围`[1,1024]`;S在Atlas A2/A3上额外要求16对齐。
270- 
271-### 4.5 属性参数
272- 
273-| 参数 | 适用 | 校验内容 |
274-| --- | --- | --- |
275-| `softmax_scale` | S/M/Q | 必须是有限浮点数;NaN、正负Inf均拦截 |
276-| `cmp_ratio` | S/M/Q | 必须大于0;无`cmp_kv`的SWA固定为1;S在A2/A3上的CSA固定为4、HCA固定为128;与`cmp_residual_kv`的逐元素数值关系由用户保证 |
277-| `ori_mask_mode` | S/M/Q | 仅支持0、3、4;决定窗口和稀疏索引联动;M中存在`cmp_kv`且取3/4时要求`cmp_mask_mode=3`;存在`ori_sparse_indices`时必须为0 |
278-| `cmp_mask_mode` | S/M/Q | 仅支持0、3;SWA固定为0;S不含`ori_sparse_indices`的HCA/CSA固定为3;三个算子存在`ori_sparse_indices`时必须为0;影响TopK Length和Residual的存在性 |
279-| `ori_win_left` | S/M/Q | 只能为-1或非负数;非Sliding Window要求为-1;S在A2/A3上固定为127 |
280-| `ori_win_right` | S/M/Q | 只能为-1或非负数;非Sliding Window要求为-1;S在A2/A3上固定为0 |
281-| `layout_q` | S/M/Q | 仅支持`BSND``TND`;决定Q、稀疏索引、TopK Length和LSE的rank/shape,以及`cu_seqlens_q`存在性 |
282-| `layout_kv` | S/M/Q | 仅支持`BSND``TND``PA_BBND`;非PA时必须与`layout_q`相同;决定KV rank、KV长度Tensor和Block Table存在性 |
283-| `topk_value_mode` | S/M/Q | 当前仅支持1,其他值直接拦截 |
284-| `return_softmax_lse` | S/M/Q | bool类型由算子Schema保证;为`false`时完全跳过`softmax_lse`检查;为`true`且输出存在时检查有效LSE dtype、format和shape;输出不存在时不拦截 |
285-| `quant_mode` | M/Q | M必须存在且仅支持1、2,并决定KV head dim为608或584;Q仅支持1 |
286-| `rope_head_dim` | M | 仅支持64 |
287- 
288-### 4.6 输出参数
289- 
290-| 参数 | 适用 | 存在性 | dtype、format、rank/shape | 一致性检查 |
291-| --- | --- | --- | --- | --- |
292-| `attn_out`(文档中也称`attention_out`) | S/M/Q | 必须存在 | S:FP16/BF16;M/Q:BF16;仅支持ND;rank随`layout_q`为4或3;任一维度必须大于0 | shape必须与`q`完全一致;S还要求dtype与`q`一致 |
293-| `softmax_lse` | S/M/Q | 可选;存在或不存在均合法 | `return_softmax_lse=false`时不检查;开启且存在时要求`float32`、ND,BSND shape为`(b,kv_n,q_s,q_n/kv_n)`,TND shape为`(kv_n,q_t,q_n/kv_n)` | 关闭或不存在时跳过全部检查;存在且开启时要求`q_n`能被`kv_n`整除,所有相关轴与Q/KV一致;输出内容不在Host侧检查 |
294- 
295-### 4.7 前置 Metadata 接口参数说明
296- 
297-本文新增 Checker 设计为运行在三个主算子的 Tiling 入口,启用后仅接收前置Metadata接口生成的
298-`metadata` Tensor;当前 M/Q 和 DAV_3510 上的 S 已接入实际入口,DAV_2201 上的 S 仍走旧
299-Checker。因此:
300- 
301-- `*_sparse_flash_mla_metadata`接口中的`num_heads_q``num_heads_kv``head_dim``batch_size``max_seqlen_*``has_ori_kv``has_cmp_kv`等参数,不会在主算子 Checker 中再次逐项读取。
302-- 主算子 Checker只验证`metadata``int32`、ND、shape `(1024,)`
303-- Metadata是否由与主算子完全一致的布局、序列长度、压缩率、Mask、窗口和KV存在状态生成,只能由调用方保证。
304- 
305-## 5. 算子特有拦截
306- 
307-### 5.1 sparse_flash_mla
308- 
309-DAV_2201(arch22/Atlas A2/A3)由原有`SMLATilingCheck`执行;DAV_3510(Atlas A5)由新增
310-`SparseFlashMlaChecker`执行。
311- 
312-按输入组合识别计算模式:
313- 
314-| 模式 | 输入组合 | 公共约束 | Atlas A2/A3附加约束 |
315-| --- | --- | --- | --- |
316-| SWA | 仅`ori_kv` | `cmp_ratio=1``cmp_mask_mode=0` | `cmp_ratio=1` |
317-| ORI_SPARSE | `ori_kv + ori_sparse_indices + ori_topk_length` | `ori_mask_mode=cmp_mask_mode=0`、`cmp_ratio=1` | `cmp_ratio=1` |
318-| HCA | `ori_kv + cmp_kv`,无`cmp_sparse_indices` | `cmp_mask_mode=3`、必须有`cmp_residual_kv` | `cmp_ratio=128` |
319-| CSA | `ori_kv + cmp_kv + cmp_sparse_indices` | `cmp_mask_mode=3`、必须有`cmp_residual_kv` | `cmp_ratio=4`,TopK仅支持512或1024 |
320-| ORI_CMP_SPARSE | `ori_kv + ori_sparse_indices + ori_topk_length + cmp_kv + cmp_sparse_indices + cmp_topk_length` | `ori_mask_mode=cmp_mask_mode=0` | `cmp_ratio=4`,cmp TopK仅支持512或1024 |
321- 
322-Ascend 950 的 HCA/CSA 仅要求`cmp_ratio > 0`,CSA TopK 要求大于0。
323- 
324-### 5.2 mixed_quant_sparse_flash_mla
325- 
326-当前由新增`MixedQuantSparseFlashMlaChecker`执行。
327- 
328-特有属性:
329- 
330-- `quant_mode`仅支持1或2。
331-- `rope_head_dim`仅支持64。
332-- Q/输出仅支持 BF16。
333-- KV 仅支持 FP8_E4M3FN。
334-- `quant_mode=1`时 KV head dim 为608。
335-- `quant_mode=2`时 KV head dim 为584。
336--`cmp_mask_mode=3``cmp_ratio != 1`时,必须传入`cmp_residual_kv`
337- 
338-### 5.3 quant_sparse_flash_mla
339- 
340-当前由新增`QuantSparseFlashMlaChecker`执行。
341- 
342-特有属性和 Tensor:
343- 
344-- `quant_mode`仅支持1。
345-- Q、ori KV、cmp KV 仅支持 HIFLOAT8。
346-- 输出仅支持 BF16。
347-- `q_descale`当前版本必传。
348-- `ori_kv_descale`当前版本必传。
349-- `cmp_kv_descale``cmp_kv`同时存在或同时不存在。
350-- 三个 descale Tensor 均要求`float32``ND`、shape为`(1,)`
351--`cmp_mask_mode=3``cmp_ratio != 1`时,必须传入`cmp_residual_kv`
352- 
353-## 6. 新旧拦截链路关系
354- 
355-当前拦截入口如下:
356- 
357-| 算子/架构 | 实际拦截类 |
358-| --- | --- |
359-| `sparse_flash_mla` / DAV_2201(Atlas A2/A3) | 原有`SMLATilingCheck` |
360-| `sparse_flash_mla` / DAV_3510(Atlas A5) | 新增`SparseFlashMlaChecker` |
361-| `mixed_quant_sparse_flash_mla` | 新增`MixedQuantSparseFlashMlaChecker` |
362-| `quant_sparse_flash_mla` | 新增`QuantSparseFlashMlaChecker` |
363- 
364-SMLA 解析完成后根据`npuArch`实例化 Checker:DAV_2201 使用旧 Checker,DAV_3510 使用新增
365-Checker;MQSMLA 和 QSMLA 继续实例化新增 Checker。CMake 不额外添加 M/Q 的旧
366-`*_check*.cpp`,新增 Checker 中不再保留 arch22/A2-A3 专用规则,原`sparse_variant_checker`
367-也不再编译。
368- 
369-解析器仍负责从 Tiling Context 获取 Tensor、属性、布局、shape 和硬件信息;解析完成后执行对应 Checker。
370- 
371-## 7. 检查边界
372- 
373-Host Tiling 阶段只检查可获得的描述信息和属性值,包括:
374- 
375-- Tensor 是否存在
376-- dtype、format、rank、shape
377-- 布局和硬件平台
378-- 标量属性范围
379-- 参数之间的存在性和形状联动
380- 
381-以下内容无法在当前阶段完整读取或验证,由调用方保证:
382- 
383-- `cu_seqlens_*`的首元素、末元素及单调性
384-- `seqused_*`的逐元素范围
385-- `cmp_residual_kv`每个元素是否位于`[0, cmp_ratio)`
386-- 稀疏索引是否为-1或合法 token 索引
387-- TopK Length 的逐元素范围
388-- Block Table 内的 block id 是否有效
389-- `metadata`是否由完全一致的前置参数生成
390- 
391-## 8. 拦截日志规范
392- 
393-三个算子的新增 Checker 禁止使用通用的 `OP_LOGE`,错误日志按照失败原因选择结构化宏:
394- 
395-| 错误类别 | 使用的日志宏 |
396-| --- | --- |
397-| 参数缺失、参数不应存在或参数存在性联动错误 | `OP_LOGE_FOR_INVALID_ARGUMENT_WITH_REASON` |
398-| 单个/多个属性值或取值范围错误 | `OP_LOGE_FOR_INVALID_VALUE*``OP_LOGE_FOR_INVALID_VALUES_WITH_REASON` |
399-| dtype 错误或多个 Tensor dtype 不一致 | `OP_LOGE_FOR_INVALID_DTYPE*``OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON` |
400-| Tensor format 错误 | `OP_LOGE_FOR_INVALID_FORMAT` |
401-| dim num(rank)错误 | `OP_LOGE_FOR_INVALID_SHAPEDIM` |
402-| shape、轴长度或 shape size 错误 | `OP_LOGE_FOR_INVALID_SHAPE*``OP_LOGE_FOR_INVALID_SHAPES*``OP_LOGE_FOR_INVALID_SHAPESIZE*` |
403- 
404-日志同时提供参数名、实际值/形状以及正确值或失败原因,便于直接定位拦截条件;所有失败原因和说明文本均以大写字母开头。
@@ -490,7 +490,6 @@ metadata校验
490 </td>490 </td>
491 <td>491 <td>
492 <ul>492 <ul>
493- <li>只有ori_kv稀疏场景下,cmp_mask_mode为0和ori_mask_mode必须为0</li>
494 <li>SWA场景下,ori_mask_mode为0、3、4</li>493 <li>SWA场景下,ori_mask_mode为0、3、4</li>
495 </ul>494 </ul>
496 </td>495 </td>
@@ -513,10 +512,7 @@ metadata校验
513 </td>512 </td>
514 <td>513 <td>
515 <ul>514 <ul>
516- <li>当ori_kv/cmp_kv/cmp_sparse_indices/ori_sparse_indices传入时,cmp_mask_mode为0和ori_mask_mode必须为0</li>515+ <li>cmpKvOptional未传入时,cmpMaskMode必须为0 </li>
517- <li>当cmp_kv不传时,ori_mask_mode为3、4</li>
518- <li>当ori_mask_mode为3时,cmp_mask_mode必须为3</li>
519- <li>当ori_mask_mode为4时,cmp_mask_mode必须为3</li>
520 </ul>516 </ul>
521 </td>517 </td>
522 </tr>518 </tr>
@@ -527,7 +527,6 @@ metadata校验
527 </td>527 </td>
528 <td>528 <td>
529 <ul>529 <ul>
530- <li>只有ori_kv稀疏场景下,cmp_mask_mode为0和ori_mask_mode必须为0</li>
531 <li>SWA场景下,ori_mask_mode为0、3、4</li>530 <li>SWA场景下,ori_mask_mode为0、3、4</li>
532 </ul>531 </ul>
533 </td>532 </td>
@@ -550,7 +549,6 @@ metadata校验
550 </td>549 </td>
551 <td>550 <td>
552 <ul>551 <ul>
553- <li>当ori_kv/cmp_kv/cmp_sparse_indices/ori_sparse_indices传入时,cmp_mask_mode为0和ori_mask_mode必须为0</li>
554 <li>SWA场景下cmp_mask_mode必须为0 </li>552 <li>SWA场景下cmp_mask_mode必须为0 </li>
555 </ul>553 </ul>
556 </td>554 </td>