news 2026/10/7 8:29:21

FlashAttention 的反向传播设计:在 SRAM 内部重计算 Softmax 的显存收益

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
FlashAttention 的反向传播设计:在 SRAM 内部重计算 Softmax 的显存收益

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$。

根据多元微积分的链式法则:

  1. 关于 $V$ 的梯度:
    $$d V = P^T d O$$
    这里直接依赖于前向计算出的完整概率矩阵 $P \in \mathbb{R}^{N \times N}$。
  2. 关于中间概率 $P$ 的梯度:
    $$d P = d O V^T$$
  3. 关于原始注意力得分 $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}$。
  4. 关于 $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$ 矩阵相比,显存体积直接暴降了数万倍!

而在反向传播执行时:

  1. 同样按照块尺寸(Block Size)将 $Q, K, V, d O$ 的微小子块分批加载到片上高速 SRAM 中;
  2. 利用保存的标量 $m$ 和 $l$,直接在片上由加载的 $Q_{\text{tile}}$ 与 $K_{\text{tile}}$ 重新计算出当前局部的 $S_{\text{tile}} = Q_{\text{tile}} K_{\text{tile}}^T / \sqrt{d}$;
  3. 原地利用公式 $P_{\text{tile}} = \exp(S_{\text{tile}} - m) / l$在寄存器内瞬间重构出该子块的精确注意力概率;
  4. 概率矩阵重构出来的刹那,立即与加载的 $d O_{\text{tile}}$ 和 $V_{\text{tile}}$ 完成矩阵乘加,更新梯度累加器;
  5. 计算完毕后立即丢弃该子块,绝对不向外层显存做任何多余搬运。

三、现代 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]; } } } }

四、显存收益与“以计算换访存”的数学算力账本

很多第一次了解重计算机制的工程师会问:多算了一次矩阵乘法与指数计算,算法整体不是变慢了吗?

在现代处理器的物理世界里,这笔账有着极其惊人的反直觉结论:

  1. GPU 核心计算速度极快,但显存带宽极其狭窄:
    • 现代 GPU(如 H100)的片上浮点算力高达数百 TFLOPS,但全局显存带宽只有不到 3 TB/s;
    • 在片上 SRAM 内部重新计算一次 $Q K^T$ 和指数,消耗的时间通常不足几微秒;
    • 而如果要把庞大的 $P$ 矩阵从全局显存读进写出,总线搬运的延迟高达数十微秒!
  2. 算子执行时间不仅没变慢,反而快了 2 到 4 倍:
    • 因为消除了全局显存的密集写回与重读,反向传播从“访存瓶颈(Memory-Bound)”瞬间转变为“纯计算密集(Compute-Bound)”;
    • 同时,由于显存占用从 $O(N^2)$ 断崖式下降到 $O(N)$,系统支持的单卡训练 Batch Size 和最大序列长度直接暴增了 5 到 10 倍!

FlashAttention 反向传播的设计,深刻揭示了现代高性能算子优化的终极法则:不要吝啬廉价的计算周期,去全力拯救昂贵而脆弱的内存总线。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/7 8:29:21

拦截 AI 生成的弱加密算法:MD5/SHA1 自动替换为安全散列

拦截 AI 生成的弱加密算法&#xff1a;MD5/SHA1 自动替换为安全散列在很多安全审计报告中&#xff0c;一个高频出现的低级漏洞让安全团队大跌眼镜&#xff1a;在 2026 年的今天&#xff0c;某些核心系统的密码存储或数据签名逻辑中&#xff0c;竟然依然赫然写着 crypto/md5 或 …

作者头像 李华
网站建设 2026/10/7 8:29:02

让AI自己管理跨机器Skill:一句话同步安装实战

1. 二十个 skill 散在三台电脑&#xff0c;这件事到底难在哪先说清楚这个场景。我手头有三台机器&#xff1a;一台主力台式机放在家里&#xff0c;一台笔记本随身带着跑客户现场&#xff0c;还有一台放在公司工位的开发机。三台机器上各自装了一堆 AI 编程助手的 skill——有的…

作者头像 李华
网站建设 2026/10/7 8:29:01

深圳科飞时速推出桌面级AI应用软件 -初元AI 24天内迭代三个版本,面向零基础用户提供建站与业务软件生成能力

【深圳&#xff0c;2026年10月】深圳科飞时速创始人徐宝林近日宣布&#xff0c;其团队推出的桌面级AI应用生成器“初元AI”已在24天内连续发布三个版本&#xff1a;V1.0于9月11日发布&#xff0c;V1.1于9月19日发布&#xff0c;V1.2于9月30日发布。定位让零基础用户通过自然语言…

作者头像 李华
网站建设 2026/10/7 8:28:59

利用大模型自动分析锁等待链:从 sys.innodb_lock_waits 提炼瓶颈事务

利用大模型自动分析锁等待链&#xff1a;从 sys.innodb_lock_waits 提炼瓶颈事务在核心高并发交易数据库中&#xff0c;行级锁等待堆积&#xff08;Row Lock Contention&#xff09;是引发服务熔断的最凶险元凶。当某个业务模块在事务内执行长耗时远程调用&#xff08;RPC&…

作者头像 李华