跳转至

Block Sparse Flash Attention

会议: NeurIPS2026
arXiv: 2512.07011
代码: https://github.com/Danielohayon/Block-Sparse-Flash-Attention
领域: LLM 效率
关键词: 块稀疏注意力、精确分数门控、阈值校准、长上下文、预填充加速

一句话总结

BSFA 在 FlashAttention-2 内先精确计算所有因果可见的 QK 分数,再用离线校准的块最大值阈值跳过低分块的 V 读取、PV 和 softmax 状态更新,在 Llama-3.1-8B 的 LongBench 上以 39.78% 对比稠密基线 40.24% 的分数,获得最长 10 个样本上 1.13× 的端到端预填充加速。

研究背景与动机

长上下文推理的首个输出 token 要等待整段输入完成预填充。FlashAttention-2 通过分块与在线 softmax 避免把完整注意力矩阵写入显存,但每个因果可见块仍需计算 QK、指数与归一化、PV。QK 负责判断相关性,PV 负责把对应内容加入输出;二者的主矩阵乘法规模相同,因此只消除注意力矩阵的显存存储并不能消除随上下文平方增长的算术工作。对长文档问答而言,很多位置最终几乎没有注意力权重,这为跳过部分值侧工作留下了空间。

许多稀疏方法在看到真实分数前就决定哪些块值得计算。例如 MInference 搜索稀疏模式,FlexPrefill 和 XAttention 利用局部或压缩分数近似选择重要区域。这能同时节省 QK 与 PV,却也会错过远处、分散或不能被代理指标准确刻画的信息;额外的选择过程在短上下文还可能比所省工作更贵。BSFA 的选择不是继续改进一个更便宜的预测器,而是保留所有 QK 计算,用已经算出的真实分数判断后半段是否值得执行。代价是不能消除 QK 的二次复杂度,收益是选择信息更可靠、门控能直接嵌在原内核里。

这里还存在一个部署问题:如果每次输入都在内核外排序所有块以求精确 top-k,额外开销会抵消稀疏收益。作者利用同一模型的层、头与位置具有相对稳定的分数分布,预先校准查表阈值,把在线排序换成一次比较。核心 idea:先付出完整 QK 的成本获得精确块最大值,再用按层、头和位置校准的阈值近似目标块预算,只对保留块继续执行值聚合与在线 softmax。

方法详解

整体框架

BSFA 是无需更新模型权重的预填充注意力替换内核。离线阶段从少量校准输入中生成不同块预算的阈值表;推理阶段输入 Q、K、V 与选定预算对应的阈值切片,输出按保留块重新归一化的注意力结果。它不在全局显存中构造完整分数矩阵,也不永久删除 KV 缓存中的 token。

流程依次是“离线阈值校准”“精确分数门控”“选择性流式聚合”:第一步提供部署参数,后两步在每个注意力调用内部完成。图中的虚线只传递离线校准产物,不表示训练梯度或额外的在线选择网络;对角因果区域直接进入聚合,非对角区域才由门控决定保留或跳过。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    C["校准样本"] --> T["离线阈值校准"]
    X["推理 Q、K、V"] --> G["精确分数门控"]
    T -.->|阈值表| G
    G -->|保留块与对角区域| A["选择性流式聚合"]
    G -->|跳过块:状态不变| N["下一键块"]
    A --> N
    N -->|尚有键块| G
    N -->|遍历结束| O["归一化输出"]

关键设计

1. 离线阈值校准:把在线 top-k 排序换成查表比较

对每个校准样本、每一层、每个注意力头和查询块位置,先计算该查询块与所有可见非对角键块的最大分数,再排序寻找第 k 个最高分数作为该样本的阈值。随后对校准样本所得阈值取平均,得到部署时使用的阈值表。不同层和头对局部上下文、起始注意力汇聚位置或远程内容的偏好不同,查询位置又决定了可见历史的数量,因此使用单个全模型阈值难以表达这些差别。

这一过程以目标非对角块数 \(k\) 为预算,而不是直接规定全序列统一稀疏率。作者使用 16 个 RULER 校准样本,并将校准任务类别与延迟评估所用类别分开;同一批阈值还用于 LongBench,而不是针对 LongBench 重新调参。多个预算可以存成多张切片,部署时切换速度与质量档位,无需再运行校准。

需要区分校准目标与实际执行:平均阈值不是新输入的第 k 大分数,阈值以上的块数可能多于或少于 k。因此论文的“fixed-k”应理解为目标预算设计,而不是每次输入严格保留 k 块的保证;它减少工作量波动的意图不能被写成完全消除了线程负载不均衡。

较早的查询位置若没有足够的非对角候选块,阈值设为 \(-\infty\),不强行裁剪这些短历史。附录还说明,超过最大校准长度的位置复用末端位置阈值;这是一项经验外推策略,并不是已经证明任意长上下文都保持相同分布。改变模型或分块形状时,阈值的语义也会改变,不能机械复用。

2. 精确分数门控:先看完整块分数,再判断值侧工作

推理时先把查询块留在片上,逐块读取键并计算缩放点积分数,再对整个查询—键分数块取最大值。这里的最大值覆盖块内所有查询与键位置,不是每个 token 独立选一次 top-k;一对位置足够相关,就能使该块保留。门控使用的是 softmax 之前的分数,而不是已经归一化的注意力概率。

\[ s_{\max}^{(i,j)}=\max_{p,q}\left[\frac{\mathbf Q_i\mathbf K_j^{\top}}{\sqrt d}\right]_{pq},\qquad \text{keep}(i,j)\iff s_{\max}^{(i,j)}\ge T_{\ell,h,i}^{(k)}. \]

式中,\(i,j\) 是概念上的查询块和键块位置,\(\ell,h\) 指层和头。精确指 QK 分数没有通过均值池化、抽样或另一模型近似,并不指整个稀疏输出与 dense attention 精确相等。所有因果可见的 QK 块仍要算;门控省的是其后的值侧处理,不是提前省掉键侧打分。

对角因果区域始终执行,并对未来位置应用因果掩码,以保留当前局部上下文。正文与图 2 的保留条件是“大于等于”,简化算法 1 则写成严格“大于”;这里沿用正文表达,等号边界需以实现为准。算法使用简化的对角块索引,而 A100 实际查询与键块尺寸不同,因此不能把它的单一 \(j=i\) 索引直接当成不等尺寸 CUDA tile 的完整映射。

块最大值的保守性有明确原因:只有整个块都没有高分位置时才被丢弃,可以保护分散的高相关 token。不过,某个极高分查询也会替块内其他查询保留该块;反过来,许多单项低分但合计有用的块仍可能被删除。最大值提供可靠的打分依据,却不是关于被删概率质量的严格误差上界。

3. 选择性流式聚合:被跳过块既不读 V,也不进入分母

若块通过门控,执行与 FlashAttention 类似的在线 softmax:更新逐查询的运行最大值,计算稳定指数,必要时重缩放既有累积输出与归一化量,再读取 V 做 PV 累加。遍历结束后,按累积归一化量缩放输出。整个过程中,分数 tile、状态和输出累积量保持在片上,而不是写回完整注意力矩阵。

若块没有通过门控,内核不读取对应 V、不执行 PV,也不计算该块的指数或更新运行最大值、归一化量和输出累积量。这一点比“只把该块的 PV 置零”更强:它连该块在 softmax 分母中的贡献也一并移除。最终输出是在保留块集合上重新归一化的近似注意力,而不是在全体键的稠密 softmax 概率下做一次精确值聚合。

这解释了为何模型质量仍需实际评估。低分块的概率质量若确实很小,删去它们后的归一化扰动也通常较小;但精确 QK 本身并不能证明删去的质量总量为零。论文采用经验校准和任务准确率验证,而没有给出所有输入上的无损保证。

每个被跳过 tile 的 QK 与 PV 主矩阵乘法各需 \(2B_MB_Nd\) FLOPs,所以不做 PV 可省去这两个矩阵乘法合计成本的一半,另有 softmax 开销节省。K 与 V 等大小,因此不读 V 可省该 tile 的 50% KV 读取量。这里的 50% 是被裁剪 tile 的局部 KV 流量或两次主矩阵乘法的理论分账,不是整层、全部注意力或整个推理流水线都减少 50%;保留块、全部 QK、投影、MLP 与内核调度仍然存在。

一个完整示例

下面是解释机制的假设示例,不是论文实验数据。设某个查询块能看到 5 个非对角候选块,部署预算为 \(k=2\),校准阈值为 3;这次输入的块最大分数依次为 1、4、2、5、3.5。

内核仍读取全部 5 个键块并计算其 QK,随后跳过最大分数为 1 和 2 的两个块,保留 4、5、3.5 对应的三个值块。保留数为 3 而不是 2,正说明离线平均阈值只近似预算。对角因果区域额外保留,不挤占这个非对角预算。

当分数为 2 的块被跳过时,之前累积的最大值、softmax 分母和加权输出全部保持原样;之后分数为 5 的块若抬高运行最大值,就对已有状态做稳定重缩放。最终只在三个保留的非对角块与对角因果区域上归一化,没有把被跳过块计入分母。若查询位置太早、候选不足,则使用 \(-\infty\) 阈值保留可用历史。

损失函数 / 训练策略

BSFA 没有新增损失函数、不训练选择器,也不微调 LLM 权重。离线校准估计的是分数分布的阈值,不是有监督学习目标;部署后依靠原模型的 QK 和预先存储的阈值执行选择。

主实验使用 A100 80GB、CUDA 12.1、FP16,查询块 128、键块 64。H100 移植改为查询块 128、键块 224,并因块尺寸变化重新校准;它沿用 FA-2 风格算法,不是已完成 FA-3 集成。A6000 在相同分块语义下复用 A100 阈值;Qwen2.5-7B 则单独校准,不能据此推断不同模型可共享阈值。

实验关键数据

主实验

表中准确率来自完整 LongBench 的 27 个任务、约 4,750 个样本;加速来自最长的 10 个 narrativeqa 样本、每个约 65K token。两列不是同一评估子集。除特别说明外,模型为 Llama-3.1-8B,设备为 A100,基线为 Dense FlashAttention-2;括号内下降是相对基线比例,不是百分点。

方法 参数 LongBench 准确率 相对下降 端到端 TTFT 加速
Dense FlashAttention-2 全部 40.24% — 1.00×
BSFA k=32 39.39% 2.1% 1.16×
BSFA k=64 39.78% 1.1% 1.13×
BSFA k=96 39.88% 0.9% 1.10×
BSFA k=128 40.03% 0.5% 1.08×
MInference default 40.06% 0.45% 1.17×
XAttention default 39.95% 0.7% 1.16×
FlexPrefill γ=0.99 38.90% 3.3% 1.06×

这些结果支持 BSFA 的可调质量—延迟档位,但不支持“表 1 中唯一达到 99% 相对准确率的方法”:MInference 与 XAttention 在此表也达到该范围。BSFA 的额外优势要结合短、中长度输入的回退幅度来看。BLASST 的 LongBench 50% 稀疏配置为 39.23%、1.11×,但运行在 H100,不应混入上表解释为同卡严格排名。

消融实验

论文主要提供预算与长度分析,而非逐模块移除消融。下表合并原文表 3 的内核计时与附录表 6 的密度数据,均对应 RULER 64K、Llama-3.1-8B、A100。预测密度以目标预算计算,实测密度为实际保留比例;二者的差异也直接反映阈值门控不保证精确 k。

配置 RULER 准确率 预测 / 实测密度 端到端加速 注意力内核加速
稠密参考 84.96% 1.00 / 1.00 1.00× 1.00×
BSFA k=192 83.08% 0.35 / 0.36±0.05 1.13× 1.38×
BSFA k=256 83.39% 0.44 / 0.45±0.05 1.09× 1.30×
BSFA k=384 84.24% 0.62 / 0.61±0.05 1.05× 1.19×

原文表 3 将稠密行标为 SDPA,附录表 6 标为 Dense FlashAttention-2;这里保留“稠密参考”表述,不自行认定这两个标签的实现身份完全相同。内核计时排除 Q/K/V/O 投影、RoPE、GQA 展开、MLP 与 LM head,因此 1.38× 不能写成端到端 1.38×。

另一项分析使用独立的 LongBench 混合长度子集,共 160 个样本、295–65,461 token。以下仅列端到端 TTFT 加速,不把完整 LongBench 准确率当成这些分桶的准确率。

方法 <4K 4–8K 8–16K 16–32K 32–65K
BSFA k=64 0.96× 0.97× 1.01× 1.05× 1.19×
MInference default 0.17× 0.31× 0.48× 0.58× 0.99×
FlexPrefill γ=0.99 0.69× 0.80× 0.93× 0.99× 1.12×
XAttention default 0.85× 0.90× 0.97× 1.01× 1.09×

关键发现

  • 降低块预算提高速度,但 RULER 64K 的 k=192 仍比稠密分数低 1.88 个百分点。不存在“QK 精确,所以质量无损”的推论。
  • BSFA 在 <8K 也有小幅回退,并非所有长度都加速;其优势是比需要额外在线选块的基线更接近稠密延迟。最长 10 个样本上的 1.13× 与混合子集 32–65K 桶的 1.19× 不能互相替换。
  • Needle-in-a-Haystack 64K 单键检索以 k=32 达到 99% 基线准确率、1.24× 加速,不能推广到需要分散证据的所有任务;其他四个主要稀疏基线未在该任务测量。

亮点与洞察

  • 保留打分,裁剪消费分数的工作:无需预测尚未看到的相关性,把选择风险转移到“低分块是否真的可忽略”。这个设计保留了精确评分依据,但诚实的收益上限也较低。
  • 校准的是预算,不是固定窗口:层、头、位置阈值能选择远程内容,LongBench 同预算滑窗 k=64 仅得 13.43%。但阈值平均后的实际块数可变,部署时应同时监测密度而非只记录 k。

局限与展望

  • 全部 QK 保留意味着算术复杂度仍是二次;收益受 PV 可跳过比例和非注意力开销共同限制,不能视作次二次注意力方案。
  • 阈值依赖模型、位置和分块形状。跨数据集迁移已有实验支持,但超长位置复用末端阈值、分布变化和低分质量累计仍缺少严格误差保证。
  • 当前针对预填充,未证明解码加速或 KV 缓存容量缩减。H100 结果对比 FA-2 风格基线,不能宣称超过 FA-3;未来可研究与 FA-3 调度或稀疏输出校正的兼容性。
  • 基线实现并非完全同质:SpargeAttention 使用 INT8 Q/K 量化,FlexPrefill 使用 bf16 且对比匹配精度的稠密基线,BLASST 使用 H100。这些条件限制了跨方法速度和质量归因。

相关工作与启发

  • vs FlashAttention-2:FA-2 精确处理全部可见块,BSFA 保留其分块与在线归一化框架但排除部分块,因此是近似注意力替换,而非另一种无损 IO 优化。
  • vs BLASST:二者都在分数之后门控;BSFA 使用按层、头、位置校准的绝对阈值,BLASST 使用相对运行最大值的门控参考。论文将 BSFA 描述为固定预算,但实际执行仍由阈值决定,不能把差异绝对化为“固定对可变”。
  • vs MInference / FlexPrefill / XAttention:这些方法先用模式或代理分数限制 QK,理论可省更多计算;BSFA 少省 QK,换取直接观察完整分数和较低在线选择开销,两类方案的合理选择依赖长度与质量要求。
  • vs KV 缓存压缩 / Δ-Attention:缓存压缩解决容量或解码带宽,Δ-Attention 校正稀疏输出偏移,均不同于 BSFA 的预填充值块选择;潜在组合需独立验证。

评分

  • 新颖性: 4/5 — 将精确块分数与细粒度离线阈值结合,贡献清晰,但与已有分数后门控方法相邻。
  • 实验充分度: 4/5 — 覆盖长文任务、预算、长度、模型与设备,但准确率和计时子集分离、跨设备基线及实现差异需保留。
  • 写作质量: 3/5 — 核心算法直观,但 fixed-k 措辞、阈值等号边界与稠密参考标签有不一致。
  • 价值: 4/5 — 适合重视质量且输入长度多变的 FA-2 预填充部署,收益务实而非数量级突破。