SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse–Linear Attention

algorithm 2509.24006
sparse-attentionlinear-attentiondiffusion-transformervideo-generationgpu-kernel

SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse–Linear Attention #

Jintao Zhang, Haoxu Wang, Kai Jiang, Shuo Yang, Kaiwen Zheng, Haocheng Xi, Ziteng Wang, Hongzhou Zhu, Min Zhao, Ion Stoica, Joseph E. Gonzalez, Jun Zhu, Jianfei Chen Tsinghua University, UC Berkeley | 2025-09 | Code:

§1 TL;DR #

DiT 注意力权重天然分解为高秩稀疏(<10%)+ 低秩密集(>90%)两部分。SLA 据此将注意力块三分类为 critical(FlashAttention $O(N^2)$)、marginal(线性注意力 $O(N)$)、negligible(跳过),融合为单一 Triton kernel,在 Wan2.1-1.3B 上实现 95% 注意力计算削减、13.7× kernel 加速、2.2× 端到端加速且无质量损失。


§2 Q1 / Q2 / Q3 #

Q1 痛点 #

DiT 视频生成中注意力是首要计算瓶颈:序列长度 10K–100K,复杂度 $O(N^2 d)$。两条加速路线均遇瓶颈:

根本原因:完整注意力权重的 stable rank 较高,线性注意力被限制在秩 ≤ $d$ 的子空间内无法逼近;而稀疏注意力丢弃了大量非零但值小的权重。

Q2 方法 #

核心洞察:注意力权重 $P$ 可分解为

$$P = \underbrace{P \odot M}_{\text{高秩稀疏}} + \underbrace{P \odot (1-M)}_{\text{极低秩密集}}$$

top 8% 的权重贡献了几乎全部 stable rank,而 bottom 92% 的 stable rank 接近 1。这意味着高秩部分必须精确计算,低秩部分可以用线性注意力高效近似。

SLA 算法

  1. Compressed attention prediction:对 $Q, K$ 做 mean pooling,将 $N \times N$ 注意力降至 $(N/b_q) \times (N/b_{kv})$ 的小矩阵 $P_c = \text{Softmax}(\text{pool}(Q) \cdot \text{pool}(K)^\top / \sqrt{d})$。
  2. 三分类:$M_c[i,j] = 1$ (top $k_h$% = critical), $M_c[i,j] = -1$ (bottom $k_l$% = negligible), $M_c[i,j] = 0$ (marginal)。
  3. 混合计算:critical 块用 FlashAttention(OnlineSoftmax),marginal 块累加预计算的 $h_j = (\phi(K_j))^\top V_j$ 实现 $O(Nd^2)$ 线性注意力,negligible 块跳过。
  4. 融合输出:$O = O^s + \text{Proj}(O^l)$,其中 $\text{Proj}$ 是可学习线性变换(初始化为零),弥合 softmax 与线性注意力的分布差异。
  5. 微调:替换原始注意力为 SLA,在预训练一致数据上微调 2000 步(<0.1% 预训练成本)。
  6. 核心技术壁垒:块级 compressed attention prediction 的精度与效率平衡。mean pooling 将复杂度从 $O(N^2)$ 降至 $O((N/b)^2 d)$,但必须准确到不将 critical 块误分为 marginal/negligible。论文选择 $b_q = b_kv = 64$,使得预测成本 < 0.5% 全注意力 FLOPs。若 pooling 精度不足(例如方差大的块内 mean 失真),整个方法将退化为低精度稀疏注意力。

    Q3 结果 #

    指标
    注意力 FLOPs 削减95%(52.75T → 2.74T)
    注意力 kernel 加速(前向)13.7× vs FlashAttention2
    注意力 kernel 加速(反向)~6.2× vs FlashAttention2
    端到端视频生成加速2.2×(185s → 88s)
    注意力延迟97s → 11s
    VBench 质量(VA/VT)76.96/83.92 vs Full Attention 76.78/82.88
    微调成本2000 步,batch 64
    评测 GPURTX 5090
    评测模型Wan2.1-1.3B,序列长度 ~30K

    §3 架构 / 方法图 #

    Figure 4: SLA architecture overview

    Paper Figure 4, verbatim (caption: "Overview of SLA. The left figure illustrates the high-level idea: attention weights are classified into three categories and assigned to computations of different complexity. The right figure shows the detailed forward algorithm of SLA using the predicted compressed attention weights.").

    左图展示 SLA 的核心思想:$N \times N$ 注意力矩阵被划分为红色 critical 块($O(N^2)$ FlashAttention)、黄色 marginal 块($O(N)$ 线性注意力)、灰色 negligible 块(跳过)。右图是完整前向流程:$Q, K$ 经 mean pooling 得到压缩表示,计算 $P_c = \text{Softmax}(Q_c K_c^\top / \sqrt{d})$,通过 TopK/BottomK 生成三分类掩码 $M_c$,$M_c$ 引导三路并行计算后融合输出。

    flowchart TB subgraph Input Q["Q ∈ ℝ^(N×d)"] K["K ∈ ℝ^(N×d)"] V["V ∈ ℝ^(N×d)"] end subgraph MaskPrediction["Compressed Mask Prediction — O((N/b)²d)"] POOL_Q["mean_pool(Q) → Q_c ∈ ℝ^(N/b×d)"] POOL_K["mean_pool(K) → K_c ∈ ℝ^(N/b×d)"] Pc["P_c = Softmax(Q_c · K_c⊤ / √d)"] Mc["M_c: TopK→critical(1), BottomK→negligible(−1), rest→marginal(0)"] Q --> POOL_Q K --> POOL_K POOL_Q --> Pc POOL_K --> Pc Pc --> Mc end subgraph LinearPrecomp["Linear Attention Precomputation — O(T_n · d²)"] PHI["φ(K), φ(Q) via feature map (softmax)"] HZ["h_j = φ(K_j)⊤ · V_j, z_j = rowsum(φ(K_j)⊤)"] K --> PHI Q --> PHI V --> HZ PHI --> HZ end subgraph FusedKernel["Fused Triton Kernel — per query block i"] BRANCH{"M_c[i,j] ?"} SPARSE["critical: FlashAttention + OnlineSoftmax"] LINEAR["marginal: H_i += h_j, Z_i += z_j"] SKIP["negligible: skip"] Mc --> BRANCH BRANCH -->|"= 1"| SPARSE BRANCH -->|"= 0"| LINEAR BRANCH -->|"= −1"| SKIP HZ --> LINEAR end subgraph Output Os["O^s = diag(l_i)⁻¹ · O^s_acc"] Ol["O^l = φ(Q_i)·H_i / (φ(Q_i)·Z_i)"] PROJ["Proj(O^l) — learnable ℝ^d→ℝ^d, zero-init"] FINAL["O = O^s + Proj(O^l)"] SPARSE --> Os LINEAR --> Ol Ol --> PROJ Os --> FINAL PROJ --> FINAL end

    §4 作者证明 #

    符号表 #

    符号含义维度
    $Q, K, V$Query, Key, Value 矩阵$\mathbb{R}^{N \times d}$
    $P$注意力权重矩阵$\mathbb{R}^{N \times N}$, 行和为 1
    $M$稀疏掩码$\{0,1\}^{N \times N}$
    $P_c$压缩注意力权重$\mathbb{R}^{(N/b_q) \times (N/b_{kv})}$
    $M_c$压缩掩码(三分类)$\{-1, 0, 1\}^{(N/b_q) \times (N/b_{kv})}$
    $\phi(\cdot)$线性注意力的特征映射$\mathbb{R}^d \to \mathbb{R}^d$
    $H_i$线性注意力的 KV 累积$\mathbb{R}^{d \times d}$
    $Z_i$线性注意力的归一化累积$\mathbb{R}^{d \times 1}$
    $b_q, b_{kv}$Q 块和 KV 块大小标量,默认 64
    $k_h, k_l$critical / negligible 阈值比例标量,默认 5% / 10%
    $\text{Proj}$可学习线性投影$\mathbb{R}^{d \times d}$,零初始化

    方程物理意义 #

    分解方程 (Eq. 1):$P = P \odot M + P \odot (1-M)$。这是一个恒等分解——物理意义在于将注意力权重分为需要精确计算的高秩部分和可低秩近似的部分。分解本身无近似,近似发生在用线性注意力替代 $P \odot (1-M)$ 对应的输出时。

    线性注意力重排 (Eq. 3–4):$O = \phi(Q) \cdot (\phi(K)^\top V) / (\phi(Q) \cdot \text{rowsum}(\phi(K)^\top))$。通过先计算 $\phi(K)^\top V \in \mathbb{R}^{d \times d}$ 再乘 $\phi(Q)$,将复杂度从 $O(N^2 d)$ 降至 $O(N d^2)$。当 $d \ll N$ 时(DiT 中 $d \approx 128$, $N \approx 30000$),加速比约 $N/d \approx 234\times$。

    融合输出 (Eq. 6):$O = O^s + \text{Proj}(O^l)$。$\text{Proj}$ 零初始化确保训练起始时 SLA 等价于纯稀疏注意力,线性注意力的贡献通过微调逐步引入,避免训练初期的分布冲击。

    6 项验证检查 #

    #检查项结果
    1符号一致性:$P, M, P_c, M_c$ 在 §3–§5 中定义和使用一致✓ 一致
    2边界退化:$k_h = 100\% \Rightarrow$ SLA = FlashAttention;$k_h = 0\%, k_l = 0\% \Rightarrow$ SLA = 线性注意力;$k_l = 100\% \Rightarrow$ SLA = 零输出。前两者符合预期,第三者为退化情况✓ 正确退化
    3复杂度验证:sparse $O(k_h \cdot N^2 d / b^2 \cdot b^2) = O(k_h N^2 d)$,linear $O(N d^2)$,预测 $O((N/b)^2 d)$。总计 $O(k_h N^2 d + N d^2)$,与论文声称一致
    4近似误差:论文给出实证 L1 误差(Fig. 1 右):跳过 bottom 45% → <3% 误差,跳过 bottom 92% → ~33% 误差。无形式化误差界,但实证数据支持 marginal 区间可安全用线性注意力替代△ 仅实证
    5数值稳定性:sparse 路径使用 OnlineSoftmax(running max $m_{ij}$ 防止 exp 溢出);linear 路径使用 softmax 作为 $\phi$(保证非负,避免 div-by-zero,分母加 1e-5 epsilon)
    6反向传播完整性:Algorithm 2 提供了 $dQ, dK, dV, dQ^\phi, dK^\phi$ 的完整梯度。sparse 路径梯度遵循 FlashAttention backward;linear 路径梯度通过预计算 $dH_i, dZ_i$ 实现,与前向的预计算对称

    §5 实验与数据 #

    注意力权重分布与稀疏度-精度关系 #

    Figure 1: Attention weight distribution and sparsity–accuracy tradeoff

    Paper Figure 1, verbatim (caption: "The left figure shows a typical distribution of attention weights sampled from the Wan2.1 model. The right figure shows the accuracy of sparse attention with different sparsity.").

    左图揭示注意力权重的三层结构:仅 8.1% 的权重超过均值 $1/N$(critical),约 45% 低于 $1/(100N)$(negligible),中间 ~47% 为 marginal。右图定量展示稀疏度-误差关系的非线性:跳过 45% 仅 <3% 误差,但从 45% 到 92% 误差急剧攀升至 33%。这条曲线直接激发了三分类策略——45% 可安全跳过,47% 可低秩近似,仅 8% 需要精确计算。

    稀疏-低秩分解可视化 #

    Figure 3: Decomposition of attention weights into sparse-few and low-rank-many

    Paper Figure 3, verbatim (caption: "Decomposition of attention weights. We sample attention weights from the Wan2.1 model: the left figure shows the full weights, the middle the top 8%, and the right the bottom 92%.").

    三张热力图直观验证了分解的合理性。完整 $P$(左)呈现复杂的块状结构。top 8%(中)保留了几乎全部的空间结构和 stable rank,呈稀疏块状分布。bottom 92%(右)颜色高度均匀,stable rank 接近 1——这是线性注意力可以高效逼近的理想条件。这张图是论文最重要的理论支撑:它将"为什么线性注意力单独不行"和"为什么稀疏注意力难以超 90%"统一解释为同一个现象的两面。

    主实验:质量与效率对比 (Table 1) #

    MethodVA ↑VT ↑IQ ↑OC ↑AQ ↑SC ↑VR ↑FLOPs ↓Sparsity ↑
    Full Attention76.7882.8862.523.356.193.00.05952.75T0%
    Sparge-F0.0020.02626.04.635.785.1−0.2167.91T85%
    Sparge-T73.8377.8761.922.755.493.10.0147.38T84%
    VMoBa32.3335.7958.018.846.289.9−0.1757.91T85%
    VSA55.3764.6160.622.451.983.6−0.0695.92T89%
    SLA76.9683.9262.223.655.993.10.0482.74T95%

    SLA 在最高稀疏度(95%)下,7 项 VBench 指标中 4 项超越 Full Attention(VA, VT, OC, SC),其余 3 项差距 <1%。计算量仅为 Full Attention 的 5.2%,为次优基线 Sparge-T 的 37%。注意 Sparge-F(training-free)的灾难性失败(VA: 0.002)说明微调对视频 DiT 的稀疏注意力是必需的。

    消融实验 (Table 2, 关键行) #

    MethodVA ↑VT ↑FLOPs ↓Sparsity
    Linear Only0.0420.0990.10T100%
    Sparse Only64.0070.507.91T85%
    L+S (naive sum)29.6541.155.37T90%
    SLA (softmax φ)76.9683.922.73T95%
    SLA (elu+1 φ)75.5081.012.74T95%

    三个关键消融结论:(1) Linear Only 完全失败,验证了线性注意力无法处理 DiT 的高秩注意力;(2) L+S(直接相加)反而比 Sparse Only 更差(VA: 29.65 vs 64.00),说明朴素融合引入干扰——$\text{Proj}$ 层和微调是必需的;(3) softmax 作为 $\phi$ 优于 elu+1(VA: 76.96 vs 75.50),反直觉但可能因为 softmax 产生的非负归一化输出更匹配注意力权重的分布特性。

    Kernel 速度与端到端延迟 #

    Figure 6: Kernel speed and end-to-end latency on RTX 5090

    Paper Figure 6, verbatim (caption: "Attention kernel speed and end-to-end generation latency of SLA and baselines on Wan2.1-1.3B with RTX5090. FlashAttn refers to FlashAttn2, the fastest available version on RTX5090.").

    左图/中图展示 kernel 吞吐量(FLOPS),SLA 前向达 FlashAttention2 的 13.7×,反向达 6.2×。反向加速低于前向的原因:反向需要重载 $Q, K, V$ 和输出梯度,内存流量更大,且 sparse 路径需要重计算 $P_{ij}$。右图展示端到端延迟分解:注意力从 97s 降至 11s(8.8× 加速),但非注意力部分(77s)不受 SLA 影响,限制了端到端加速至 2.2×。这表明若要进一步加速,需要联合优化非注意力组件。


    §6 论证链 #

    Step论点证据论文位置
    1注意力权重在 DiT 中呈现"少量高秩 + 大量极低秩"的双峰结构Fig. 1:8.1% 权重 > $1/N$,45% 权重 < $1/(100N)$;Fig. 3:top 8% 保留全部 stable rank,bottom 92% stable rank ≈ 1§3.1–§3.2
    2这种结构导致稀疏注意力和线性注意力各自的瓶颈稀疏:跳过 >90% 权重时 L1 误差 >30%(Fig. 1 右);线性:受限于秩 $d$,无法拟合高秩部分(§3.2 分析)§3.1–§3.2
    3SLA 的三分类策略(critical/marginal/negligible)利用了双峰结构,将稀疏度从 ~85% 推到 95%critical(top 5%)用 FlashAttention 保精度,marginal(~50%)用线性注意力补偿,negligible(~45%)安全跳过。线性注意力成本 < 0.5% 全注意力 FLOPs§4
    4线性注意力不是对 marginal 权重的近似,而是可学习补偿Proj 零初始化 → 训练初期 SLA = 纯稀疏 → 微调逐步学习 $O^l$ 的贡献。朴素 L+S 相加效果更差(VA: 29.65 vs SLA: 76.96),证明需要端到端学习§4.2, Table 2
    5融合 Triton kernel 将算法优势转化为实际 GPU 加速前向 13.7× vs FlashAttention2,反向 6.2×。LUT-based 跳过避免了逐块判断的分支开销§6.3, Fig. 6
    62000 步微调足以使 SLA 在所有质量指标上匹配或超越 Full AttentionTable 1:7 项 VBench 中 4 项超越 Full Attention,3 项差距 <1%。微调成本 <0.1% 预训练§6.2, Table 1

    §7 实现 cross-reference #

    代码来源:

    核心文件结构 #

    文件功能
    sparse_linear_attention/core.pySparseLinearAttention 模块:orchestrates mask prediction, sparse kernel, linear attention, Proj
    sparse_linear_attention/kernel.pyTriton kernel:_attn_fwd(前向)、_attn_bwd_dq(dQ 反向)、_attn_bwd_dkdv(dK/dV 反向),以及 _attention autograd wrapper
    sparse_linear_attention/utils.pymean_pool(Triton pooling kernel)、get_block_map(mask prediction + TopK)
    SageSLA/SageAttention 集成变体

    关键实现细节 #

    1. Proj 零初始化(core.py:57–59

    
    def init_weights_(self):
        with torch.no_grad():
            nn.init.zeros_(self.proj_l.weight)
            nn.init.zeros_(self.proj_l.bias)
    

    proj_l 初始化为全零,使得训练起始时 $O^l$ 的贡献为零,SLA 退化为纯 sparse attention。这避免了线性注意力在早期引入的分布冲击,让模型通过微调逐步学会利用线性注意力的补偿信号。

    2. 代码与论文的关键差异——两路分离而非三路融合(core.py:74–87

    论文 Algorithm 1 描述了在单一内循环中根据 $M_c[i,j]$ 分支处理 critical/marginal/negligible 三类的融合 kernel。但开源代码将 sparse 和 linear 分为两个独立计算:

    
    o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, ...)
    # Linear attention: computed over ALL q, k, v (not just marginal blocks)
    o_l = calc_linear(q, k, v)
    o = (o_s + o_l).to(dtype)
    

    sparse attention 在 Triton kernel 中仅处理 TopK 块(LUT-driven),linear attention 在 PyTorch 中对全部 tokens 计算 $(K^\top V)$ 然后 $Q \cdot (K^\top V)$。没有论文描述的"仅对 marginal 块计算线性注意力"的优化——线性注意力覆盖所有 tokens。这简化了实现但放弃了论文声称的三分类节省。由于 $O(Nd^2)$ 在 $N \approx 30K, d \approx 128$ 下仅占全注意力的 ~0.5%,对总延迟影响极小。

    3. smooth-K 技巧(utils.py:60

    
    arg_k = k - torch.mean(k, dim=-2, keepdim=True)  # smooth-k from SageAttention
    

    mask prediction 时对 $K$ 做零均值化(借鉴同一作者的 SageAttention),提升 pooled attention score 的区分度。论文未提及此技巧。

    4. Triton kernel 使用 log2 加速 OnlineSoftmax(kernel.py:58–69

    
    qk = tl.dot(q, k) * (qk_scale * 1.4426950408889634)  # 1/ln(2)
    p = tl.math.exp2(qk)
    alpha = tl.math.exp2(m_i - new_m)
    

    将 $\exp(x)$ 替换为 $2^{x / \ln 2}$,利用 exp2 在 GPU 上的更快实现。乘以 $1/\ln 2 \approx 1.4427$ 在 scale 阶段完成,无额外开销。

    5. 前向 kernel 仅处理 sparse path,不含 linear(kernel.py:22–83

    _attn_fwd kernel 的内循环通过 LUT 索引直接跳转到 critical 块,不处理 marginal 块:

    
    for block_idx in tl.range(topk):
        idx_n = tl.load(LUT_ptr + block_idx)
        # load K, V at block idx_n and compute FlashAttention
    

    这是 LUT-driven 而非 scan-based 的 sparse attention:不扫描全部 $T_n$ 块再判断 mask,而是直接遍历预先排好的 topk 个非零块索引。避免了分支预测失败和无效内存访问。

    核心技术壁垒(§7 专属段落) #

    SLA 最难复现的部分不是 Triton kernel(kernel 逻辑与 FlashAttention2 近似,区别仅在 LUT-driven 跳转),而是 compressed attention prediction 的精度保证。mean pooling 是一个有损压缩——块内 Q/K 方差越大,pooled score 与真实 block-level attention 的偏差越大。论文未给出 pooling 精度的理论分析,也未讨论哪些注意力模式下 pooling 会失效。实际部署中,若模型的注意力模式与 Wan2.1 显著不同(例如更长序列、不同 head dimension),$k_h = 5\%$ 的阈值可能不再最优,需要重新调参。开源代码中的 smooth-K 技巧(零均值化 K)是一个未在论文中讨论但对精度有帮助的工程 trick。


    §8 关键疑问与开放方向 #

    1. 序列长度外推:论文在 $N \approx 30K$ 上验证。更长序列(100K+)下 pooling 精度是否下降?stable rank 分布是否变化?更长序列的 top 5% 可能包含更多"伪 critical"块。
    2. 模型泛化性:仅在 Wan2.1-1.3B 和 LightningDiT 上测试。大模型(14B+)或不同架构(非 DiT、autoregressive LLM)的注意力权重是否也呈现相同的双峰 stable rank 结构?
    3. 硬件移植:Triton kernel 仅在 RTX 5090 评测。H100/A100 的不同 shared memory 大小和 Tensor Core 架构可能需要调整 $b_q, b_{kv}$。AMD GPU 需要完全重写 kernel。
    4. 后续工作 SLA2(arXiv 2602.12675):同一团队的后续工作已识别出 block-level TopK/BottomK 分类的启发式路由和分解不匹配问题,表明 compressed attention prediction 是已知的薄弱环节。
    5. Proj 层的理论理解:论文承认线性注意力在 SLA 中不是对 marginal 权重的近似而是"可学习补偿",但未解释这个补偿信号具体学到了什么。Proj 收敛后的权重矩阵结构是否有可解释的模式?