M4-049M4: Sequences & TransformersEfficient Attention & FlashAttentionMedium
Mastery:
Efficient Attention & FlashAttention: 解释 FlashAttention 的 backward 如何避免重算整块。
📐 Mathematical Definition
⚡ Executive Summary
Core Concept: 保存每行的 (m, ℓ) 与输出 O,反向时用它们重算 S、P(不存 L×L 矩阵),显存 O(L)。
📌 Key Takeaways
- •只保存 O(L) 的统计量,不保存 O(L²) 的 P 矩阵
- •反向时重算 S 与 P(需要 Q/K/V 分块)
- •用'分块重算 + 累积'避免物化
📐 Mathematical Derivations
数学机理:<strong>标准反向的问题</strong>——注意力反向需要 dP(softmax 输出 P 的梯度)与 dS;而 dP 的计算需要 P=softmax(S),故标准实现<strong>保存 P(L×L)</strong>,显存 O(L²)。<strong>Flash 的做法</strong>——前向只保存每行的<strong>两个标量</strong> (m_i, ℓ_i)(running max 与 running sum)与输出 O(O(L·d));反向时<strong>重新计算</strong> P:P=exp(S−m)/ℓ,其中 S 由 Q/K 分块重算(S=QKᵀ,只需当前分块)。这样显存从 O(L²) 降到 <strong>O(L·d)</strong>(与序列长度线性)。<strong>反向的分块累积</strong>——与在线 softmax 类似,反向也可分块进行:遍历 K/V 分块,计算 dQ、dK、dV 的贡献并累积(dQ 需在所有 K 块上累积、dK/dV 在 Q 块上累积);其中也需处理'重缩放'(因为 m、ℓ 在分块中变化)。<strong>额外 FLOPs</strong>——重算 S 需要一次额外的 QKᵀ 与 softmax(约等于前向的注意力计算),故反向的额外开销约 +33%(前向 1 + 重算 1 + 反向 1 ≈ 3 vs 原 2)。<strong>关键收益</strong>——(a) 显存 O(L²)→O(L),使长序列训练可行;(b) HBM 读写大幅减少(不需读写 L×L 矩阵),故<strong>反向也更快</strong>(尽管 FLOPs 略增)。这是'用计算换显存,同时因减少 IO 而净提速'的经典案例。
🏭 Production Trade-offs
深度剖析与工程权衡:① <strong>与梯度检查点的对比</strong>——Flash Attention 的反向重算可视为'注意力内部的梯度检查点'(保存少量统计量、重算中间结果);两者思想一致,但 Flash 的重算是<strong>kernel 内部</strong>的、更细粒度且避免了 HBM 往返。② <strong>'IO 减少 > FLOPs 增加'</strong>——这是 memory-bound 场景的核心规律;若某优化减少 IO 但增加 FLOPs,仍可能净提速(只要算力有冗余)。面试中能说出这一规律很有说服力。③ <strong>与确定性</strong>——分块计算的浮点求和顺序与标准实现不同,故结果有微小差异;训练时通常可接受,但需注意'可复现性'要求(见 M3 可复现性题)。④ <strong>FlashAttention-2/3 的改进</strong>——FA2 优化了并行度与 warp 调度(减少非 matmul 的 FLOPs、改善 occupancy);FA3 针对 Hopper 架构用 TMA 与 warp-specialization 进一步提升。⑤ <strong>与 KV cache 的交互</strong>——训练时 Flash 不存 P;但<strong>推理 decode</strong> 时需读取 KV cache(已存),此时瓶颈是 KV 的读取带宽(见 Flash-Decoding 题)。⑥ <strong>面试要点</strong>——被问'Flash 的反向怎么省显存',应给出'<strong>只存 (m, ℓ, O) 三个 O(L) 量、反向重算 S 与 P、分块累积</strong>',并说明'额外 FLOPs ≈ +33% 但因减少 IO 而净提速';能联系到梯度检查点是加分。
⚠️ Common Interview Pitfalls
- ✕以为反向必须保存 P 矩阵
- ✕忽略重算带来的额外 FLOPs(约 +33%)
🎯 Interviewer Follow-ups
- ?为什么反向的显存也是 O(L)?
- ?重算的额外 FLOPs 有多少?
📚
Associated Knowledge Base Guides & Mindmaps
Explore the comprehensive technical article, exam cards, and global architecture tree.