FlashAttention 的反向传播设计:在 SRAM 内部重计算 Softmax 的显存收益
在探讨 FlashAttention 时,大多数人的目光往往集中在其前向传播(Forward Pass)中如何通过 Tiling 分块与 Online Softmax 消除 $N \times N$ 中间注意力得分矩阵的显存写回。
然而,在深度学习全生命周期中,反向传播(Backward Pass)才是对显存容量与总线带宽更为残酷的终极考验。
在传统的 PyTorch 标准实现中,训练长序列大模型时发生 OOM(显存溢出)的元凶几乎全部来自前向阶段保存的中间激活值(Saved Tensors)。为了在反向传播中计算关于查询矩阵 $Q$、键矩阵 $K$ 和值矩阵 $V$ 的梯度($d Q, d K, d V$),系统必须把尺寸为 $N \times N$ 的完整 Softmax 注意力概率矩阵 $P$ 硬生生保存在高昂的全局显存(HBM)中,直到反向传播执行完毕才能释放。
FlashAttention 在反向传播上的突破性设计,展现了极其精湛的算子工程哲学:拒绝在显存中保留任何 $N \times N$ 的中间激活值;仅保留极其轻量的行统计标量,并在反向传播时,直接在片上高速 SRAM 内部原地重计算(Recompute)Softmax!
本文我们将严格推导注意力反向传播的链式求导公式,剖析重计算机制的数学必然性与巨大的显存降维收益。
一、传统 Attention 反向传播的显存绝境
我们先回顾标准 Self-Attention 的前向输出与梯度推导:
前向计算公式:
$$S = \frac{Q K^T}{\sqrt{d}}, \quad P = \text{softmax}(S), \quad O = P V$$
在反向传播阶段,上层网络反向传递回输出张量的梯度 $d O \in \mathbb{R}^{N \times d}$。我们需要求解 $d Q, d K, d V$。
根据多元微积分的链式法则:
- 关于 $V$ 的梯度:
$$d V = P^T d O$$
这里直接依赖于前向计算出的完整概率矩阵 $P \in \mathbb{R}^{N \times N}$。 - 关于中间概率 $P$ 的梯度:
$$d P = d O V^T$$ - 关于原始注意力得分 $S$ 的梯度:
Softmax 算子的雅可比矩阵求导展开后极其复杂,其矩阵形式为:
$$d S = P \circ (d P - D)$$
其中 $\circ$ 表示逐元素相乘(Hadamard Product),$D \in \mathbb{R}^{N}$ 是一个一维归约向量,其第 $i$ 行元素为 $D_i = \sum_j (d P){i,j} \cdot P{i,j}$。 - 关于 $Q$ 和 $K$ 的梯度:
$$d Q = \frac{1}{\sqrt{d}} d S \cdot K, \quad d K = \frac{1}{\sqrt{d}} d S^T \cdot Q$$
审视这一连串数学公式:
无论是计算 $d V$ 还是推导 $d S$,$P$ 矩阵都如影随形。
当上下文长度 $N = 32768$(32K)时,单头注意力矩阵 $P$ 拥有超过10.7 亿个元素。以 FP16 存储,单个注意力头仅这一项就需要霸占 2.14GB 显存;一个 32 个注意力头的模型层,单层前向激活值就要消耗近 70GB 显存!
传统框架为了能够在有限的 GPU 上跑训练,不得不引入耗时巨大的激活值检查点(Activation Checkpointing)把整层网络重新跑一遍,带来了沉重的算力惩罚。
二、FlashAttention 的破局点:以片上重计算换显存解脱
FlashAttention 反向传播的核心破局点极其纯粹:
前向传播时,绝对不向全局显存写回 $N \times N$ 的 $P$ 矩阵;取而代之的是,前向传播仅仅把每一行最后收敛的两个标量:全局最大值 $m \in \mathbb{R}^N$ 与全局归一化分母 $l \in \mathbb{R}^N$,写回到全局内存中。
这两个标量向量的显存占用是纯粹的 $O(N)$ 线性复杂度:
对于 $N = 32768$ 的序列,一行两个 float32 标量,总共仅占用区区256KB显存!与原本数十吉字节的 $P$ 矩阵相比,显存体积直接暴降了数万倍!
而在反向传播执行时:
- 同样按照块尺寸(Block Size)将 $Q, K, V, d O$ 的微小子块分批加载到片上高速 SRAM 中;
- 利用保存的标量 $m$ 和 $l$,直接在片上由加载的 $Q_{\text{tile}}$ 与 $K_{\text{tile}}$ 重新计算出当前局部的 $S_{\text{tile}} = Q_{\text{tile}} K_{\text{tile}}^T / \sqrt{d}$;
- 原地利用公式 $P_{\text{tile}} = \exp(S_{\text{tile}} - m) / l$在寄存器内瞬间重构出该子块的精确注意力概率;
- 概率矩阵重构出来的刹那,立即与加载的 $d O_{\text{tile}}$ 和 $V_{\text{tile}}$ 完成矩阵乘加,更新梯度累加器;
- 计算完毕后立即丢弃该子块,绝对不向外层显存做任何多余搬运。
三、现代 C++ 模拟片上 SRAM 反向重计算闭环
我们用现代 C++ 代码来清晰展现反向传播在 SRAM 分块内部的重计算与梯度聚合流程:
#include <vector> #include <cmath> #include <algorithm> #include <span> #include <iostream> void flash_attention_backward_tile( std::span<const float> Q_tile, // [Br x d] std::span<const float> K_tile, // [Bc x d] std::span<const float> V_tile, // [Bc x d] std::span<const float> dO_tile, // [Br x d] std::span<const float> row_m, // [Br] 前向保存的行最大值 std::span<const float> row_l, // [Br] 前向保存的归一化分母 std::span<float> dQ_tile, // [Br x d] 梯度累加 std::span<float> dK_tile, // [Bc x d] 梯度累加 std::span<float> dV_tile, // [Bc x d] 梯度累加 size_t Br, size_t Bc, size_t d, float scale ) { // 1. 在片上 SRAM 内部原地重计算局部 S_tile = Q_tile * K_tile^T * scale std::vector<float> S_tile(Br * Bc, 0.0f); std::vector<float> P_tile(Br * Bc, 0.0f); for (size_t r = 0; r < Br; ++r) { float m = row_m[r]; float l = row_l[r]; for (size_t c = 0; c < Bc; ++c) { float dot = 0.0f; for (size_t k = 0; k < d; ++k) { dot += Q_tile[r * d + k] * K_tile[c * d + k]; } dot *= scale; S_tile[r * Bc + c] = dot; // 原位重计算精准的 Softmax 概率值! P_tile[r * Bc + c] = std::exp(dot - m) / l; } } // 2. 原地计算 dV_tile = P_tile^T * dO_tile for (size_t c = 0; c < Bc; ++c) { for (size_t k = 0; k < d; ++k) { float sum = 0.0f; for (size_t r = 0; r < Br; ++r) { sum += P_tile[r * Bc + c] * dO_tile[r * d + k]; } dV_tile[c * d + k] += sum; } } // 3. 计算中间梯度 dP_tile = dO_tile * V_tile^T std::vector<float> dP_tile(Br * Bc, 0.0f); for (size_t r = 0; r < Br; ++r) { for (size_t c = 0; c < Bc; ++c) { float dot = 0.0f; for (size_t k = 0; k < d; ++k) { dot += dO_tile[r * d + k] * V_tile[c * d + k]; } dP_tile[r * Bc + c] = dot; } } // 4. 计算 dS_tile 与 dQ, dK (利用行内积归约项 D_i) for (size_t r = 0; r < Br; ++r) { // 计算 D_i = sum_c (dP_tile * P_tile) float Di = 0.0f; for (size_t c = 0; c < Bc; ++c) { Di += dP_tile[r * Bc + c] * P_tile[r * Bc + c]; } // 计算 dS = P * (dP - D) 并累加到 dQ for (size_t c = 0; c < Bc; ++c) { float p_val = P_tile[r * Bc + c]; float ds_val = p_val * (dP_tile[r * Bc + c] - Di) * scale; for (size_t k = 0; k < d; ++k) { dQ_tile[r * d + k] += ds_val * K_tile[c * d + k]; dK_tile[c * d + k] += ds_val * Q_tile[r * d + k]; } } } }四、显存收益与“以计算换访存”的数学算力账本
很多第一次了解重计算机制的工程师会问:多算了一次矩阵乘法与指数计算,算法整体不是变慢了吗?
在现代处理器的物理世界里,这笔账有着极其惊人的反直觉结论:
- GPU 核心计算速度极快,但显存带宽极其狭窄:
- 现代 GPU(如 H100)的片上浮点算力高达数百 TFLOPS,但全局显存带宽只有不到 3 TB/s;
- 在片上 SRAM 内部重新计算一次 $Q K^T$ 和指数,消耗的时间通常不足几微秒;
- 而如果要把庞大的 $P$ 矩阵从全局显存读进写出,总线搬运的延迟高达数十微秒!
- 算子执行时间不仅没变慢,反而快了 2 到 4 倍:
- 因为消除了全局显存的密集写回与重读,反向传播从“访存瓶颈(Memory-Bound)”瞬间转变为“纯计算密集(Compute-Bound)”;
- 同时,由于显存占用从 $O(N^2)$ 断崖式下降到 $O(N)$,系统支持的单卡训练 Batch Size 和最大序列长度直接暴增了 5 到 10 倍!
FlashAttention 反向传播的设计,深刻揭示了现代高性能算子优化的终极法则:不要吝啬廉价的计算周期,去全力拯救昂贵而脆弱的内存总线。