Back to blog

SpargeAttn: Accurate and Training-free Sparse Attention Accelerating Any Model Inference

训练无关的通用稀疏注意力,加速所有模型推理(语言、图像、视频生成)

SpargeAttn: Accurate and Training-free Sparse Attention Accelerating Any Model Inference

一、论文概述

项目内容
标题SpargeAttn: Accurate Sparse Attention Accelerating Any Model Inference
作者Jintao Zhang, Chendong Xiang, Haofeng Huang, Jia Wei, Haocheng Xi, Jun Zhu, Jianfei Chen
机构Tsinghua University (thu-ml)
论文arXiv:2502.18137
代码github.com/thu-ml/SpargeAttn
发布2025-02-25 (ICML 2025)
许可—

二、核心思想

SpargeAttn 是一种通用的、训练无关的稀疏注意力算子,可应用于各类生成式模型(语言建模、文本到图像、文本到视频),在不损失端到端性能的前提下显著加速推理。

问题定义

注意力机制的时间复杂度为 O(N2d)O(N^2 d),随着序列长度增长(视频生成中可达 45K-128K),注意力成为推理延迟的主要部分。虽然注意力图通常具有稀疏性(softmax 产生大量接近零的值),但现有稀疏注意力方法面临两大挑战:

  • L1. 通用性不足: 已有方法多针对特定任务设计(如语言模型使用滑动窗口或注意力 sink),而不同任务的注意力模式差异很大(见 Fig. 2)。
  • L2. 可用性与精度的矛盾: 准确预测稀疏区域需要高开销的计算,而低开销预测又难以保证精度。例如 MInference 需要极长序列(100K)才能实现明显加速。

解决方案概述

SpargeAttn 通过三个核心技术解决上述问题:

  1. 选择性 Token 压缩的稀疏掩码预测:通过将每个 block 内的 Q/K 按 token 相似度有选择地压缩为单个 token,构建稀疏掩码,准确预测注意力图中应计算的块。该方法跨任务通用。
  2. Sparse Warp Online Softmax:在 GPU warp 级别设计的在线 softmax 稀疏算法,利用全局最大值与局部最大值的差异进一步跳过部分 PVPV 乘法,且无额外开销。
  3. 与 SageAttention 集成:将稀疏方法整合到 8-bit 量化的 SageAttention 框架中实现进一步加速。

三、技术架构

整体框架

SpargeAttn Workflow

SpargeAttn 包含两阶段在线滤波器来实现稀疏 FlashAttention:

  • Stage 1(Step 1-2):快速准确地预测注意力图中的稀疏块,跳过对应的 QiKj⊤Q_i K_j^\top 和 P~ijVj\widetilde{P}_{ij} V_j 计算。
  • Stage 2(Step 3):通过稀疏在线 softmax 进一步跳过部分 P~ijVj\widetilde{P}_{ij} V_j 计算。

核心公式

FlashAttention 在线 Softmax 基础:

Sij=QiKj⊤/d,(mij,P~ij)=σ~(mi,j−1,Sij)lij=exp⁡(mi,j−1−mij)li,j−1+rowsum(P~ij)Oij=diag(exp⁡(mi,j−1−mij))Oi,j−1+P~ijVj(1)\begin{aligned} S_{ij} &= Q_i K_j^\top / \sqrt{d}, \quad (m_{ij}, \widetilde{P}_{ij}) = \tilde{\sigma}(m_{i,j-1}, S_{ij}) \\ l_{ij} &= \exp(m_{i,j-1} - m_{ij}) l_{i,j-1} + \mathrm{rowsum}(\widetilde{P}_{ij}) \\ O_{ij} &= \mathrm{diag}\left(\exp(m_{i,j-1} - m_{ij})\right) O_{i,j-1} + \widetilde{P}_{ij} V_j \end{aligned} \tag{1}

其中 mijm_{ij} 和 lijl_{ij} 是 bq×1b_q \times 1 向量,初始化为 −∞-\infty 和 00。σ~\tilde{\sigma} 是类 softmax 算子:mij=max⁡{mi,j−1,rowmax(Sij)}m_{ij} = \max\{m_{i,j-1}, \mathrm{rowmax}(S_{ij})\}, P~i,j=exp⁡(Sij−mij)\widetilde{P}_{i,j} = \exp(S_{ij} - m_{ij})。最终输出 Oi=diag(lij)−1OijO_i = \mathrm{diag}(l_{ij})^{-1} O_{ij}。

稀疏掩码定义:

令 MgM_g 和 MpvM_{pv} 为维度 ⌈N/bq⌉×⌈N/bk⌉\lceil N/b_q \rceil \times \lceil N/b_k \rceil 的二进制掩码:

QiKj⊤,P~ijVj are skipped if Mg[i,j]=0(2)Q_i K_j^\top, \widetilde{P}_{ij} V_j \text{ are skipped if } M_g[i,j] = 0 \tag{2} P~ijVj is skipped if Mpv[i,j]=0(3)\widetilde{P}_{ij} V_j \text{ is skipped if } M_{pv}[i,j] = 0 \tag{3}

选择性 Token 压缩:

对每个 block 计算平均 token 和块内自相似度:

q={qi}={mean(Qi,axis=0)}(4)q = \{q_i\} = \{\mathrm{mean}(Q_i, \mathrm{axis}=0)\} \tag{4} k={kj}={mean(Kj,axis=0)}(5)k = \{k_j\} = \{\mathrm{mean}(K_j, \mathrm{axis}=0)\} \tag{5} sqi=CosSim(Qi),skj=CosSim(Kj)(6)s_{qi} = \mathrm{CosSim}(Q_i), \quad s_{kj} = \mathrm{CosSim}(K_j) \tag{6}

其中 CosSim(X)=XX⊤∣max⁡(XX⊤)∣\mathrm{CosSim}(X) = \frac{XX^\top}{|\max(XX^\top)|} 用于衡量 block 内 token 的相似程度。

注意力近似矩阵:

S^[i]=qik⊤;S^[:,j]=−∞, if skj<θ(7)\hat{S}[i] = q_i k^\top; \quad \hat{S}[:,j] = -\infty, \text{ if } s_{kj} < \theta \tag{7}

TopCdf 掩码选择:

对每行 P^[i]\hat{P}[i],选择累积概率超过阈值 τ⋅∑P^[i]\tau \cdot \sum \hat{P}[i] 的位置:

Mg[i,:]=TopCdf(P^[i],τ)(8)M_g[i,:] = \mathrm{TopCdf}(\hat{P}[i], \tau) \tag{8}

同时强制保留低相似度 block 的所有计算:

Mg[i,:]=1, if sqi<θ;Mg[:,j]=1, if skj<θ(9)M_g[i,:] = 1, \text{ if } s_{qi} < \theta; \quad M_g[:,j] = 1, \text{ if } s_{kj} < \theta \tag{9}

Sparse Warp Online Softmax:

定义局部最大值 mlocal=rowmax(Sij)m_{\mathrm{local}} = \mathrm{rowmax}(S_{ij}),当 max⁡(mlocal−mij)<λ\max(m_{\mathrm{local}} - m_{ij}) < \lambda 时,可近似认为 P~ijVj≈0\widetilde{P}_{ij} V_j \approx 0,从而跳过 PVPV 乘法。

在 warp 级别,设 warp 索引为 iwi_w,其覆盖的行范围为 Iw=[iwbqcw:(iw+1)bqcw]I_w = [\frac{i_w b_q}{c_w} : \frac{(i_w+1)b_q}{c_w}],若满足:

max⁡(mlocal[Iw]−mij[Iw])<λ\max(m_{\mathrm{local}}[I_w] - m_{ij}[I_w]) < \lambda

则跳过 P~ij[Iw]Vj\widetilde{P}_{ij}[I_w] V_j 的计算,直接令 Oij[Iw]≈Oi,j−1[Iw]O_{ij}[I_w] \approx O_{i,j-1}[I_w]。

Algorithm 1: SpargeAttn 实现流程

Input: Q(FP16), K(FP16), V(FP16) ∈ ℝ^(N×d), block size bq, bk_v,
       GPU Warps count cw, hyper-parameters τ, θ, λ

1. Divide Q to Tm=N/bq blocks {Qi}; divide K,V to Tn=N/bk_v blocks {Ki},{Vi}

2. Q̂i, K̂j, δQ, δK = Quant(Qi, Kj)              // per-block quantization (SageAttention)

3. qi = mean(Qi, axis=0);  kj = mean(Kj, axis=0)

4. Ŝ = qkᵀ;  sqi = CosSim(Qi);  skj = CosSim(Kj)
   Ŝ[:,j] = -∞, if skj < θ

5. P̂[i] = Softmax(Ŝ[i]);  M[i,:] = TopCdf(P̂[i], τ)
   M[i,:] = 1, if sqi < θ;  M[:,j] = 1, if skj < θ

6. for i = 1 to Tm do
7.   Load Q̂i and δQ[i] into a SM
8.   for j = 1 to Tn do
9.     if M[i,j] != 0 then
10.      Compute Si_j = Q̂i K̂jᵀ / √δQ[i]δK[j] + δBias
11.      Compute (mi_j, P̃i_j) via online softmax
12.      if Mp_v[i,j] != 0 then
13.        Compute P̃i_j Vj and accumulate to output
14.      end
15.    end
16.  end
17. end

模型组件

组件说明关键参数
Block Mask MgM_g第一阶段掩码,决定跳过哪些 QiKj⊤Q_i K_j^\top 和 PVPV 计算τ∈(0,1)\tau \in (0,1), θ∈(−1,1)\theta \in (-1,1)
PV Mask MpvM_{pv}第二阶段掩码,仅跳过 PVPV 乘法λ<0\lambda < 0
Selective Token Compression按块内 token 相似度压缩 Q/K block 到单 tokenCosSim 归一化
TopCdf基于累积分布函数的稀疏选择阈值 τ⋅∑P^\tau \cdot \sum \hat{P}
HilbertCurve Permutation空间填充曲线排列 Q/K/V 以提升块内相似性块大小 4
SageAttention Integration8-bit 量化集成,per-block quantizationFP16 → INT8

超参数确定策略

三个超参数 τ∈(0,1)\tau \in (0,1), θ∈(−1,1)\theta \in (-1,1), λ<0\lambda < 0 通过 L1 误差约束自适应确定:

L1=∑∣O−O′∣/∑∣O∣L1 = \sum |O - O'| / \sum |O|

其中 OO 为完整注意力输出,O′O' 为稀疏注意力输出。给定两层阈值 (l1,l2)(l_1, l_2):

  • 若 L1<l1L1 < l_1:增大 τ\tau(提高稀疏度)
  • 若 l1<L1<l2l_1 < L1 < l_2:保持当前参数
  • 若 L1>l2L1 > l_2:减小 τ\tau(降低稀疏度以保证精度)

各模型使用的 (l1,l2)(l_1, l_2) 值:

  • Llama3.1: (0.08,0.09)(0.08, 0.09)
  • CogvideoX / Mochi: (0.05,0.06)(0.05, 0.06)
  • Stable-Diffusion3.5 / Flux: (0.07,0.08)(0.07, 0.08)

HilbertCurve Permutation

图像和视频模型受益于强空间先验:相邻像素往往相似。为提升稀疏预测准确性,使用 Hilbert 空间填充曲线将 Q,K,V∈RT×H×W×dQ, K, V \in \mathbb{R}^{T \times H \times W \times d} 重排为 RL×d\mathbb{R}^{L \times d}(L=T×H×WL=T \times H \times W),使得空间相邻的 token 在序列中也相邻,从而提高 block 内自相似性,增加可跳过的计算比例。

四、核心创新

创新点说明理论/实验依据
选择性 Token 压缩按块内 token 相似度有选择地压缩 Q/K block 为单 token,而非简单取平均Eq.(4)-(9):CosSim 筛选确保低相似度 block 不被过度压缩
两阶段在线稀疏滤波第一阶段 MgM_g 跳过 QK⊤QK^\top 和 PVPV,第二阶段 MpvM_{pv} 仅跳过 PVPVTable 6:合并后 sparsity 达 54%,高于单独使用任一阶段
Warp-level 稀疏 Softmax在 GPU warp 粒度上利用 mlocal−mij<λm_{\mathrm{local}} - m_{ij} < \lambda 条件跳过 PVPV零额外开销——在线 softmax 已计算 mijm_{ij}
HilbertCurve 排列利用空间填充曲线提升图像/视频模型的 block 自相似性和稀疏度Table 4:HilbertCurve 的 Sim-q=0.572, Sim-k=0.479 最优
跨模型通用性训练无关,无需针对特定模型调整在 LLM、T2I、T2V 五种模型上均有效

五、代码实现分析

GitHub: thu-ml/SpargeAttn

  • 使用 CUDA 实现
  • 基于 FlashAttention 的 tiling 策略
  • 可选集成 SageAttention 的 per-block 8-bit 量化
  • 支持 Llama3.1、CogvideoX、Mochi、Flux、Stable-Diffusion3.5 等多种模型

关键实现文件(按 GitHub repo 结构推断):

  • sparge_attn.py / CUDA kernel:核心稀疏注意力实现
  • 两阶段掩码生成逻辑
  • HilbertCurve 排列工具

六、实验结果

基准测试

评估模型:

  • 文本: Llama3.1 (8B)
  • 视频: CogvideoX (2B), Mochi
  • 图像: Flux (.1-dev), Stable-Diffusion3.5 (large)

评估数据集:

  • WikiText(ppl)、Longbench、InfiniteBench En.MC、Needle-in-A-Haystack
  • Open-Sora prompt sets(视频)
  • COCO annotations(图像,FID / CLIP / ImageReward)

端到端指标对比(Table 1):

模型指标Full AttnSpargeAttnMInfer 30%FlexPref γ=0.95
Llama3.1WikiText ppl5.595.5814.43330.90
Llama3.1Longbench62.6862.6650.4430.43
Llama3.1Needle Haystack10010022.5010.00
CogvideoXVQA-a54.8554.6539.5452.48
CogvideoXVQA-t67.4567.2044.2866.74
CogvideoXFScore1.8141.8131.3751.801
MochiVQA-a54.2754.1534.6653.01
MochiVQA-t67.0166.9244.7266.66
MochiFScore1.8071.8051.1381.802
FluxFID ↓13.7113.7416.6713.86
FluxCLIP ↑30.4730.4630.3730.46
FluxIR ↑0.5740.5730.5730.574
SD3.5FID ↓11.7511.8012.3911.76
SD3.5CLIP ↑31.5131.5031.4931.50
SD3.5IR ↑0.5790.5790.5790.579

SpargeAttn 在全部模型上几乎无损,而 MInference 和 FlexPrefill 在多项指标上显著下降。

注意力内核速度对比(Fig. 9):

在 RTX4090 上,序列长度 22K,head dim 128:

  • SpargeAttn + FA2: 相比完整 FlashAttention 2 加速明显
  • SpargeAttn + Sage: 结合 8-bit 量化,速度进一步提升
  • SpargeAttn + Sage2: 结合 SageAttention2,达到最高速度
  • 随稀疏度增加,SpargeAttn 速度持续提升,且在相同稀疏度下全面超越基线

端到端生成延迟(Table 2):

模型GPUOriginalSageAttnSpargeAttn
CogvideoXRTX409087 s68 s53 s
MochiL401897 s1544 s1037 s
Llama3.1 (24K)RTX40904.01 s3.53 s2.6 s
Llama3.1 (128K)L4052 s42 s29.98 s

SpargeAttn 在 Mochi 上实现 1.83x 端到端加速,在 Llama3.1 128K 长序列上实现 1.74x 加速。

消融实验

稀疏掩码预测开销(Table 3):

Sequence LenPrediction (ms)Full Attention (ms)Overhead
8k0.2516.6493.78%
16k0.48726.831.82%
32k0.972106.680.911%
64k2.599424.240.612%
128k8.7641696.20.516%

预测开销随序列长度增长而递减,在 128K 时仅占 0.516%。

Permutation 方法对比(Table 4):

MethodSim-q ↑Sim-k ↑L1 ↓Sparsity ↑
Random0.3210.0190.04140.048
Rowmajor0.5510.3900.03070.363
Timemajor0.5140.3670.03420.338
HilbertCurve0.5720.4790.03890.392

HilbertCurve 在块自相似性和稀疏度上均最优。

Self-similarity Judge 消融(Table 5):

MethodVQA-a ↑VQA-t ↑FScore ↑
W/o self-sim Judge34.66444.7221.138
With self-sim Judge54.17967.2191.807

自相似度判断对质量至关重要,移除后各项指标大幅下降。

两阶段稀疏贡献分析(Table 6):

StrategySparsity
only MgM_g51.2%
only MpvM_{pv}27.7%
Mg+MpvM_g + M_{pv}54%

两阶段组合获得最高稀疏度,说明 MpvM_{pv} 在 MgM_g 基础上进一步跳过部分 PVPV 计算。

长序列稀疏度变化(Table 7):

在 Llama3.1 上,保持恒定精度约束下,稀疏度随序列长度增加而提高(长序列有更多可跳过的块)。

七、相关工作

三类稀疏注意力方法:

  1. Pattern-based methods(依赖固定模式):

    • H2O, InfLLM, DUOAttention — 滑动窗口
    • SampleAttention, MOA, StreamingLLM — 滑动窗口 + attention sink
    • DitFastAttn — 滑动窗口 + 注意力图相似性(仅限简单 diffusion transformer,不兼容语言模型和 MMDiT)
  2. Dynamic sparse methods(输入驱动,更通用):

    • SparQAttn, LokiAttn — 通道压缩(降低注意力维度)
    • MInference, FlexPrefill — Token 压缩(block 压缩到单 token)
    • SeerAttention — 需训练额外参数
  3. Training-based methods(需重新训练):

    • Reformer, FastAttention

其他加速方向(正交方法):kernel 优化(FlashAttention)、量化、分布式、线性时间注意力。

八、总结

核心贡献

  1. 首个训练无关的通用稀疏注意力算子:SpargeAttn 可在语言、图像、视频生成模型上统一应用,无需针对特定模型调整。
  2. 选择性 Token 压缩:通过块内 token 自相似度判断进行有选择的压缩,避免 MInference 式激进压缩导致的精度丢失。
  3. Warp-level 稀疏在线 Softmax:零额外开销的第二阶段稀疏,进一步跳过 PVPV 计算。
  4. HilbertCurve 排列:利用空间填充曲线提升图像/视频模型的稀疏度。
  5. 全面的实验验证:在 5 种模型(LLM + 2 T2V + 2 T2I)上验证,端到端指标几乎无损,加速比 2.5x-5x。

技术影响

SpargeAttn 为长序列推理提供了一种即插即用的加速方案,可与量化(SageAttention)、FlashAttention 等正交方法叠加使用。

局限性

  • 超参数 τ,θ,λ\tau, \theta, \lambda 需针对不同模型手动设定(通过 L1 误差约束自适应调节,但仍需先验调参)
  • 对于短序列场景,预测开销占比相对较高(8K 时 3.78%)
  • 稀疏度依赖于注意力图的内在稀疏性,在某些注意力均匀分布的场景下效果可能受限

九、参考资源

图片索引

图片说明文件名
Figure 1SpargeAttn 在 Mochi 上实现 1.83x 加速(L40 GPU)figure-1-speedup-mochi.png
Figure 2不同任务(视频/图像/语言)的注意力图采样模式figure-2-attention-patterns.png
Figure 3SpargeAttn 工作流程(两阶段在线滤波)figure-3-workflow.png
Figure 4各类模型中 Query 和 Key 的注意力模式示例figure-4-qk-patterns.png
Figure 5不同 Token 排列方法对比(1×6×6 空间,块大小 4)figure-5-permutation-comparison.png
Figure 6Flux 和 SD3.5 上的定性对比示例figure-6-image-comparison-flux-sd35.png
Figure 7Mochi 上的定性对比示例figure-7-video-comparison-mochi.png
Figure 8Llama3.1 上 NeedleInAHaystack 对比示例figure-8-nah-comparison-llama.png
Figure 9RTX4090 上不同稀疏度的内核速度对比figure-9-kernel-speed-comparison.png
Figure 10Llama3.1 上另一组 NeedleInAHaystack 对比figure-10-nah-comparison-llama-long.png
Figure 11Mochi 上另一组视频对比示例figure-11-video-comparison-mochi-additional.png