CANN ops-transformer SparseAttnSharedkvMetadata 算子深度解析:稀疏共享KV注意力负载均衡元数据生成原理与使用指南
2026/9/19 7:46:48 网站建设 项目流程

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_versionaic_core_numaiv_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 数总和(前缀和)INT32ND
cu_seqlens_ori_kv可选输入layout_kv为 TND 时,表示不同 Batch 中 ori_kv 的有效 token 数,语义同为前缀和。当前layout_kv仅支持 PA_ND,故设置此参数无效INT32ND
cu_seqlens_cmp_kv可选输入layout_kv为 TND 时,表示不同 Batch 中 cmp_kv 的有效 token 数,语义同为前缀和。当前layout_kv仅支持 PA_ND,故设置此参数无效INT32ND
seqused_q可选输入表示不同 Batch 中 q 实际参与运算的 token 数,维度为 B。目前暂不支持指定该参数INT32ND
seqused_kv可选输入表示不同 Batch 中 ori_kv 实际参与运算的 token 数,维度为 BINT32ND

从 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_kvcu_seqlens_cmp_kvlayout_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 = 36AIV_CORE_NUM = 72,并通过static_assert保证 1024 个 INT32 足以容纳该结构。

FA(FlashAttention)元数据每核 8 个字段,索引常量定义如下(源码位置):

索引常量含义
FA_CORE_ENABLE_INDEX0该核是否启用(1 启用 / 0 禁用)
FA_BN2_START_INDEX1该核处理的 BN2(Batch×Head)起点
FA_M_START_INDEX2该核处理的 Q 分块(M/GS1)起点
FA_S2_START_INDEX3该核处理的 KV 分块(S2)起点
FA_BN2_END_INDEX4BN2 终点(右开区间)
FA_M_END_INDEX5M 终点
FA_S2_END_INDEX6S2 终点
FA_FIRST_FD_DATA_WORKSPACE_IDX_INDEX7该核第一份 FD 归约数据在 workspace 中的位置

FD(FlashDecode)元数据同样每核 8 个字段(源码位置):

索引常量含义
FD_CORE_ENABLE_INDEX0该 Vector 核是否参与归约
FD_BN2_IDX_INDEX1归约任务所属的 BN2 索引
FD_M_IDX_INDEX2归约任务所属的 GS1 索引
FD_WORKSPACE_IDX_INDEX3归约数据在 workspace 中的存放位置
FD_WORKSPACE_NUM_INDEX4该归约任务的 S2 核间切分份数
FD_M_START_INDEX5该 Vector 核处理的 M 轴起点
FD_M_NUM_INDEX6该 Vector 核处理的 M 轴行数

属性参数

属性分为必需属性与可选属性两类。必需属性在 算子原型 中通过REQUIRED_ATTR声明(除文档表格列出的num_heads_qnum_heads_kvhead_dim外,还包括框架注入的soc_versionaic_core_numaiv_core_num),AI CPU Kernel 的Prepare()会强制读取必需属性,读取失败即返回参数非法。

参数名输入/输出/属性描述数据类型
num_heads_q必需属性Q 的多头数,目前仅支持 64INT32
num_heads_kv必需属性K 和 V 的多头数,目前仅支持 1INT32
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 计算的数量,仅支持默认值 127INT32
ori_win_right可选属性q 和 ori_kv 计算中 q 对未来 token 计算的数量,仅支持默认值 0INT32
layout_q可选属性q 的数据排布格式,默认 BSND,目前支持 BSND 与 TNDSTRING
layout_kv可选属性ori_kv 与 cmp_kv 的数据排布格式,目前仅支持默认值 PA_NDSTRING
has_ori_kv可选属性是否传入 ori_kv,默认 trueBOOL
has_cmp_kv可选属性是否传入 cmp_kv,默认 trueBOOL
device可选属性用于获取设备信息,默认值为 NoneSTRING

上述可选属性在 算子原型 中的默认值均为代码层面注册的 fallback(如batch_size=0cmp_ratio=-1ori_mask_mode=4cmp_mask_mode=3ori_win_left=127layout_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_qori_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作为preTokennextToken=0,见SparseMode枚举(sparse_attn_sharedkv_metadata_aicpu.h);
  • 计算groupSize = num_heads_q / num_heads_kv(SharedKV 场景下即 64);
  • 若传入了cmp_kvcmp_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 长度的经验修正系数;
  • 最终totalCosttotalBlockNumkvHeadNum加权累加,得到全图总负载。

第三步:多级分配(AssignBlocksToCore)

CalcSplitPlan()totalCost / aicCoreNum为每核负载上限costLimit,逐个核调用AssignBlocksToCore(),依次执行四级分配策略(源码):

  1. 按整 Batch 分配(AssignByBatch):若整个 BN2 的负载加上当前核已有负载仍在容差(FA_TOLERANCE_RATIO=2,见 aicpu.h)范围内,则整批划给当前核;
  2. 按行分配(AssignByRow):否则退化为按 S1G 行分配,逐行累加直至接近负载上限;
  3. 按块分配(AssignByBlock):再以单个 S2 块为粒度补齐(仅supportFd场景);
  4. 强制分配(ForceAssign):兜底保证每个启用核至少获得一块任务,避免核空闲。

分配过程中同步记录每个核的bN2EndgS1Ends2End(右开区间)以及maxCost,供GenMetadata()写出 FA 元数据。

第四步:FlashDecode 归约任务的负载均衡(SplitFD)

当存在跨核行(某一行 KV 被切分到多个核)时,会产生需要 Vector 核归约的 FD 任务。RecordFDInfo()在切分点处记录归约任务的 BN2、GS1、workspace 位置与 S2 切分份数;随后SplitFD()(源码)按归约数据总量fdS2SplitNum × fdMSizeaivCoreNum个 Vector 核间做二次负载均衡:先按平均负载计算每个任务占用的核数(向下取整、至少 1),再对每个任务的 M 轴行数做均分(向上取整),最终产出每个 Vector 核的fdMStartfdMNum

第五步:写出 metadata(GenMetadata)

GenMetadata()把上述SplitResult填充进输出张量(内存中按SasMetadata结构布局):

  • FA 部分:每个 AIC 核写入启用标志与 BN2/M/S2 的起止索引;未被使用的核(i >= usedCoreNum)写入禁用标志FA_CORE_ENABLE_INDEX=0
  • FD 部分:每个 AIV 核写入归约任务索引(fdBN2IdxfdMIdxfdWorkspaceIdxfdS2SplitNum)与 M 轴划分(fdMStartfdMNum),未参与的核写入禁用标志。

调用方式

调用模式

  • 单算子模式:直接调用该算子的 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=64num_heads_kv=1,并按实际数据设置max_seqlen_q/max_seqlen_kvcmp_topk=512cmp_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询