接口功能:该算子为AICPU算子,
[object Object]算子为[object Object]算子的前序算子,负责根据输入的序列长度信息和注意力配置参数,生成负载均衡的分核元数据(metadata)。该元数据包含每个AICore上FlashAttention计算任务的Batch、Head、Query分块和KV分块的索引,以及每个VectorCore上FlashDecode归约任务的索引信息。该算子不建议单独使用,建议与aclnnSparseFlashMla算子配合使用,形成完整的工作流。
场景简称:SWA(Sliding Window Attention)、CSA(Compressed Sparse Attention)、HCA(Heavily Compressed Attention)。
计算公式:
该算子为AICPU调度算子,不涉及数值计算。核心流程为:解析各Batch的Q/KV序列长度 → 根据mask模式计算每个S1G块的有效S2范围 → 基于开销模型进行负载均衡分核 → 输出分核元数据。
输出metadata tensor的shape为(1024,),数据类型为INT32,内部结构如下:
FA Metadata区域(AIC_CORE_NUM × 8个INT32),每个AICore的FA阶段任务信息:
[object Object]undefined
FD Metadata区域(AIV_CORE_NUM × 8个INT32),每个AIVCore的FD归约任务信息:
[object Object]undefined
符号说明
[object Object]undefined
每个算子分为,必须先调用[object Object]接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用[object Object]执行实际计算。
参数说明
[object Object]- [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:numHeadsQ/numHeadsKv支持1、2、4、8、16、32、64、128,oriMaskMode仅支持4,cmpMaskMode仅支持3,oriWinLeft仅支持127,oriWinRight仅支持0。
- [object Object]Ascend 950PR/Ascend 950DT[object Object]:numHeadsQ/numHeadsKv不支持1,oriMaskMode仅支持4,cmpMaskMode仅支持3,oriWinLeft仅支持127,oriWinRight仅支持0。
返回值
第一段接口完成入参校验,出现以下场景时报错:
[object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]、[object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:
[object Object][object Object]Ascend 950PR/Ascend 950DT[object Object]:
[object Object]
确定性计算
- aclnnSparseFlashMlaMetadata默认采用确定性实现,相同输入多次调用结果一致。
使用约束
- layoutQOptional和layoutKvOptional组合仅支持"BSND"/"BSND"、"TND"/"TND"、"BSND"/"PA_BBND"、"TND"/"PA_BBND";非PA_BBND场景下layoutQOptional和layoutKvOptional必须一致。
- layoutQOptional为TND时,
[object Object]必须传入。 - layoutKvOptional为PA_BBND时,
[object Object]必须传入。BSND场景可选传入[object Object]覆盖每个batch的oriKv有效长度;TND场景使用[object Object]表达oriKv序列边界。 - layoutKvOptional为TND时,
[object Object]必须传入;若hasCmpKv为true,[object Object]也必须传入。 [object Object]为所有layoutKvOptional下的可选输入,显式传入时用于覆盖cmp侧逻辑有效长度。[object Object]为[object Object]和[object Object]的可选输入,在CSA、HCA、cmpRatio不等于1且cmpMaskMode为3场景必传,用于恢复cmp侧mask使用的压缩前长度。- 该算子为AICPU算子,在Host侧CPU上执行,不占用NPU计算资源。