Back to blog

FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision

利用异步性和低精度实现快速准确的注意力机制

FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision

一、论文概述

项目内容
标题FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision
作者Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao
机构Colfax International, Princeton
论文arXiv:2407.08608
代码GitHub
发布2024年7月11日 (v1), 2024年7月12日 (v2)
主题cs.LG, cs.AR

二、核心思想

问题定义

注意力机制作为 Transformer 的核心层,是大语言模型和长上下文应用的瓶颈。FlashAttention 通过最小化内存读写来加速注意力,但尚未充分利用新硬件的特性。FlashAttention-2 在 H100 GPU 上仅达到 35% 的利用率(vs GEMM 的 80-90%)。

核心观察

  1. 异步性: Hopper GPU 的 Tensor Core 和 TMA 可异步执行,允许重叠计算和数据移动
  2. 低精度: FP8 提供 2× 吞吐量,但需要小心处理量化误差(特别是异常值特征)
  3. 非 GEMM 瓶颈: 指数函数吞吐量比 GEMM 低 256×,但 softmax 可占 50% 的周期

解决方案概述

FlashAttention-3 提出三种技术加速 Hopper GPU 上的注意力:

  1. 生产者-消费者异步: Warp 特化软件流水线,利用 TMA 和 Tensor Core 的异步执行
  2. GEMM-Softmax 重叠: 将 softmax 操作隐藏在异步 WGMMA 指令下
  3. FP8 低精度: 块量化 + 不相干处理,利用 FP8 Tensor Core

核心性能

指标数值
FP16 前向加速1.5-2.0× vs FlashAttention-2
FP16 后向加速1.5-1.75× vs FlashAttention-2
FP16 峰值性能740 TFLOPs/s (75% 利用率)
FP8 峰值性能~1.2 PFLOPs/s
FP8 精度2.6× 低于标准 FP8 注意力的数值误差

三、技术架构

核心洞察

Pingpong 调度

Figure 1: Pingpong 调度:2 个 warpgroup 重叠 softmax 和 GEMM。一个 warpgroup 的 softmax 应在另一个 warpgroup 运行 GEMM 时调度。相同颜色表示相同迭代。

关键问题: H100 有 989 TFLOPS FP16 GEMM 但仅 3.9 TFLOPS 特殊函数(如指数)。对于 head dimension 128 的 FP16 前向:

  • GEMM FLOPS 比指数操作多 512×
  • 但指数吞吐量低 256×
  • 所以指数可占 GEMM 50% 的周期
  • FP8 时更糟:GEMM 吞吐翻倍,指数不变

核心公式

多头注意力

给定查询 QQ, 键 KK, 值 V∈RN×dV \in \mathbb{R}^{N \times d}:

S=αQKT∈RN×N,P=softmax(S)∈RN×N,O=PV∈RN×dS = \alpha QK^T \in \mathbb{R}^{N \times N}, \quad P = \text{softmax}(S) \in \mathbb{R}^{N \times N}, \quad O = PV \in \mathbb{R}^{N \times d}

其中 α=1/d\alpha = 1/d,softmax 逐行应用。

反向传播

dV=PTdO,dP=dOVT,dS=dsoftmax(dP),dQ=αdSK,dK=αdSTQdV = P^T dO, \quad dP = dO V^T, \quad dS = \text{dsoftmax}(dP), \quad dQ = \alpha dS K, \quad dK = \alpha dS^T Q

三大核心技术

1. 生产者-消费者异步 (Warp 特化)

Warp 特化:

  • CTA 中的 warps 分为生产者和消费者角色
  • 生产者仅发出数据移动(TMA)
  • 消费者仅执行计算(WGMMA)
  • 通过 setmaxnreg 动态重分配寄存器

Pingpong 调度:

  • 使用 bar.sync 强制 warpgroup 1 的 GEMMs 在 warpgroup 2 之前调度
  • 结果:warpgroup 1 的 softmax 在 warpgroup 2 执行 GEMMs 时调度
  • 然后角色互换
  • 性能提升:570 TFLOPS → 620-640 TFLOPS

2. GEMM-Softmax 流水线

2 阶段流水线

Figure 2: 2 阶段 WGMMA-softmax 流水线。

核心思想: 打破 softmax 和 GEMMs 之间的顺序依赖,通过跨迭代流水线化。

Algorithm 2 (消费者 warpgroup 前向):

  • 初始化 Oi=(0)O_i = (0), ℓi,mi=(0),(−∞)\ell_i, m_i = (0), (-\infty)
  • 计算 Scur=QiK0TS_{\text{cur}} = Q_i K_0^T (WGMMA)
  • 主循环 (lines 8-16):
    • 计算 Snext=QiKjTS_{\text{next}} = Q_i K_j^T (WGMMA, commit 但不 wait)
    • 计算 Oi=Oi+P~curVj−1O_i = O_i + \tilde{P}_{\text{cur}} V_{j-1} (WGMMA, commit 但不 wait)
    • Wait SnextS_{\text{next}} WGMMA, 计算 P~next\tilde{P}_{\text{next}}
    • Wait P~curVj−1\tilde{P}_{\text{cur}} V_{j-1} WGMMA, rescale OiO_i

关键: 第二个 WGMMA (P~curVj−1\tilde{P}_{\text{cur}} V_{j-1}) 与下一次迭代的 softmax (SnextS_{\text{next}}) 重叠。

寄存器压力: 需要额外寄存器存储 SnextS_{\text{next}},大小 Br×Bc×sizeof(float)B_r \times B_c \times \text{sizeof(float)}。

3 阶段流水线: 进一步重叠第二个 WGMMA 与 softmax,但需要更多寄存器。

3. FP8 低精度

布局挑战:

  • FP8 WGMMA 仅支持 k-major 格式
  • 但 VV 通常在 head dimension 上连续
  • 解决方案:内核内转置(使用 LDSM/STSM 指令)

寄存器布局差异:

  • FP32 累加器布局(Figure 3)与 FP8 操作数 A 布局(Figure 4)不同
  • 使用字节置换指令转换:{d0 d1 d4 d5 d2 d3 d6 d7}\{d0\ d1\ d4\ d5\ d2\ d3\ d6\ d7\}

精度优化:

  1. 块量化: 每个块保持一个标量(而非每张量),自然融合到 FlashAttention 的块操作中
  2. 不相干处理: 将 QQ 和 KK 乘以随机正交矩阵 MM 再量化到 FP8
    • MMT=IMM^T = I,所以 (QM)(KM)T=QKT(QM)(KM)^T = QK^T
    • 每个条目是原始条目的随机和,“分散”异常值
    • 选择 MM 为随机 ±1\pm 1 对角矩阵和 Hadamard 矩阵的乘积
    • 计算复杂度 O(dlog⁡d)O(d \log d),可融合到 rotary embedding

GPU 硬件特性

内存层次

硬件级别并行代理数据位置容量 @ 带宽
ChipGridGMEM80 GiB @ 3.35 TB/s
GPCThreadblock ClustersL250 MiB @ 12 TB/s
SMThreadblock (CTA)SMEM228 KiB per SM, 31 TB/s per GPU
ThreadThreadRMEM256 KiB per SM

线程层次

  • Threads → Warps (32 threads) → Warpgroups (4 warps) → Threadblocks (CTAs) → Clusters → Grids

Hopper 异步特性

  • TMA: 专用硬件单元,GMEM↔SMEM 异步拷贝
  • WGMMA: Warpgroup 级异步 GEMM,可直接从 SMEM 源操作数
  • setmaxnreg: 动态重分配 warpgroup 间寄存器

算法伪代码

Algorithm 1 (无重叠的前向):

Producer warpgroup:
  Load Q_i from HBM to SMEM
  For j = 0 to T_c - 1:
    Wait buffer stage consumed
    Load K_j, V_j from HBM to SMEM
    Commit notification

Consumer warpgroup:
  Initialize O_i = 0, l_i = 0, m_i = -inf
  Wait Q_i loaded
  For j = 0 to T_c - 1:
    Wait K_j loaded
    S_i(j) = Q_i K_j^T (SS-GEMM)
    m_i = max(m_i, rowmax(S_i(j)))
    P~_i(j) = exp(S_i(j) - m_i)
    l_i = exp(m_i_old - m_i) l_i + rowsum(P~_i(j))
    Wait V_j loaded
    O_i = diag(exp(m_i_old - m_i))^{-1} O_i + P~_i(j) V_j (RS-GEMM)
    Release buffer stage
  O_i = diag(l_i)^{-1} O_i, L_i = m_i + log(l_i)

四、核心创新

创新点说明理论/实验依据
Warp 特化生产者/消费者分离,利用 TMA/WGMMA 异步性570→620 TFLOPS
Pingpong 调度跨 warpgroup 重叠 softmax 与 GEMM隐藏指数操作延迟
2/3 阶段流水线跨迭代重叠 GEMM 与 softmax570→661 TFLOPS
FP8 块量化每块一个标量,自然融合2.6× 精度提升
不相干处理随机正交矩阵分散异常值2.6× 精度提升
内核内转置使用 LDSM/STSM 指令转置 V避免额外预处理内核

五、实验结果

测试配置

配置值
GPUNVIDIA H100 SXM5 80GB
实现CUTLASS 原语 (WGMMA, TMA)
序列长度512, 1k, …, 16k
总 token 数16k
Hidden dimension2048
Head dimension64, 128, 256

性能基准

FP16 前向加速

配置FA-3 vs FA-2FA-3 vs 标准注意力
Head dim 641.5-1.8×3-10×
Head dim 1281.5-2.0×3-12×
Head dim 2561.5-1.9×5-16×

关键发现:

  • 中长序列 (1k+) 时 FA-3 超越 cuDNN 优化实现
  • 达到 740 TFLOPS/s (75% 利用率)
  • 后向加速 1.5-1.75×

FP8 前向

  • 达到接近 1.2 PFLOPS/s
  • Head dim 256 时性能对比见 Figure 7

消融实验

配置时间TFLOPS/s
FlashAttention-33.538 ms661
无 GEMM-Softmax 流水线, 有 Warp 特化4.021 ms582
有 GEMM-Softmax 流水线, 无 Warp 特化4.105 ms570

关键发现: 两种技术都重要,组合效果最佳。

数值精度

方法RMSE
FP16
Baseline FP163.2e-4
FlashAttention-2 FP161.9e-4
FlashAttention-3 FP161.9e-4
FP8
Baseline FP8 (per-tensor)2.4e-2
FlashAttention-3 FP89.1e-3
无块量化9.3e-3
无不相干处理2.4e-2

关键发现:

  • FP16: FA-3 与 FA-2 数值误差相同,比标准实现低 1.7×(因 softmax 保持 FP32)
  • FP8: FA-3 比基线 FP8 低 2.6× 数值误差

测试数据分布

模拟 LLM 异常值特征: N(0,1)+N(0,100)⋅Bernoulli(0.001)\mathcal{N}(0, 1) + \mathcal{N}(0, 100) \cdot \text{Bernoulli}(0.001)

即每个条目服从零均值、标准差 1 的正态分布,但 0.1% 的条目加上标准差 10 的独立项。

六、关键算法细节

Warp 特化实现

  • 使用 setmaxnreg 进行寄存器 (de)allocations
  • TMA 加载 QiQ_i 和 {Kj,Vj}\{K_j, V_j\}
  • WGMMA 执行消费者主循环中的 GEMMs
  • SS 前缀表示第一个操作数来自 SMEM,RS 表示来自 RMEM

FP8 布局转换

问题: FP32 累加器和 FP8 操作数 A 的寄存器布局不同

解决方案: 字节置换指令,将序列改为 {d0 d1 d4 d5 d2 d3 d6 d7}\{d0\ d1\ d4\ d5\ d2\ d3\ d6\ d7\},每 8 字节复制

V 的内核内转置:

  • 使用 LDSM/STSM 指令(128 字节粒度)
  • 生产者 warpgroup 执行
  • 可在前一个 VV 块和当前 KK 块的两个 WGMMA 阴影中执行

不相干处理

选择 MM 为: M=D1HD2HD3M = D_1 H D_2 H D_3

其中 DiD_i 为随机 ±1\pm 1 对角矩阵,HH 为 Hadamard 矩阵。

计算复杂度 O(dlog⁡d)O(d \log d),可融合到 rotary embedding。

七、与 FlashAttention-2 的对比

方面FlashAttention-2FlashAttention-3
目标硬件AmpereHopper
异步利用有限充分 (TMA, WGMMA)
Warp 特化无生产者/消费者分离
GEMM-Softmax 重叠无2/3 阶段流水线
FP8 支持无块量化 + 不相干处理
H100 利用率35%75%
FP16 前向基准1.5-2.0×
FP8N/A~1.2 PFLOPS

八、相关工作

相关工作与本文关系
FlashAttention前身,引入 tiling 策略
FlashAttention-2直接前作,FA-3 的基准
ThunkerKittenHopper 特定指令简化实现
cuDNN 9厂商优化实现,FA-3 在长序列上超越
CUTLASS提供 WGMMA 和 TMA 抽象
QuIP不相干处理的灵感来源

九、总结

核心贡献

  1. 生产者-消费者异步: Warp 特化软件流水线,利用 Hopper 异步执行
  2. GEMM-Softmax 重叠: 将 softmax 隐藏在异步 WGMMA 下
  3. FP8 注意力: 块量化 + 不相干处理,1.2 PFLOPS 且精度高
  4. 75% H100 利用率: 从 35% 提升到 75%
  5. 开源: 宽松许可证,集成 PyTorch 和 Hugging Face

技术影响

  • 异步编程范式: Warp 特化成为 GPU 高性能计算的标准技术
  • 低精度注意力: FP8 注意力可行且准确
  • 长上下文加速: 1.5-2.0× 加速解锁更长上下文应用
  • 硬件-算法协同设计: 充分利用 Hopper 特性

局限性

  • 仅针对 Hopper GPU,需适配其他架构
  • FP8 内核未集成持久内核设计
  • 未研究低精度注意力在大规模训练中的效果
  • 3 阶段流水线寄存器压力大,需权衡 tile 大小

十、参考资源