显存瓶颈:中间激活才是真正的显存大户
训练一个 24 层 Transformer 时,你可能会发现即使把 batch size 设为 1,显存依然不够用。模型参数本身只占一小部分,真正的显存消耗来自前向传播时每一层输出的激活值(activation)。假设输入是 [batch=16, seq_len=1024, hidden=1024],单个激活张量在 float32 下占 64 MB,24 层就是 1.5 GB,这还只是单次前向传播。反向传播时为了计算梯度,需要用到每一层的激活值,因此这些中间结果必须全部保留,显存占用随层数线性增长。
当模型加深或序列变长时,这个线性增长很快就会撞上 GPU 显存上限。常见的应对办法是减小 batch size,但 batch 太小会导致梯度估计噪声大、训练不稳定,收敛变慢。换更小的模型则直接牺牲表达能力。购买更大显存的 GPU 成本高昂,而且对于动辄数十亿参数的大模型,即使 80 GB 的 H100 也未必够用。
激活检查点(Activation Checkpointing,也叫梯度检查点 Gradient Checkpointing)提供了一条不同的路径:不保存所有中间激活,只保存少数关键位置的检查点,反向传播需要时再重新计算被丢弃的部分。这是一种以时间换空间的策略,用额外的计算量换取显存的大幅下降。
核心机制:丢弃激活,反向传播时重算
激活检查点的基本思想可以用图书馆借书来类比。传统训练方式把前向传播所有层的激活值都保存在显存里,就像把需要的所有书都堆在桌子上,取用方便但空间有限。激活检查点则只保留少数关键参考书(检查点),需要其他书时再回书架取(重新计算)。代价是取书多花时间,但桌子能放下更多资料。
具体来说,假设一个 N 层的网络,传统方式需要保存每一层的激活值,显存占用为 O(N)。激活检查点策略每隔一定层数设置一个检查点,只保存检查点处的激活值,丢弃其他中间激活。反向传播时,如果需要某个被丢弃的激活值,就从最近的检查点开始重新执行前向计算来恢复它。
陈天奇等人在 2016 年的论文《Training Deep Nets with Sublinear Memory Cost》中提出了这一系统化方法。论文指出,通过合理设置检查点间隔,可以将显存占用从 O(N) 降到 O(√N),而额外计算成本只是每个 mini-batch 多一次前向传播。论文还展示了极端情况下显存可降到 O(log N),但额外计算成本会增加到 O(N log N)。
在 PyTorch 中,torch.utils.checkpoint.checkpoint 是标准的实现接口。它接受一个函数(通常是模型的一个模块)和输入,在前向传播时执行该函数但不保存中间激活,只保存输入作为检查点。反向传播时,从保存的输入开始重新执行前向计算,再计算梯度。关键参数 use_reentrant 在 PyTorch 2.5 之后必须显式传递,推荐设为 False,因为它支持关键字参数和 torch.autograd.grad(),性能也更好。
贯穿场景:长序列 Transformer 训练
为了具体说明,我们假设一个实际场景:在一张 24 GB 显存的 GPU(如 RTX 3090)上训练一个 BERT-large 规模的模型,序列长度为 1024,batch size 为 16。这个配置在传统方式下显存会溢出,因为激活值占用超过 10 GB,加上参数、优化器状态和梯度,总显存需求远超 24 GB。
使用激活检查点后,我们可以在模型的每个 Transformer 块上应用 checkpoint。前向传播时,每个块只保存输入张量,丢弃块内部的中间激活。反向传播时,每个块需要梯度时重新计算自己的前向过程。这样显存占用大幅下降,但训练时间会增加,因为每个块在反向传播时都多了一次前向计算。
下面是一个简化的流程图,展示了激活检查点在单个 Transformer 块上的数据流:
flowchart TD
A[输入 x] --> B[保存 x 作为检查点]
B --> C[执行块前向计算]
C --> D[丢弃中间激活]
D --> E[输出 y 继续前向]
E --> F[反向传播开始]
F --> G{需要该块梯度?}
G -- 是 --> H[从检查点 x 重新前向]
H --> I[使用重算的激活计算梯度]
G -- 否 --> J[跳过重算]
I --> K[更新参数]
这个流程的关键在于:检查点 x 必须保留,它是重计算的起点;中间激活被丢弃,显存因此减少;反向传播时重算前向,计算量增加。
显存节省与时间开销的量化分析
激活检查点的收益和代价可以通过复杂度分析来理解。设模型有 N 层,每层激活大小为 |a|。传统方式显存占用为 O(N × |a|),计算时间为 O(N)(前向)+ O(N)(反向)。激活检查点如果每 √N 层设一个检查点,显存占用降为 O(√N × |a|),计算时间变为 O(N)(前向)+ O(N × √N)(反向,因为每段需要重算),总时间约为 O(N × √N)。
实际训练中,显存节省通常在 50%~70%,时间增加约 20%~30%。OneFlow 文档中给出了一个 BERT 模型在 RTX 3090 上的对比实验:不开启激活检查点时平均显存占用 9141 MB,训练完成用时 25 分 16 秒;开启后显存降到 5978 MB,用时增加到 33 分 36 秒。显存下降了约 35%,时间增加了约 33%。陈天奇论文中的实验显示,1000 层残差网络在 ImageNet 上显存从 48 GB 降到 7 GB,额外运行时间约 30%。
这些数字表明,激活检查点能显著降低显存峰值,但代价是训练时间增加。对于显存受限的场景,这是值得的,因为否则根本无法训练。
与梯度累积、混合精度训练的组合
在实际训练大模型时,激活检查点通常不是单独使用,而是与梯度累积和混合精度训练配合。梯度累积解决的是 batch size 太小导致梯度估计不稳定的问题:通过多个小 batch 的梯度累加来模拟大 batch 的效果。激活检查点则进一步降低单个 batch 的显存占用,使得在有限显存下可以使用更大的 batch size 进行梯度累积,提高训练稳定性。
混合精度训练(如 FP16 或 BF16)通过降低激活值和梯度的精度来减少显存占用。激活检查点与混合精度可以叠加:检查点保存的激活值可以以低精度存储,进一步节省显存。但需要注意,重计算时如果使用低精度,可能会引入数值误差,影响梯度精度。因此,在关键检查点位置可能需要保存更高精度的激活值。
在 Hugging Face Transformers 中,可以通过 TrainingArguments 的 gradient_checkpointing=True 和 fp16=True 同时开启,配合 gradient_accumulation_steps 设置梯度累积步数。这种组合是当前大模型训练的标准配置。
实现细节与注意事项
在 PyTorch 中使用激活检查点,需要注意几个关键点。首先,torch.utils.checkpoint.checkpoint 的第一个参数是一个可调用对象,通常用 create_custom_forward 包装模块,以便传入多个输入。其次,use_reentrant=False 必须显式传递,否则在 PyTorch 2.5+ 会报错。
对于 nn.Sequential 模型,可以使用 checkpoint_sequential 简化操作,它自动将模型分成多个段,每段作为一个检查点。
PyTorch 还提供了选择性检查点(Selective Checkpointing),通过 create_selective_checkpoint_contexts 和策略函数,可以精细控制哪些操作保存结果、哪些重新计算。例如,矩阵乘法(torch.ops.aten.mm)通常保存结果,而其他操作可以重算。这种细粒度控制能进一步优化显存和时间的平衡。
另一个重要注意事项是随机数生成器状态。如果模型中有 dropout 等随机操作,重计算时需要使用相同的随机数序列,否则结果不一致。preserve_rng_state=True(默认)会保存和恢复 RNG 状态,确保重计算与原始前向一致。
何时失效:适用边界与失败模式
激活检查点并非万能,它有一些明确的适用边界和失败模式。首先,重计算增加了计算量,如果模型本身计算量已经很大,时间开销可能不可接受。对于计算密集型操作(如大矩阵乘法),重计算的成本较高;而对于内存密集型操作(如激活函数),重计算成本较低。
其次,激活检查点主要针对中间激活的显存占用。如果显存瓶颈来自参数、优化器状态或梯度本身,激活检查点的效果有限。例如,使用 Adam 优化器时,优化器状态通常是参数大小的 2~3 倍,这部分显存无法通过激活检查点减少。
在长序列训练中,激活检查点尤其有效,因为序列长度直接影响激活大小。但序列过长时,重计算的时间开销也会显著增加,因为每次重算都要遍历长序列。此外,如果模型层数较少(如少于 15 层),激活检查点的收益不明显,因为中间激活总量不大,而重计算开销相对较高。
另一个失败模式是数值精度问题。在混合精度训练中,如果检查点保存的激活值精度过低,重计算出的梯度可能不准确,导致训练不稳定。因此,在关键位置可能需要保存 FP32 激活,这会减少显存节省。
最后,激活检查点与某些并行策略(如张量并行、流水线并行)的交互需要谨慎。在流水线并行中,每个设备只负责部分层,激活检查点可以进一步降低每台设备的显存,但会增加通信和重计算的协调复杂度。
决策参考:何时使用激活检查点
下表总结了激活检查点与几种替代方案的对比,帮助你在实际训练中做出选择。
| 方案 | 显存占用 | 训练时间 | 实现复杂度 | 适用场景 |
|---|---|---|---|---|
| 传统训练 | 高(O(N)) | 低 | 低 | 显存充足时 |
| 减小 batch size | 中 | 中(收敛慢) | 低 | 显存不足但可接受小 batch |
| 激活检查点 | 中低(O(√N)) | 中(+20%~30%) | 中 | 显存不足且需要较大 batch |
| 混合精度训练 | 低(约减半) | 低(可能加速) | 中 | 显存不足且硬件支持 |
| 激活检查点 + 混合精度 | 最低 | 中高 | 中高 | 显存极度受限时 |
从表中可以看出,激活检查点适合在显存不足且无法通过减小 batch size 解决问题的场景。如果显存只是略超,可以先尝试混合精度训练;如果显存严重不足,则激活检查点与混合精度结合是更优选择。
可观测性与调试
在生产环境中,使用激活检查点时需要监控几个关键指标。显存占用是最直接的收益指标,可以通过 nvidia-smi 或 PyTorch 的 torch.cuda.max_memory_allocated() 监控。训练时间(吞吐量)是主要代价,需要记录每个 step 的耗时。
如果训练速度明显低于预期,可能原因是重计算过于频繁或检查点间隔过小。可以通过调整检查点间隔(即每个 checkpoint 包含的层数)来平衡显存和时间。间隔越小,显存节省越多,但重计算开销越大。
另一个需要监控的是数值稳定性。如果损失曲线出现异常波动,可能是重计算导致的数值误差。可以对比开启和关闭激活检查点时的梯度范数,确保差异在可接受范围内。
仍未解决的问题
激活检查点虽然成熟,但仍有一些开放问题。如何自动选择最优的检查点位置和间隔,目前多依赖经验或启发式方法。选择性检查点的策略函数如何设计,才能在不同模型和硬件上达到最优平衡,也缺乏通用指导。此外,在分布式训练中,激活检查点与模型并行、流水线并行的联合优化,仍是一个活跃的研究方向。
对于训练超大模型(如数千亿参数),激活检查点通常与 ZeRO、张量并行等技术结合使用。在这些复杂系统中,激活检查点的重计算与通信重叠、显存碎片等问题,都需要更精细的调度策略。
总的来说,激活检查点是一个简单而有效的显存优化手段,它用计算时间换取了显存空间,使得在有限硬件上训练更大模型成为可能。理解其原理、开销和边界,能帮助你在实际训练中做出更合理的决策。