系统工程FlashAttention
反向传播实现
会员专享实现 Flash Attention 的梯度计算,通过 Recomputation 实现内存高效的训练。
在权益中心获取代码在前面的章节中,我们实现了完整的 Flash Attention 前向传播,支持任意序列长度、Causal Masking 和 GQA。但这还不够——要在训练中使用 Flash Attention,我们需要实现反向传播(Backward Pass)来计算梯度。
本章将探讨如何在保持 IO 效率的同时计算 、、。
为什么需要自定义反向传播?
PyTorch 自动微分的局限性
PyTorch 的 autograd 可以自动为大部分操作生成反向传播代码,但对于 Flash Attention 这样的融合 Kernel,自动微分会遇到问题:
# 标准 Attention 的前向传播
def standard_attention(Q, K, V):
S = Q @ K.T / sqrt(d) # (1) 计算分数
P = softmax(S, dim=-1) # (2) Softmax
O = P @ V # (3) 加权求和
return O自动微分的行为:
- PyTorch 会保存中间结果 和 (Attention Matrix)
- 在反向传播时使用这些中间结果计算梯度
- 内存开销: 用于存储 Attention Matrix
但 Flash Attention 的核心优势就是不物化(materialize) Attention Matrix!
如果使用自动微分,我们会:
- 前向传播:不保存 Attention Matrix,内存 ✅
- 反向传播:需要 Attention Matrix,被迫重新计算,失去了优势 ❌
Recomputation 策略
Flash Attention 采用了一个巧妙的权衡:Recomputation(重计算)
核心思想:
- 前向传播时:只保存少量的中间统计量 ( 和 )
- 反向传播时:使用这些统计量重新计算 Attention 分数,而非从 HBM 加载
传统方法 (Materialization):
Forward: 计算 S, P → 保存到 HBM (O(N²) 内存)
Backward: 从 HBM 读取 S, P → 计算梯度
Flash Attention (Recomputation):
Forward: 计算 S, P → 只保存 L, M (O(N) 内存)
Backward: 重新计算 S, P → 计算梯度权衡分析:
- ✅ 内存节省:
- ⚠️ 额外计算:需要重新计算 Attention(约 1.5x FLOPS)
- ✅ IO 效率:仍然比标准实现快,因为避免了 HBM 往返
为什么 Recomputation 仍然更快?
虽然增加了 FLOPS,但现代 GPU 是 IO-bound 而非 compute-bound。重新计算 Attention 的 FLOPS 开销小于从 HBM 加载 矩阵的 IO 开销。
这就是为什么 Flash Attention 即使在反向传播中也能保持性能优势。
登录以继续阅读
这是一篇付费内容,请登录您的账户以访问完整内容。
CookLLM文档