返回 AI 基础设施 思维导图
中文·English
🖥️ AI 基础设施ID: flashattention-kernel

FlashAttention 与 Online Softmax

FlashAttention Kernel & Online Softmax
🎯核心定义
FlashAttention = 对标准注意力做 GPU 访存优化:通过 SRAM 切块(Tiling)与 Online Softmax 增量归一化,把 HBM 读写从 O(N2)O(N^2) 降到 O(N2/M)O(N^2 / M)(大 NN 时近似 O(N)O(N),MM 为 SRAM 块容量)。标准注意力 S=QKdS = \frac{QK^\top}{\sqrt{d}}、softmax、O=SVO = SV,N×NN \times NSS 矩阵需写回 HBM 再读回——这是 O(N2)O(N^2) 的访存瓶颈;FlashAttention 把 Q/K/VQ/K/V 切成块常驻 SRAM,逐块计算局部 softmax 并用运行最大值 + 重缩放合并: 处理元素 xix_i 时,mnew=max(m,xi)m_{new} = \max(m, x_i),lnew=lemmnew+eximnewl_{new} = l \cdot e^{m - m_{new}} + e^{x_i - m_{new}},输出按 emmnewe^{m - m_{new}} 重缩放,最终 1l\frac{1}{l} 归一化,数学结果与标准 softmax 完全一致。
💡使用场景
长序列注意力(4K/32K/1M 上下文)训练与推理的标准 kernel;面试高频“FlashAttention 为什么快”“Online Softmax 是什么”“FA-1/2/3 区别”;也是 PD 分离、KV 量化等推理优化的底层支撑。
解决的核心痛点
标准注意力受 HBM 带宽支配:长序列下 SS 矩阵 O(N2)O(N^2) 读写主导耗时,且显存 O(N2)O(N^2) 存不下。FlashAttention-1(2022)用切块 + Online Softmax + 反传重算(不存 SS)把 HBM 访问降一个数量级、显存从 O(N2)O(N^2) 降到 O(N)O(N),64K 序列上比 PyTorch eager 快 2-4 倍;FA-2(2023)做多头并行、把 N2N^2 的 softmax 开销移出并行域,约再快 2 倍;FA-3(2024)用 Hopper 的 warp specialization + TMA 异步拷贝,并把 FP8 引入 kernel。
🎯5 个高频面试考点 (Exam Points)
1
推导标准注意力的 HBM 瓶颈: 为什么 S=QKdS = \frac{QK^\top}{\sqrt{d}} 与 softmax 之间的 N×NN \times N 矩阵往返是 O(N2)O(N^2) 访存;长序列(如 128K)下为什么带宽主导?
2
写出 Online Softmax 增量公式: mnew=max(m,xi)m_{new} = \max(m, x_i)lnew=lemmnew+eximnewl_{new} = l \cdot e^{m - m_{new}} + e^{x_i - m_{new}};为什么需要 running max + 重缩放而不能直接累加 exie^{x_i}?
3
Tiling 细节: 块大小受什么限制(每 SM SRAM 如 H100 约 228KB)?HBM 访问如何从 O(N2)O(N^2) 降到 O(N2/M)O(N^2 / M);causal mask 如何在 kernel 内处理?
4
FA-1/2/3 演进: 反传重算为何把显存从 O(N2)O(N^2) 降到 O(N)O(N);FA-2 的多头并行与 upcast 优化;FA-3 的 warp specialization/TMA 异步与 FP8 各解决什么?
5
FlashAttention 与 KV cache/PagedAttention 的关系: kernel 层优化为何不与推理引擎层优化冲突;为什么 decode 阶段带宽瓶颈(逐 token 读权重)不能靠 FA 解决?
更新于 2026-08-12
🎯
检验攻克程度:针对「FlashAttention 与 Online Softmax」专属刷题排雷
做单选排雷题、推导选项机制,答错自动收录进专属错题本。
🚀 开始本考点专项刷题
上一个知识点推理量化(跨模块)下一个知识点KV 缓存优化

🔗 更多 AI 基础设施 知识点卡片

激活显存估算Agent 运行时(跨模块)弹性伸缩与成本优化检查点与故障恢复