导读:本期聚焦于深圳SEO公司创作的《全参数训练显存不够怎么办?梯度检查点与ZeRO如何组合破局》,敬请观看详情。把大模型塞进单卡训练时,显存往往先被激活值和优化器状态吃光,而不是模型权重本身。很多开发者习惯直接减小 batch size 或开启 CPU offload,却忽略了两种可以叠加的显存优化手段:梯度检查点与 ZeRO。前者在前向传播时丢弃部分中间激活,反向传播需要时重新计算,用约三分之一的额外算力换回数倍的激活显存;后者把优化器状态、梯度和模型参数分片到多张 GPU,让每张卡只保存完整状态的一部分。本文从显存分布拆解入手,给出两者的实现要点和组合策略,包括如何调整检查点分段、ZeRO 各阶段差异以及通信代价,帮助你在不牺牲全局 batch size 的前提下完成全参数训练。

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

全参数训练显存不够怎么办?梯度检查点与ZeRO如何组合破局

一、先看清显存消耗的完整结构

全参数训练中的显存占用可以分成静态和动态两部分。静态部分由模型权重、梯度和优化器状态组成,它们在整个训练过程中一直存在。动态部分主要是前向传播产生的激活值,以及通信库、注意力计算、梯度裁剪等过程产生的临时缓冲区。如果使用 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 实现会处理随机数状态保存与恢复,但如果自定义层中使用了不受自动管理的外部随机源,需要额外确保重算过程与原始前向过程在数值上一致,否则可能出现反向梯度不一致的问题。

梯度检查点ZeRO显存优化修改时间:2026-09-27 22:40:09

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/0927/62732.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。