BLASST:一个标量阈值实现的动态块稀疏注意力(MLSys 2026)

8634 字
43 分钟
BLASST:一个标量阈值实现的动态块稀疏注意力(MLSys 2026)

AI 生成内容声明

背景:FlashAttention 解决了 IO,但没有解决 O(n2)O(n^2)#

自注意力最核心的瓶颈是二次复杂度:序列长度为 nn 时,注意力分数矩阵 S=QK\mathbf{S} = \mathbf{Q}\mathbf{K}^\topn2n^2 个元素,随后的 softmax(S)V\mathrm{softmax}(\mathbf{S})\mathbf{V} 乘加也是 O(n2)O(n^2) 的。随着模型上下文窗口从 4K 一路扩张到 128K(DeepSeek-R1、Qwen3)、甚至 1M,这个二次项从”可以忍受”变成了”系统级灾难”。TensorRT-LLM 在 Skip Softmax Attention 博客里给了一组很直观的内核级基线数据:B200 上 64K 序列、BF16 的 prefill 注意力内核单次耗时 67.8 ms,decode 每生成一个 token 的注意力内核耗时 1.2 ms——而这只是一个注意力算子,还没算 MLP 和 MoE。

FlashAttention 系列做的事情,是在不改变数学结果的前提下,把注意力的内存访问模式改造成分块融合的在线形式,让矩阵乘法吃满 Tensor Core、避免中间矩阵往返 HBM。它把常数因子榨到了极致,但算法上仍然是完整计算整个注意力矩阵:n2n^2 个分数的指数、行和、PV\mathbf{P}\mathbf{V} 乘加,一个不落。前文《FlashAttention-4:面向 Blackwell 的算法-流水线协同设计》讲过,这类内核的优化目标是让注意力跑出”矩阵乘法级别”的速度——但天花板依然是 O(n2)O(n^2) 的计算量本身。

那么问题变成了:这些 n2n^2 个分数里,有多少是真正有用的?

稀疏注意力:方向正确,落地困难#

大量工作证明了同一个经验事实:训练好的 Transformer 的注意力矩阵天然是重尾分布——少数 token 对吸引绝大部分注意力质量,绝大多数分数在 softmax 归一化后趋近于零。既然如此,直接跳过那些”注定趋近于零”的块,就能同时省下指数计算、PV\mathbf{P}\mathbf{V} 乘加和 Value 的访存。这条路线统称稀疏注意力,2024 年以来成果很多,但论文《BLASST: Dynamic BLocked Attention Sparsity via Softmax Thresholding》(MLSys 2026 Oral,Rice / UC Davis / NVIDIA / Meta 合作)把这些方法的落地障碍归纳为五条:

  1. 昂贵的预计算MInference 需要先做一次 pattern 搜索确定每个头的稀疏模式(垂直条带、斜线、块状),XAttention 需要先算反对角线和的最小值做动态规划,Quest 的 KV 块打分要额外跑一遍估计量。这些预计算的延迟开销在中等长度上下文上经常吃掉理论加速。
  2. 需要训练或改架构DuoAttention、SeerAttention 要微调出稀疏门控;DeepSeek 的 DSANSA 要在训练时引入 indexer 网络。已部署的存量模型用不上。
  3. 只加速一个阶段。MInference、XAttention、SpargeAttentionFlexPrefill 只做 prefill;H2OSnapKVRocketKV、Quest 只做 decode。而实际服务的瓶颈取决于负载:长 prompt 卡 TTFT,长生成卡 TPOT,只优化一半拿不到端到端收益。
  4. 缺少新硬件上的内核。论文在 H200/B200 上实测,很多方法的理论加速在 Hopper/Blackwell 的特性(Tensor Core 指令、TMA、warp 特化)下无法兑现。
  5. 框架侵入性。需要改模型结构或注意力接口,难以集成进 TensorRT-LLM、vLLM、SGLang 这类服务框架。

BLASST 的卖点就在这五条的”反面”上:训练无关、零预计算、prefill 与 decode 都加速、原生支持 MHA/GQA/MQA/MLA 四种注意力变体、在 Hopper/Blackwell 上给出优化内核、以单个标量参数的形式接入现有框架。论文报告的成果是:精度与稠密基线持平的前提下,71.9% 稀疏度时 prefill 加速 1.52 倍,73.2% 稀疏度时 decode 加速 1.48 倍。NVIDIA 已把它以 “Skip Softmax Attention” 的名字落地进 TensorRT-LLM,FlashInfer 侧的移植也在进行中(issue #2306)。

核心思想:被 softmax 压死的块,根本不需要算#

BLASST 的核心思想一句话:在 FlashAttention 分块在线 softmax 的过程中,用已经算出来的 running max 作为”全局最大分数的代理”,判断某个块的分数是否低到 softmax 之后可以忽略;如果是,就跳过这个块的指数计算、Value 加载和 PV\mathbf{P}\mathbf{V} 乘加。

BLASST 概览:沿注意力矩阵每一行分块处理,更新 running max、计算块内最大值,并在块最大值低于 running max 超过阈值时跳过后续计算
BLASST 概览:沿注意力矩阵每一行分块处理,更新 running max、计算块内最大值,并在块最大值低于 running max 超过阈值时跳过后续计算

图源:BLASST 论文 Figure 1(arXiv:2512.12087)

这个判定不引入任何新计算。FlashAttention 的每个线程为了数值稳定本来就要维护 running max,BLASST 只是把它”顺便”拿来做剪枝决策。

一个直观类比#

可以把每行注意力看成一场蛋糕分配:softmax 把整行的总质量(1)按 ee 的指数缩放分给每个 token,分数越接近最大值的 token 分得越多。在线 softmax 维护的 running max 相当于”目前已知的擂主分数”。处理下一个块时,先看这个块的块内最高分

  • 如果块内最高分比擂主低了超过 ln(λ)\ln(\lambda),那么指数缩放之后,这个块里每个 token 分到的质量都小于 λ\lambda,整个块的贡献可以忽略——直接跳过。
  • 如果块内最高分刷新了擂主记录,那这个块永远不会被跳过(新擂主与自己的差为 0,而 ln(λ)<0\ln(\lambda) < 0,因为 λ<1\lambda < 1)。

这个”擂主规则”顺带保证了安全性的基石:真正的全局最大值所在的块永远不会被剪掉。这为后面的误差分析提供了 Z1Z \geq 1 的关键前提。

与已有方案的本质区别#

最接近的前作是清华的 SpargeAttention(ICML 2025)。BLASST 论文明确列出了三点差异:SpargeAttention 只优化 prefill,BLASST 两阶段都做;SpargeAttention 需要额外的预测步骤(先算压缩注意力图选出块),BLASST 的决策直接复用在线 softmax 已经算出的统计量,零额外开销;BLASST 的 decode 内核还会跳过 Value 从 HBM 的加载,直接削减内存带宽这个 decode 阶段的真正瓶颈。

XAttention 用反对角线和作为块重要性的代理分数,MInference 用预计算的模式库。BLASST 与它们的分水岭在于:它用的不是任何”代理”,而是真实 softmax 计算过程中的真实统计量。分数多低算”低”,由 softmax 本身的数学性质(指数缩放)严格保证,而不是由某个启发式打分来近似。

原理详解#

铺垫:FlashAttention 的在线 softmax#

要理解 BLASST 的判定式,先要清楚 FlashAttention 分块计算时每个线程手里维护着什么。对第 ii 个查询块 Qi\mathbf{Q}_i,内核按顺序遍历 KV 块 j=1,,Tcj = 1, \dots, T_c,维护三个状态:

  • mi(j1)m_i^{(j-1)}:前 j1j-1 个 KV 块中所有分数的最大值(running max,初值 -\infty);
  • i(j1)\ell_i^{(j-1)}:softmax 分母的累计值(row sum,初值 0);
  • Oi(j1)\mathbf{O}_i^{(j-1)}:注意力输出的累计值(初值 0)。

处理第 jj 个块时,先算分数块 Sij=QiKj\mathbf{S}_{ij} = \mathbf{Q}_i \mathbf{K}_j^\top,取块内最大值 m~i(j)=rowmax(Sij)\tilde{m}_i^{(j)} = \mathrm{rowmax}(\mathbf{S}_{ij}),更新 running max:

mi(j)=max(mi(j1),  m~i(j))m_i^{(j)} = \max\left(m_i^{(j-1)},\; \tilde{m}_i^{(j)}\right)

其中 m~i(j)\tilde{m}_i^{(j)} 是当前块分数的局部最大值,mi(j1)m_i^{(j-1)} 是历史 running max,mi(j)m_i^{(j)} 是更新后的新 running max。然后计算指数化的注意力权重、更新分母和输出:

P~ij=exp(Sijmi(j))\tilde{\mathbf{P}}_{ij} = \exp\left(\mathbf{S}_{ij} - m_i^{(j)}\right)i(j)=emi(j1)mi(j)i(j1)+rowsum(P~ij)\ell_i^{(j)} = e^{m_i^{(j-1)} - m_i^{(j)}}\, \ell_i^{(j-1)} + \mathrm{rowsum}\left(\tilde{\mathbf{P}}_{ij}\right)Oi(j)=emi(j1)mi(j)Oi(j1)+P~ijVj\mathbf{O}_i^{(j)} = e^{m_i^{(j-1)} - m_i^{(j)}}\, \mathbf{O}_i^{(j-1)} + \tilde{\mathbf{P}}_{ij} \mathbf{V}_j

最后输出 Oi=Oi(Tc)/i(Tc)\mathbf{O}_i = \mathbf{O}_i^{(T_c)} / \ell_i^{(T_c)}。公式 (3)(4) 中的重缩放因子 emi(j1)mi(j)e^{m_i^{(j-1)} - m_i^{(j)}} 的意义是:如果新块刷新了最大值,之前累计的分母和输出是用旧最大值做的指数化,需要整体乘一个小于 1 的因子”贬值”到新尺标下;如果没刷新,这个因子就是 e0=1e^0 = 1,历史累计原样保留。这套机制的完整推导在前文 FlashAttention-4 文章里讲过,这里只需记住一点:每一步每个线程都会算 m~i(j)\tilde{m}_i^{(j)}mi(j)m_i^{(j)},这是”免费”的副产物。

BLASST 的判定准则#

BLASST 在算法上的全部改动,就是在更新完 running max 之后插入一次比较:

m~i(j)mi(j)<ln(λ)\tilde{m}_i^{(j)} - m_i^{(j)} < \ln(\lambda)

如果成立,跳过这个块。其中 λ(0,1)\lambda \in (0, 1) 是唯一的超参数(阈值),ln(λ)\ln(\lambda) 因此恒为负数。注意比较的对象是更新后的 running max mi(j)m_i^{(j)}(而不是更新前的 mi(j1)m_i^{(j-1)}),这样新擂主块自身的差值为 0>ln(λ)0 > \ln(\lambda),天然不会被跳过。

为什么这个判据成立?对块内任意分数 ss,都有 sm~i(j)s \leq \tilde{m}_i^{(j)}(块内最大值),而 (5) 式给出 m~i(j)mi(j)<ln(λ)\tilde{m}_i^{(j)} - m_i^{(j)} < \ln(\lambda),两边取指数:

exp(smi(j))exp(m~i(j)mi(j))<λ\exp\left(s - m_i^{(j)}\right) \leq \exp\left(\tilde{m}_i^{(j)} - m_i^{(j)}\right) < \lambda

即该块每个 token 在”以 running max 为基准”的指数化权重都小于 λ\lambda。而 softmax 的分母 Z=kexp(skM)1Z = \sum_k \exp(s_k - M) \geq 1(达到全局最大值 MM 的那个元素贡献 exp(0)=1\exp(0) = 1),所以每个被跳过的 token 的真实 softmax 权重都满足 pjk<λp_{jk} < \lambda。取 λ=103\lambda = 10^{-3},每个被跳过的 token 最多分到千分之一的质量——一个 128 token 的块整体贡献不超过 0.128。

这里有一个容易忽略的设计细节:为什么比较在 ln\ln 空间做,而不是直接比 exp(m~m)<λ\exp(\tilde{m} - m) < \lambda 因为判定本身也要花钱。分数在指数化之前是普通浮点数,做一次减法加一次比较只需要一条 FADD 加一条比较指令;而先算指数再比较,每个元素都要多一次 MUFU.EX2 指令——判定开销和它想省掉的计算同数量级,就不划算了。在 log 空间比较,等于把”指数”这一步留给了被判定为值得计算的块。

跳过了什么? 对通过判定的块(值得算),照常执行公式 (2)-(4);对被跳过的块,省掉三样东西:

  1. softmax 的指数计算(CUDA Core 负载):exp\exp 每个元素要 MUFU.EX2(指数)加 FMUL(乘缩放)加 FADD(加偏移)多条指令,行和归约还要一串 FADD;
  2. P~ijVj\tilde{\mathbf{P}}_{ij}\mathbf{V}_j 的矩阵乘(Tensor Core 负载):一次 MMA 全部省掉;
  3. Value 块 Vj\mathbf{V}_j 从 HBM 到 SRAM 的加载(内存带宽):这在 decode 阶段是主要瓶颈。

没有跳过什么? Sij=QiKj\mathbf{S}_{ij} = \mathbf{Q}_i\mathbf{K}_j^\top(BMM1)必须算——块内最大值 m~i(j)\tilde{m}_i^{(j)} 只能从它得出,这是判定的信息来源。这个”BMM1 不可省”的性质是 BLASST 整个性能模型的核心,后面内核一节会反复回到它。

误差界:跳过的质量有多少#

论文附录 B 给出了输出近似的误差上界,推导干净利落,值得完整走一遍。考虑单个 query token 的注意力输出:

y=j=1Tck=1Bcexp(sjkM)vjkZy = \frac{\sum_{j=1}^{T_c}\sum_{k=1}^{B_c} \exp(s_{jk} - M)\, v_{jk}}{Z}

其中 BcB_c 是 KV 块大小,sjks_{jk} 是第 jj 块第 kk 个 token 的注意力分数,M=maxj,ksjkM = \max_{j,k} s_{jk} 是全局最大分数,Z=j,kexp(sjkM)Z = \sum_{j,k} \exp(s_{jk} - M) 是 softmax 归一化常数,vjkv_{jk} 是 Value 向量。

第一步,单块质量界。jj 被跳过意味着 m~(j)m(j)<lnλ\tilde{m}^{(j)} - m^{(j)} < \ln\lambda。由于 running max 恒不超过全局最大(m(j)Mm^{(j)} \leq M),块内每个分数满足:

exp(sjkM)exp(m~(j)M)exp(m~(j)m(j))<λ\exp\left(s_{jk} - M\right) \leq \exp\left(\tilde{m}^{(j)} - M\right) \leq \exp\left(\tilde{m}^{(j)} - m^{(j)}\right) < \lambda

对块内 BcB_c 个 token 求和,被跳过的单个块的未归一化质量满足 kexp(sjkM)<Bcλ\sum_k \exp(s_{jk} - M) < B_c \lambda

第二步,softmax 权重界。 前文说过 Z1Z \geq 1(全局最大值的元素贡献 1),所以被跳过 token 的 softmax 权重 pjk=exp(sjkM)/Z<λp_{jk} = \exp(s_{jk} - M)/Z < \lambda

第三步,输出误差。S\mathcal{S} 为被跳过的块的集合,Vmax=maxj,kvjkV_{\max} = \max_{j,k} \|v_{jk}\|。近似输出 y^\hat{y} 与真实输出 yy 的差就是被跳过 token 的贡献之和:

yy^=jSk=1Bcpjkvjk(jSk=1Bcpjk)δVmax\|y - \hat{y}\| = \left\|\sum_{j \in \mathcal{S}}\sum_{k=1}^{B_c} p_{jk}\, v_{jk}\right\| \leq \underbrace{\left(\sum_{j \in \mathcal{S}}\sum_{k=1}^{B_c} p_{jk}\right)}_{\delta}\, V_{\max}yy^δVmax<SBcλVmax\|y - \hat{y}\| \leq \delta\, V_{\max} < |\mathcal{S}|\, B_c\, \lambda\, V_{\max}

其中 δ\delta 是被跳过 token 的 softmax 质量总和,S|\mathcal{S}| 是被跳过的块数。这个界把输出误差分解成三个可控因素:跳过的块数、每块的 token 数、阈值本身,每个因素都是线性出现的。实践中还有一个二阶修正:y^\hat{y} 只在未跳过的块上重归一化(分母从 ZZ 变成 ZZSZ - Z_{\mathcal{S}}),带来的误差是 O(δ2Vmax)O(\delta^2 V_{\max}) 量级;因为 δ1\delta \ll 1,这一项可以忽略。这个推导解释了为什么”以块为单位跳”是可控的:误差随跳过的块数线性累积,只要 λ\lambda 足够小,即使跳过 70% 的块,总质量损失也被压在很小的数上。

为什么是块级而不是 token 级#

分数矩阵里的低分 token 远比低分块多,直觉上 token 级剪枝能拿到更高的稀疏度。但 token 级剪枝在 GPU 上是坏主意,原因有二。其一,判定粒度决定判定成本:块级判定每 Bc×BcB_c \times B_c(典型 128×128)的分数子矩阵只花一次比较,token 级判定每个元素都要比较,且 token 级跳过会破坏 MMA 的规整形状——Tensor Core 的矩阵乘按 tile 执行,稀疏成锯齿状的 P\mathbf{P} 无法喂给 MMA。其二,访存连续性:Value 的加载和 PV\mathbf{P}\mathbf{V} 的乘加都以块为单位才有规整的地址模式和预取机会。块级是”硬件友好”与”稀疏收益”之间的折中点,这也是为什么 MInference、XAttention、SpargeAttention 这些工作不约而同选择块级粒度的原因。

校准:λ\lambda 与上下文长度成反比#

有了判定准则,剩下的问题是最关键的那一个:λ\lambda 取多少? 取大了精度崩,取小了没加速。BLASST 的回答分两步:先发现一个反比规律,再设计一个自动化校准流程。

固定阈值为什么不行#

第一个实验结论:精度损失由稀疏度决定,而不是由阈值决定。论文在 Llama-3.1-8B 上跑 RULER 的困难子集(NIAH_MULTI、VT、FWE),画出相对精度损失随稀疏度的曲线:不同任务、不同长度(8K-64K)的曲线高度重合——稀疏度到 60-70% 之前精度几乎不掉,之后急剧下滑。这说明部署时应该控制的变量是稀疏度,阈值只是实现手段。

不同任务与上下文长度下,相对精度损失随稀疏度的变化曲线高度重合:稀疏度 60-70% 前精度几乎不掉
不同任务与上下文长度下,相对精度损失随稀疏度的变化曲线高度重合:稀疏度 60-70% 前精度几乎不掉

图源:BLASST 论文 Figure 2(左)(arXiv:2512.12087)

但第二个实验结论马上泼了冷水:同一个阈值在不同长度下产生的稀疏度天差地别。论文的表 6 数据很有说服力——目标 50% 稀疏度,固定 λ=1×103\lambda = 1\times10^{-3} 时,实测稀疏度从 4K 的 23.09% 一路漂到 64K 的 74.63%;目标 70% 时固定 λ=3×103\lambda = 3\times10^{-3},4K 只有 42.35%,64K 却高达 84.63%。固定阈值意味着:请求的上下文长度一变,稀疏度(进而精度和速度)就失控,这在生产环境不可接受。

同一阈值在不同上下文长度下产生的稀疏度差异显著,说明必须对阈值做长度校准
同一阈值在不同上下文长度下产生的稀疏度差异显著,说明必须对阈值做长度校准

图源:BLASST 论文 Figure 2(右)(arXiv:2512.12087)

λ=a/L\lambda = a/L:反比关系的直觉与依据#

论文发现的规律是:最优阈值与上下文长度成反比:

λ=aL\lambda = \frac{a}{L}

其中 aa 是模型相关的比例系数(取决于目标稀疏度),LL 是上下文长度。这个反比关系有扎实的直觉支撑:softmax 把每一行的总质量归一化为 1,序列越长,这 1 份质量就要分给越多的 token,平均每个 token 的分数就越低。要让”被跳过的块”的质量占比保持不变(即固定稀疏度),衡量”低”的标尺 λ\lambda 就必须随长度同步缩小。

论文表 12 里还有一个看似奇怪的数据需要解释:同一模型同一目标稀疏度下,prefill 的比例系数 a900a \approx 90011001100,而 decode 的 aa 只有 4.6–11.7,相差两个数量级。原因藏在块级判定的聚合方式里:prefill 的一个查询块包含 BcB_c 个 query 行(典型值 128),内核只有在块内所有行的判定都同意跳过时才跳过整个块(kernel 一节会讲到 warp 内 VOTE 加 warpgroup 级 ATOMIC 的聚合实现)。这是对 BcB_c 行取逻辑与的条件——任何一行关注了这个块,整个块就必须保留。因此同一阈值下,prefill 的块稀疏度天然远低于单行稀疏度,想达到 50% 的块稀疏度就必须放宽淘汰线(更大的 λ\lambda);decode 每步只有一个 query,块判定就是单行判定,块稀疏度与行稀疏度直接相等,小阈值就能剪掉一半。这个机制差异也带来一个额外收益:decode 的实测稀疏度直接反映了”当前这一行注意力有多集中”,检索型任务(niah_single)注意力更集中,因此同样的目标稀疏度对应更小的 aa(4.6),与表 12 的数据吻合。

校准算法:一次前向,一个指数模型#

比例系数 aa 怎么定?论文给出一个只需要一次前向的校准流程(Algorithm 2)。思路是反着用上述反比关系:对校准数据集 D\mathcal{D} 中的每条样本 (xi,Li)(x_i, L_i) 和候选阈值 λj\lambda_j,测量该阈值下实际达到的稀疏度 sijs_{ij},记录数据点 (λjLi,  sij)(\lambda_j \cdot L_i,\; s_{ij})。由于同一个分数矩阵可以同时算出所有候选阈值对应的稀疏度(不同阈值只是移动淘汰线),整轮校准只需要一次前向传播。收集完数据点后,拟合指数模型:

λL=αexp(βs)\lambda \cdot L = \alpha \cdot \exp(\beta \cdot s)

其中 ss 是稀疏度,α\alphaβ\beta 是拟合参数。指数形式反映的是注意力分数的重尾分布:阈值小幅上调就能剪掉大量低分块(ss 快速增长),之后进入”只剩高分块”的平台期,再加大阈值收益递减。推理时给定目标稀疏度 SS,阈值直接按 λ=αexp(βS)/L\lambda = \alpha \exp(\beta S) / L 计算——注意它保持了对长度的反比依赖,同时允许运行时动态调整目标稀疏度而无需重新校准

论文的实现细节:从 RULER 采样约 1000 条序列,覆盖 4K/8K/16K/32K/64K 五种长度,拟合出 α\alphaβ\beta;校准后的 λ=a/L\lambda = a/L 在 50% 目标下把实测稀疏度的平均偏差控制在 1.2% 以内(表 6)。跨任务稳定性方面(表 12),在六类子数据集上分别校准,prefill 的 aa 稳定在 900–1100 区间;decode 的 aa 波动稍大(检索型任务 niah_single 为 4.6,多键任务为 11.4),但论文指出这源于检索型任务注意力更集中、分数分布更陡,且混合数据集上单次校准即可覆盖各类负载——实际部署不需要按任务重调。

Kernel 设计:一个分支,两条流水线#

算法层面的改动只有一行 if,但把这一行 if 放进 Hopper/Blackwell 的内核而不伤性能,是 BLASST 工程上的主体工作。设计目标两条:判定逻辑零开销针对各阶段的真实瓶颈做裁剪

判定本身:几乎零成本的秘密#

块级判定需要三类指令:每个线程基于比较结果设置谓词寄存器;warp 内做一次 VOTE(__all_sync 类指令)确认整个 warp 是否一致同意跳过;然后每个 warp 派一个线程向共享内存做一次 ATOMIC,汇总出 warpgroup 级的块判定。全部加起来只有几条指令,且被刻意安排在已有的计算指令之间——prefill 里藏在 Tensor Core MMA 发射间隙后,decode 里藏在 HBM 加载等待里。论文用一组数据证明了这个”零开销”不是空话:0% 稀疏度(阈值设到极小,永不触发跳过)时,内核速度是稠密基线的 0.96–1.00 倍,判定逻辑的开销完全被隐藏。

Prefill 内核:砍计算#

Prefill 阶段的注意力是 compute-bound 的:TensorRT-LLM 博客里的基线数据显示,B200 上 64K 序列 BF16 prefill 内核跑到 1038 TFLOP/s(有效算力利用率已经很高),H200 上 610 TFLOP/s——瓶颈在算力而非带宽。所以 prefill 内核的剪枝重点是跳过计算:被判定跳过的块不执行指数、行和归约和 PV\mathbf{P}\mathbf{V} 的 MMA,省下的全是 CUDA Core 和 Tensor Core 的指令。

但有一个反直觉的设计决策:prefill 内核不跳过 Value 的 HBM 加载。论文给出三个理由:(1) prefill 不是带宽瓶颈,省带宽换不来时间;(2) 预取流水线依赖可预测的访存模式,条件化加载会打乱预取节奏;(3) 条件加载的判断-跳转延迟本身可能超过省下的带宽时间。既然计算被砍掉了,省出的执行单元会让后续操作提前发射,整个流水线被”压缩”——论文图 3 的调度图显示,50% 稀疏度下,FlashAttention-4 的 prefill 流水线从 18 个时间单位压缩到 14 个。

FlashAttention-4 与 BLASST 的 prefill 流水线调度对比:跳过块的 MMA 与指数计算后,流水线被压缩
FlashAttention-4 与 BLASST 的 prefill 流水线调度对比:跳过块的 MMA 与指数计算后,流水线被压缩

图源:BLASST 论文 Figure 3(arXiv:2512.12087);上为正常 FlashAttention-4 调度,下为 BLASST 跳过部分块后的调度

Decode 内核:砍带宽#

Decode 阶段恰好相反:每步只有一个 query,QK\mathbf{Q}\mathbf{K}^\topPV\mathbf{P}\mathbf{V} 的计算量微不足道,真正的瓶颈是从 HBM 读 KV 缓存。H200 上 64K 序列 decode 内核的实测带宽 4.37 TB/s,B200 上 7.10 TB/s,都已贴近各自 HBM 的理论带宽上限。所以 decode 内核的剪枝重点是跳过被判定块对应的 Value 加载——直接减少 HBM 流量,省下的时间与稀疏度成正比。

真正的难点在流水线结构上。朴素实现里,Value 的加载在 BMM1 之前就发射了(预取),但 BLASST 的判定要等 BMM1 算完、块内最大值出来之后才能做。如果改成”算出判定再决定是否加载 Value”,就引入了一条 scoreboard 依赖:下一批 Value 加载必须等上一批判定完成,加载流水线被串行化,产生气泡——省下的带宽被气泡吃掉。

BLASST 的解法叫 batched load scheduling(批量加载调度):不逐块”算一个判一个”,而是把 BB 个连续的 KjQ\mathbf{K}_j^\top\mathbf{Q}(BMM1)一次性算完,各自把分数块 Sj\mathbf{S}_j 存在共享内存的小缓冲区里(decode 的 query 长度为 1,每个分数块只有 128 个数,BB 个缓冲区的开销很小);然后对这 BB 个块的判定结果做一次批量处理,只为通过判定的块统一发射 Value 加载。依赖从”逐块的串行判定-加载”变成”批量判定-批量加载”,消除了气泡。论文图 4 的调度图显示,跳过三个块的场景下 Value 加载阶段从 38 个时间单位缩短到 31 个。

FlashAttention-4 与 BLASST 的 decode 流水线调度对比:BLASST 通过批量加载调度,只对被判定保留的块统一发射 Value 加载
FlashAttention-4 与 BLASST 的 decode 流水线调度对比:BLASST 通过批量加载调度,只对被判定保留的块统一发射 Value 加载

图源:BLASST 论文 Figure 4(arXiv:2512.12087);上为正常 FlashAttention-4 decode 调度,下为 BLASST 批量加载调度跳过部分块后的调度

对 MLA 这类架构还有一个补充优化:MLA 的 decode 因为 KV 状态被压缩到低秩潜空间,可能反而变成 compute-bound,此时 decode 内核在跳过 Value 加载之外还会跳过 softmax 计算,两头都省。

性能上限:为什么是 1.8 倍#

TensorRT-LLM 博客的结论部分点明了一个重要性质:因为 BMM1(QK\mathbf{Q}\mathbf{K}^\top)永远要算,kernel 级加速的理论上限约 2 倍(QK\mathbf{Q}\mathbf{K}^\topPV\mathbf{P}\mathbf{V} 各占一半 FLOPs,加上 softmax 开销后实际封顶约 1.8 倍)。这意味着 BLASST 是常数因子加速:它不改变二次复杂度,只是把常数砍掉一块。想获得数量级加速,仍然要靠 KV 压缩、训练型稀疏架构(NSA/DSA)或线性注意力这些改变渐近线的方法;BLASST 的定位是在”存量模型 + 现成框架”上拿到一个确定、稳定、低风险的 1.3–1.8 倍。

实验:精度与速度的双重证据#

精度:50% 稀疏度几乎无损,偶尔还涨#

论文主表(Table 2)在 Llama-3.1-8B 和 Qwen3-8B 上测试三种部署场景:仅 prefill 稀疏(RULER-32K、LongBench)、仅 decode 稀疏(MATH500、AIME 2024、GPQA)、两阶段同时稀疏。50% 目标稀疏度下所有指标与稠密基线持平,75% 下损失仍然很小。其中多个指标甚至超过稠密基线:Qwen3-8B 在 50% 稀疏度下 MATH500 得 96.23(稠密 95.87),AIME 2024 得 76.50(稠密 75.00)。论文对这种现象给出两个解释:长上下文任务中信息天然稀疏,剪掉低分块相当于把概率质量重新集中到真正相关的 token 上,是一种隐式去噪;长程推理任务中部分中间步骤冗余甚至有害(“overthinking”),跳过它们的低分块恰好过滤了这些干扰。

对比基线(Table 3、Table 4)进一步说明阈值判定的质量。Prefill 对比中,BLASST(约 50% 稀疏度)RULER 平均 92.87,超过 XAttention 的 92.44、FlexPrefill 的 87.72,大幅领先 MInference 的 84.15,且与稠密基线的 93.21 几乎无差。Decode 对比中,Qwen3-8B 六项任务平均 BLASST 得 68.97,超过稠密基线的 68.57,而 RocketKV 只有 66.91、Quest 只有 60.75——注意 Quest 在 RULER-32K 上只有 56.23(稠密 91.90),说明基于估计量的 token 级剪枝在检索任务上会漏掉关键信息,而 BLASST 用真实 softmax 统计量做判定,不犯这类错误。

RULER-16K 上高稀疏度区间的精度-稀疏度权衡:BLASST 相比 XAttention 退化更平稳
RULER-16K 上高稀疏度区间的精度-稀疏度权衡:BLASST 相比 XAttention 退化更平稳

图源:BLASST 论文 Figure 8(arXiv:2512.12087)

Kernel 速度:随稀疏度单调上升#

内核级加速(Table 5,基线为 FlashAttention-3 BF16):

稀疏度(B200 prefill)38.9%49.2%63.0%71.9%80.8%94.2%
加速比1.25×1.33×1.43×1.52×1.61×1.77×
稀疏度(B200 decode)36.9%46.7%61.2%73.2%82.6%92.0%
加速比1.18×1.25×1.34×1.48×1.64×1.79×

Hopper 上的趋势一致:H200 prefill 在 71.0% 稀疏度达到 1.52×,decode 在 70.5% 达到 1.40×。两个规律值得注意。第一,加速比随稀疏度近乎线性增长,且两条曲线都明确穿过”约 50% 稀疏度 ≈ 1.25–1.33ד的点——这正好是精度无损区间的上限,也就是实际部署的甜点。第二,decode 在低稀疏度(24% 附近)只有 1.08×,说明小稀疏度下省下的 Value 带宽有限;prefill 则从 38.9% 起就有 1.25×,因为计算端的节约更直接。

端到端:TTFT 与 TPOT 的真实收益#

TensorRT-LLM 博客给出了单卡 H200/B200 上的端到端数据(Qwen3-30B-A3B-Instruct-2507,LongBench V1:平均输入 10K、输出 6 token、并发 64;LongBench V2 medium:平均输入 130K、输出 200 token、并发 1)。LongBench V1 场景下,H200 的 TTFT 从 9419.61 ms 降到 8107.73 ms(0.8 稀疏度,约 1.16×),TPOT 从 1731.80 ms 降到 1507.82 ms;B200 的 TTFT 从 4854.55 ms 降到 4150.44 ms。注意这个场景 TTFT 里混入了并发 64 下的排队效应,真实的内核加速被稀释了。看并发 1 的 LongBench V2 更干净:H200 TTFT 从 16486.70 ms 降到 12507.95 ms(0.9 稀疏度,约 1.32×),B200 从 6990.59 ms 降到 5276.67 ms;但 TPOT 几乎不动(H200 从 9.34 ms 到 8.42–8.61 ms 区间徘徊)——单请求 decode(batch=1)省不下带宽,这正是下一节局限部分要展开的问题。

Qwen3-30B-A3B-Instruct 在 LongBench V1 上、H200/B200 的端到端加速随目标稀疏度变化
Qwen3-30B-A3B-Instruct 在 LongBench V1 上、H200/B200 的端到端加速随目标稀疏度变化

图源:BLASST 论文 Figure 5(arXiv:2512.12087)

更大模型、MLA 与组合性#

附录实验补齐了可扩展性的证据:Qwen3-30B-A3B 在 LongBench V1 上 70% 稀疏度仍保持 47.21(稠密 47.77),V2 在 60-70% 稀疏度反而升到 39.53(稠密 36.28);Llama-3.1-70B 在 RULER-hard 上 80% 稀疏度仍保留 97% 以上精度;DeepSeek-R1(MLA,NVFP4)在 GPQA Diamond / MMLU Pro / LiveCodeBench 上 60% 稀疏度精度几乎无差——证明 MLA 的潜空间注意力同样适用阈值判定。组合性实验(Table 7)显示 BLASST 可以与 XAttention(prefill)、RocketKV(KV 压缩)正交叠加:XAttention+BLASST 组合 RULER-16K 92.89(稠密 93.22),BLASST+RocketKV 92.60,损失都远小于叠加各自的损失之和。最后,200K 上下文的 RepoQA 实验(Table 8)显示超长上下文天然稀疏度更高:Qwen3-Coder-30B 在 200K 下 prefill 达到 57.5% 稀疏度,精度从 0.850 降到 0.841,再叠加 decode 稀疏(40.8%)也只降到 0.838。

工程落地:TensorRT-LLM 与 FlashInfer#

BLASST 的工程化名字叫 Skip Softmax Attention,已进入 TensorRT-LLM 主线。接入方式只有一个配置项:

from tensorrt_llm import LLM
from tensorrt_llm.llmapi import SkipSoftmaxAttentionConfig
sparse_attention_config = SkipSoftmaxAttentionConfig(
threshold_scale_factor={"prefill": 587.18, "decode": 16.52}
)
llm = LLM(model="Qwen/Qwen3-30B-A3B-Instruct-2507",
sparse_attention_config=sparse_attention_config)

threshold_scale_factor 就是校准一节里的比例系数 aa(实际阈值 = 系数 / 上下文长度),prefill 与 decode 可以分开配置。以 Qwen3-30B-A3B 为例,50% 目标稀疏度对应的系数是 prefill 587.18、decode 16.52;70% 对应 3293.04 和 118.62。也可以通过 trtllm-serve 的 YAML 选项在 OpenAI 兼容端点上一键开启;NVIDIA Model Optimizer 提供自动校准,给定目标稀疏度直接产出系数。

落地形态上有两点值得一提。第一,内核是逐架构移植的:Hopper 的 prefill 基于 fmha_v2、decode 基于 XQA 内核改造,Blackwell 基于 trtllm-gen 的 warp 特化流水线改造——这也是为什么论文强调”现代硬件上的优化内核”是其他稀疏注意力方法普遍缺失的一环。第二,兼容性来自”只是内核的近似计算”这个定位:因为它不改模型结构、不加额外的选择步骤,与 FP8 注意力、KV cache 复用、chunked prefill、in-flight batching 等既有特性天然兼容,端到端收益可预测。FP8 注意力的增益比 BF16 小(博客原文指出),因为 FP8 本身已经把带宽和计算各压缩了一半,留给剪枝的空间变小。

开源侧,FlashInfer 的 issue #2306 记录了这一特性的移植计划(Hopper 侧基于 fmha_v2 分支进行中,Blackwell 侧复用 trtllm-gen 内核);论文配套的可复现实验仓库在 cameronshinn/blasst-ae-mlsys26(Apache 2.0,Docker 环境,H200/B200 内核基准测试脚本)。

局限与未解决的问题#

加速上限 1.8 倍。 BMM1 不可跳过的结构性约束决定了 BLASST 是常数因子优化:哪怕稀疏度推到 95%,内核加速也封顶在 1.8 倍左右,复杂度仍是二次的。百万级上下文的场景,它救不了命。

小 batch decode 收益有限。 LongBench V2 并发 1 的端到端数据显示 TPOT 几乎不动。原因在 roofline:batch=1 的 decode 里,Value 加载只是 KV 加载的一部分(Key 还是要全部读),且单请求的带宽利用率本来就低(H200 decode 内核在 batch=1 时有效带宽远低于峰值)。博客坦言”小 batch 长上下文服务的 decode 优化还在进行中”。decode 的加速在中等以上 batch(论文内核基准是 batch 128-148)才明显。

阈值依赖校准。 每个模型、每个阶段、每个目标稀疏度都需要标定系数;decode 的 aa 在检索型任务上漂移(4.6 与 11.7 之差意味着稀疏度几倍的差异)。自动化校准把成本压到一次前向,但”校准数据分布 ≠ 线上分布”的工程风险仍在。

稀疏度超过 75% 后精度悬崖。 60-70% 是经验安全区,再往上精度急剧下滑。论文给出的两条出路:一是稀疏度感知训练——微调时就把 BLASST 接在前向里(被跳过的块自然收不到梯度),让模型学会把信息集中在高分块上,实验显示可以把同等稀疏度下的精度损失最多降低 1.7 倍,甚至低稀疏度下反超稠密基线;二是与其他方法正交组合,例如与 XAttention、RocketKV 叠加,或与细粒度通道/头剪枝方法结合。不过前者引入训练成本,与”训练无关”的卖点有所偏离,属于可选的进阶路线。

稀疏度感知训练把精度-稀疏度前沿外推:训练时接 BLASST 的模型在同等稀疏度下保持更高精度
稀疏度感知训练把精度-稀疏度前沿外推:训练时接 BLASST 的模型在同等稀疏度下保持更高精度

图源:BLASST 论文 Figure 6(arXiv:2512.12087)

判定方向的固有偏差。 判定用 running max 做全局最大值的代理,处理顺序会影响早期判定的质量——序列开头 running max 还很低,块更容易”过关”;论文对 tile 行序重排的研究(附录图 9)结论是影响依赖数据集、总体可忽略,但这个问题在极端长上下文下没有系统验证。另外,被跳过的块贡献的误差是单侧的(只减不加),多层累积后对输出的系统性偏移有多大,论文没有给出端到端的理论分析。

tile 行序重排对精度-稀疏度权衡的影响:顺序 Cummax 与逆序 Cummax 差异可忽略
tile 行序重排对精度-稀疏度权衡的影响:顺序 Cummax 与逆序 Cummax 差异可忽略

图源:BLASST 论文 Figure 9(arXiv:2512.12087)

与训练型稀疏的关系。 BLASST 与 NSA/DSA 解决的是同一问题的两个层次:后者在训练时学会”哪些位置值得看”,前者在推理时用统计量判断”哪些块不值得算”。二者理论上可以叠加(在 DSA 的 indexer 选出的块之上再做 softmax 阈值筛选),但目前没有公开实验验证。

小结#

BLASST 的价值在于它把稀疏注意力的”决策成本”降到了零:不引入代理分数、不引入预测步骤、不需要训练、不需要预计算,只用在线 softmax 本来就要维护的 running max,加一次比较,就换来了 prefill 1.52×、decode 1.48× 的确定加速。它的误差界严格、校准自动化、内核针对 Hopper/Blackwell 逐架构优化,并且已经以 Skip Softmax Attention 的名字在 TensorRT-LLM 里可用了——对一个存量的稠密注意力模型,这是目前接入成本最低、风险最可控的长上下文加速手段之一。

另一方面要清醒:1.8 倍的内核上限提醒我们这只是常数因子战争中的一役。长上下文的真正解法仍然要回答”哪些信息值得保留与计算”这个结构性问题——KV 压缩、训练型稀疏、线性注意力与 BLASST 这类统计量剪枝的叠加,才是 2026 年推理优化最值得跟踪的组合拳。

参考资料#

  1. BLASST: Dynamic BLocked Attention Sparsity via Softmax Thresholding(arXiv)
  2. BLASST 论文页(MLSys 2026 Oral)
  3. Accelerating Long-Context Inference with Skip Softmax Attention(TensorRT-LLM 官方博客)
  4. Accelerating Long-Context Inference with Skip Softmax in NVIDIA TensorRT-LLM(NVIDIA Developer Blog)
  5. FlashInfer issue #2306:BLASST 内核移植计划
  6. BLASST MLSys 2026 Artifact 复现仓库
  7. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
  8. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision
  9. SpargeAttention: Accurate Sparse Attention Accelerating Any Model Inference
  10. XAttention: Block Sparse Attention with Antidiagonal Scoring
  11. MInference 1.0: Accelerating Pre-filling for Long-Context LLMs via Dynamic Sparse Attention
  12. Quest: Query-Aware Sparsity for Efficient Long-Context LLM Inference
  13. RocketKV: Accelerating Long-Context LLM Inference via Two-Stage KV Cache Compression
  14. DeepSeek-V3.2-Exp(DeepSeek Sparse Attention)

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

BLASST:一个标量阈值实现的动态块稀疏注意力(MLSys 2026)
https://pinghaoyang.com.cn/aigc/posts/blasst/
作者
平昊阳
发布于
2026-08-15
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
平昊阳
乘长风,破巨浪, 展鸿图于未央!
--
总访问量
--
访客数
公告
欢迎来到我的个人博客!欢迎关注交流吖!
更多相关公告,见
社交-留言」。
音乐
封面

音乐

暂未播放

0:000:00
暂无歌词
站点统计
文章
66
分类
16
标签
93
总字数
477,284
运行时长
0
最后活动
0 天前

文章目录