- [object Object]Ascend 950PR/Ascend 950DT[object Object]:支持
- [object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:支持
- [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]:支持
- [object Object]Atlas 200I/500 A2 推理产品[object Object]:不支持
- [object Object]Atlas 推理系列产品[object Object]:不支持
- [object Object]Atlas 训练系列产品[object Object]:不支持
接口功能:
[object Object]接口用于生成一个任务列表,包含每个AIcore的Attention计算任务的起止点的Batch、Head、以及Q和K的分块的索引,供后续sparse_flash_mla算子使用。[object Object]是基于[object Object]的[object Object]扩展接口,用于调用[object Object]算子完成共享KV(Key和Value使用同一份输入)的稀疏注意力计算。该接口支持以下三类计算模式:- SWA(Sliding Window Attention):仅使用
[object Object],对原始KV做滑动窗口注意力。 - CSA(Compressed Sparse Attention):同时使用
[object Object]、[object Object]和[object Object],对原始KV窗口和TopK选择出的压缩KV共同做注意力。 - HCA(Heavily Compressed Attention):同时使用
[object Object]和[object Object],对原始KV窗口和连续压缩KV段共同做注意力。
[object Object]是[object Object]的metadata前置接口,用于在主接口执行前生成metadata。metadata记录AICore/AIVCore的任务切分结果,主接口必须传入该metadata。典型调用流程如下:- 准备
[object Object]、[object Object]、[object Object]、序列长度、[object Object]、[object Object]等输入。 - 调用
[object Object]生成[object Object]。 - 调用
[object Object],将上一步得到的[object Object]传入主算子。
- SWA(Sliding Window Attention):仅使用
计算公式:
其中,由
[object Object]的滑动窗口部分和[object Object]的压缩部分共同组成,实际参与计算的KV范围由[object Object]、[object Object]、[object Object]、[object Object]、[object Object]以及[object Object]决定。
[object Object]
调用sparse_flash_mla接口之前,先调用前置接口sparse_flash_mla_metadata,完成sparse_flash_mla负载均衡的计算。
- metadata:
[object Object]的输出,shape固定为[object Object],dtype为[object Object]。
- attention_out:
[object Object]的第一个输出,shape和[object Object]一致,dtype和[object Object]一致。 - softmax_lse:
[object Object]的第二个输出。[object Object]时返回FLOAT32标量占位Tensor;[object Object]时返回FLOAT32的log-sum-exp结果。
公共约束:
- 适用场景:该接口支持训练、推理场景下使用。
- 调用方式:该接口支持单算子模式和TorchAir图模式调用。
[object Object]、[object Object]、[object Object]的数据类型必须一致,支持FLOAT16和BFLOAT16。[object Object]和[object Object]组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下[object Object]和[object Object]必须一致。[object Object]支持"BSND"和"TND";[object Object]时,[object Object]必须为4维;[object Object]时,[object Object]必须为3维且必须传入[object Object]。[object Object]支持"BSND"、"TND"和"PA_BBND";[object Object]或[object Object]时,[object Object]和[object Object]必须为4维;[object Object]时,[object Object]和[object Object]必须为3维。[object Object]时必须传入[object Object];传入[object Object]时,还必须传入[object Object]。- 参数
[object Object]、[object Object]及[object Object]要求其值为当前Batch与前序Batch有效token数的累加值,后一个元素的值必须大于等于前一个元素的值。 - 参数
[object Object]、[object Object]、[object Object]要求其值表示每个Batch中的有效token数。 - 参数
[object Object]需满足[object Object]<[object Object]。 [object Object]时必须传入[object Object]和[object Object];传入[object Object]时,还必须传入[object Object]。[object Object]为所有[object Object]下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。[object Object]及[object Object]所表示的mask模式的详细介绍见。[object Object]固定为1024个INT32元素,[object Object]仅支持1,[object Object]、[object Object]和[object Object]当前版本不支持传入非空Tensor。[object Object]和[object Object]允许存在行间padding类非连续内存,接口会通过aclnn获取stride信息传给底层算子。
规格约束:
公共参数约束:
[object Object]仅支持512,[object Object]仅支持1。[object Object]支持1、2、4、8、16、32、64、128。[object Object]仅支持4,[object Object]仅支持127,[object Object]仅支持0。- PageAttention的block_size支持1到1024。
产品型号约束如下:
- [object Object]Ascend 950PR/Ascend 950DT[object Object]:
[object Object]/[object Object]不支持1。 - [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:cmp_ratio在SWA场景保持默认值1,CSA支持传入4,HCA支持传入128;block_size支持16的倍数,且不超过1024;ori_sparse_indices当前暂不支持,cmp_sparse_indices的最后一维压缩KV TopK长度,支持0、512、1024。默认值为0。
- [object Object]Ascend 950PR/Ascend 950DT[object Object]:
SWA:
- 仅传入
[object Object]时,[object Object]不参与压缩KV计算,需保持默认值1。 - 不传入
[object Object]、[object Object]和[object Object]。 [object Object]传0,[object Object]传0。
- 仅传入
CSA:
[object Object]仅支持3。[object Object]必须传入,最后一维支持泛化;[object Object]对应传非0。[object Object]必须传入,长度必须等于batch大小。
HCA:
[object Object]仅支持3。- 不传入
[object Object];[object Object]传0。 [object Object]必须传入,长度必须等于batch大小。
- 默认支持确定性计算。
- 默认支持batch invariance。
下面示例用单进程顺序模拟两个CP rank,说明全局TND数据与每个rank入参之间的关系。假设全局有2个序列,[object Object]:
rank1虽然只计算seq1的[object Object],但[object Object]和[object Object]需要传到当前位置结束为止的前缀。此时[object Object],kernel推导出的q起点正好是CP切分点。每个本地batch都需要满足[object Object]。