ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

CANN ops-transformer SparseAttnSharedkvMetadata 算子深度解析:稀疏共享KV注意力负载均衡元数据生成原理与使用指南

2026/9/19 7:46:58 拓冰建站 浏览量
CANN ops-transformer SparseAttnSharedkvMetadata 算子深度解析:稀疏共享KV注意力负载均衡元数据生成原理与使用指南 CANN ops-transformer SparseAttnSharedkvMetadata 算子深度解析稀疏共享KV注意力负载均衡元数据生成原理与使用指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformerSparseAttnSharedkvMetadata 是 CANN ops-transformer 算子库中SparseAttnSharedkv稀疏注意力算子的前置元数据算子它不执行任何实际 Attention 计算而是在 AI CPU 上根据输入序列长度、稀疏关键 token 数量、mask 模式等参数为后续的 FlashAttentionFA与 FlashDecodeFD计算生成每个 AI Core 应处理哪个 Batch、哪段 Q、哪段 K的负载均衡任务切分方案。本文基于 sparse_attn_sharedkv_metadata/README.md 展开并结合仓库中的算子原型、AI CPU Kernel 实现与 Shape 推导源码完整说明该算子的功能定位、全部参数语义、约束条件、底层切分调度算法与调用方式帮助读者在 Atlas A3 训练/推理系列产品上正确配置并理解 SparseAttnSharedkv 的负载均衡元数据生成机制。功能定位Attention 计算之前的任务划分器在稀疏共享 KVSharedKV注意力场景中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、HeadBN2、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, )数据类型固定为 INT32AI 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 数。维度为 B1每个元素表示当前 batch 与之前所有 batch 的 token 数总和前缀和INT32NDcu_seqlens_ori_kv可选输入当layout_kv为 TND 时表示不同 Batch 中 ori_kv 的有效 token 数语义同为前缀和。当前layout_kv仅支持 PA_ND故设置此参数无效INT32NDcu_seqlens_cmp_kv可选输入当layout_kv为 TND 时表示不同 Batch 中 cmp_kv 的有效 token 数语义同为前缀和。当前layout_kv仅支持 PA_ND故设置此参数无效INT32NDseqused_q可选输入表示不同 Batch 中 q 实际参与运算的 token 数维度为 B。目前暂不支持指定该参数INT32NDseqused_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_qTND 时cu_seqlens_q[b1] - cu_seqlens_q[b] 属性max_seqlen_q。KV 侧有效序列长度S2优先级seqused_kvTND 时cu_seqlens_ori_kv[b1] - 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.hShape 推导固定为(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 足以容纳该结构。FAFlashAttention元数据每核 8 个字段索引常量定义如下源码位置索引常量值含义FA_CORE_ENABLE_INDEX0该核是否启用1 启用 / 0 禁用FA_BN2_START_INDEX1该核处理的 BN2Batch×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 中的位置FDFlashDecode元数据同样每核 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_q、num_heads_kv、head_dim外还包括框架注入的soc_version、aic_core_num、aiv_core_numAI CPU Kernel 的Prepare()会强制读取必需属性读取失败即返回参数非法。参数名输入/输出/属性描述数据类型num_heads_q必需属性Q 的多头数目前仅支持 64INT32num_heads_kv必需属性K 和 V 的多头数目前仅支持 1INT32head_dim必需属性注意力头的维度INT32batch_size可选属性输入样本批量大小默认值为 None原型默认 0实际以seqused_q/cu_seqlens_q推断为准INT32max_seqlen_q可选属性所有 batch 中 q 的最大有效 token 数INT32max_seqlen_kv可选属性所有 batch 中 ori_kv 的最大有效 token 数INT32ori_topk可选属性通过 QLI 算法从 ori_kv 中筛选出的关键稀疏 token 个数。目前暂不支持指定该参数默认值为 None原型默认 0INT32cmp_topk可选属性通过 QLI 算法从 cmp_kv 中筛选出的关键稀疏 token 个数目前仅支持 512默认值为 None原型默认 0INT32cmp_ratio可选属性对 ori_kv 的压缩率数据范围支持 4/128默认值为 None原型默认 -1INT32ori_mask_mode可选属性q 和 ori_kv 计算的 mask 模式仅支持默认值 4band 模式INT32cmp_mask_mode可选属性q 和 cmp_kv 计算的 mask 模式仅支持默认值 3rightDownCausal 模式INT32ori_win_left可选属性q 和 ori_kv 计算中 q 对过去 token 计算的数量仅支持默认值 127INT32ori_win_right可选属性q 和 ori_kv 计算中 q 对未来 token 计算的数量仅支持默认值 0INT32layout_q可选属性q 的数据排布格式默认 BSND目前支持 BSND 与 TNDSTRINGlayout_kv可选属性ori_kv 与 cmp_kv 的数据排布格式目前仅支持默认值 PA_NDSTRINGhas_ori_kv可选属性是否传入 ori_kv默认 trueBOOLhas_cmp_kv可选属性是否传入 cmp_kv默认 trueBOOLdevice可选属性用于获取设备信息默认值为 NoneSTRING上述可选属性在 算子原型 中的默认值均为代码层面注册的 fallback如batch_size0、cmp_ratio-1、ori_mask_mode4、cmp_mask_mode3、ori_win_left127、layout_qBSND、layout_kvPA_ND、has_ori_kv/has_cmp_kvtrue与文档描述一致AI CPU Kernel 侧通过GetAttrValueOpt在属性存在时覆盖默认值。约束说明该算子支持推理场景下使用该算子支持aclgraph 模式图模式调用结合参数表可知当前实现还有若干参数级约束num_heads_q仅支持 64、num_heads_kv仅支持 1、cmp_topk仅支持 512、cmp_ratio仅支持 4/128、ori_mask_mode仅支持 4band、cmp_mask_mode仅支持 3rightDownCausal、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 / ParamsInitPrepare()从上下文读取 5 个输入张量、强制读取num_heads_q/num_heads_kv/head_dim三个必需属性再以GetAttrValueOpt读取全部可选属性最后调用ParamsInit()做派生量初始化按ori_mask_mode映射出稀疏模式DEFAULT_MASK(0)时preTokenINT64_MAX, nextTokenINT64_MAX无 mask 语义RIGHT_DOWN_CAUSAL(3)时preTokenINT64_MAX, nextToken0其余即当前仅支持的BAND4使用ori_win_left作为preToken、nextToken0见SparseMode枚举sparse_attn_sharedkv_metadata_aicpu.h计算groupSize num_heads_q / num_heads_kvSharedKV 场景下即 64若传入了cmp_kv且cmp_topk 0判定为 SCFA稀疏压缩注意力模式否则为 CFA全压缩注意力模式确定基本块尺寸SCFA 下 M 方向基本块mBaseSize groupSize否则mBaseSize 256S2 方向基本块s2BaseSize 512。第二步划分基本块并统计开销BalanceSchedule / CalcCostInfoBalanceSchedule()首先调用CalcSplitInfo()对每个 Batch 计算S1GQ×groupSize 方向基本块数s1GBaseNum ceil(s1Valid * groupSize / mBaseSize)S2KV 方向基本块数s2BaseNum ceil(s2Size / s2BaseSize)同时记录 S1G 尾块大小s1GTailSize、S2 尾块大小s2TailSize并标记是否存在全空 KV 序列isKvSeqAllZero。随后CalcCostInfo()遍历所有 Batch×HeadBN2组合统计整批开销对每个 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_q128 cmp_topk1024还附加了一段基于 token 长度的经验修正系数最终totalCost、totalBlockNum按kvHeadNum加权累加得到全图总负载。第三步多级分配AssignBlocksToCoreCalcSplitPlan()以totalCost / aicCoreNum为每核负载上限costLimit逐个核调用AssignBlocksToCore()依次执行四级分配策略源码按整 Batch 分配AssignByBatch若整个 BN2 的负载加上当前核已有负载仍在容差FA_TOLERANCE_RATIO2见 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。第五步写出 metadataGenMetadataGenMetadata()把上述SplitResult填充进输出张量内存中按SasMetadata结构布局FA 部分每个 AIC 核写入启用标志与 BN2/M/S2 的起止索引未被使用的核i usedCoreNum写入禁用标志FA_CORE_ENABLE_INDEX0FD 部分每个 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 等序列长度信息 │ ▼ SparseAttnSharedkvMetadataAI CPU本算子 │ metadata每核 FA/FD 任务的起止索引 ▼ SparseAttnSharedkvNPU 计算按 metadata 执行稀疏 FlashAttention 与 FlashDecode 规约用户在接入时只需保证设备为 Atlas A3 系列、num_heads_q64、num_heads_kv1并按实际数据设置max_seqlen_q/max_seqlen_kv、cmp_topk512、cmp_ratio4 或 128等稀疏参数其余负载均衡细节全部由本算子自动完成。相关源码索引算子文档experimental/attention/sparse_attn_sharedkv_metadata/README.md算子原型experimental/attention/sparse_attn_sharedkv_metadata/op_graph/sparse_attn_sharedkv_metadata_proto.hAI CPU Kernel 实现experimental/attention/sparse_attn_sharedkv_metadata/op_kernel_aicpu/sparse_attn_sharedkv_metadata_aicpu.cppAI CPU Kernel 头文件数据结构定义experimental/attention/sparse_attn_sharedkv_metadata/op_kernel_aicpu/sparse_attn_sharedkv_metadata_aicpu.hShape/数据类型推导experimental/attention/sparse_attn_sharedkv_metadata/op_host/sparse_attn_sharedkv_metadata_infershape.cppACLNN 单算子接口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),仅供参考