
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 内部原地重计算RecomputeSoftmax本文我们将严格推导注意力反向传播的链式求导公式剖析重计算机制的数学必然性与巨大的显存降维收益。一、传统 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::spanconst float Q_tile, // [Br x d] std::spanconst float K_tile, // [Bc x d] std::spanconst float V_tile, // [Bc x d] std::spanconst float dO_tile, // [Br x d] std::spanconst float row_m, // [Br] 前向保存的行最大值 std::spanconst float row_l, // [Br] 前向保存的归一化分母 std::spanfloat dQ_tile, // [Br x d] 梯度累加 std::spanfloat dK_tile, // [Bc x d] 梯度累加 std::spanfloat 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::vectorfloat S_tile(Br * Bc, 0.0f); std::vectorfloat 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::vectorfloat 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 反向传播的设计深刻揭示了现代高性能算子优化的终极法则不要吝啬廉价的计算周期去全力拯救昂贵而脆弱的内存总线。