文件最后提交记录最后更新时间
10 天前
10 天前
17 天前
17 天前
17 天前
10 天前
README

SparseFlashMlaGradMetadata

产品支持情况

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

功能说明

  • 算子功能:SparseFlashMlaGradMetadata算子旨在生成一个任务列表,包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及Q和K的分块的索引,供后续SparseFlashMlaGrad算子使用。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
cu_seqlens_q 可选输入 表示不同Batch中Query的有效Sequence Length,shape为(B+1, ),仅layout_q为TND场景需传入。 INT32 ND
cu_seqlens_ori_kv 可选输入 表示不同Batch中ori_kv的有效Sequence Length,shape为(B+1, ),仅layout_kv为TND场景需传入。 INT32 ND
cu_seqlens_cmp_kv 可选输入 表示不同Batch中cmp_kv的有效Sequence Length,shape为(B+1, ),仅layout_kv为TND场景需传入。 INT32 ND
seqused_q 可选输入 表示不同Batch中Query实际参与运算的Sequence Length,shape为(B, )。 INT32 ND
seqused_ori_kv 可选输入 表示不同Batch中ori_kv实际参与运算的Sequence Length,shape为(B, )。 INT32 ND
seqused_cmp_kv 可选输入 表示不同Batch中cmp_kv实际参与运算的Sequence Length,shape为(B, )。 INT32 ND
cmp_residual_kv 可选输入 表示不同Batch中cmp_kv压缩后Sequence Length的余数,配合cmp_ratio实现cmp_kv部分的mask和负载计算。cmp_mask_mode=3且cmp_ratio≠1时必须传入,shape为(B, )。 INT32 ND
ori_topk_length 可选输入 表示不同q token对应的ori_kv部分关键稀疏token的个数,shape为(B, S1, N2)或(T1, N2)。 INT32 ND
cmp_topk_length 可选输入 表示不同q token对应的cmp_kv部分关键稀疏token的个数,shape为(B, S1, N2)或(T1, N2)。 INT32 ND
num_heads_q 属性 表示Query的head个数,当前支持[1, 128]。 INT32 -
num_heads_kv 属性 表示Key和Value对应的多头数,当前仅支持1。 INT32 -
head_dim 属性 表示注意力头的维度,当前仅支持512。 INT32 -
batch_size 可选属性 表示Batch数量,默认值为0。 INT32 -
max_seqlen_q 可选属性 表示Query的最长Sequence Length,默认值为0。 INT32 -
max_seqlen_ori_kv 可选属性 表示ori_kv的最长Sequence Length,默认值为0。 INT32 -
max_seqlen_cmp_kv 可选属性 表示cmp_kv的最长Sequence Length,默认值为0。 INT32 -
ori_topk 可选属性 表示ori_kv中筛选出的关键稀疏token的个数,0表示非稀疏场景,默认值为0。 INT32 -
cmp_topk 可选属性 表示cmp_kv中筛选出的关键稀疏token的个数,0表示非稀疏场景,默认值为0。 INT32 -
cmp_ratio 可选属性 表示对cmp_kv的压缩率,默认值为1,当前支持[1, 128]。 INT32 -
ori_mask_mode 可选属性 表示q和ori_kv计算的mask模式,0表示No mask,3表示rightDownCausal模式,4表示sliding window模式,默认值为0。 INT32 -
cmp_mask_mode 可选属性 表示q和cmp_kv计算的mask模式,0表示No mask,3表示rightDownCausal模式,默认值为0。 INT32 -
ori_win_left 可选属性 表示q和ori_kv计算中q对过去token计算的数量,-1表示无穷大,默认值为-1。 INT32 -
ori_win_right 可选属性 表示q和ori_kv计算中q对未来token计算的数量,-1表示无穷大,默认值为-1。 INT32 -
layout_q 可选属性 表示Query的排列格式,支持BSND、TND,默认值为BSND。 STRING -
layout_kv 可选属性 表示Key的排列格式,支持BSND、TND,默认值为BSND。 STRING -
has_ori_kv 可选属性 用于标识是否含有ori_kv,默认值为true。 BOOL -
has_cmp_kv 可选属性 用于标识是否含有cmp_kv,默认值为true。 BOOL -
metadata 输出 表示负载均衡结果输出,shape固定为[1024]。 INT32 ND

约束说明

  • SparseFlashMlaGradMetadata算子需要与SparseFlashMlaGrad算子配套使用。
  • B(Batch)表示输入样本批量大小。
  • 参数cu_seqlens_q、cu_seqlens_ori_kv及cu_seqlens_cmp_kv要求其值为当前Batch与前序Batch有效token数的累加值,后一个元素的值必须大于等于前一个元素的值。
  • 参数seqused_q、seqused_ori_kv、seqused_cmp_kv要求其值表示每个Batch中的有效token数。
  • 参数cmp_residual_kv需满足cmp_residual_kv[i] < cmp_ratio。
  • ori_mask_mode及cmp_mask_mode所表示的mask模式的详细介绍见sparse_mode参数说明

调用说明

调用方式 样例代码 说明
aclnn API test_aclnn_sparse_flash_mla_grad_metadata 通过aclnnSparseFlashMlaGradMetadata接口调用SparseFlashMlaGradMetadata算子。
PyTorch API test_torch_sparse_flash_mla_grad_metadata 通过sparse_flash_mla_grad_metadata接口调用SparseFlashMlaGradMetadata算子。