CANN ops-transformer SparseAttnSharedkvMetadata 算子深度解析:稀疏共享KV注意力负载均衡元数据生成原理与使用指南
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
SparseAttnSharedkvMetadata 是 CANN ops-transformer 算子库中SparseAttnSharedkv稀疏注意力算子的前置元数据算子,它不执行任何实际 Attention 计算,而是在 AI CPU 上根据输入序列长度、稀疏关键 token 数量、mask 模式等参数,为后续的 FlashAttention(FA)与 FlashDecode(FD)计算生成"每个 AI Core 应处理哪个 Batch、哪段 Q、哪段 K"的负载均衡任务切分方案。本文基于 sparse_attn_sharedkv_metadata/README.md 展开,并结合仓库中的算子原型、AI CPU Kernel 实现与 Shape 推导源码,完整说明该算子的功能定位、全部参数语义、约束条件、底层切分调度算法与调用方式,帮助读者在 Atlas A3 训练/推理系列产品上正确配置并理解 SparseAttnSharedkv 的负载均衡元数据生成机制。
功能定位:Attention 计算之前的"任务划分器"
在稀疏共享 KV(SharedKV)注意力场景中,Q 有多头(目前仅支持 64 头),而 K/V 共享且仅 1 头,每个 Q 头只需要从全局 KV 中挑选出若干关键稀疏 token(通过 QLI 算法筛选)参与计算,并配合 band、causal 等稀疏 mask 模式。这类稀疏计算的最大难点在于:不同 Batch、不同 Q 分块对应的有效 KV 范围差异巨大,如果按固定的均匀方式把任务分给各个 AI Core,必然出现严重的核间负载不均衡。
SparseAttnSharedkvMetadata 正是为解决这个问题而生。它的职责可以概括为:
- 在 AI CPU 上完成"分核规划":根据 batch 内各条序列的实际有效 token 数、稀疏 topk 参数、mask 窗口等输入,把整个 Attention 计算任务切分成基本块(Block),统计每个块的估算开销(Cost),再通过多级分配策略把块集合均衡地分配给各个 AI Core;
- 输出
metadata张量:为每个 Cube 核(AI Core,负责 FlashAttention 计算)记录其负责的 Batch、Head(BN2)、Q 分块(M/GS1)、KV 分块(S2)的起止索引,同时为每个 Vector 核(AIV,负责 FlashDecode 规约)记录归约任务的索引与 M 轴划分范围; - 生成的 metadata 直接作为
SparseAttnSharedkv算子的输入,指导其按既定范围执行稀疏 Attention,从而最大化计算资源利用率,避免各 Core 间负载不均衡。
从仓库源码看,该算子的实现横跨三层:
| 层次 | 文件 | 作用 |
|---|---|---|
| 算子原型(图定义) | op_graph/sparse_attn_sharedkv_metadata_proto.h | 注册算子输入、输出、属性,声明必需的soc_version、aic_core_num、aiv_core_num等属性 |
| 主机侧 Shape 推导 | op_host/sparse_attn_sharedkv_metadata_infershape.cpp | 将输出 metadata 的 Shape 固定为(SAS_META_SIZE, ),数据类型固定为 INT32 |
| AI CPU 核函数 | op_kernel_aicpu/sparse_attn_sharedkv_metadata_aicpu.cpp | 实现完整的块划分、开销估算、负载均衡分配与 metadata 生成逻辑 |
此外,AI CPU Kernel 的入口Compute()中实际调用的分核数据结构、元数据索引常量定义位于 experimental/attention/sparse_attn_sharedkv/op_kernel/sparse_attn_sharedkv_metadata.h,它同时被 metadata 算子(AI CPU 侧写入)与 SparseAttnSharedkv 算子(NPU 侧读取)共用,保证了元数据布局的一致性。
产品支持情况
当前仓库的 README 明确给出了产品支持矩阵:
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | × |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | × |
| Atlas 200I/500 A2 推理系列产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
也就是说,该算子(以及其服务的 SparseAttnSharedkv)目前仅在 Atlas A3 训练/推理系列产品上可用,其余产品线暂不支持。使用时需要确认运行环境为 Atlas A3 系列,并在图/单算子调用中正确传入与设备匹配的soc_version等属性。
参数说明
输入参数
该算子共有 5 个可选输入,全部为 INT32 类型、ND 格式,用于描述每条序列的有效 token 数:
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
cu_seqlens_q | 可选输入 | 当layout_query为 TND 时,表示不同 Batch 中 q 的有效 token 数。维度为 B+1,每个元素表示当前 batch 与之前所有 batch 的 token 数总和(前缀和) | INT32 | ND |
cu_seqlens_ori_kv | 可选输入 | 当layout_kv为 TND 时,表示不同 Batch 中 ori_kv 的有效 token 数,语义同为前缀和。当前layout_kv仅支持 PA_ND,故设置此参数无效 | INT32 | ND |
cu_seqlens_cmp_kv | 可选输入 | 当layout_kv为 TND 时,表示不同 Batch 中 cmp_kv 的有效 token 数,语义同为前缀和。当前layout_kv仅支持 PA_ND,故设置此参数无效 | INT32 | ND |
seqused_q | 可选输入 | 表示不同 Batch 中 q 实际参与运算的 token 数,维度为 B。目前暂不支持指定该参数 | INT32 | ND |
seqused_kv | 可选输入 | 表示不同 Batch 中 ori_kv 实际参与运算的 token 数,维度为 B | INT32 | ND |
从 AI CPU Kernel 实现 的Prepare()与GetQueryBatchSize()/GetS1SeqSize()/GetS2SeqSize()可以看出这 5 个输入在底层的实际用途与优先级:
- BatchSize 推断优先级:先看
seqused_q是否传入(取其第 0 维作为 B);未传且layout_query == "TND"时用cu_seqlens_q的维度减 1 得到 B;否则回退到属性batch_size。 - Q 侧有效序列长度(S1)优先级:
seqused_q>(TND 时)cu_seqlens_q[b+1] - cu_seqlens_q[b]> 属性max_seqlen_q。 - KV 侧有效序列长度(S2)优先级:
seqused_kv>(TND 时)cu_seqlens_ori_kv[b+1] - cu_seqlens_ori_kv[b]> 属性max_seqlen_kv。 cu_seqlens_ori_kv、cu_seqlens_cmp_kv在layout_kv仅支持 PA_ND 的前提下不参与实际逻辑,仅在 TND 分支中作为备用数据源保留。
输出参数
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
metadata | 输出 | 每个 Cube 核上 FlashAttention 计算任务的 Batch、Head 以及 Q 和 K 的分块索引,以及每个 Vector 核上 FlashDecode 的规约任务索引 | INT32 | - |
metadata输出的一维长度为常量SAS_META_SIZE = 1024(见 sparse_attn_sharedkv_metadata.h),Shape 推导固定为(1024,),数据类型固定为DT_INT32(见 infershape 实现)。其内部布局为SasMetadata结构:faMetadata[AIC_CORE_NUM][8]紧接fdMetadata[AIV_CORE_NUM][8],其中AIC_CORE_NUM = 36、AIV_CORE_NUM = 72,并通过static_assert保证 1024 个 INT32 足以容纳该结构。
FA(FlashAttention)元数据每核 8 个字段,索引常量定义如下(源码位置):
| 索引常量 | 值 | 含义 |
|---|---|---|
FA_CORE_ENABLE_INDEX | 0 | 该核是否启用(1 启用 / 0 禁用) |
FA_BN2_START_INDEX | 1 | 该核处理的 BN2(Batch×Head)起点 |
FA_M_START_INDEX | 2 | 该核处理的 Q 分块(M/GS1)起点 |
FA_S2_START_INDEX | 3 | 该核处理的 KV 分块(S2)起点 |
FA_BN2_END_INDEX | 4 | BN2 终点(右开区间) |
FA_M_END_INDEX | 5 | M 终点 |
FA_S2_END_INDEX | 6 | S2 终点 |
FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX | 7 | 该核第一份 FD 归约数据在 workspace 中的位置 |
FD(FlashDecode)元数据同样每核 8 个字段(源码位置):
| 索引常量 | 值 | 含义 |
|---|---|---|
FD_CORE_ENABLE_INDEX | 0 | 该 Vector 核是否参与归约 |
FD_BN2_IDX_INDEX | 1 | 归约任务所属的 BN2 索引 |
FD_M_IDX_INDEX | 2 | 归约任务所属的 GS1 索引 |
FD_WORKSPACE_IDX_INDEX | 3 | 归约数据在 workspace 中的存放位置 |
FD_WORKSPACE_NUM_INDEX | 4 | 该归约任务的 S2 核间切分份数 |
FD_M_START_INDEX | 5 | 该 Vector 核处理的 M 轴起点 |
FD_M_NUM_INDEX | 6 | 该 Vector 核处理的 M 轴行数 |
属性参数
属性分为必需属性与可选属性两类。必需属性在 算子原型 中通过REQUIRED_ATTR声明(除文档表格列出的num_heads_q、num_heads_kv、head_dim外,还包括框架注入的soc_version、aic_core_num、aiv_core_num),AI CPU Kernel 的Prepare()会强制读取必需属性,读取失败即返回参数非法。
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 |
|---|---|---|---|
num_heads_q | 必需属性 | Q 的多头数,目前仅支持 64 | INT32 |
num_heads_kv | 必需属性 | K 和 V 的多头数,目前仅支持 1 | INT32 |
head_dim | 必需属性 | 注意力头的维度 | INT32 |
batch_size | 可选属性 | 输入样本批量大小,默认值为 None(原型默认 0,实际以seqused_q/cu_seqlens_q推断为准) | INT32 |
max_seqlen_q | 可选属性 | 所有 batch 中 q 的最大有效 token 数 | INT32 |
max_seqlen_kv | 可选属性 | 所有 batch 中 ori_kv 的最大有效 token 数 | INT32 |
ori_topk | 可选属性 | 通过 QLI 算法从 ori_kv 中筛选出的关键稀疏 token 个数。目前暂不支持指定该参数,默认值为 None(原型默认 0) | INT32 |
cmp_topk | 可选属性 | 通过 QLI 算法从 cmp_kv 中筛选出的关键稀疏 token 个数,目前仅支持 512,默认值为 None(原型默认 0) | INT32 |
cmp_ratio | 可选属性 | 对 ori_kv 的压缩率,数据范围支持 4/128,默认值为 None(原型默认 -1) | INT32 |
ori_mask_mode | 可选属性 | q 和 ori_kv 计算的 mask 模式,仅支持默认值 4(band 模式) | INT32 |
cmp_mask_mode | 可选属性 | q 和 cmp_kv 计算的 mask 模式,仅支持默认值 3(rightDownCausal 模式) | INT32 |
ori_win_left | 可选属性 | q 和 ori_kv 计算中 q 对过去 token 计算的数量,仅支持默认值 127 | INT32 |
ori_win_right | 可选属性 | q 和 ori_kv 计算中 q 对未来 token 计算的数量,仅支持默认值 0 | INT32 |
layout_q | 可选属性 | q 的数据排布格式,默认 BSND,目前支持 BSND 与 TND | STRING |
layout_kv | 可选属性 | ori_kv 与 cmp_kv 的数据排布格式,目前仅支持默认值 PA_ND | STRING |
has_ori_kv | 可选属性 | 是否传入 ori_kv,默认 true | BOOL |
has_cmp_kv | 可选属性 | 是否传入 cmp_kv,默认 true | BOOL |
device | 可选属性 | 用于获取设备信息,默认值为 None | STRING |
上述可选属性在 算子原型 中的默认值均为代码层面注册的 fallback(如batch_size=0、cmp_ratio=-1、ori_mask_mode=4、cmp_mask_mode=3、ori_win_left=127、layout_q="BSND"、layout_kv="PA_ND"、has_ori_kv/has_cmp_kv=true),与文档描述一致,AI CPU Kernel 侧通过GetAttrValueOpt在属性存在时覆盖默认值。
约束说明
- 该算子支持推理场景下使用;
- 该算子支持aclgraph 模式(图模式)调用;
- 结合参数表可知,当前实现还有若干参数级约束:
num_heads_q仅支持 64、num_heads_kv仅支持 1、cmp_topk仅支持 512、cmp_ratio仅支持 4/128、ori_mask_mode仅支持 4(band)、cmp_mask_mode仅支持 3(rightDownCausal)、ori_win_left仅支持 127、ori_win_right仅支持 0、layout_q仅支持 BSND/TND、layout_kv仅支持 PA_ND,且seqused_q、ori_topk暂不支持指定。配置时若超出上述范围,算子行为不受保证。
底层原理:AI CPU 上的负载均衡切分算法
SparseAttnSharedkvMetadata 的核心算法全部实现在 sparse_attn_sharedkv_metadata_aicpu.cpp 中,主流程Compute()依次执行Prepare()→BalanceSchedule()→GenMetadata()三步。结合源码可以还原出完整的算法链路。
第一步:参数解析与初始化(Prepare / ParamsInit)
Prepare()从上下文读取 5 个输入张量、强制读取num_heads_q/num_heads_kv/head_dim三个必需属性,再以GetAttrValueOpt读取全部可选属性,最后调用ParamsInit()做派生量初始化:
- 按
ori_mask_mode映射出稀疏模式:DEFAULT_MASK(0)时preToken=INT64_MAX, nextToken=INT64_MAX(无 mask 语义);RIGHT_DOWN_CAUSAL(3)时preToken=INT64_MAX, nextToken=0;其余(即当前仅支持的BAND=4)使用ori_win_left作为preToken、nextToken=0,见SparseMode枚举(sparse_attn_sharedkv_metadata_aicpu.h); - 计算
groupSize = num_heads_q / num_heads_kv(SharedKV 场景下即 64); - 若传入了
cmp_kv且cmp_topk > 0判定为 SCFA(稀疏压缩注意力)模式,否则为 CFA(全压缩注意力)模式; - 确定基本块尺寸:SCFA 下 M 方向基本块
mBaseSize = groupSize,否则mBaseSize = 256;S2 方向基本块s2BaseSize = 512。
第二步:划分基本块并统计开销(BalanceSchedule / CalcCostInfo)
BalanceSchedule()首先调用CalcSplitInfo()对每个 Batch 计算:
- S1G(Q×groupSize 方向)基本块数:
s1GBaseNum = ceil(s1Valid * groupSize / mBaseSize); - S2(KV 方向)基本块数:
s2BaseNum = ceil(s2Size / s2BaseSize); - 同时记录 S1G 尾块大小
s1GTailSize、S2 尾块大小s2TailSize,并标记是否存在全空 KV 序列(isKvSeqAllZero)。
随后CalcCostInfo()遍历所有 Batch×Head(BN2)组合统计整批开销:
- 对每个 S1G 行,通过
CalcS2TokenRange()依据 mask 模式与窗口参数推算出该行需要访问的 ori_kv token 区间(band 模式即[s1First - ori_win_left, s1Last + ori_win_right]),并据此计算出 win(窗口 band)与 cmp(压缩)两段 S2 块的起止范围; WinCalcCost()与CmpCalcCost()使用与硬件对齐粒度(M 按 16 对齐、S2 按 64 对齐)的线性代价模型估算块开销,代价系数分别为 M 轴 6、S2 轴 10;- 对 SCFA 高优先级场景(
num_heads_q==128 && cmp_topk==1024),还附加了一段基于 token 长度的经验修正系数; - 最终
totalCost、totalBlockNum按kvHeadNum加权累加,得到全图总负载。
第三步:多级分配(AssignBlocksToCore)
CalcSplitPlan()以totalCost / aicCoreNum为每核负载上限costLimit,逐个核调用AssignBlocksToCore(),依次执行四级分配策略(源码):
- 按整 Batch 分配(AssignByBatch):若整个 BN2 的负载加上当前核已有负载仍在容差(
FA_TOLERANCE_RATIO=2,见 aicpu.h)范围内,则整批划给当前核; - 按行分配(AssignByRow):否则退化为按 S1G 行分配,逐行累加直至接近负载上限;
- 按块分配(AssignByBlock):再以单个 S2 块为粒度补齐(仅
supportFd场景); - 强制分配(ForceAssign):兜底保证每个启用核至少获得一块任务,避免核空闲。
分配过程中同步记录每个核的bN2End、gS1End、s2End(右开区间)以及maxCost,供GenMetadata()写出 FA 元数据。
第四步:FlashDecode 归约任务的负载均衡(SplitFD)
当存在跨核行(某一行 KV 被切分到多个核)时,会产生需要 Vector 核归约的 FD 任务。RecordFDInfo()在切分点处记录归约任务的 BN2、GS1、workspace 位置与 S2 切分份数;随后SplitFD()(源码)按归约数据总量fdS2SplitNum × fdMSize在aivCoreNum个 Vector 核间做二次负载均衡:先按平均负载计算每个任务占用的核数(向下取整、至少 1),再对每个任务的 M 轴行数做均分(向上取整),最终产出每个 Vector 核的fdMStart与fdMNum。
第五步:写出 metadata(GenMetadata)
GenMetadata()把上述SplitResult填充进输出张量(内存中按SasMetadata结构布局):
- FA 部分:每个 AIC 核写入启用标志与 BN2/M/S2 的起止索引;未被使用的核(
i >= usedCoreNum)写入禁用标志FA_CORE_ENABLE_INDEX=0; - FD 部分:每个 AIV 核写入归约任务索引(
fdBN2Idx、fdMIdx、fdWorkspaceIdx、fdS2SplitNum)与 M 轴划分(fdMStart、fdMNum),未参与的核写入禁用标志。
调用方式
调用模式
- 单算子模式:直接调用该算子的 ACLNN 接口,接口原型见 aclnn_sparse_attn_sharedkv_metadata.h:先调用
aclnnSparseAttnSharedkvMetadataGetWorkspaceSize查询 workspace 大小并创建执行器,再调用aclnnSparseAttnSharedkvMetadata提交执行; - aclgraph 模式:以图模式将该算子作为
SparseAttnSharedkv的前序算子挂在图中。
与 SparseAttnSharedkv 的配合
该算子作为 SparseAttnSharedkv 的前置算子使用,其输出的metadata张量会作为 SparseAttnSharedkv 的输入,驱动后续真正的稀疏 Attention 计算。完整调用示例参见 SparseAttnSharedkv 调用示例。典型的数据流为:
cu_seqlens_q / seqused_kv 等序列长度信息 │ ▼ SparseAttnSharedkvMetadata(AI CPU,本算子) │ metadata(每核 FA/FD 任务的起止索引) ▼ SparseAttnSharedkv(NPU 计算,按 metadata 执行稀疏 FlashAttention 与 FlashDecode 规约)用户在接入时只需保证:设备为 Atlas A3 系列、num_heads_q=64、num_heads_kv=1,并按实际数据设置max_seqlen_q/max_seqlen_kv、cmp_topk=512、cmp_ratio(4 或 128)等稀疏参数,其余负载均衡细节全部由本算子自动完成。
相关源码索引
- 算子文档:experimental/attention/sparse_attn_sharedkv_metadata/README.md
- 算子原型:experimental/attention/sparse_attn_sharedkv_metadata/op_graph/sparse_attn_sharedkv_metadata_proto.h
- AI CPU Kernel 实现:experimental/attention/sparse_attn_sharedkv_metadata/op_kernel_aicpu/sparse_attn_sharedkv_metadata_aicpu.cpp
- AI CPU Kernel 头文件(数据结构定义):experimental/attention/sparse_attn_sharedkv_metadata/op_kernel_aicpu/sparse_attn_sharedkv_metadata_aicpu.h
- Shape/数据类型推导:experimental/attention/sparse_attn_sharedkv_metadata/op_host/sparse_attn_sharedkv_metadata_infershape.cpp
- ACLNN 单算子接口:experimental/attention/sparse_attn_sharedkv_metadata/op_host/op_api/aclnn_sparse_attn_sharedkv_metadata.h
- 元数据布局与索引常量(FA/FD 共用头文件):experimental/attention/sparse_attn_sharedkv/op_kernel/sparse_attn_sharedkv_metadata.h
- 后置稀疏注意力算子:experimental/attention/sparse_attn_sharedkv/README.md
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考