全参数训练出现 OOM,通常不能只归咎于模型权重太大。训练时显存里同时存在权重、梯度、优化器状态、激活值和各类临时缓冲区,其中 Adam 优化器状态和激活值往往比权重本身更占空间。比如一个 7B 参数的模型,仅权重用 FP16 保存需要约 14GB,但混合精度训练下 Adam 的 FP32 主副本、一阶矩和二阶矩还会额外占用约 84GB,如果不做任何优化,单卡 24GB 根本跑不起来。面对这种情况,有两个优化方向可以叠加使用:梯度检查点负责压缩激活值,ZeRO 负责把优化器状态、梯度甚至参数切分到多卡。它们不冲突,组合后能在基本不降低全局 batch size 的前提下完成全参数训练。

一、先看清显存消耗的完整结构
全参数训练中的显存占用可以分成静态和动态两部分。静态部分由模型权重、梯度和优化器状态组成,它们在整个训练过程中一直存在。动态部分主要是前向传播产生的激活值,以及通信库、注意力计算、梯度裁剪等过程产生的临时缓冲区。如果使用 PyTorch 默认的自动求导机制,每一层前向计算得到的中间结果都会被保留下来,用于反向传播时直接计算梯度。这种做法的计算速度很快,但激活值会随着层数、序列长度、隐藏维度和 batch size 线性增长,成为大模型训练中最容易触发 OOM 的部分。
优化器状态同样不容小觑。以常用的 Adam 为例,它在混合精度训练下通常需要为每个参数保存 FP32 的主副本、一阶矩和二阶矩,也就是 3 个 FP32 张量。对于一个 7B 参数的模型,这部分占用约为 84GB。再加上 FP16 的权重和梯度各 14GB,总显存需求会轻松超过 100GB。多卡数据并行虽然可以把部分计算分摊,但传统 DDP 每张卡上仍然保存完整的优化器状态、梯度和权重,显存压力并没有因为卡数增加而降低。
因此,解决问题的思路就很清晰:要么减少动态激活值的保存量,要么减少每张卡上静态优化器状态和梯度的重复冗余。前者对应梯度检查点,后者对应 ZeRO。两者针对不同的显存瓶颈,组合使用时收益可以叠加。
二、梯度检查点:用算力换激活显存
梯度检查点的核心思想是不保存完整的前向激活,只保存若干层输入作为检查点。反向传播需要某一层的激活时,从最近的检查点开始重新执行前向计算,得到所需的激活后再计算梯度。这样做会增加一次前向计算,但能把激活显存从与层数成正比压缩到与检查点数量成正比。对于 Transformer 结构,常见的做法是以每一层或每几层为一个检查点单元,反向时只重算该单元内部的激活。
PyTorch 提供了 torch.utils.checkpoint.checkpoint,可以直接包裹某个前向函数。下面是一个 Transformer Block 的示例,把真正的计算逻辑放到 _forward 中,再通过 checkpoint 调用:
import torch
from torch.utils.checkpoint import checkpoint
class Block(torch.nn.Module):
def __init__(self):
super().__init__()
self.norm = torch.nn.LayerNorm(768)
self.attn = torch.nn.MultiheadAttention(768, 12, batch_first=True)
self.mlp = torch.nn.Sequential(
torch.nn.Linear(768, 3072),
torch.nn.GELU(),
torch.nn.Linear(3072, 768),
)
def forward(self, x):
return checkpoint(self._forward, x, use_reentrant=False)
def _forward(self, x):
h = self.norm(x)
attn_out, _ = self.attn(h, h, h, need_weights=False)
h = h + attn_out
h = h + self.mlp(self.norm(h))
return h
这个示例中,每个 Block 的输入会被保存为一个检查点,Block 内部的 LayerNorm、多头注意力和 MLP 的中间激活都不会长期保留。反向传播到达该 Block 时,会根据输入重新跑一次 _forward,再计算各参数梯度。需要注意 use_reentrant=False 在 PyTorch 1.11 之后可用,它可以减少一部分显存碎片,但与个别自定义算子可能存在兼容性问题。如果遇到报错,可以先切换回 use_reentrant=True。
梯度检查点的代价是额外计算,通常在 Transformer 模型上会增加约 30% 到 40% 的前向计算时间。为了平衡收益,检查点粒度不需要过细。如果对每一个线性层都做检查点,重算次数会显著增加;如果以整个编码器为检查点,激活压缩效果又不明显。实践中以单个 Transformer Block 或两个 Block 为一个检查点单元,往往能获得较好的折中。
三、ZeRO:把优化器状态、梯度和参数分片
ZeRO 是 DeepSpeed 提出的显存优化方法,核心思路是在数据并行的多卡之间对训练状态进行分片,而不是让每张卡保存完整副本。根据分片对象的不同,ZeRO 分为三个阶段。ZeRO-1 只分片 Adam 优化器状态,每张卡只保存与自身参数分片对应的 FP32 主副本、一阶矩和二阶矩;ZeRO-2 在 ZeRO-1 基础上继续分片梯度;ZeRO-3 进一步分片模型权重本身。阶段越高,单卡显存越低,但通信量也会相应增加。
对于大多数全参数训练 OOM 问题,瓶颈通常不在模型权重,而在优化器状态和梯度。因此开启 ZeRO-2 往往就能把单卡静态显存从完整保存降到接近原来的 1/N,其中 N 是 GPU 数量。ZeRO-3 虽然单卡显存最低,但前向和反向过程中需要频繁通过 all-gather 收集参数分片,通信压力较大,训练速度可能明显下降。除非模型权重本身就放不下,否则不必一开始就上 ZeRO-3。
下面是一段 DeepSpeed 的 ZeRO-2 配置示例,同时启用了优化器 CPU offload,用于进一步降低显存:
{
"train_batch_size": 32,
"gradient_accumulation_steps": 4,
"fp16": {
"enabled": true
},
"zero_optimization": {
"stage": 2,
"offload_optimizer": {
"device": "cpu"
},
"allgather_partitions": true,
"allgather_bucket_size": 500000000,
"overlap_comm": true,
"reduce_scatter": true,
"reduce_bucket_size": 500000000,
"contiguous_gradients": true
}
}
ZeRO-2 在反向传播时会通过 reduce-scatter 把梯度先聚合再分片,每张卡只保留自己负责的那部分梯度,然后更新对应的优化器状态和参数。更新完成后如果需要完整参数,再通过 all-gather 收集。由于通信被拆成多个 bucket 并行执行,并且可以与计算重叠,实际增加的通信开销在高速互联环境下通常可以接受。
四、组合策略与调试顺序
梯度检查点和 ZeRO 优化的是显存的两个不同维度,前者削减激活值,后者削减优化器状态、梯度和参数。两者可以同时开启,组合后的显存峰值通常会低于只使用其中一种方案。但两者也都不是免费的:梯度检查点增加前向重算,ZeRO 增加跨卡通信。因此在调试时要避免一开始就同时开启所有选项,否则训练速度可能变得非常慢,反而难以判断瓶颈。
推荐的顺序是:先固定一个不导致 OOM 的全局 batch size,再开启梯度检查点,观察显存下降程度和训练吞吐变化。如果仍然 OOM,再逐步开启 ZeRO-1、ZeRO-2。对于 7B 到 13B 规模的模型,梯度检查点配合 ZeRO-2 通常已经能在多卡上稳定训练。ZeRO-3 更适合模型权重本身接近单卡容量上限的场景,或者需要进一步降低每卡显存来增大单卡 micro batch size 的情况。
在 Hugging Face Transformers 中,开启梯度检查点和 DeepSpeed 可以这样配置:
from transformers import AutoModelForCausalLM, TrainingArguments, Trainer
model = AutoModelForCausalLM.from_pretrained("path/to/model")
model.gradient_checkpointing_enable()
training_args = TrainingArguments(
output_dir="./checkpoints",
per_device_train_batch_size=1,
gradient_accumulation_steps=8,
fp16=True,
deepspeed="./ds_config.json",
save_steps=500,
)
trainer = Trainer(model=model, args=training_args, train_dataset=dataset)
trainer.train()
调试过程中还要留意一些容易忽略的显存消耗。比如 PyTorch 的显存分配器可能保留大量已释放但未归还给 CUDA 的缓存,长时间训练时显存碎片也会增加。可以设置 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True 来缓解碎片问题。如果使用了 FlashAttention 或 SDPA,激活值本身已经比传统注意力实现低很多,梯度检查点的收益会相对变小,但仍然有效。
此外,梯度检查点会重新执行前向计算,这意味着 Dropout、随机数生成等操作在重算时会消耗额外的随机状态。PyTorch 的 checkpoint 实现会处理随机数状态保存与恢复,但如果自定义层中使用了不受自动管理的外部随机源,需要额外确保重算过程与原始前向过程在数值上一致,否则可能出现反向梯度不一致的问题。