跳转至

激活值优化

🧠 梯度检查点的原理是什么?用多少计算换多少显存?

💡 梯度检查点(Gradient Checkpointing)的核心思想是在前向传播时丢弃大部分中间激活,反向传播时再通过额外的前向计算重新生成它们,从而用大约 20-30% 的额外计算量,换取 50-80% 的激活显存节省。

🔬 原理:

  • 训练时,反向传播需要用到前向传播产生的中间结果(即激活值)来计算梯度。通常这些激活会全部保留在显存中,等待反向使用。

  • 梯度检查点将网络分成若干段(checkpoint 段)。在前向时,只保留每一段输入的激活(即检查点),而段内各层的中间激活全部丢弃。

  • 反向传播到该段时,利用保留的检查点输入,重新执行该段的前向计算,即时生成所需的中间激活,然后再计算梯度。这样,同一段内的激活峰值只存在一个段的大小,而不是整个网络。

📊 计算-显存权衡:

  • 计算增量:每个检查点段需要重新前向一次。若将模型均匀分为 √N 个段(最优策略),则额外计算量约为 1/√N。实际通常将每个 Transformer Block 作为一个检查点,则额外前向次数等于 Block 数(但仅重计算激活,不更新参数),额外计算量约为 20-30%。

  • 显存节省:原本需要存储所有 Block 的全部激活,现在只需存储每个 Block 的输入(少量),中间激活几乎全部释放。激活显存可降至原来的 O(√N) 或更低,节省 50-80% 很常见。

✅ 因此,梯度检查点是用可控的计算代价,大幅压缩激活显存,是长序列或大 batch 训练不可或缺的技术。


🧱 为什么梯度检查点通常加在 Transformer 的每个 Block 上?

💡 Transformer Block 具有天然的模块化和顺序依赖,且每个 Block 内部激活量巨大但输入很小,将检查点设在 Block 边界能以最小存储成本保留重计算所需的“种子”,最大化显存节省效率。

🔍 原因:

  • 每个 Block 输入很小:一个 Block 的输入是 hidden states,形状 [batch, seq, hidden],远小于 Block 内部产生的 QKV、注意力矩阵、FFN 中间激活等。保存这个输入就能重算整个 Block,存储成本极低。

  • 模块化天然分段:Block 的输入输出是清晰的分界点,不需要额外分割。在框架中只需标记 Block 的边界,就能轻松实现检查点。

  • 最优计算/存储比:若将检查点设在更小的粒度(如每个 Linear 层),重计算次数增多,但节省的显存增量减少;设在更大粒度(如多个 Block)则保存的检查点输入增多,显存节省下降。逐 Block 是工程上的甜点。

📌 因此,几乎所有框架(PyTorch、Hugging Face、Megatron)默认将 gradient checkpointing 应用于每个 Transformer Block,达到极佳平衡。


⏱️ 梯度检查点对训练速度的影响有多大?

💡 开启梯度检查点通常会使训练速度降低 15%~30%,具体取决于模型结构、检查点粒度和计算/通信比重。但结合其他优化后,吞吐下降往往可接受,甚至因能增大 batch 而总体吞吐不降反升。

📊 影响因素:

  • 额外前向计算:重计算带来约 20-30% 的额外 FLOPs,若 GPU 计算能力充足,这部分可能会完全转化为延迟增加。

  • 计算与通信重叠:在分布式训练中,重计算可以在反向传播时并行,填充原本的流水线气泡,实际速度损失可能小于理论值。

  • 批次大小变化:开启检查点后显存宽裕,可以增大 batch size 或序列长度,提高 GPU 利用率,有时总吞吐反而提升。

  • 框架实现:PyTorch 的 checkpoint 函数有额外开销(如 Python 调用栈、保存/恢复张量)。使用 torch.utils.checkpoint.checkpoint 可能比手动重计算慢一些,但已充分优化。

🔧 实测经验:在训练 7B 模型时,开启逐 Block 检查点,单个迭代时间增加约 20%。但 batch size 可以从 16 扩大到 64,总 tokens/s 提升近 2 倍。因此,不能仅看迭代延迟,要从系统吞吐角度评估。

✅ 因此,梯度检查点的速度代价通常可接受,尤其当它允许使用更大 batch 时,整体训练效率往往不降反升。


⚡ FlashAttention 如何通过分块和重计算减少显存?

💡 FlashAttention 通过将注意力矩阵分块(tiling),在片上 SRAM 内完成 Softmax 的增量计算,并依赖重计算(recomputation)丢弃中间大矩阵,从而避免存储完整的 [seq, seq] 注意力分数矩阵,将显存复杂度从 O(N²) 降至 O(N)。

🔬 核心机制:

  1. 分块(Tiling):将 Q、K、V 矩阵沿序列维度切分成小块。计算时,逐块加载到 SRAM 中,使用在线 Softmax 算法(numerically stable)累积归一化因子,从而在块内完成注意力的所有数学运算,而无需将整个注意力分数矩阵 S = QK^T(尺寸 [N, N])写出到 HBM。

  2. 重计算(Recomputation):反向传播时,不保留前向的注意力分数矩阵(那是显存杀手)。而是利用前向保存的 Q、K、V 和 Softmax 的统计量(row_max, row_sum),在反向时重新计算一次 QK^T 和 Softmax,得出梯度所需的值。这样,显存中始终不需要完整存储 S

📊 显存节省:标准注意力需要存储 SP = softmax(S),两者均为 [batch, heads, N, N] 形状,FP16 下占 2 * batch * heads * N * N * 2 字节。FlashAttention 将这些完全消除,只保留 Q、K、V 和少量标量,显存从 O(N²) 降为 O(N)。对于长序列(如 N=4096),节省可达几十倍。

✅ 因此,FlashAttention 通过巧妙的 IO 优化和重计算,在保证精确性的同时,让长序列训练的显存需求回归线性。


📏 FlashAttention 对于不同序列长度的显存节省效果如何?

💡 序列长度越长,节省越惊人:当 N=512 时节省约 5-10 倍注意力内存;N=4096 时节省可达 40 倍;N=32K 时甚至百倍以上,直接决定了长上下文训练是否可能。

📈 具体分析(假设 batch=1, heads=32, FP16):

  • 标准注意力:存储注意力分数矩阵 S [N, N] 和 softmax P [N, N],显存 ≈ 2 * heads * N * N * 2 字节。

  • FlashAttention:无需存储这两个矩阵,额外显存仅保留 Q、K、V 及少量统计量,与 N 成线性关系,可忽略不计(相比其他激活)。

  • 比例:节省的注意力矩阵部分几乎被完全剔除。N=2048,标准注意力分数矩阵约 32 * 2048² * 2 ≈ 0.5 GB(单个矩阵),两个矩阵共 1 GB,而 FlashAttention 几乎为 0。整体激活显存中,注意力贡献从主导变为可忽略。

📊 实际效果(参考):

查看内嵌表格

可见,对于 32K 序列,标准注意力仅分数矩阵就需要 128 GB,完全不可行。FlashAttention 使之成为可能。

✅ 所以,FlashAttention 的显存节省对短序列温和,对长序列却是革命性的,是实现 128K 训练的核心技术。


🚀 FlashAttention-2/3 在显存优化上又有哪些进步?

💡 FlashAttention-2 和 3 在保持 O(N) 显存的基础上,主要通过改进并行策略和减少非矩阵乘法运算,进一步降低 前向/反向的激活内存需求,并提高了计算效率,但显存复杂度的实质突破仍在 FlashAttention-1。

🔍 各版本进步:

  • FlashAttention-2
  • 优化了线程块调度,减少 warp 间通信,提高了 SM 利用率。
  • 改进了反向传播的重计算流程,减少了需要存储的中间标量数量,进一步压缩了反向所需的工作内存(较少量,但仍线性)。
  • 总的显存占用比 v1 略低(约 10-20%),但对于极长序列,效果有限,主要优势在于速度。

  • FlashAttention-3

  • 针对 Hopper 架构(H100)设计,利用 TMA(Tensor Memory Accelerator)异步数据拷贝,更高效地管理 SRAM 和 HBM 之间的数据流。
  • 进一步融合了某些激活函数和 dropout,减少临时缓冲区。
  • 显存上,它更加精细地管理重计算所需的统计量,使其与 batch 和 head 维度更好地扩展。但对用户而言,显存节省相比 v2 没有数量级突破,更多是速度(提速 1.5-2x)和计算效率。

📌 总结:显存优化在 FlashAttention-1 已完成从 O(N²) 到 O(N) 的跨越,后续版本是工程精益化。选择时优先用最新版,因其速度和显存都有温和改善。


💾 激活卸载(Activation Offloading)如何工作?代价是什么?

💡 激活卸载在前向传播后将部分激活张量从 GPU 显存拷贝到 CPU 内存,反向传播时再异步取回,以 PCIe 带宽和延迟为代价,换取 GPU 显存的扩大。

⚙️ 工作流程:

  1. 前向传播:正常计算每一层的激活,在产生后,将选定的激活张量(如 attention 输出、FFN 中间结果)从 GPU 复制到 CPU 固定内存(pinned memory)。

  2. 显存释放:GPU 上的该激活立即被释放,显存被腾出。

  3. 反向传播需要时:提前从 CPU 异步预取需要的激活张量回 GPU。通过 overlap,将传输隐藏在计算之后。

  4. 选择性卸载:通常只卸载部分层(如每隔几层)或较大张量,以避免传输量过大。

⚖️ 代价:

  • 速度下降:PCIe 带宽(如 PCIe 4.0 x16 理论 ~32 GB/s)远低于 GPU HBM 带宽(>1 TB/s)。如果计算时间不能完全隐藏传输时间,GPU 会停顿等待数据,导致训练变慢。实测可能导致 20%-50% 的吞吐下降,取决于卸载量和模型。

  • CPU 内存压力:需要大量 CPU 内存(数十 GB 至数百 GB)来存放激活。尤其大 batch 长序列时,CPU 内存也可能不足。

  • 实现复杂度:需要框架支持异步拷贝和良好调度。

🤝 与梯度检查点结合:激活卸载可与梯度检查点互补。梯度检查点减少激活总量,卸载将剩余的激活转移到 CPU,两者结合能在极端条件下进一步扩展序列长度或 batch。

✅ 因此,激活卸载是显存极限扩展的最后手段之一,以明显的速度代价换取突破显存墙的可能。


🔬 梯度检查点和 FlashAttention 同时使用,能省多少显存?以 GPT-3 为例说明。

💡 两者协同可节省大量显存:梯度检查点消除绝大多数中间激活,FlashAttention 避免生成 O(n²) 的注意力矩阵,总激活显存从原来的数十 GB 降至数 GB。以 GPT-3 (175B) 为例,在序列长度 2048、batch=1 时,激活显存从 ~60 GB 降至 ~2-3 GB。

📐 GPT-3 具体估算:

  • GPT-3 结构:96 层,hidden=12288,96 个注意力头,头维度 128。

  • 标准训练(无优化)时,每层激活主要包括:

  • Q, K, V:3 × [batch, seq, hidden] → 约 3×2048×12288×2/1024^3 ≈ 0.14 GB/层。
  • 注意力分数矩阵 S [batch, heads, seq, seq]:96×2048×2048×2/1024^3 ≈ 0.75 GB/层。
  • softmax 后的 P:同样 0.75 GB/层。
  • FFN 中间激活:约 0.3 GB/层。 每层总计约 2 GB,96 层需 约 192 GB 激活(这还没算梯度等,已远超单卡显存)。实际上标准实现还会更高。

  • 仅用梯度检查点:每层只保留输入 hidden states(~0.14 GB),其余丢弃。反向重计算每层时,仍需临时生成注意力矩阵。激活显存降至约 13.5 GB (96 × 0.14 GB),但重计算时注意力矩阵依然会出现峰值(临时 O(n²) 矩阵),可能导致显存尖峰。

  • 仅用 FlashAttention:不存储注意力分数矩阵,激活显存减少 2×0.75 GB/层。每层激活剩下约 0.5 GB,96 层约 48 GB。仍然很高。

  • 梯度检查点 + FlashAttention:梯度检查点将每层保留的激活减少为输入 hidden states(0.14 GB),而 FlashAttention 确保重计算时不再产生 O(n²) 的注意力矩阵,仅需 Q,K,V 和输出,无大矩阵驻留。因此每层激活峰值仅约 0.14 GB(输入)+ 重计算时的 Q,K,V(~0.14 GB,临时),总显存占用约 0.28 GB/层,96 层约 27 GB,但得益于重计算段内释放,实际平均驻留激活仅需存储每段输入,峰值大约为几个层的激活。若将每个 Transformer Block 作为一个检查点段,驻留激活仅为该 Block 输入(~0.14 GB)加上当前 Block 重计算时的临时激活(~0.14 GB QKV 等),全局激活峰值可控制在 2-3 GB 左右。

✅ 因此,两者联手使得 GPT-3 级别的训练成为可能,从 >100 GB 降至几 GB,消解了注意力矩阵的二次显存噩梦,是长序列训练的标配。


🧪 在训练超长序列时,激活值显存成为瓶颈,有哪些专用优化方案?

💡 超长序列 (≥32K) 训练时,激活显存通常成为主要瓶颈,常用以下专用方案组合应对:

查看内嵌表格

🔧 典型组合:对于 128K 序列训练,常用 FlashAttention + 梯度检查点 + 序列并行 + CPU Offload。序列并行将 LayerNorm 及残差连接的序列维度切分,结合张量并行,每卡激活进一步减半。Ring Attention 则利用多卡分摊整个长序列的注意力存储。

✅ 没有单一银弹,必须根据序列长度、硬件和模型架构选择恰当的组合。


🔗 序列并行与激活值重计算的组合效果如何?

💡 序列并行 (SP) 切分序列维度,激活值重计算 (如梯度检查点) 减少时间维度的激活驻留,两者组合可正交地降低单卡激活显存,效果近似相乘:序列长度降 1/N,激活驻留降 1/K,单卡激活可降至原来的 1/(N*K) 量级。

🔍 工作机理:

  • 序列并行将激活张量(如 LayerNorm 输出、Dropout 输入)在 seq 维度上切成 N 份(N=TP 度),每张卡只存储 seq/N 长度的激活,激活内存直接除以 N。

  • 梯度检查点则让每卡在反向时重计算部分激活,从而在时间上减少了同时驻留的激活层数,将激活内存降至约 O(√L) 或更少(L 为层数)。

  • 两者作用在不同维度(序列长度、模型深度),因此可以叠加,且互不干扰。最终单卡激活峰值可降至原来的 1/(N * C),其中 C 是检查点节省系数。

📊 实测效果:在训练 65B 模型、序列长度 4096、TP=8 且开启序列并行和逐 Block 梯度检查点时,单卡激活峰值比无优化时减少超过 10 倍。这允许使用更大 batch 或更长序列。

✅ 因此,序列并行与重计算是超长序列训练的黄金搭档,分别从宽度和深度压缩激活。


🧮 为什么注意力矩阵的显存是 O(L²),而其他激活值是 O(L)?

💡 注意力机制需要计算并存储形状为 [batch, heads, L, L] 的注意力分数矩阵和 softmax 结果,其元素数与序列长度 L 的平方成正比;而其他激活(如 Q、K、V、FFN 输出)的形状是 [batch, L, hidden],仅随 L 线性增长。

📐 数学表达:

  • 注意力分数矩阵 S = Q × K^T,Q 形状 [B, H, L, d],K 形状 [B, H, L, d],相乘得到 [B, H, L, L]。存储该矩阵需要 B * H * L * L 个元素,所以是 O(L²)。

  • softmax 后的 P 同样 O(L²)。

  • 其他激活:例如 Q、K、V 本身、FFN 的中间输出 [B, L, 4D]、LayerNorm 前后的 hidden states [B, L, D],它们都是沿序列维度为 L,另一维度固定,因此存储量 O(L)。

🔬 为什么注意力是二次? 因为每个 token 要与所有 token 计算相关性,所以产生一个 L×L 的矩阵。这是 Transformer 的核心特性,也是长文本的瓶颈。FlashAttention 通过不显式存储这个矩阵解决了显存瓶颈,但仍需进行 O(L²) 次计算。


🚫 有没有方法可以完全避免保存注意力矩阵?

💡 是的,FlashAttention 及其变体通过分块算法和在线 softmax 完全避免了在 HBM 中保存整个注意力矩阵,仅需在片上 SRAM 中进行小块计算,从而将显存复杂度降为 O(L)。

🔬 具体方法:

  1. FlashAttention (v1/v2/v3):将 Q、K、V 分块,逐块加载到 SRAM,采用 numerically stable 的 online softmax 算法增量计算输出 O,并保留归一化统计量(m, l)。前向时无需写出完整的 S 和 P;反向时基于保存的统计量和 Q、K、V 重新计算局部 S 和 P,同样不产生完整矩阵。这是目前工业界标准方案,完全避免 O(L²) 显存。

  2. Sparse Attention / 近似注意力:如 Reformer (LSH attention)、Linformer、BigBird、Longformer 等,通过稀疏连接或低秩近似将注意力矩阵从稠密变为稀疏或低秩,计算和存储复杂度降至 O(L log L) 或 O(L)。不过会牺牲部分模型质量。

  3. Kernel-based 线性注意力:将注意力分解为核函数,可将复杂度降至 O(L),例如 Performer、Linear Transformers,但实际效果往往不如 FlashAttention + 原版注意力。

✅ 目前,FlashAttention 是实现“完全避免保存注意力矩阵”且无损精度的最佳方案,已成为 PyTorch 等框架的默认注意力后端。


🎯 什么是“选择性梯度检查点”?比全量检查点有什么优势?

💡 选择性梯度检查点是指仅对部分操作或层设置检查点,而不是对每个 Block 都重计算。它允许在显存节省和计算开销之间更精细地权衡,避免不必要的重计算。

🔍 实现方式:

  • 在全量检查点中,每个 Transformer Block 都被标记为检查点段,前向时丢弃块内所有中间激活,反向时重算整个 Block。

  • 选择性检查点则根据算子或层的“激活大小 vs 计算成本”来决定是否保留激活。例如:

  • 保留计算代价高但激活小的层(如某些 FFN),不丢弃其激活。
  • 丢弃激活占用极大但计算较便宜的层(如注意力中的 softmax、dropout)。

  • PyTorch 提供 torch.utils.checkpoint.checkpoint 的区域控制,可通过 use_reentrant=False 实现细粒度;或使用 torch.utils.checkpoint.create_selective_checkpoint

📊 优势:

  • 减少额外计算:全量检查点增加 20-30% 计算,选择性可能只增加 10-15%。

  • 显存节省接近全量:因为注意力矩阵等大激活仍被丢弃,而小激活保留,节省效果略低但可接受。

  • 更灵活的显存-速度平衡:可根据硬件特点(如计算富余、显存紧张)调整。

✅ 因此,选择性梯度检查点适合在显存压力不大、或某些层重计算代价过高的场景,避免“一刀切”。


🧠 在推理阶段,激活值显存还需要特别优化吗?

💡 推理时没有反向传播,因此不存在训练中的梯度、优化器和中间激活保留问题,激活显存压力大幅减轻。但长序列推理时,KV Cache 取代激活成为主要显存瓶颈,而非传统中间激活。所以推理阶段通常不专门优化激活,而是优化 KV Cache。

🔍 区别:

  • 训练激活:前向产生的所有中间结果(QKV、注意力矩阵、FFN 输出等)都需为反向保留,占用巨大。

  • 推理前向:仅需产生输出,中间结果用后即扔,不需要长期保存。因此激活显存很小,仅需当前层的少量临时缓冲区。

  • 推理新瓶颈:KV Cache 存储历史 Key、Value,随序列长度线性增长,且常驻显存,必须优化。

📌 推理阶段的显存优化聚焦于:

  • 使用 GQA/MQA 减少 KV 头数。

  • KV Cache 量化(INT8/FP8)。

  • PagedAttention 管理内存分配。

  • 共享前缀等。

✅ 因此,推理时几乎不需要关注激活显存,重点在 KV Cache。


⏫ 训练时显存中的“二次峰值”和激活值有什么关系?

💡 “二次峰值”通常指训练过程中显存出现的两个尖峰:第一次在正向传播末尾(全激活保留),第二次在反向传播开始后(重计算或生成梯度前)。激活值的存储与重计算策略直接塑造了这个双峰形态。

🔎 详细解释:

  1. 首次峰值:前向传播完毕时 所有层的前向激活都保留在显存中(若不使用梯度检查点),此时激活显存达到最大。这就是第一个峰值,完全由激活值贡献。

  2. 二次峰值:反向传播初期 当开始反向传播,一方面需要释放部分前向激活(或利用它们),另一方面要为梯度计算分配缓冲区,同时若启用梯度检查点,会进行重计算,产生新的临时激活。这些临时激活叠加尚未释放的激活检查点以及部分梯度,可能形成第二个显存峰值。它通常高于前向峰值,因为额外增加了梯度、临时重计算激活等。

🔬 与激活值的关系:

  • 如果使用梯度检查点,前向峰值会大大降低(只保留检查点输入),反向重计算时临时激活又会产生一个峰值(重计算当前块的激活)。这个反向重计算峰值就是“二次峰值”的主要来源,其高度等于重计算所需的局部激活大小(通常低于无检查点的前向峰值,但可能高于检查点缩减后的前向占用)。

  • 若不开检查点,前向峰值包含全部激活,反向时激活逐步释放,通常显存逐渐下降,可能不会出现明显的二次峰值,但峰值绝对值很高。

📊 影响:优化显存峰值需同时关注前向和反向,确保二次峰值不超出容量。FlashAttention 可减少重计算时的临时注意力矩阵,从而降低二次峰值。

✅ 因此,训练显存波动是激活、梯度、优化器状态交织的结果,二次峰值与激活值重计算密切相关。