资讯动态

CANN ops-transformer sparseMode 稀疏模式全解析:FlashAttentionScore 的十种注意力掩码机制

发布时间:2026/9/19 13:10:03 来源:尧图企业网站定制
CANN ops-transformer sparseMode 稀疏模式全解析FlashAttentionScore 的十种注意力掩码机制【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer导读sparseMode稀疏模式是 CANN ops-transformer 算子库中控制注意力掩码attention mask生成方式的核心属性直接决定 FlashAttentionScore 系列融合算子在 NPU 上如何遮蔽QK^T矩阵、如何处理因果掩码causal mask、prefix 前缀、band 滑动窗口以及 varlen 长序列外切等场景。本文以仓库 docs/zh/context/sparse_mode_introduction.md 为骨架结合 FlashAttentionScore 算子的算子原型op_graph、tiling 参数校验op_host与 CPU/NPU 参考实现tests/pytest源码系统讲解 sparseMode 09 共十种模式的语义、掩码形状、参数约束与典型使用场景帮助读者在模型迁移和算子调用时正确配置 sparse_mode、preTokens、nextTokens、prefix 等关键属性。sparseMode 概述与模式速查在大模型领域sparseMode 通常指模型架构或计算公式中参数或激活的稀疏性设计与稠密模式DenseMode相对。在 CANN ops-transformer 中sparseMode 是 FlashAttentionScore 等注意力算子如 aclnnFlashAttentionScore / aclnnFlashAttentionVarLenScore 系列的一个 int 类型属性用于声明当前注意力计算的掩码类型从而决定算子内部是使用传入的 attenMask还是依据 preTokens / nextTokens 等属性现场推导掩码区域。在算子原型 flash_attention_score_proto.h 中sparse_mode被声明为Int类型、默认值为 0官方注释给出的枚举含义为0 defaultMsk1 allMask2 leftUpCausal3 rightDownCausal4 band5 prefix6 prefixCompress7 rightDownCausalBand8 bandLeftUpCausal在 flash_attention_score_def.cpp 中同样通过Attr(sparse_mode).AttrType(OPTIONAL).Int(0)确认其为可选属性且默认 0。仓库文档给出的完整模式速查表如下sparseMode含义备注0defaultMask 模式-1allMask 模式-2leftUpCausal 模式-3rightDownCausal 模式-4band 模式-5prefix 非压缩模式varlen 场景不支持6prefix 压缩模式-7varlen 外切场景rightDownCausal 模式仅 varlen 场景支持8varlen 外切场景leftUpCausal 模式仅 varlen 场景支持9treeMask 模式非量化支持 GQA 和 MLA 场景全量化仅支持 MLA 场景注意sparse_mode的默认属性值 0 对应 defaultMask而 proto 中pre_tockens/next_tockens的默认值均为2147483647INT32_MAX意味着默认情况下不做范围限制。attenMask 工作原理QK^T 遮蔽机制在理解各 sparseMode 之前需要先明确 attenMask 的基本工作原理在 Mask 为 True 的位置遮蔽 query(Q) 与 key(K) 的转置矩阵乘积的值即对QK^T矩阵中对应位置行索引来自 Q 的序列位置列索引来自 K 的序列位置进行遮蔽。图中的原理示意可以这样理解QK^T矩阵的尺寸为 Sq × SkvSq 为 query 序列长度Skv 为 key/value 序列长度attenMask 与QK^T同形状Mask 为 True图中深色区域的位置在 softmax 之前被替换为极小值从而在概率分布上等效于该位置的注意力分数被完全屏蔽。该原理在测试参考实现 test_utils.py 的generate_npu_mask/generate_cpu_mask中有逐模式的矩阵构造逻辑与之对应下文将结合各模式逐一展开。sparseMode0defaultMask 模式sparseMode 为 0 时代表 defaultMask 模式也是属性默认值。该模式的行为完全由是否传入 attenMask以及 preTokens / nextTokens 的取值决定共有五种子场景场景一不传 mask全计算如果 attenMask 未传入则不做任何 mask 操作attenMask 取值为 None同时忽略 preTokens 和 nextTokens 取值QK^T矩阵全量参与计算等效于稠密注意力。场景二nextTokens0preTokens≥Sqcausal 下三角当 nextTokens 取值为 0、preTokens 大于等于 Sq 时表示 causal 场景的稀疏attenMask 应传入下三角矩阵只有 preTokens 和 nextTokens 之间的部分需要计算即每个 query 只能看到自身及之前含部分历史的 key对应自回归解码的因果注意力。此时传入的 attenMask 形状为下三角矩阵场景三preTokensSqnextTokensSkv 且均 ≥0band当 preTokens 小于 Sq、nextTokens 小于 Skv且两者都大于等于 0 时表示 band带状场景只有 preTokens 和 nextTokens 之间的部分需要计算相当于滑动窗口注意力每个 query 只关注前后一定范围内的 key。此时 attenMask 应传入 band 形状矩阵场景四nextTokens 为负数当 nextTokens 为负数时以 preTokens9nextTokens-3 为例preTokens 和 nextTokens 之间的部分仍然需要计算。此时遮蔽区域相当于在 causal 基础上向未来方向额外放行部分 keynextTokens 为负意味着允许 query 关注自身之后最多 |nextTokens| 个位置的 key。注意nextTokens 为负数时preTokens 取值必须大于等于 nextTokens 的绝对值且 nextTokens 的绝对值小于 Skv。场景五preTokens 为负数当 preTokens 为负数时以 nextTokens7preTokens-3 为例preTokens 和 nextTokens 之间的部分需要计算。遮蔽区域相当于在 causal 基础上向过去方向收窄query 只能看到更近的历史 key。注意preTokens 为负数时nextTokens 取值必须大于等于 preTokens 的绝对值且 preTokens 的绝对值小于 Sq。源码佐证在 flash_attention_score_tiling_varlen.cpp 的SparseNoMaskModeCheck中当nextTokens 0时直接报错要求其大于等于 0而 defaultMask 模式在 CheckPretokenAndNexttoken 分支之外通过专门的 default 模式检查处理正负取值组合保证 band 区域存在有效数据块。测试参考实现 test_utils.py 中 default 模式掩码构造公式为atten_mask_u torch.triu(torch.ones(s1, s2), diagonalnext_tokens 1) atten_mask_l torch.tril(torch.ones(s1, s2), diagonal-pre_tokens - 1) mask (atten_mask_u atten_mask_l).bool()即上三角部分与下三角部分取并集中间band 内区域保持可见。sparseMode1allMask 模式sparseMode 为 1 时代表 allMask即传入完整的 attenMask 矩阵完全由用户自定义掩码内容。该场景下忽略 nextTokens、preTokens 取值。从源码看allMask 模式在 CheckPretokenAndNexttoken 中会检查 preTokens / nextTokens 是否足够大小于 s1Size-1 或 s2Size-1 时打印告警并将其重置为 INT32_MAX随后将稀疏类型直接映射为对应枚举表示掩码完全交由 attenMask 输入决定。sparseMode2leftUpCausal 模式sparseMode 为 2 时代表 leftUpCausal 模式对应以左上顶点为参数起点划分的下三角场景即标准的因果掩码每个 query 只能看到自身及之前位置的 key。该场景下忽略 preTokens、nextTokens 取值。传入的 attenMask 为优化后的压缩下三角矩阵2048×2048即无论实际序列多长都使用固定 2048×2048 的压缩下三角掩码下文 sparseMode3/4/6 同理示意如下源码中的处理逻辑是leftUpCausal 在 CheckPretokenAndNexttoken 分支内强制设置preTokens s1Size、nextTokens 0并将稀疏类型映射为SparseEnum::CAUSAL模板始终按 preTokens 等于 s1Size 的 causal 形态执行。对应的测试掩码构造为torch.triu(torch.ones(s1, s2), diagonal1)test_utils.py。sparseMode3rightDownCausal 模式sparseMode 为 3 时代表 rightDownCausal 模式对应以右下顶点为参数起点划分的下三角场景即每个 query 只能看到自身及之后位置的 key适用于 Skv ≥ Sq 的因果对齐场景例如 prefix 已被吸收进 KV 缓存的情况。该场景下忽略 preTokens、nextTokens 取值attenMask 同样为优化后的压缩下三角矩阵2048×2048。从源码看rightDownCausal 分支flash_attention_score_tiling_varlen.cpp会逐个 batch 校验actual_seq_qlen actual_seq_kvlen然后设置preTokens s2Size、nextTokens 0并映射为SparseEnum::RIGHT_DOWN_CAUSAL。测试参考实现中的构造公式为torch.triu(torch.ones(s1, s2), diagonals2 - s1 1)test_utils.py即对角线的位置由 Sq 与 Skv 的差值决定。与 PyTorch 生态的对应关系测试用例注释中明确指出sparseMode3 等价于 GPU 侧的causaltrue见 test_case.py在模型迁移时可据此映射。sparseMode4band 模式sparseMode 为 4 时代表 band 场景即计算 preTokens 和 nextTokens 之间的部分对应滑动窗口注意力。与 sparseMode0 的 band 子场景不同band 模式的参数起点为右下角rightDown 对齐且 preTokens 和 nextTokens 之间需要有交集。attenMask 为优化后的压缩下三角矩阵2048×2048。源码校验逻辑CheckBandPretokenAndNexttoken包含三条关键约束preTokens 0若nextTokens 0则必须满足preTokens nextTokens 0每个 batch 需满足actual_seq_len_q - next_tokens actual_seq_len_kv。当preTokens s2Size nextTokens 0时band 退化为 rightDownCausal稀疏类型映射为RIGHT_DOWN_CAUSAL否则映射为BAND_COMPRESS。测试参考实现的 band 构造公式为test_utils.pyatten_mask_u torch.triu(torch.ones(s1, s2), diagonalnext_tokens 1 s2 - s1) atten_mask_l torch.tril(torch.ones(s1, s2), diagonal-pre_tokens - 1 s2 - s1) mask (atten_mask_u atten_mask_l).bool()与 default 模式相比对角线偏移量增加了s2 - s1的修正项这正是参数起点为右下角的体现。sparseMode5prefix 非压缩模式sparseMode 为 5 时代表 prefix 非压缩场景用于前缀prefix / few-shot 上下文可见、其余部分因果的注意力结构常见于需要让每个 query 都能看到同一段系统提示词的场景。其结构为在 rightDownCausal 的基础上左侧加上一个长为 Sq、宽为 N 的矩形区域N 的值由可选输入 prefix 获取。例如在 batch2 场景下prefix 传入数组[4, 5]表示每个 batch 轴的前缀长度可以不同第一个 batch 前 4 列全可见第二个 batch 前 5 列全可见参数起点为左上角。该场景下忽略 preTokens、nextTokens 取值。attenMask 矩阵数据格式须为 BNSS 或 B1SS即 shape 为 (B, N, S1, S2) 或 (B, 1, S1, S2) 的四维矩阵。应传入的 attenMask 示意如下在源码中prefix 非压缩模式的掩码构造逻辑为test_utils.pymask torch.triu(torch.ones(b, 1, s1, s2), diagonals2 - s1 1) for i in range(0, b): mask[i, :, :, :prefix[i]] 0即先构造 rightDownCausal 下三角再将每个 batch 前缀长度prefix[i]之前的列全部置为可见。tiling 侧通过VarLenGetPrefixNList从 prefix 输入解析每个 batch 的 N 值并映射为SparseEnum::PREFIXflash_attention_score_tiling_varlen.cpp。注意sparseMode5 在 varlen 场景下不支持见速查表备注且 prefix 的 N 值取值范围与 Sq/Skv 的关系见下文使用建议一节。sparseMode6prefix 压缩模式sparseMode 为 6 时代表 prefix 压缩场景。此时 attenMask 为优化后的压缩下三角矩形的矩阵3072×2048其结构为上半部分为 [2048, 2048] 的下三角矩阵对应 causal 部分下半部分为 [1024, 2048] 的矩形矩阵其中矩形矩阵左半部分全 0、右半部分全 1对应 prefix 全可见部分。该场景下忽略 preTokens、nextTokens 取值。测试参考实现的构造逻辑test_utils.py与文档描述完全一致upper torch.triu(torch.ones(2048, 2048), diagonal1) lower torch.cat((torch.zeros(1024, 1024), torch.ones(1024, 1024)), dim1) mask torch.cat((upper, lower), dim0).bool()tiling 侧在VarLenSparseModeProcess中对PREFIX_COMPRESS进行专门处理flash_attention_score_tiling_varlen.cpp。当用户指定 attenMask 输入时prefix 压缩场景同样需要传入对应形状的掩码BNSS/B1SS 格式且更推荐直接使用压缩格式以节省内存。sparseMode7varlen 外切场景rightDownCausal 切分sparseMode 为 7 时表示varlen 且为长序列外切场景即长序列在模型脚本中进行多卡切分 query 的 sequence length。使用前提是外切前为 sparseMode3rightDownCausal的场景。当前模式下用户需要设置 preTokens 和 nextTokens起点为右下顶点且必须保证参数正确否则会存在精度问题。以一个具体示例说明在第二个 batch 对 query 进行切分key 和 value 不切分4×6 的 mask 矩阵被切分成 2×6 和 2×6 的 mask分别在卡 1 和卡 2 上计算卡 1最后一块 mask 为 band 类型配置preTokens6保证大于等于最后一个 Skv、nextTokens-2actual_seq_qlen传入{3,5}actual_seq_kvlen传入{3,9}卡 2mask 类型切分后不变仍为 sparseMode3actual_seq_qlen传入{2,7,11}actual_seq_kvlen传入{6,11,15}。说明sparseMode7 时band 表示的是最后一个非空 tensor 的 Batch 的 sparse 类型如果只有一个 batch用户需按照 band 模式的要求来配置参数sparseMode7 时用户需要输入2048×2048 的下三角 mask作为该融合算子的输入。基于 sparseMode3 进行外切产生的 band 模式 sparse 参数应符合以下条件preTokens last_Skvlast_Skv 为最后一个非空 batch 的有效 KV 长度last_Sq - last_Skv nextTokens 0当前模式下不支持可选输入 pse位置编码。非 band 模式的 batch 应满足Sq Skv。源码侧对应的校验函数为 CheckRightDownCausalBandPretokenAndNexttoken要求preTokens lastS2最后一个有效 s2、nextTokens 0、每个 batchactual_seq_qlen actual_seq_kvlen且 band batch 满足actual_seq_len_q - next_tokens actual_seq_len_kv校验通过后映射为SparseEnum::RIGHT_DOWN_CAUSAL_BAND。测试参考实现test_utils.py中band batch 使用与 sparseMode4 相同的对角线修正公式非 band batch 使用diagonals2 - s1 1的下三角。sparseMode8varlen 外切场景leftUpCausal 切分sparseMode 为 8 时同样表示 varlen 且为长序列外切场景但使用前提是外切前为 sparseMode2leftUpCausal的场景。当前模式下用户需要设置 preTokens 和 nextTokens起点为右下顶点且必须保证参数正确否则会存在精度问题。示例在第二个 batch 对 query 进行切分key 和 value 不切分5×4 的 mask 矩阵被切分成 2×4 和 3×4 的 mask分别在卡 1 和卡 2 上计算卡 1mask 类型切分后不变仍为 sparseMode2actual_seq_qlen传入{3,5}actual_seq_kvlen传入{3,7}卡 2第一块 mask 为 band 类型配置preTokens4保证大于等于第一个 Skv、nextTokens1actual_seq_qlen传入{3,8,12}actual_seq_kvlen传入{4,9,13}。说明sparseMode8 时band 表示的是第一个非空 tensor 的 Batch 的 sparse 类型如果只有一个 batch用户需按照 band 模式的要求来配置参数sparseMode8 时用户需要输入2048×2048 的下三角 mask作为该融合算子的输入。基于 sparseMode2 进行外切产生的 band 模式的 sparse 参数应符合以下条件preTokens first_Skvfirst_Skv 为第一个非空 batch 的有效 KV 长度nextTokens first_Sq - first_Skv根据实际情况进行配置当前模式下不支持可选输入 pse。源码侧对应的校验分支为 BAND_LEFT_UP_CAUSAL要求 band batch 满足actual_seq_len_q - next_tokens actual_seq_len_kv且preTokens firstS2第一个有效 s2通过后映射为SparseEnum::BAND_LEFT_UP_CAUSAL。测试参考实现中非 band batch 使用diagonal1的标准 leftUpCausal 下三角test_utils.py。sparseMode9treeMask 模式推测解码sparseMode 为 9 时代表 treeMask 模式用于推测解码speculative decoding场景下的树形注意力掩码。用户需传入自定义的树形 maskmask 中值为 1 的位置会被遮蔽。树形 mask 矩阵特征对角线位置s1s2值为 0表示 token 关注自身上三角位置s1s2值为 1表示不关注未来 token下三角位置s1s2值为 0 或 1由树结构决定部分注意力关系即允许当前 token 关注其在推测树中的父节点/兄弟节点对应的历史 token。attenMask 输入格式根据 inputLayout 的不同treeMask 有两种输入格式inputLayout 为 BSH、BSND 或 BNSD 时attenMask 的 shape 为(B, S1, S1)每个 batch 传入 S1×S1 大小的 tree maskinputLayout 为 TND 时attenMask 为 1D 紧凑格式shape 为(ΣS1i²,)即每个 batch 的 S1i×S1i mask 拼接传入约束说明非量化场景支持 GQA 和 MLA全量化场景仅支持 MLA不支持左 padding、pseShift、sharedPrefix输出 dtype 不支持 INT8每个 batch 需满足Q_S ≤ KV_S。源码级实现佐证sparseMode 在算子内部的流转为便于读者在仓库中进一步追踪实现这里将 sparseMode 从属性声明到掩码最终生效的关键环节串起来属性声明sparse_mode定义于 flash_attention_score_proto.hATTR(sparse_mode, Int, 0)host 侧在 flash_attention_score_def.cpp 注册为可选属性tiling 解析与校验在 flash_attention_score_tiling_varlen.cpp 的GetSparseInfo中会先校验 sparseMode 取值范围varlen 场景下 prefix 及大于BAND_LEFT_UP_CAUSAL的取值会被拒绝再依据 sparseMode 与 preTokens/nextTokens 的组合推导出内部SparseEnum类型如CAUSAL、RIGHT_DOWN_CAUSAL、BAND_COMPRESS、PREFIX、RIGHT_DOWN_CAUSAL_BAND、BAND_LEFT_UP_CAUSAL等并完成全部参数合法性检查掩码生成参考不同模式的 attenMask 矩阵构造公式集中在测试参考实现 test_utils.py 的generate_npu_mask/generate_cpu_mask中可作为用户自行构造掩码时的对照标准算子文档联动aclnnFlashAttentionScore 系列 API 文档均将 sparseMode 的详细语义指向本文所对应的 sparse 模式说明例如 aclnnFlashAttentionScore.md。使用建议与注意事项汇总结合算子 API 文档aclnnFlashAttentionScore.md与 sparseMode 文档给出以下实战建议尽量使用默认值用户不特意指定时建议 sparseMode 传入 0defaultMask掩码与参数一致性配置为 0、4 时须保证 attenMaskOptional 与 preTokens、nextTokens 的范围一致配置为 1、2、3、5、6 时用户配置的 preTokens、nextTokens 不会生效为 1、2、3、4、5、6 时应传入对应正确的 attenMaskOptional否则将导致计算结果错误当 attenMaskOptional 输入为 None 时sparseMode、preTokens、nextTokens 参数不生效固定为全计算内存优化当所有 attenMaskOptional 的 shape 小于 2048 且相同时建议使用 default 模式以减少内存使用量sparseMode2/3/4/6 使用固定 2048×2048或 3072×2048的压缩掩码格式prefix 取值prefix 稀疏计算场景sparseMode5 或 6当 Sq Skv 时prefix 的 N 值取值范围为[0, Skv]当 Sq Skv 时N 值取值范围为[Skv-Sq, Skv]band 交集要求band 场景下 preTokens 和 nextTokens 之间必须要有交集varlen 外切7/8必须先确认外切前的原始场景是 sparseMode3对应 7或 sparseMode2对应 8并按上文约束设置 preTokens / nextTokens、actual_seq_qlen / actual_seq_kvlen且这两种模式均不支持可选输入 pse无效行计算如需避免整行 mask 导致的精度损失可关注innerPrecise2无效行计算配置该配置会带来性能下降算子可在可判断的场景如 sparseMode3 且 Sq Skv自动开启。总结sparseMode 是 CANN ops-transformer 注意力算子掩码体系的总开关04 覆盖从全计算、全自定义、两种因果到 band 滑动窗口的基础稀疏形态5/6 解决 prefix 前缀可见性非压缩与压缩两种实现7/8 支撑 varlen 长序列多卡外切的正确切分9 面向推测解码的树形掩码。理解各模式的参数起点左上/右下、压缩掩码形状2048×2048 / 3072×2048以及 varlen 场景的 preTokens/nextTokens 边界条件是在 NPU 上正确配置 FlashAttentionScore 系列算子的关键相关实现与校验逻辑可在 flash_attention_score_proto.h、flash_attention_score_tiling_varlen.cpp 与 test_utils.py 中进一步查阅。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价