SparseLightningIndexerKLLossGradMetadata
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
功能说明
- 算子功能:
SparseLightningIndexerKLLossGradMetadata算子旨在根据SparseLightningIndexerKLLossGrad算子的输入shape、layout、mask和压缩比例信息,计算并输出分核切分metadata,供后续SparseLightningIndexerKLLossGrad算子使用。
参数说明
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| cu_seqlens_q | 可选输入 | 表示不同Batch中Query的有效Sequence Length,shape为(B+1, ),仅layout_q为TND场景下必传,第一个值固定为0。 | INT32 | ND |
| cu_seqlens_k | 可选输入 | 表示不同Batch中Key的有效Sequence Length,shape为(B+1, ),仅layout_k为TND场景下必传,第一个值固定为0。 | INT32 | ND |
| seqused_q | 可选输入 | 表示不同Batch中Query实际参与运算的Sequence Length,shape为(B, )。 | INT32 | ND |
| seqused_k | 可选输入 | 表示不同Batch中Key实际参与运算的Sequence Length,shape为(B, )。 | INT32 | ND |
| cmp_residual_k | 可选输入 | 表示不同Batch中cmp_kv压缩后Sequence Length的余数,配合cmp_ratio实现cmp_kv部分的mask和负载计算,shape为(B, )。cmp_ratio不为1且mask_mode为3场景下必传。 | INT32 | ND |
| num_heads_q | 属性 | 表示Query的head个数,当前支持[1, 128]。 | INT32 | - |
| num_heads_k | 属性 | 表示Key的head个数,当前仅支持1。 | INT32 | - |
| head_dim | 属性 | 表示注意力头的维度,当前仅支持128。 | INT32 | - |
| topk | 可选属性 | 表示从Query中筛选出的关键稀疏token的个数,当前支持[1, 2048]和4096、8192。 | INT32 | - |
| batch_size | 可选属性 | 表示Batch数量,默认值为0。 | INT32 | - |
| max_seqlen_q | 可选属性 | 表示Query的最长Sequence Length,默认值为0。 | INT32 | - |
| max_seqlen_k | 可选属性 | 表示Key的最长Sequence Length,默认值为0。 | INT32 | - |
| layout_q | 可选属性 | 表示Query的排列格式,支持BSND、TND,默认值为BSND。 | STRING | - |
| layout_k | 可选属性 | 表示Key的排列格式,支持BSND、TND,默认值为BSND。 | STRING | - |
| mask_mode | 可选属性 | 表示sparse模式,0表示No mask,3表示rightDownCausal模式,默认值为0。 | INT32 | - |
| cmp_ratio | 可选属性 | 表示Key的压缩率,取值范围[1, 128],默认值为1,表示无压缩。 | INT32 | - |
| metadata | 输出 | 表示负载均衡结果输出,shape固定为[64]。 | INT32 | ND |
- Ascend 950PR/Ascend 950DT :topk仅支持[1, 2048]。
- Atlas A3 训练系列产品/Atlas A3 推理系列产品 :不支持seqused_q、seqused_k、cmp_residual_k,num_heads_q仅支持8/16/32/64,topk仅支持512/1024/2048/4096/8192。
- Atlas A2 训练系列产品/Atlas A2 推理系列产品 :不支持seqused_q、seqused_k、cmp_residual_k,num_heads_q仅支持8/16/32/64,topk仅支持512/1024/2048/4096/8192。
约束说明
- SparseLightningIndexerKLLossGradMetadata算子需要与SparseLightningIndexerKLLossGrad算子配套使用。
- B(Batch)表示输入样本批量大小。
- 参数cu_seqlens_q、cu_seqlens_k要求其值为当前Batch与前序Batch有效token数的累加值,后一个元素的值必须大于等于前一个元素的值。
- 参数seqused_q、seqused_k要求其值表示每个Batch中的有效token数。
- 参数cmp_residual_k需满足cmp_residual_k[i] < cmp_ratio。
- mask_mode所表示的mask模式的详细介绍见sparse_mode参数说明。
调用说明
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| aclnn API | test_aclnn_sparse_lightning_indexer_kl_loss_grad_metadata | 通过aclnnSparseLightningIndexerKLLossGradMetadata接口调用SparseLightningIndexerKLLossGradMetadata算子。 |
| PyTorch API | test_torch_sparse_lightning_indexer_kl_loss_grad_metadata | 通过sparse_lightning_indexer_kl_loss_grad_metadata接口调用SparseLightningIndexerKLLossGradMetadata算子。 |