Skip to content
🔗 分享本题
查看我的学习进度 →

FlashAttention动漫知识图:传统Attention在HBM上两次读写中间矩阵,FlashAttention分块将计算全部留在SRAM,通过在线softmax和重计算避免写出QK^T和Attention矩阵

🧠 图解记忆:注意力计算的瓶颈不是算不快,而是数据搬太多——把计算移到SRAM里做,少搬数据就快了。

FlashAttention = I/O 感知的注意力算法:把计算拆成小方块放在 GPU SRAM 上做,避免反复从 HBM(高带宽显存)读写中间结果。

核心问题:标准 Attention 的 I/O 瓶颈

标准 Softmax(QK^T)V 的计算流程:
1. Q @ K^T → M×M 中间矩阵(写入 HBM)
2. softmax(M×M) → 又一个中间矩阵(写入 HBM)
3. Softmax_M × V → 最终输出

问题:
- 对于一个 seq_len=8192 的序列,QK^T 就是 8192×8192 ≈ 67M 个 float16 = 128MB
- 长序列下中间矩阵远大于 SRAM,只能存在 HBM
- HBM 读写的速度远慢于 SRAM(约 10~20 倍差距)
- 每次 block 计算完都要写回 HBM、下次再读回来 → 大量时间浪费在搬运数据

FlashAttention 的三大优化

1. Tiling(分块计算)

把 Q、K、V 分成小 blocks,每个 block 大小适配 SRAM:
  Block-Q (n×d) + Block-K (m×d)^T → Block-QK^T → softmax → Block-O (n×d)

每个 block 的计算完全在 SRAM 内完成,只有最终输出写回 HBM
→ 大幅减少 HBM 读写次数

2. Online Softmax

传统 softmax 需要两遍遍历数据(先求 max,再归一化)。
FlashAttention 用递推方式一次搞定:

在线 softmax 递推公式:
给定前 i 个块的 max(m_i) 和 sum(s_i),加入第 i+1 块后:
  m_new = max(m_i, m_{i+1})
  s_new = s_i * exp(m_i - m_new) + s_{i+1} * exp(m_{i+1} - m_new)
  o_new = (o_i * exp(m_i - m_new) + o_{i+1} * exp(m_{i+1} - m_new)) / s_new

这样每处理一个 block 就能更新全局统计量,不需要保存整个中间矩阵!

3. Re-computation(重计算)

因为不保存中间矩阵,反向传播时需要重新计算 QK^T。
但重计算的成本远低于存储和传输中间矩阵的成本。

权衡:多一次正向计算 ↔ 省掉 M×M 显存存储
对于长序列场景,后者收益巨大

性能对比

指标标准 AttentionFlashAttention
显存占用O(n²)(存完整中间矩阵)O(n)(只存分块结果)
训练速度基准1.5~2× 更快
支持的最大序列长度受限于显存大幅提升
数值精度标准 softmax数值稳定(在线算法保证)

面试高频追问

  • FlashAttention-2 vs FlashAttention-1 有什么区别? FA-2 进一步优化了 CUDA kernel,减少了寄存器压力并改进了 load/store 调度,实测再提速 ~25%
  • FlashAttention 能用于 Transformer Encoder 吗? 可以——任何使用注意力机制的场景都可以受益,不限于 Decoder-only
  • FlashAttention 对硬件有要求吗? 需要较新的 GPU(如 A100/H100 的更大 SRAM),对消费级显卡也有优化版本

面试话术:

"FlashAttention 的核心思想是 I/O 感知设计——不是让计算更快,而是让数据搬运更少。传统 Attention 要把 QK^T 这个巨大的中间矩阵写到显存再读回来,而 FlashAttention 把 Q、K、V 切成小块全放在 SRAM 里算完,用在线 Softmax 递推替代全矩阵操作,最后只把输出写回显存。实际效果是训练速度提升 1.5~2 倍,而且显存占用从 O(n²) 降到 O(n),这让超长上下文训练成为可能。现在几乎所有主流 LLM 框架都内置了 FlashAttention。"