文件最后提交记录最后更新时间
12 天前
1 个月前
7 天前
3 天前
3 天前
14 天前
1 个月前
12 天前
README

QuantSparseFlashMla

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品 ×
Atlas A2 训练系列产品/Atlas A2 推理系列产品 ×
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×

功能说明

  • 算子功能:

    QuantSparseFlashMla算子旨在完成全量化和稀疏场景下的MLA(Multi-head Latent Attention)注意力计算。该接口支持以下三类计算模式:

    • SWA(Sliding Window Attention):仅传入ori_kv,对原始KV做滑动窗口注意力。
    • CSA(Compressed Sparse Attention):同时传入ori_kvcmp_kvcmp_sparse_indices,对原始KV窗口和topK选择出的压缩KV共同做注意力。
    • HCA(Heavily Compressed Attention):同时传入ori_kvcmp_kv,对原始KV窗口和连续压缩KV段共同做注意力。

    QuantSparseFlashMlaMetadataQuantSparseFlashMla的分核信息,在主算子执行前生成。当前版本主算子必须传入该metadata。典型调用流程如下:

    1. 根据调用场景准备对应输入。
    2. 调用QuantSparseFlashMlaMetadata生成metadata,作为QuantSparseFlashMla的入参。
    3. 调用QuantSparseFlashMla,生成计算结果。
  • 计算公式: QuantSparseFlashMla采用MLA(Multi-head Latent Attention)对KV共享输入的稀疏注意力进行计算,其原理是对输入的KV进行选择性压缩与量化处理,再将Query与拼接后的KV计算结果通过Softmax得到注意力权重。

    MLA的计算公式一般定义如下,其中K~=V~\tilde{K}=\tilde{V}为基于入参控制的实际参与计算的KV,由ori_kv的滑动窗口部分和cmp_kv的压缩部分共同组成,实际参与计算的KV范围由cmp_ratioori_mask_modecmp_mask_modeori_win_leftori_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。 HIFLOAT8 ND
ori_kv 可选输入 对应公式中K和V的一部分,表示原始不经压缩的量化KV。 HIFLOAT8 ND
cmp_kv 可选输入 对应公式中K和V的一部分,表示经过压缩的量化KV。 HIFLOAT8 ND
q_descale 输入 q对应的量化参数。 FLOAT ND
ori_kv_descale 可选输入 ori_kv对应的量化参数。 FLOAT ND
cmp_kv_descale 可选输入 cmp_kv对应的量化参数。 FLOAT ND
ori_sparse_indices 可选输入 表示从ori_kv中离散取数的索引,无效位置填-1。 INT32 ND
cmp_sparse_indices 可选输入 表示从cmp_kv中离散取数的索引,无效位置填-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 可选输入 QuantSparseFlashMlaMetadata生成的任务切分结果。 INT32 ND
quant_mode 属性 表示量化模式。量化模式1表示Q、K、V 为per-token量化,Q、K、V 数据类型为HIFLOAT8。 INT -
softmax_scale 可选属性 对应公式中的softmax_scale。 FLOAT -
cmp_ratio 可选属性 表示cmp_kv相对于压缩前KV长度的压缩倍率,用于恢复cmp侧mask使用的压缩前KV长度。默认值为1。 INT -
ori_mask_mode 可选属性 表示q和ori_kv计算的mask模式。
0: No Mask。
3: RightDownCausal模式。
4: Band模式。
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索引取值模式。 INT -
return_softmax_lse 可选属性 表示是否返回softmax_lse。默认值False。 BOOL -
attn_out 输出 对应公式中的输出O。 BFLOAT16 ND
softmax_lse 可选输出 返回softmax的log-sum-exp结果。 FLOAT ND

约束说明

  • 该接口支持推理场景下使用。

  • 该接口支持aclgraph模式。

  • 该接口当前支持三种计算场景:SWA(Sliding Window Attention)场景仅传入ori_kv;CSA(Compressed Sparse Attention)场景同时传入ori_kvcmp_kvcmp_sparse_indices;HCA(Heavily Compressed Attention)场景同时传入ori_kvcmp_kv

  • 通用规格约束如下:

    • N2仅支持1,D仅支持512。
    • 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支持-1或非负数,ori_win_right支持-1或非负数,只有ori_mask_mode为4时,ori_win_leftori_win_right可以>=0。
    • layout_qlayout_kv组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下layout_qlayout_kv必须一致;PA_BBND场景下block_size支持1到1024。
    • 全平台均不支持传入非空Tensor。
  • layout_q为TND时,功能使用限制如下:

    • q的shape需要为[T1,N1,D]。
    • ori_sparse_indices的shape维度为[Q_T, KV_N, K1],其中K1为对ori_kv一次离散选取的token数。
    • cmp_sparse_indices的shape需要为[Q_T, KV_N, K2],其中K2为对cmp_kv一次离散选取的token数。
    • cu_seqlens_q必须传入,输入维度为B+1,大小为参数中每个元素的值表示当前batch与之前所有batch的token数总和,即前缀和,因此后一个元素的值必须>=前一个元素的值且首元素必须为0。
  • layout_q为BSND时,功能使用限制如下:

    • q的shape需要为[B, Q_S, Q_N, D]。
    • ori_sparse_indices的shape需要为[B, Q_S, KV_N, K1],其中K1为对ori_kv一次离散选取的token数。
    • cmp_sparse_indices的shape需要为[B, Q_S, KV_N, K2],其中K2为对cmp_kv一次离散选取的token数。
  • PageAttention场景下,功能使用限制如下:

    • ori_kvcmp_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取值为1到1024,KV_N仅支持1。
    • ori_block_tablecmp_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",需满足上述layout_qlayout_kv组合约束。

    • 当输入为PA_BBND时,seqused_ori_kvori_block_table必须传入;当输入为BSND时,seqused_ori_kv可用于表达每个batch的ori_kv有效长度;当输入为TND时,ori_kv有效长度由cu_seqlens_ori_kv表达。
    • 当输入为BSND时,ori_kvcmp_kv的layout都必须为BSND,ori_kv的shape为[B, S2, N2,D],cmp_kv的shape为[B, S3, N2,D]。
    • 当输入为TND时,cu_seqlens_ori_kv必须传入;若存在cmp_kvcu_seqlens_cmp_kv也必须传入。
  • return_softmax_lse为False时返回占位Tensor;为True时返回softmax的log-sum-exp结果。

  • 传入Tensor不支持为空。

  • seqused_cmp_kv为所有layout_kv下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度;未传时由cmp_kv shape、cu_seqlens_cmp_kv或PA block table相关语义推导。

  • cmp_residual_kv为算子的可选入参;传入后用于按cmp_len * cmp_ratio + residual恢复cmp侧mask使用的压缩前KV长度,其中cmp_len优先来自显式传入的seqused_cmp_kv

  • qori_kvcmp_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中的输入样本序列长度的累加和。

调用说明

调用方式 样例代码 说明
aclnn API test_aclnn_quant_sparse_flash_mla 通过aclnnQuantSparseFlashMla调用QuantSparseFlashMla算子
PyTorch API quant_sparse_flash_mla 通过cann_ops_transformer.quant_sparse_flash_mla调用QuantSparseFlashMla算子