MixedQuantSparseFlashMla
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | × |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | × |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
功能说明
-
算子功能:
MixedQuantSparseFlashMla算子旨在完成量化和稀疏场景下的MLA(Multi-head Latent Attention)注意力计算,支持SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)三类Attention计算场景。与SparseFlashMla的区别在于,本算子支持KV的per-token-group量化输入。该算子的三种典型场景:- SWA(Sliding Window Attention):仅传入
ori_kv,对原始KV做滑动窗口注意力。 - CSA(Compressed Sparse Attention):同时传入
ori_kv、cmp_kv和cmp_sparse_indices,对原始KV窗口和topK选择出的压缩KV共同做注意力。 - HCA(Heavily Compressed Attention):同时传入
ori_kv和cmp_kv,对原始KV窗口和连续压缩KV段共同做注意力。
调用时需要使用
MixedQuantSparseFlashMlaMetadata生成的任务列表metadata,在主算子执行前生成,当前版本主算子必须传入该metadata。典型调用流程如下:- 根据调用场景准备
q、ori_kv、cmp_kv等对应输入。 - 调用
MixedQuantSparseFlashMlaMetadata生成metadata,作为MixedQuantSparseFlashMla的入参。 - 调用
MixedQuantSparseFlashMla,将上一步得到的metadata传入主算子,生成计算结果。
- SWA(Sliding Window Attention):仅传入
-
计算公式:
MixedQuantSparseFlashMla采用MLA对KV共享输入的稀疏注意力进行计算,其原理是对输入的KV进行选择性压缩与量化处理,再将Query与拼接后的KV计算结果通过Softmax得到注意力权重。MLA的计算公式一般定义如下,其中K~=V~\tilde{K}=\tilde{V}为基于入参控制的实际参与计算的KV,由
ori_kv的滑动窗口部分和cmp_kv的压缩部分共同组成,实际参与计算的KV范围由cmp_ratio、ori_mask_mode、cmp_mask_mode、ori_win_left、ori_win_right以及cmp_sparse_indices决定。O=softmax(Q@K~T⋅softmax_scale)@V~O = \text{softmax}(Q@\tilde{K}^T \cdot \text{softmax\_scale})@\tilde{V}
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| q | 输入 | 表示对应公式中的Q。 | BFLOAT16 | ND |
| ori_kv | 可选输入 | 表示对应公式中K和V的一部分,为原始不经压缩的量化KV,Key和Value共享同一份数据。由nope、rope、scale、padding拼接而成,详见quant_mode。 | 详见quant_mode | ND |
| cmp_kv | 可选输入 | 表示对应公式中K和V的一部分,为经过压缩的量化KV,Key和Value共享同一份数据。由nope、rope、scale、padding拼接而成,详见quant_mode。 | 详见quant_mode | ND |
| ori_sparse_indices | 可选输入 | 表示原始KV topK索引,无效位置填-1。 | INT32 | ND |
| cmp_sparse_indices | 可选输入 | 表示压缩KV topK索引,无效位置填-1。 | INT32 | ND |
| ori_block_table | 可选输入 | 表示PageAttention场景下ori_kv使用的block映射表。 | INT32 | ND |
| cmp_block_table | 可选输入 | 表示PageAttention场景下cmp_kv使用的block映射表。 | INT32 | ND |
| cu_seqlens_q | 可选输入 | 表示TND布局下不同batch中q的累积序列长度。 | INT32 | ND |
| cu_seqlens_ori_kv | 可选输入 | 表示TND布局下不同batch中ori_kv的累积序列长度。 | INT32 | ND |
| cu_seqlens_cmp_kv | 可选输入 | 表示TND布局下不同batch中cmp_kv的累积序列长度。 | INT32 | ND |
| seqused_q | 可选输入 | 表示不同batch中q实际参与计算的token数。 | INT32 | ND |
| seqused_ori_kv | 可选输入 | 表示不同batch中ori_kv实际参与计算的token数。 | INT32 | ND |
| seqused_cmp_kv | 可选输入 | 表示不同batch中cmp_kv实际参与计算的token数。 | INT32 | ND |
| cmp_residual_kv | 可选输入 | 表示压缩KV余数,用于恢复cmp侧mask使用的压缩前KV长度。 | INT32 | ND |
| ori_topk_length | 可选输入 | 表示ori_sparse_indices实际参与计算的长度。 | INT32 | ND |
| cmp_topk_length | 可选输入 | 表示cmp_sparse_indices实际参与计算的长度。 | INT32 | ND |
| sinks | 可选输入 | 表示各注意力头设置独立可学习虚拟偏移项,用于维持长文本推理时的稳定性。 | FLOAT | ND |
| metadata | 可选输入 | 表示MixedQuantSparseFlashMlaMetadata生成的分核信息。 | INT32 | ND |
| quant_mode | 属性 | 表示量化模式。量化模式1表示K、V nope为per-token-group量化,K、V依次由rope(64,bfloat16)、nope(448,FLOAT8_E4M3FN)、scale(7,bfloat16)、pad(18B)拼接而成;量化模式2表示K、V nope为per-token-group量化,K、V依次由nope(448,FLOAT8_E4M3FN)、rope(64,bfloat16)、scale(7,FLOAT8_E8M0)、pad(1B)拼接而成。当前仅支持1和2,量化模式2仅支持layout_kv为PA_BBND。 | INT | - |
| rope_head_dim | 可选属性 | 表示rope头的维度,仅支持64。 | INT | - |
| softmax_scale | 可选属性 | 表示对应公式中的softmax_scale,默认值为1.0。 | FLOAT | - |
| cmp_ratio | 可选属性 | 表示cmp_kv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度;仅传入ori_kv时不参与压缩KV计算。支持1到128。 | INT | - |
| ori_mask_mode | 可选属性 | 表示q和ori_kv计算的mask模式。 0: No mask。 3: rightDownCausal模式。 4: sliding window模式。 |
INT | - |
| cmp_mask_mode | 可选属性 | 表示q和cmp_kv计算的mask模式。 0: No mask。 3: rightDownCausal模式。 |
INT | - |
| ori_win_left | 可选属性 | 表示q和ori_kv计算中q对历史token计算的数量,-1表示无穷大,即全部参与运算。默认值为-1。 | INT | - |
| ori_win_right | 可选属性 | 表示q和ori_kv计算中q对未来token计算的数量,-1表示无穷大,即全部参与运算。默认值为-1。 | INT | - |
| layout_q | 可选属性 | 表示输入q的数据排布格式,支持"BSND"和"TND"。 | STRING | - |
| layout_kv | 可选属性 | 表示输入ori_kv和cmp_kv的数据排布格式,支持"BSND"、"TND"和"PA_BBND"。 | STRING | - |
| topk_value_mode | 可选属性 | 表示topK索引取值模式,默认值为1。 | INT | - |
| return_softmax_lse | 可选属性 | 表示是否返回softmax的lse结果,默认值为False。 | BOOL | - |
| attn_out | 输出 | 表示对应公式中的输出O。 | BFLOAT16 | ND |
| softmax_lse | 可选输出 | 表示对query乘key的结果先取max得到softmax_max,query乘key的结果减去softmax_max后取exp再取sum得到softmax_sum,最后对softmax_sum取log再加上softmax_max得到的结果。 | FLOAT | ND |
约束说明
- 该接口支持推理场景下使用。
- 该接口支持aclgraph模式。
- 该接口当前支持三种计算场景: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。
常见字段释义
| 命名 | 含义 |
|---|---|
| b | 输入样本batch大小 |
| q_s | 输入q的序列长度 |
| ori_kv_s | 输入ori_kv的序列长度 |
| cmp_kv_s | 输入cmp_kv的序列长度 |
| q_n | 输入q的头数 |
| kv_n | 输入ori_kv/cmp_kv的头数 |
| q_d | 输入q的注意力头的维度 |
| kv_d | 输入ori_kv/cmp_kv的注意力头的维度 |
| q_t | 输入q所有batch序列长度的累加和 |
| ori_kv_t | 输入ori_kv所有batch序列长度的累加和 |
| cmp_kv_t | 输入cmp_kv所有batch序列长度的累加和 |
| ori_kv_k | 输入ori_sparse_indices中topK选出的token个数 |
| cmp_kv_k | 输入cmp_sparse_indices中topK选出的token个数 |
| ori_kv_s_max | 输入ori_kv的最大序列长度 |
| cmp_kv_s_max | 输入cmp_kv的最大序列长度 |
| ori_kv_block_size | 输入ori_kv在PagedAttention场景下的block大小 |
| cmp_kv_block_size | 输入cmp_kv在PagedAttention场景下的block大小 |
| ori_kv_block_nums | 输入ori_kv在PagedAttention场景下的block数量 |
| cmp_kv_block_nums | 输入cmp_kv在PagedAttention场景下的block数量 |
- 通用规格约束如下:
- kv_n仅支持1,q_d仅支持512。其中,
ori_kv和cmp_kv的kv_d由nope、rope、scale、padding拼接而成,详见quant_mode。 cmp_ratio表示cmp_kv相对于压缩前KV长度的压缩倍率;仅传入ori_kv时,cmp_ratio不参与压缩KV计算,需保持默认值1;支持1到128。ori_mask_mode支持0、3和4,cmp_mask_mode支持0和3,ori_win_left和ori_win_right支持-1或非负数,-1表示对应方向不受限。rope_head_dim仅支持64。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。
- kv_n仅支持1,q_d仅支持512。其中,
- 当
layout_q为TND时,功能使用限制如下:q的shape需要为[q_t, q_n, q_d]。ori_sparse_indices的shape需要为[q_t, kv_n, ori_kv_k]。cmp_sparse_indices的shape需要为[q_t, kv_n, cmp_kv_k]。cu_seqlens_q必须传入,输入维度为b+1,每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须>=前一个元素的值且首元素必须为0。
- 当
layout_q为BSND时,功能使用限制如下:q的shape需要为[b, q_s, q_n, q_d]。cmp_sparse_indices的shape需要为[b, q_s, kv_n, ori_kv_k]。cmp_sparse_indices的shape需要为[b, q_s, kv_n, cmp_kv_k]。
- PageAttention场景下,功能使用限制如下:
ori_kv和cmp_kv的shape分别为[ori_kv_block_nums, ori_kv_block_size, kv_n, kv_d]和[cmp_kv_block_nums, cmp_kv_block_size, kv_n, kv_d],其中ori_kv_block_nums和cmp_kv_block_nums为PagedAttention场景下的block数量,ori_kv_block_size和cmp_kv_block_size为一个block的token数,取值为1到1024。ori_block_table和cmp_block_table的shape为2维,其中第一维长度为b,第二维长度不小于所有batch中最大的ori_kv_s和cmp_kv_s对应的block数量,即ori_kv_s_max / ori_kv_block_size和cmp_kv_s_max / cmp_kv_block_size向上取整。
metadata为算子实际需要使用的分核结果,目前该参数必传,shape大小固定为[1024]。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, ori_kv_s, kv_n, kv_d],cmp_kv的shape为[b, cmp_kv_s, kv_n, kv_d]。 - 当输入为TND时,
cu_seqlens_ori_kv必须传入;若存在cmp_kv,cu_seqlens_cmp_kv也必须传入。
- 当输入为PA_BBND时,
return_softmax_lse为False时返回占位Tensor;为True时返回softmax的log-sum-exp结果。- 除
ori_topk_length和cmp_topk_length等预留输入可不传或传入空Tensor外,其余已传入Tensor不支持为空。 seqused_cmp_kv为所有layout_kv下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由cmp_kvshape、cu_seqlens_cmp_kv或PA block table相关语义推导。cmp_residual_kv为算子的可选入参;传入后用于按cmp_len * cmp_ratio + residual恢复cmp侧mask使用的压缩前KV长度,其中cmp_len优先来自显式传入的seqused_cmp_kv。
调用说明
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn API | test_aclnn_mixed_quant_sparse_flash_mla | 通过aclnnMixedQuantSparseFlashMla调用MixedQuantSparseFlashMla算子 |
| PyTorch API | mixed_quant_sparse_flash_mla | 通过cann_ops_transformer.mixed_quant_sparse_flash_mla调用MixedQuantSparseFlashMla算子 |