模型层 开放阅读

Activation Checkpointing

Gradient Checkpointing

概念 ID
gradient-checkpointing
更新时间
2026-05-29
来源数量
待补

Activation Checkpointing

1. 引言:大模型时代的显存困境

当模型参数量迈入千亿乃至万亿级别,单张加速器的显存容量已不再是宽裕的资源,而成为制约训练可行性的最直接瓶颈。大语言模型、多模态基础模型以及大规模推荐系统的训练,不仅需要存储海量参数和优化器状态,还必须在前向传播过程中保留中间激活张量,以供反向传播计算梯度。以典型的 Transformer 结构为例,假设模型层数为 L,隐藏维度为 d,序列长度为 s,批次大小为 b,则单个 Transformer 层产生的激活张量总字节数约为 b × s × d × 常数。对于 L=96d=12288s=4096 的千亿参数模型,仅激活部分即可轻易占据数百 GB 显存,而目前最先进的单 GPU 高带宽存储器 (HBM) 容量仅约 80 GB 至 192 GB 不等。即使引入张量并行、流水线并行等模型并行手段将参数切分至多个加速器,单卡的激活显存压力依然巨大,且在增大批次大小或处理更长序列时呈线性增长。

显存与算力增速的剪刀差进一步加剧了这一矛盾。过去十年间,GPU 的浮点运算能力提升了数十倍,而 HBM 容量和带宽的增速则远低于此。产业界长期面临一个现实:算力相对“充裕”,显存却无比“稀缺”。在此背景下,任何一种能够在不显著增加总计算量的前提下压缩激活内存占用的技术,都将直接转化为模型扩展和工程迭代的自由度。Activation Checkpointing(激活检查点,亦常称作梯度检查点)正是这样一种以计算换内存的核心技术。它通过在前向传播中策略性地丢弃部分中间激活,并在反向传播时选择性重新计算,将激活所需的峰值显存降低至亚线性甚至对数级别,代价仅为引入一定比例的额外前向计算。自 1990 年代自动微分领域的理论奠基以来,经深度学习框架的工程化改造,这一技术已成为当前千亿/万亿参数大模型训练的标配组件,深刻重塑了分布式训练系统的内存管理体系。

在展开论述之前,有必要厘清一个常见误解:Activation Checkpointing 并非简单地“少存一些中间结果”。它立足于一个被严格证明的图灵奖级理论——在计算图中,通过精心选择保留哪些中间节点的值,可获得内存占用的最优下界以及对应的重计算复杂度上界。它背后涉及组合最优、层次化分治和现代编译优化等多领域的交叉,也是自动微分、分布式系统和机器学习系统工程三者汇流的典范案例。本节将以此为起点,综合剖析该技术从数学原理到工业落地的完整图景。

2. Activation Checkpointing 的核心思想与直观解释

Activation Checkpointing 的根本思想可以归纳为一句话:不保留所有中间结果,只保留若干“检查点”,需要时从最近检查点快速重演计算,临时复现被丢弃的激活。

在深度学习的一次标准训练迭代中,前向传播逐层计算输出,同时将每一层的输入、输出以及某些中间缓存的张量保存在显存中,形成计算图的反向依赖。反向传播时,根据链式法则,每一层参数的梯度都依赖于该层的输入激活与上游传递而来的输出梯度。因此,若没有特殊的内存管理策略,L 层模型就必须同时在显存中驻留 L 组激活张量,导致激活内存占用量与网络深度呈严格线性关系。

Activation Checkpointing 挑战了这一朴素逻辑:它允许开发者在计算图的特定位置设置“检查点”。前向传播经过检查点时,正常保存该节点的输出张量;在非检查点层,激活张量在完成前向传递后即被释放。反向传播进行到某个区段时,由于缺少直接可用的中间激活,梯度计算模块会自动识别当前区段的起始检查点,并以该检查点处保存的输出为输入,重新执行该段内若干层的前向计算,从而按需“复现”出该段内每一层所需的中间激活。梯度计算完成之后,这些临时激活随即被回收,整个过程中显存中同时驻留的激活总量由“全部层”变为“若干检查点加当前区段内的激活”。

直观上,这好比一个登山者在长途跋涉中选择性搭建营地。传统方法要求登山者在每一步都留下帐篷(激活),导致装备物资(显存)快速耗尽;而 Activation Checkpointing 则只在高海拔关键处(检查点)安置大型营地,其余路段轻装前进。需要返回某段获取信息时,再从最近营地快速沿原路复攀一次,事毕即撤收。显然,营地越密集则额外攀登越少,但需携带的保障物资越多;营地越稀疏则物资需求越低,但重攀代价越大。这一权衡精确地刻画了激活检查点技术的参数配置空间。若将“营地”数目记为 c,则内存大致降低至 O(L/c),而额外计算开销约为 O(c) 的前向传播,由此可导出经典的平方根均衡策略,即在给定内存预算下最小化额外计算量。

3. 理论渊源:从自动微分的检查点存储到深度学习

Activation Checkpointing 并非深度学习时代的原创发明,其理论基础可追溯至计算数学领域对自动微分(Automatic Differentiation, AD)内存复杂度的长期研究。在反向模式自动微分(即反向传播的数学本质)中,为了计算标量函数对众多输入变量的梯度,系统必须保存一条记录所有中间变量的“磁带”,该磁带的长度通常与原始计算的操作步数成正比。对于迭代求解、时间序列仿真或深度链式神经网络等包含长序列计算的问题,磁带的内存占用迅速膨胀,成为内存瓶颈。

1992 年,Andreas Griewank 在其开创性论文《Achieving Logarithmic Growth of Temporal and Spatial Complexity in Reverse Automatic Differentiation》中,提出了一种革命性的检查点调度策略:利用“分治递归”的思想,将长度为 n 的计算过程划分为不均匀的段落,并在各阶段末端设置检查点,使得反向传播过程的总重计算次数仅为原始前向计算的数倍,而内存占用仅需 O(\log n) 个检查点。这一成果证明:反向模式自动微分的内存复杂度可以从线性降至对数级别,而时间开销仅增加一个常数因子。该策略被后续文献称为“ binomial checkpointing”或“ Griewank checkpointing”,奠定了检查点技术的数学最优性基础。

在 Griewank 之前,Volin 与 Ostrovskii(1985)以及 Morgenstern(1979)等已研究了类似“蛇形”或“分段”重计算方法,但未达到对数内存的理论保证。Griewank 将检查点问题形式化为有向无环图上的峰值内存最小化问题,并给出了最优检查点放置的动态规划算法。随着时间的推移,这一理论被吸收进自动微分工具(如 ADIFOR、TAPENADE)中,用于处理大规模数值模拟的梯度计算。进入深度学习时代后,由于神经网络的高度层状结构和反向传播的特性,自动微分中的检查点存储技术被自然地迁移过来,形成了我们今天所熟知的 Activation Checkpointing。因此,该技术的本质是对计算图进行内存感知的调度,它的数学根基远比简单的“丢弃再重算”深厚得多。

4. 形式化定义与计算图模型

为建立精确的分析框架,我们需要将深度学习模型的前向计算抽象为一个有向无环图 G = (V, E),其中节点 v \in V 表示操作(算子),边表示张量数据的流动。前向传播相当于按拓扑序执行每个节点,每个节点在执行结束后会生成一个或多个输出张量。反向传播时,梯度沿逆向拓扑序流动,每个节点需要其直接前驱节点的输出激活来计算局部梯度。传统训练模式下,所有节点的输出激活均须保留在内存中,直到反向传播不再需要为止。因此,内存占用的峰值表现为所有“活跃”张量的大小之和,其下界与图的关键路径长度相关。

Activation Checkpointing 的核心在于引入一个“检查点集合” C \subseteq V,并规定:在前向执行过程中,仅当节点属于 C 时,其输出才被长期保留;其余节点执行后立即释放输出。反向传播按层逆序推进,当需要计算某个非检查点节点 v 的梯度时,必须从一个距离最近的前向检查点 c 出发,沿前向边重算 cv 之间所有被丢弃的中间节点,从而临时再现 v 所需的输入激活。这一定义自然地引出了两个关键性能指标:内存峰值 M,即任意时刻所存检查点与当前重计算区段内激活之和的最大值;重计算开销因子 R,定义为由于回退重算而引入的总前向计算量与原始前向计算量的比值减一。

设总操作步数为 n,检查点数目为 k,在不同调度策略下,MR 呈现出不同的权衡函数。均匀检查点策略中,每隔 n/k 步设置一个检查点,导致 M = O(n/k)R = O(k),在固定内存预算下,取 k \propto sqrt(n) 时达到 M = O(sqrt(n))R = O(sqrt(n)) 的平方根平衡。Griewank 的分治策略则可实现 M = O(\log n)R 仅为 O(\log n)O(1) 部分取决于具体实现,实际上经典 Griewank 算法的时间开销约为 R \approx \log_2 n 次重计算。显然,后者在大规模网络中更具优越性。这一形式化框架使我们能够将 Activation Checkpointing 视为一个资源受限的图调度问题,进而借助动态规划或机器学习方法自动发现最优检查点放置。

5. 基本调度策略:均匀检查点与平方根平衡

最简单的检查点调度策略是均匀分布检查点。假设模型由 L 个相同的顺序块组成(例如 Transformer 层),我们每隔 s 层设置一个检查点,总共保留 k = L/s 个检查点。前向过程中,所有层的输出初始均被计算,但仅检查点层的输出被驻留内存,其余层输出即时释放。反向传播开始时,从最后一个检查点开始,对每个检查点区间 [i, i+s),从该检查点的保存输出重算该区间内所有层的激活,然后执行该区间所有层的反向传播,重算的激活在该区间反向结束后释放。由此,峰值内存为单个区间的激活加上 k 个检查点的输出,即 M = O(s + k)。在检查点大小与层激活大小相当的假设下,给定总内存预算 M,我们令 s \approx Mk \approx M,则 L = s \times k,即 M \approx sqrt(L)。额外重计算次数为 k \times s = L,恰好等于原始前向计算量,因此额外开销 R = 1.0(即总前向计算变为原来的 2 倍)。这就是著名的平方根平衡:内存从 O(L) 降至 O(sqrt(L)),而计算时间仅翻倍。

均匀检查点的优点在于实现简单,且与深度学习框架的分层模型天然契合。PyTorch 的 torch.utils.checkpoint 在早期版本中便采用了固定间隔的检查点策略(用户可自定义分段)。对于层数不是极端庞大的网络(如数十到上百层),均匀策略足以在内存和计算开销之间取得实用平衡。然而,当网络进一步加深或单层本身包含复杂子结构时,平方根的缩减可能不再足够,因为 O(sqrt(L)) 对于数万步的序列可能仍然过于庞大。这促使人们寻求更优的调度策略,即 Griewank 所提出的对数内存方案,它允许内存占用随 L 的增加仅呈对数增长。

6. 递归分治与对数级显存:Griewank 最优调度

Griewank(1992)提出的递归分治检查点策略本质上是一种“二分”或“两进”调度,但常被称为“binomial checkpointing”。假设计算图是一个长为 n 的序列,目标是在内存中最多同时保留 d 个检查点的条件下完成反向传播,并且使重计算步骤总数尽量小。Griewank 证明,通过如下递归方式部署检查点,可实现内存 d = \lceil \log_2 n \rceil,且总重计算步数约为 n \log_2 n / 2(在原始论文中略有不同,但本质为对数级别)。

其思想可简述为:将序列划分为不等长两部分,前一部分递归处理,后一部分在反向时需要前一部分的最终状态作为检查点。具体执行时,前向过程并不一次性执行完整个序列,而是配合反向过程交错进行。经典 Griewank 调度在实际深度学习框架中实现较为复杂,因为它需要反复在不同区间内切换前向与反向的执行粒度,难以直接映射到现有的层式执行循环上。后来的研究提出了一些近似等效的实现,例如“动态检查点间隔调整”或“树形依赖记录”。

尽管完全的对数策略在工业界尚未成为主流,但其理论价值极其深远:它证明了 Activation Checkpointing 不存在根本性的扩展瓶颈。即使模型层数增长一千倍,内存需求仅增加数十倍(对数增长),且计算成本增长非常温和(对于最深度的网络,额外计算因子通常低于 2 或 3)。这为大模型开发者提供了信心——显存不足绝不是放弃加深网络的理由。近年来,一些高端训练框架(如 JAX 的 jax.checkpoint 结合 XLA 编译优化)已尝试在计算图级别实现近乎对数的检查点调度,利用编译器整体分析重计算区间,进一步逼近理论最优。

7. 深度学习框架中的实现机制

现代深度学习框架将 Activation Checkpointing 实施为一种自动微分上下文与自定义前向/反向钩子的组合。以 PyTorch 为例,torch.utils.checkpoint.checkpoint 函数接受一个任意模块的 forward 函数及输入,内部不做常规的中间张量保存,而是定义了一个自定义的 Function,其 forward 方法仅保留输入和输出(作为检查点),中间激活全部丢弃;backward 方法被调用时,使用保存的输入重新执行该模块的 forward,在临时开启梯度计算的条件下重新生成中间激活,再依次计算梯度。这个过程对用户透明,只需将希望进行检查点的网络片段包裹即可。

TensorFlow 的 tf.recompute_gradtf.GradientTape 结合自定义循环提供了类似功能,但需要显式控制梯度磁带的范围。JAX 的 jax.checkpoint(又名 jax.remat)则允许对纯函数进行装饰,并在 XLA 编译器层面将重计算指令嵌入 HLO 图,从而获得更大的优化空间,如与算子融合相结合,减少重计算的内存和计算开销。JAX 的设计使得检查点策略可以通过变换函数组合轻松定制。

关键实现细节包括:如何确保在重计算阶段原始子图所依赖的输入张量仍然存活(即检查点的保持);如何处理不可重算的随机操作(如 Dropout);如何避免因重计算导致 batch normalization 统计量变化等。实践中,Dropout 层通常被要求在检查点区域内关闭或使用固定种子重放,BatchNorm 则常被排除在检查点区段外,或使用全局统计量替代。这些工程化细节使得 Activation Checkpointing 从理论走向可靠的生产环境。

8. 与混合精度训练、梯度累积的协同作用

在大模型训练中,Activation Checkpointing 很少单独使用,而是与混合精度训练(FP16/BF16)和梯度累积等技术协同,以进一步压缩内存。混合精度训练通过将激活与参数存储为半精度(或 BF16),使得激活张量字节数减少一半;梯度累积允许在极小micro-batch下前向与反向传播,累积梯度后再更新参数,从而降低激活对批次大小的线性依赖。当三者结合时,总体显存缩减呈乘积效应。

例如,在训练 GPT-3 规模的模型时,采用 BF16 混合精度、序列并行、激活检查点(每 1 或 2 层一个检查点)以及 32 路梯度累积,单卡仅需满载数百 GB 的 HBM 即可完成原本需要 TB 级显存的训练任务。值得注意的是,混合精度下 Activation Checkpointing 的重计算部分一般以与原始相同精度执行(通常为半精度),因为重计算中的数值误差可通过损失缩放与动态调整来控制,几乎不影响收敛性。但某些敏感操作(如 LayerNorm、Softmax)可能保留高精度进行重计算,这需要框架提供灵活的策略配置。此外,梯度累积带来的多步前向/反向交错执行,也为检查点调度提供了更细粒度的调度自由度:可在多个 micro-batch 之间共享检查点,或利用累积阶段的空闲内存增加检查点密度,进一步平衡计算与内存。

9. 模型并行体系下的激活检查点设计

当模型参数量超出单卡容量,模型并行(包括张量并行和流水线并行)将参数切分至多卡,此时激活检查点的设计需要考虑并行划分的边界。对于张量并行,一层的前向计算被横向切分到多个设备,激活张量同样被分片存储。若在该层设置检查点,需确保所有设备同时保存对应的分片,或者仅在某一设备保留完整拷贝(需额外通信)。更常见的方式是,在张量并行区域内部不设检查点,而在张量并行单元的入口和出口(即通信集合点)设置检查点,重计算时只需从入口的完整激活重算整个并行区域,无需跨设备联合重计算,从而避免了复杂的通信协议。

流水线并行则天然将模型按层切分为多个阶段,检查点可自然地嵌入各阶段内的微批次调度中。1F1B(one-forward-one-backward)或交错调度中,检查点的存留与放弃需与微批次的流动相匹配。先进流水线并行系统(如 Megatron-LM、DeepSpeed)提供了集成的激活检查点配置,允许用户指定哪些流水线阶段或哪些层启用检查点,自动挂载重计算逻辑并管理跨阶段的激活传输。在 3D 并行(数据、张量、流水线)的超大规模训练中,检查点成为沟通三者的“内存控制器”,使各维度并行的局部内存开销均能被独立约束和优化。

10. 性能模型与开销量化分析

为了充分理解 Activation Checkpointing 的利弊,需建立性能模型。设单层前向计算时间为 F,反向时间为 B(通常 B ≈ 2F),带宽传输开销暂忽略。若无检查点,总迭代时间 T_0 = L(F + B),峰值内存 M_0 = L \cdot A,其中 A 为单层激活大小。采用均匀检查点,间隔 s,检查点数 k = L/s,则前向额外开销为每个区间重算 s 层,共 k 个区间,故总前向时间 F_total = L \cdot F + k \cdot s \cdot F = 2LF;反向时间 B_total = L \cdot B 不变,因此总时间 T_ckpt = 2LF + LB ≈ 2.66 T_0(假设 B=2F)。内存峰值 M_ckpt = s \cdot A + k \cdot A。在内存约束 M_budget 下,令 M_ckpt = M_budget 可反解最优 sk。若 M_budget 远小于 M_0,则 k 会很大,加倍比例接近 2。实际上,由于重计算无需保存中间激活的梯度图,重计算块的峰值内存较原始前向略低,但核心量级一致。

递归分治策略下,额外计算因子近似为 \log_2 L,而内存为 O(\log L)。当 L 非常大时(如数千层),对数策略的时间优势明显。例如 L=1024,均匀策略 R=1(总时间 2x),内存 sqrt(L)≈32;对数策略 R≈10(总时间 11x),内存 ≈10。但若内存极为紧张只能容纳 10 层激活,则对数策略成为唯一选择。实践中,千亿参数模型的层数通常在 100 层左右,均匀策略 2x 时间是可接受的,因此工业界大多采用每 1~4 层一个检查点的均匀配置,保持重计算开销在 10%~30% 之间,满足效率要求。

11. 检查点的动态选择与最优布局

给定一个具体的计算图,检查点的放置可以形式化为一个最优化问题:在内存上限约束下,最小化额外计算量。对链式图,有动态规划算法可以在 O(n^2) 时间内解出最优检查点位置。对于更复杂的 DAG(如带有残差连接的网络),问题变得 NP-hard,但现实中常采用贪心或启发式方法选择瓶颈节点作为检查点:优先保存在网络中被多次使用、且重算代价高的节点(如特征融合后的张量)。这类方法可视为“选择性检查点”(selective checkpointing),它与自动微分中的“ checkpoints on a tape”问题密切相关。

另一个维度是自适应调度:在训练过程中根据实时内存利用率动态插入或丢弃检查点。例如,当优化器状态更新占用较高内存时,系统可以临时将某些缓存的检查点释放,并在反向时从更早的检查点重算,从而在不违反内存限制的前提下最大化计算效率。DeepSpeed 的 ZeRO-Offload 和 ZeRO-Infinity 部分结合了这种动态性,将检查点与 CPU/NVMe 卸载协同,进一步拓宽了可能性边界。然而,过度的动态性可能破坏编译器优化(如图融合),因此多数静态图系统倾向编译期固定检查点策略。

12. 针对 Transformer 架构的定制化设计

Transformer 模型的独特结构为检查点优化提供了丰富机会。每一层由 Multi-Head Self-Attention(MHA)和 Feed-Forward Network(FFN)组成,中间有 LayerNorm 和残差相加。Attention 模块内部产生的 QKV 矩阵以及 attention score 矩阵是激活内存的巨大来源,其大小与序列长度的平方成正比。为此,很多实现选择在 Attention 计算内部设置检查点,即保存 QKV 投影后的结果,而丢弃 attention score 和 softmax 输出。反向重算时仅需从 QKV 快速还原 attention score,且避免了存储 O(s^2) 的大张量。对于 FFN,激活函数(如 GeLU)的输入通常被设为检查点,因为其重算代价极低。

自 Megatron-LM 起,这个策略被规范化:将每个 Transformer 层的前向函数通过 checkpoint 包装,对注意力块和 FFN 块均丢弃中间激活,仅保留输入隐藏状态和部分必要张量。在此基础上,针对序列长度极长的情况(如 32k 以上),FlashAttention 等 I/O 感知算法可结合检查点更精细地管理 SRAM 与 HBM 之间的数据移动,但本质是另一种层次的重计算:FlashAttention 在 kernel 内部进行分块重算,类似在线调度。未来,这些不同层级的内存优化将走向统一,编译器可将高层次检查点与底层 kernel 重计算融合调度,获得更极致的显存利用率。

13. 检查点与长序列、多模态等复杂场景

长序列处理是检查点技术的一大用武之地。当序列长度增至百万 token 时,注意力激活达到 TB 级,标准检查点即使每层只保留一个隐藏状态也依然过大。因此,需要发展更细粒度的或层次化的检查点策略:在序列维度上也进行分段检查点。例如,Blockwise Parallel Transformer 或 RingAttention 将序列切分为多个块,每个块分别计算注意力并通信,此时可在块间设置检查点,块内重算。这使得激活内存与序列长度解耦,更灵活地适应长上下文。

在多模态大模型中,不同模态的编码器、解码器以及融合模块的激活分布极不均衡,检查点策略需根据算子和张量大小进行异构配置。例如,图像编码器(ViT)可能注重 patch 间注意力的激活,而文本骨干则更关注序列长度的平方项。通过剖析(profiling)静态或动态的计算图,可自动识别高内存算子并施加选择性重计算。此外,引入 CPU 内存或 NVMe 作为检查点的备份存储,也能在不增加重计算的前提下降低 GPU 显存占用,即检查点与卸载的联合调度,这在千卡训练中已有初步验证。

14. 当前局限与工程挑战

尽管 Activation Checkpointing 已极为成功,但它仍面临多重挑战。第一,额外的计算开销始终存在,即便优化的均匀策略也使训练时间延长 30% 左右,对于需要数周乃至数月训练的超大模型,这部分开销转化为巨大的能源和硬件成本。第二,确定性重计算难题:当计算包含不可逆操作(如随机 Dropout、随机深度、某些池化)时,必须小心控制伪随机种子以保证重算结果完全一致,否则会引入梯度误差导致收敛问题;在分布式环境下,模型并行的通信次序也可能因重算而改变,带来不确定。第三,框架实现的性能瓶颈:基于 Python 的 torch.utils.checkpoint 在调用时产生了额外的主机端开销,且打断算子融合,降低 GPU 利用率;虽然通过 TorchScript 或 torch.compile 可部分缓解,但自动化的端到端优化仍不完美。第四,异构计算设备(如 TPU、IPU)的内存模型与 GPU 差异大,原有的检查点策略需重新评估与适配,缺乏统一抽象。克服这些限制将是下一代训练系统的重要方向。

15. 总结与未来展望

Activation Checkpointing 以极其优雅的“时间换空间”理念,从自动微分的理论被引入深度学习系统,已成为训练超大规模模型不可或缺的基石。它深植于计算图的内存调度理论,以 Griewank 的对数级内存结果为最优边界,在实践中通过均匀分段、递归分治以及选择性保留等策略,将激活内存需求降低数十倍,而额外计算代价可控。当前,它与混合精度、模型并行、梯度累积、FlashAttention 等技术深度融合,共同支撑起万亿参数时代的工程落地。

展望未来,检查点技术将继续向自动化、异构化和协同化发展。基于机器学习的计算图优化可能自动探索任意模型结构下的最优检查点布局;在异构内存池(HBM + DDR + SSD)场景中,检查点将被统一管理为多级缓存,形成以重计算为最内层、卸载为外层的层次化内存子系统;在分布式训练调度器中,检查点策略将作为第一类约束纳入全局优化,平衡各设备的内存与计算负载。同时,伴随着神经架构搜索和动态网络结构的兴起,检查点机制需支持可变计算图和动态控制流,这对编译器与运行时系统提出了更高要求。我们有理由相信,这个起源于 1990 年代的简洁思想,将在未来引领更深远的大规模机器学习系统革新,让人类在有限的物理显存中持续突破模型智能的边界。

source: 公开披露与公开资料整理 本页仅用于产业链学习、信息检索和研究辅助;不构成投资建议,不预测涨跌,不提供买卖、仓位或目标价建议。
完整概念页 复盘 13 节结构 公司投研页 沿产业链找到受益公司 投资课 把概念转成可跟踪模型