Back to blog

SageAttention3: Microscaling FP4 Attention for Inference and An Exploration of 8-Bit Training

首个 FP4 微缩放注意力推理加速与 INT8 可训练注意力探索

SageAttention3: Microscaling FP4 Attention for Inference and An Exploration of 8-Bit Training

一、论文概述

项目内容
标题SageAttention3: Microscaling FP4 Attention for Inference and An Exploration of 8-Bit Training
作者Jintao Zhang*, Jia Wei*, Pengle Zhang, Xiaoming Xu, Haofeng Huang, Haoxu Wang, Kai Jiang, Jun Zhu, Jianfei Chen
机构清华大学计算机系、智能技术与系统国家重点实验室、交叉信息研究中心、THBI Lab、清华-博世联合机器学习中心;声水科技
论文https://arxiv.org/abs/2505.11594
代码https://github.com/thu-ml/SageAttention
发表NeurIPS 2025
前置工作SageAttention (ICLR 2025), SageAttention2 (ICML 2025), SpargeAttn (ICML 2025)

核心贡献:

  1. 设计了首个 FP4 微缩放注意力(SageAttention3),利用 Blackwell GPU 的 FP4 Tensor Cores,在 RTX5090 上实现 1038 TOPS,较 FlashAttention2 提速 5×
  2. 首次探索低比特注意力用于训练(SageBwd):在指令微调任务中实现无损性能,但在预训练中收敛较慢
  3. 提出两级量化 P 矩阵策略,将 E4M3 范围利用率从 35 个值提升到 127 个值
  4. 证明 INT8 比 FP8 更适合训练,梯度精度更高且硬件支持更广

二、核心思想

问题定义

注意力计算 S=QK⊤,P=Softmax(S),O=PVS = QK^\top, P = \text{Softmax}(S), O = PV 具有二次时间复杂度 O(N2d)O(N^2d),是生成模型的瓶颈。现有低比特注意力工作(FlashAttention3、SageAttention)仅针对推理优化,未探索训练场景。同时,Blackwell GPU 引入了新的 FP4 Tensor Cores,但直接应用 FP4 量化面临严重精度损失。

解决方案概述

方向方法用途精度
SageAttention3FP4 微缩放量化 + 两级 P 缩放推理加速NVFP4 E2M1 + E4M3 缩放因子 (FP8)
SageBwdINT8 逐块量化 + FP16 关键路径训练加速(前向+反向)INT8 per-block

三大技术挑战

  1. C1 — FP4 值域严重受限:FP4 仅有 15 个可表示值,逐张量或逐 token 量化均不足以保持模型精度。通过约束量化组大小为 1×16(而非逐张量或逐通道),有效容纳每个 block 内的异常值效应。

  2. C2 — P 矩阵缩放因子范围极窄:P~\widetilde{P} 的值在 [0, 1] 范围内(online softmax 结果),直接量化到 FP4 导致缩放因子进入 0~0.167 的极窄动态范围(sP=max⁡(P~ij)/6s_P = \max(\widetilde{P}_{ij}) / 6)。而硬件要求缩放因子使用 FP8(E4M3)格式,造成显著精度损失。

  3. C3 — 训练时注意力梯度对量化误差敏感:反向传播中 dOV⊤\text{dOV}^\top 的精度直接影响 dP\text{dP} 和 dS\text{dS},误差会在 FlashAttention 反向传播过程中沿序列长度递归累积到 dQ\text{dQ} 和 dK\text{dK},序列越长误差越大。

三、技术架构

整体框架

FP4 注意力工作流程

SageAttention3 基于 FlashAttention 的 tiling 策略,将 Q 分为 Tm=N/BqT_m = N/B_q 个块 {Qi}\{\mathbf{Q}_i\},将 K、V 分为 Tn=N/BkvT_n = N/B_{kv} 个块 {Ki},{Vi}\{\mathbf{K}_i\}, \{\mathbf{V}_i\}。使用在线 softmax 避免对完整的 N×NN \times N 矩阵 S 和 P 进行大量全局内存 I/O。

核心公式

FlashAttention 在线 Softmax 基础

Sij=QiKj⊤/dmij=max⁡(mi,j−1,rowmax(Sij))P~ij=exp⁡(Sij−mij)lij=emi,j−1−mijli,j−1+rowsum(P~ij)Oij=diag(emi,j−1−mij)Oi,j−1+P~ijVjOi=diag(li,Tn)−1Oi,Tn(1)\begin{aligned} S_{ij} &= Q_i K_j^\top / \sqrt{d} \\ m_{ij} &= \max(m_{i,j-1}, \mathrm{rowmax}(S_{ij})) \\ \widetilde{P}_{ij} &= \exp(S_{ij} - m_{ij}) \\ l_{ij} &= e^{m_{i,j-1} - m_{ij}} l_{i,j-1} + \mathrm{rowsum}(\widetilde{P}_{ij}) \\ O_{ij} &= \mathrm{diag}(e^{m_{i,j-1} - m_{ij}}) O_{i,j-1} + \widetilde{P}_{ij} V_j \\ O_i &= \mathrm{diag}(l_{i,T_n})^{-1} O_{i,T_n} \end{aligned} \tag{1}

其中 mijm_{ij} 和 lijl_{ij} 是 bq×1b_q \times 1 向量,初始化为 −∞-\infty 和 00。

FP4 微缩放量化(Section 3.1)

对于矩阵 X∈RN×d\mathbf{X} \in \mathbb{R}^{N \times d},将其划分为 1×n1 \times n 块 Xij\mathbf{X}_{ij}:

sij=max⁡(∣X∣)/6,X^ij=⌈Xij/sij⌋(2)s_{ij} = \max(|\mathbf{X}|) / 6, \quad \hat{\mathbf{X}}_{ij} = \lceil \mathbf{X}_{ij} / s_{ij} \rfloor \tag{2} Xij′=sij×X^ij(3)\mathbf{X}'_{ij} = s_{ij} \times \hat{\mathbf{X}}_{ij} \tag{3}

其中 ⌈⋅⌋\lceil \cdot \rfloor 表示 FP4 rounding。每个 1×n1 \times n 块对应一个缩放因子 sijs_{ij}。

FP4 数据格式选择:

格式数据类型块大小缩放因子格式
NVFP4E2M11×16E4M3 (FP8)
MXFP4E2M11×32E8M0

选择 NVFP4 的原因:CogVideoX 全层精度对比显示 NVFP4 CosSim 为 99.52%,MXFP4 仅为 98.37%。

FP4 微缩放矩阵乘法

RTX5090 上 FP4 MM 速度约 1600 TOPS,是 FP32 GEMM(约 200 TOPS)的 8 倍:

C=FP4MM(A^,sA,B^,sB)=ϕ−1(A^,sA)×ϕ−1(B^,sB)(4)\mathbf{C} = \text{FP4MM}(\hat{\mathbf{A}}, s_A, \hat{\mathbf{B}}, s_B) = \phi^{-1}(\hat{\mathbf{A}}, s_A) \times \phi^{-1}(\hat{\mathbf{B}}, s_B) \tag{4}

注意力计算全流程

Q^,sQ=ϕ(Q),K^,sK=ϕ(K⊤)S=FP4MM(Q^,sQ,K^,sK)P~=OnlineSoftmax(S)P^,sP=ϕ(P~),V^,sV=ϕ(V)O=FP4MM(P^,sP,V^,sV)(5)\begin{aligned} \hat{\mathbf{Q}}, \mathbf{s_Q} &= \phi(\mathbf{Q}), \quad \hat{\mathbf{K}}, \mathbf{s_K} = \phi(\mathbf{K}^\top) \\ \mathbf{S} &= \text{FP4MM}(\hat{\mathbf{Q}}, \mathbf{s_Q}, \hat{\mathbf{K}}, \mathbf{s_K}) \\ \widetilde{\mathbf{P}} &= \text{OnlineSoftmax}(\mathbf{S}) \\ \hat{\mathbf{P}}, \mathbf{s_P} &= \phi(\widetilde{\mathbf{P}}), \quad \hat{\mathbf{V}}, \mathbf{s_V} = \phi(\mathbf{V}) \\ \mathbf{O} &= \text{FP4MM}(\hat{\mathbf{P}}, \mathbf{s_P}, \hat{\mathbf{V}}, \mathbf{s_V}) \end{aligned} \tag{5}

两级 P 缩放(Section 3.2)

由于 P~\widetilde{\mathbf{P}} 每块内值落在 [0, 1] 范围,缩放因子 sP=max⁡(P~ij)/6s_P = \max(\widetilde{\mathbf{P}}_{ij}) / 6 的范围仅为 0~0.167。E4M3 FP8 格式在此极窄范围内表示效率极低。两级量化策略:

第一级:逐 token 归一化到 [0, 448×6] 范围:

sP1=rowmax(P~)/(448×6),P~2=P~/sP1\mathbf{s_{P_1}} = \mathrm{rowmax}(\widetilde{\mathbf{P}}) / (448 \times 6), \quad \widetilde{\mathbf{P}}_2 = \widetilde{\mathbf{P}} / \mathbf{s_{P_1}}

第二级:标准 FP4 微缩放:

sP2,P^2=ϕ(P~2)\mathbf{s_{P_2}}, \hat{\mathbf{P}}_2 = \phi(\widetilde{\mathbf{P}}_2)

重构:

P~≈P^2×sP2×sP1,O=FP4MM(P^2,sP2,V^,sV)×sP1(6)\widetilde{\mathbf{P}} \approx \hat{\mathbf{P}}_2 \times \mathbf{s_{P_2}} \times \mathbf{s_{P_1}}, \quad \mathbf{O} = \text{FP4MM}(\hat{\mathbf{P}}_2, \mathbf{s_{P_2}}, \hat{\mathbf{V}}, \mathbf{s_V}) \times \mathbf{s_{P_1}} \tag{6}

其中 P~\widetilde{\mathbf{P}}, P~2\widetilde{\mathbf{P}}_2, sP1\mathbf{s_{P_1}} 为 FP32,sP2\mathbf{s_{P_2}}, sV\mathbf{s_V} 为 FP8,P^2\hat{\mathbf{P}}_2, V^\hat{\mathbf{V}} 为 FP4。

理论分析(Appendix A.10):

  • 直接量化:E4M3 缩放因子有 35 个可表示值,输出可表示值为 35×8=28035 \times 8 = 280
  • 两级量化:E4M3 缩放因子有 127 个可表示值,输出可表示值为 127×8=1016127 \times 8 = 1016
  • 量化间隔更细:Δ(P~2)P~2<Δ(P~)P~\frac{\Delta(\widetilde{P}_2)}{\widetilde{P}_2} < \frac{\Delta(\widetilde{P})}{\widetilde{P}},因此 E2<E1E_2 < E_1

INT8 前向传播(Algorithm 2)

INT8 逐块量化:

sX=max⁡(∣X∣)/127,X^=X/sX(7)\mathbf{s_X} = \max(|\mathbf{X}|) / 127, \quad \hat{\mathbf{X}} = \mathbf{X} / \mathbf{s_X} \tag{7}

前向过程中对 QK⊤\mathbf{Q}\mathbf{K}^\top 采用 Smoothing K + per-block INT8 量化;对 P~V\widetilde{\mathbf{P}}\mathbf{V} 采用 per-token INT8 量化(而非静态 per-block),因为静态 per-block 缩放因子 1/1271/127 不够准确。同时复用 online softmax 中的全局和局部最大值消除显式的 max 操作。

Sij=M(Q^i,K^j)×sQ×sK(8)\mathbf{S}_{ij} = \mathbb{M}(\hat{\mathbf{Q}}_i, \hat{\mathbf{K}}_j) \times \mathbf{s_Q} \times \mathbf{s_K} \tag{8} sP=exp⁡(rowmax(Sij)−mij)/127,P^ij=P~ij/sP\mathbf{s_P} = \exp(\mathrm{rowmax}(\mathbf{S}_{ij}) - m_{ij}) / 127, \quad \hat{\mathbf{P}}_{ij} = \widetilde{\mathbf{P}}_{ij} / \mathbf{s_P} Oij=diag(emi,j−1−mij)−1Oi,j−1+M(P^ij,V^j)×sP×sV\mathbf{O}_{ij} = \mathrm{diag}(e^{m_{i,j-1} - m_{ij}})^{-1} \mathbf{O}_{i,j-1} + \mathbb{M}(\hat{\mathbf{P}}_{ij}, \hat{\mathbf{V}}_j) \times \mathbf{s_P} \times \mathbf{s_V}

INT8 反向传播(Algorithm 3)

反向传播包含五个矩阵乘法:

S=QK⊤,dV=P~⊤dO,dP=dOV⊤,dQ=dSK,dK=dS⊤Q(9)\mathbf{S} = \mathbf{Q}\mathbf{K}^\top, \quad \mathbf{dV} = \widetilde{\mathbf{P}}^\top \mathbf{dO}, \quad \mathbf{dP} = \mathbf{dO}\mathbf{V}^\top, \quad \mathbf{dQ} = \mathbf{dS}\mathbf{K}, \quad \mathbf{dK} = \mathbf{dS}^\top \mathbf{Q} \tag{9}

关键发现:是否对 dOV⊤\mathbf{dO}\mathbf{V}^\top 进行量化对 dQ\mathbf{dQ}、dK\mathbf{dK} 的精度影响最大。这是因为 dOV⊤\mathbf{dO}\mathbf{V}^\top 的精度直接决定 dP\mathbf{dP} 和 dS\mathbf{dS} 的精度,而 dS\mathbf{dS} 的精度损失会在 FlashAttention 反向传播的沿序列长度递归过程中持续累积到 dQ\mathbf{dQ} 和 dK\mathbf{dK}。

策略:保持 dOV⊤\mathbf{dO}\mathbf{V}^\top 为 FP16,其余四个矩阵乘法使用 INT8 per-block 量化:

dVj←dVj+M(P^ij⊤,dO^i)×sP×sdO\mathbf{dV}_j \leftarrow \mathbf{dV}_j + \mathbb{M}(\hat{\mathbf{P}}_{ij}^\top, \widehat{\mathbf{dO}}_i) \times \mathbf{s_P} \times \mathbf{s_{dO}} dPij=M(dO,Vj⊤)// Keep in FP16\mathbf{dP}_{ij} = \mathbb{M}(\mathbf{dO}, \mathbf{V}_j^\top) \quad \text{// Keep in FP16} dSij=Pij∘(dPij−Di),dQi←dQi+M(dS^ij,K^j)×sdS×sK\mathbf{dS}_{ij} = \mathbf{P}_{ij} \circ (\mathbf{dP}_{ij} - \mathbf{D}_i), \quad \mathbf{dQ}_i \leftarrow \mathbf{dQ}_i + \mathbb{M}(\widehat{\mathbf{dS}}_{ij}, \hat{\mathbf{K}}_j) \times \mathbf{s_{dS}} \times \mathbf{s_K} dKj←dKj+M(dS^ij⊤,Q^i)×sdS×sQ\mathbf{dK}_j \leftarrow \mathbf{dK}_j + \mathbb{M}(\widehat{\mathbf{dS}}_{ij}^\top, \hat{\mathbf{Q}}_i) \times \mathbf{s_{dS}} \times \mathbf{s_Q}

其中 Di=rowsum(dO∘O)\mathbf{D}_i = \mathrm{rowsum}(\mathbf{dO} \circ \mathbf{O})。

量化误差评估指标

指标公式说明
CosSim∑OO′/∑O2∑O′2\sum OO' / \sqrt{\sum O^2 \sum O'^2}余弦相似度,越接近 100% 越好
L1$\sumO-O’
RMSE(1/n)∑(O−O′)2\sqrt{(1/n)\sum(O-O')^2}均方根误差,越小越好

完整算法流程(Algorithm 1 — FP4 注意力)

Input: Q(FP16), K(FP16), V(FP16) ∈ ℝ^(N×d), block size Bq, Bkv

Preprocessing: K ← K - mean(K)                          // Smoothing K (SageAttention)
Divide Q → {Qi} (Tm=N/Bq blocks); divide K,V → {Ki},{Vi} (Tn=N/Bkv blocks)

for i = 1 to Tm do
  qi_bar = mean(Qi)                                     // Smoothing Q (SageAttention2)
  (sQ, Q̂i) = φ(Qi - qi_bar)
  for j = 1 to Tn do
    (sK, K̂j) = φ(Kj^⊤),   (sV, V̂j) = φ(Vj)
    // Smoothed Q: FP4MM on quantized part + GEMV on mean part
    Sij = FP4MM(Q̂i, sQ, K̂j, sK) + GEMV(qi_bar, Kj^⊤)
    mij = max(mi,j-1, rowmax(Sij))
    P̃ij = exp(Sij - mij)
    lij = e^(mi,j-1 - mij) * li,j-1 + rowsum(P̃ij)
    // Two-level quantization for P
    sP1 = rowmax(P̃ij) / (448 × 6), P̃2 = P̃ij / sP1
    (sP2, P̂2ij) = φ(P̃2)
    Oij = diag(e^(mi,j-1 - mij))^(-1) * Oi,j-1
        + FP4MM(P̂2ij, sP2, V̂j, sV) × sP1
  end
  Oi = diag(li,Tn)^(-1) * Oi,Tn
end
return O = {Oi}

模型组件总览

组件说明关键参数/细节
NVFP4 微缩放E2M1 数据格式,1×16 块,E4M3 FP8 缩放因子15 个可表示值
MXFP4 微缩放E2M1 数据格式,1×32 块,E8M0 缩放因子对比基线
两级 P 量化第一级逐 token 归一化到 [0, 448×6],第二级 FP4 微缩放FP32 中间精度
Smoothing KK 矩阵去均值(来自 SageAttention)K = K - mean(K)
Smoothing QQ 块去均值(来自 SageAttention2)Qi = Qi - mean(Qi)
INT8 Per-block逐块量化,缩放因子 $\max(\mathbf{X}
INT8 Per-token逐 token 量化 P 矩阵用于前向 PV 乘法
FP16 关键路径dOV⊤\mathbf{dO}\mathbf{V}^\top 保持 FP16防止梯度累积误差

硬件实现优化(Section 3.3)

1. K 矩阵置换

FP32 累加器的寄存器布局与操作数 A 的寄存器布局不同(见 Fig. 19-20)。执行 thread shuffle 来匹配操作数 A 的布局会导致内核性能退化。

解决方案:通过置换 P tile 的列来变换累加器布局(Fig. 21),同时相应地重新排列 K 的列以保持矩阵乘法正确性。该操作可融合到 K 量化核中,无额外开销。

2. 复用 Shuffle

在 kernel 内对 P~\widetilde{\mathbf{P}} 进行微缩放量化需要找到 16 个连续行元素的最大值。但这 16 个元素分布在四个线程中,需要线程内最大值归约 + 线程间 shuffle,显著降低内核速度。

优化:将量化与 online softmax 融合——online softmax 已计算行最大值,直接复用即可。减少约 50% 冗余 shuffle 和 max 操作,整体内核提速 ~10%。

3. Producer Warp Epilogue

在传统 warp-specialized kernel 中,consumer warps 负责 MatMul 和存储,producer 仅加载输入。但由于寄存器约束,FP4 attention kernel 无法采用此方案。

创新方案:在 producer warps 之间实现 ping-pong 调度——当一个 producer 加载下一个 MatMul 的输入时,另一个 producer 同时将输出存储到全局内存。Consumer warps 仅负责将 MatMul 结果从寄存器传输到共享内存。此设计在寄存器约束下实现了 MatMul 与全局内存存储的重叠。

四、核心创新

创新点说明理论/实验依据
首个 FP4 注意力利用 Blackwell GPU FP4 Tensor Cores,推理速度达 1038 TOPSRTX5090 上 5× over FA2,CosSim 99.55%
NVFP4 vs MXFP4NVFP4 (1×16 块, E4M3 缩放) 精度远高于 MXFP4 (1×32 块, E8M0 缩放)CogVideoX 全层 CosSim: 99.52% vs 98.37% (Table 1a)
两级 P 量化将 E4M3 范围利用率从 35 个值提升到 127 个值理论证明 E2<E1E_2 < E_1;CosSim 93.32%→99.52% (Table 1b)
dOV^⊤ 保持 FP16最敏感的矩阵乘法不量化,防止梯度累积误差CosSim 97.47%→99.77%;RMSE 2.440→0.692 (Table 1c)
首个可训练低比特注意力6/7 矩阵乘法 INT8 量化,微调无损Qwen2.5-3B: GSM8K 0.607 vs BF16 0.601
INT8 > FP8 训练INT8 梯度精度更高,硬件支持更广dQ L1 error: 0.0290 vs 0.0696;支持 A100/MI250/Ascend
Smoothing Q+K结合两代 smoothing 技术提升 FP4 量化精度CosSim 从 0.9156 提升到 0.9912
HilbertCurve 排列(SpargeAttn 协同)空间填充曲线提升图像/视频模型的块自相似性Sim-q=0.572, Sim-k=0.479 (Table 4)

五、代码实现分析

实现技术栈:

  • SageAttention3: CUTLASS + CUDA(手写内核,基于 FP4MM 指令)
  • SageBwd: OpenAI Triton(便于快速迭代,但存在性能差距)

关键文件结构(GitHub: thu-ml/SageAttention):

  • FP4 内核基于 CUTLASS 的 FP4MM 指令构建
  • 8-bit 训练注意力基于 Triton 实现
  • 两阶段掩码生成逻辑
  • Smoothing Q/K 预处理

六、实验结果

内核速度对比

RTX5090(head dim=128, causal=True):

内核速度对比 head=128

方法TOPS on 5090TOPS on H100CosSim
FlashAttention2214338100.000%
FlashAttention3 (16bit)N/A470100.000%
FlashAttention3 (8bit)N/A89098.570%
SageAttention147951899.996%
SageAttention2 (8bit)64388599.995%
SageAttention3 (4bit)1038N/A99.551%

SageAttention3 较 FlashAttention2 提速 ~4.85×,较 xformers 提速 ~11×。

RTX5090(head dim=64): 趋势一致,SageAttention3 在短维度下优势更明显(Fig. 5)。

SageBwd 速度对比(RTX4090): SageBwd 前向传播最高提速 2×,反向传播提速 1.2~1.6×(Appendix Fig. 15-18)。

端到端推理指标

文本到视频模型(Table 2):

模型方法CLIPSIM↑CLIP-T↑VQA-a↑VQA-t↑FScore↑
CogVideoXFull-Precision0.18650.996870.47669.8754.780
CogVideoXSageAttn2 (8bit)0.18800.996969.41470.7504.534
CogVideoXSageAttn3 (4bit)0.18810.996969.86070.3644.035
HunyuanVideoFull-Precision0.18380.999368.99878.8911.4793
HunyuanVideoSageAttn2 (8bit)0.18360.999369.49777.0191.4741
HunyuanVideoSageAttn3 (4bit)0.18660.999370.55275.4401.232
MochiFull-Precision0.18280.999061.984061.00001.8042
MochiSageAttn2 (8bit)0.18190.999061.009360.37321.7539
MochiSageAttn3 (4bit)0.18000.999361.86359.4291.649

文本到图像模型(Table 2):

模型方法FID↓sFID↓CLIP↑IR↑
FluxFull-Precision162.812146.98031.4090.91
FluxSageAttn2 (8bit)163.107146.21331.4360.90
FluxSageAttn3 (4bit)162.121142.83931.4500.94
SD3.5Full-Precision166.421146.37931.930.93
SD3.5SageAttn2 (8bit)164.986148.55732.010.93
SD3.5SageAttn3 (4bit)166.102145.58732.010.92

SageAttention3 在所有模型上几乎无损,部分指标甚至优于全精度(如 Flux 的 FID 从 162.812 降至 162.121)。

端到端速度提升(Table 4)

(a) 推理延迟:

模型原始Sage1Sage2Sage3
CogVideoX (2B)64 s55 s46 s27 s (~2.4×)
HunyuanVideo489 s257 s240 s164 s (~3×)

(b) 训练延迟:

模型原始SageBwd
Llama (8K)2.1 s1.9 s (~1.15×)
Llama (16K)6.0 s5.2 s (~1.15×)

微调性能对比(Table 3)

模型方法GSM8KDROP(F1)MMLUHellaSwag
Qwen2.5 (1.5B)BF160.5210.7330.5690.905
Qwen2.5 (1.5B)SageBwd0.5200.7340.5740.911
Qwen2.5 (3B)BF160.6010.7850.6400.944
Qwen2.5 (3B)SageBwd0.6070.7820.6530.943
Llama3.2 (1B)BF160.2590.6410.4640.828
Llama3.2 (1B)SageBwd0.2680.6370.4580.823

多随机种子实验(Appendix Tables 5-10)进一步验证:SageBwd 在各种子下的平均性能与 BF16 高度一致,标准差几乎相同(如 Qwen2.5-1.5B GSM8K: SageBwd 0.5077±0.0090 vs BF16 0.5081±0.0089)。

预训练 vs 微调(Fig. 8)

  • 微调(700 steps, lr=3e-5, warmup 100, batch=32/128):SageBwd 损失曲线与 BF16 完全对齐,各数据集标准差几乎相同
  • 预训练(FineWeb-Edu, 400M 模型, lr=1e-3, warmup 1000, 2M tokens/step):SageBwd 可以收敛但速度较慢,目前不适用于预训练

组合使用 SageBwd + SageAttention3

先 INT8 微调再 FP4 推理的效果优于 BF16 微调 + FP4 推理:

模型方法GSM8KMMLU
Qwen2.5-1.5BBF16 微调0.49120.4688
Qwen2.5-1.5BSageBwd 微调0.52320.4934
Qwen2.5-3BBF16 微调0.58600.6000
Qwen2.5-3BSageBwd 微调0.59450.6032

可视化对比

视频/图像生成可视化对比

SageAttention3 生成的视频和图像与全精度方法视觉质量一致。附录中还提供了更多定性示例(Fig. 10-11: Flux/SD3.5 图像生成; Fig. 13-14: CogVideoX/HunyuanVideo 视频生成; Fig. 12: 两级量化 vs 直接量化的 CogVideoX 生成对比)。

七、消融研究

平滑技术消融

方法CosSim↑L1 Error↓RMSE↓
None0.91560.33590.3035
SmoothQuant0.93010.26760.2529
Hadamard0.94120.26200.2240
Smoothing_Q0.98280.11570.1259
Smoothing_K0.99120.09480.0977

结合 Smoothing_Q 和 Smoothing_K 效果最佳,CosSim 从 0.9156 提升到 0.9912。

层级量化误差累积

方法Layer1 L1Layer10 L1Layer20 L1Layer30 L1
直接使用 SageAttention30.00760.09220.11460.0571
保留 3 个最敏感层在 FP160.00760.04470.07730.0429

累积误差一般随层深增加而增大,深层偶尔出现部分误差抵消。保留最敏感的三层在 FP16 可显著降低整体误差。

不同 FP4 格式对比(Table 1a)

TypeCosSim↑L1↓RMSE↓
MXFP498.37%0.2940.994
NVFP499.52%0.0770.201

不同 P 缩放策略对比(Table 1b)

MethodCosSimL1RMSE
Direct93.32%0.1931.103
Two-level99.52%0.0770.201

dOV^⊤ 数据类型对比(Table 1c)

MethodCosSimL1RMSE
INT897.47%0.1712.440
FP1699.77%0.0390.692

INT8 vs FP8 训练对比

指标INT8 SageBwdFP8 SageBwd
dQ L1 error0.02900.0696
dK L1 error0.03170.0999
dV L1 error0.04230.0873
dQ CosSim0.99870.9880
dK CosSim0.99930.9910
dV CosSim0.99950.9955

INT8 选择原因:(1) 更高的梯度精度 (2) 更广泛的硬件支持(A100、AMD MI250、Ascend 910B)。

理论吞吐对比(Appendix A.8)

方法B300 TOPSB200 TOPSRTX5090 TOPS
FlashAttention325002500209.5
FlashAttention3 (FP8)50005000419
SageAttention3 (FP4)15000100001676

多随机种子消融(Appendix A.4, Tables 5-10)

在 Qwen2.5 (1.5B/3B) 和 Llama3.2 (1B/3B) 上,使用 5 个不同随机种子(42, 233, 1234, 5678, 11)分别微调后评估:

  • SageBwd 与 BF16 的平均性能几乎完全一致
  • 标准差非常接近(通常差异 < 0.001)
  • 证明 SageBwd 的微损效果具有统计稳健性

八、相关工作

推理加速方法

  • FlashAttention 系列:FlashAttention (ICLR 2022) 引入 tiling 减少 GPU 全局内存 I/O;FlashAttention2 (ICLR 2024) 改进并行度和 warp 分区策略;FlashAttention3 (NeurIPS 2024) 仅在 Hopper GPU 上优化,不支持 RTX 系列。
  • xformers:专用 CUDA kernel 加速注意力。
  • SageAttention 系列:SageAttention (ICLR 2025) 使用量化 + outlier smoothing;SageAttention2 (ICML 2025) 引入 thorough outlier smoothing + per-thread int4 量化。
  • FlashAttention3 FP8:虽支持 FP8 量化但未以 plug-and-play 方式应用于视频生成模型,且不支持反向传播。

线性注意力与稀疏注意力

  • 线性注意力:Linformer (ICML 2020)、Performer (ICML 2021)、Metaformer (CVPR 2022)、Transformers are RNNs (ICML 2020)、Lightning Attention-2 (arXiv 2024)、Gated Delta Networks (arXiv 2024)。
  • 稀疏注意力:Swin Transformer (ICCV 2021)、Twins (NeurIPS 2021)、Uniformer (ICLR 2022)、StreamingLLM (ICLR 2024)、InfLLM (2024)、LongLoRA (ICLR 2024)、MInference (NeurIPS 2024)、Skip-Attention (ICLR 2024)、SeerAttention (arXiv 2024)、MOA (arXiv 2024)、Sparse VideoGen (arXiv 2025)、SpargeAttn (ICML 2025)。

以上方法与 SpargeAttn 正交,可组合使用。

九、总结

核心贡献

  1. FP4 推理加速:设计首个 FP4 微缩放注意力,RTX5090 上 1038 TOPS(5× over FlashAttention2),端到端质量几乎无损
  2. INT8 训练探索:首次将低比特注意力应用于训练,微调无损、预训练收敛较慢
  3. 两级 P 量化:将 E4M3 范围利用率从 35 提升到 127,理论证明 E2<E1E_2 < E_1
  4. INT8 > FP8:证明 INT8 在训练中的优势——梯度精度更高、硬件支持更广

局限性

  1. SageBwd 在预训练任务中收敛速度较慢,目前不适用于预训练
  2. FP4 注意力依赖 Blackwell GPU 硬件支持,Hopper 及以下架构无法使用
  3. 深层网络中仍存在量化误差累积,需保留部分层在 FP16
  4. SageBwd 当前 Triton 实现与理论性能上限存在差距

未来方向

  1. 优化 Triton 内核实现以缩小 SageBwd 与理论性能差距
  2. 探索低比特注意力在预训练任务中的可行性

十、参考资源

  • arXiv: https://arxiv.org/abs/2505.11594
  • 代码: https://github.com/thu-ml/SageAttention
  • 相关论文:
    • SageAttention (ICLR 2025): Accurate 8-bit attention for plug-and-play inference acceleration
    • SageAttention2 (ICML 2025): Efficient attention with thorough outlier smoothing and per-thread int4 quantization
    • SpargeAttn (ICML 2025): Accurate and training-free sparse attention accelerating any model inference
    • FlashAttention (ICLR 2022), FlashAttention2 (ICLR 2024), FlashAttention3 (NeurIPS 2024)
    • xformers: A modular and hackable transformer modelling library

附图索引

编号文件名说明
Figure 1figure-1-speedup-rtx5090.pngRTX5090 内核速度对比及 HunyuanVideo 端到端加速
Figure 2figure-2-fp4-workflow.png微缩放 FP4 注意力工作流程
Figure 3figure-3-two-level-quantization-analysis.png两级量化分析(分布、误差、可表示值对比)
Figure 4figure-4-kernel-speed-head128.pngRTX5090 head dim=128 内核速度对比
Figure 5figure-5-kernel-speed-head64.pngRTX5090 head dim=64 内核速度对比
Figure 6figure-6-sagebwd-speed-head128.pngRTX4090 SageBwd 前向+反向速度 head=128
Figure 7figure-7-sagebwd-speed-head64.pngRTX4090 SageBwd 速度 head=64
Figure 8figure-8-training-loss-curves.png预训练和微调损失曲线
Figure 9figure-9-visible-comparison.png视频/图像生成可视化对比
Appendix Fig. 10-11—Flux/SD3.5 图像生成额外对比
Appendix Fig. 12—两级量化 vs 直接量化的 CogVideoX 生成对比
Appendix Fig. 13-14—CogVideoX/HunyuanVideo 视频生成额外对比
Appendix Fig. 15-18—SageBwd 前向/反向单独速度对比
Appendix Fig. 19-21—FP4 寄存器布局与置换示意图