DeepSpeed ZeRO-3保存Checkpoint后OOM:原理、诊断与解决方案

1. 问题现象与核心挑战

最近在微调一个70亿参数的大模型,用上了 DeepSpeed ZeRO-3 和 LLaMA-Factory 这套强力组合拳。流程跑起来很顺畅,loss 也在稳步下降,一切看起来都很美好。直到我设置了定期保存 checkpoint,问题就来了:模型成功保存了第一个 checkpoint,但在紧接着的下一个训练 step 开始,程序就直接报错退出,错误信息明确指向了 Out Of Memory (OOM)。这感觉就像你辛辛苦苦把游戏进度存了个档,结果一读档,游戏直接崩溃了,非常令人沮丧。

这个问题的诡异之处在于,训练过程本身是稳定的,内存使用也在可控范围内。但“保存”这个动作仿佛触发了一个隐藏的开关,让显存在下一个计算步骤中瞬间被榨干。如果你也遇到了类似deepspeed zero3 + llamafactory 保存checkpoint后第一step 就 OOM的情况,那么你很可能正踩在 ZeRO-3 优化策略与模型状态保存/恢复机制的一个关键交互陷阱上。这不仅仅是内存不足那么简单,它涉及到 ZeRO-3 分布式状态的分片、聚合与重新分发的完整生命周期。接下来,我会结合自己的踩坑经历,把这个问题掰开揉碎了讲清楚,并提供一套完整的诊断和解决方案。

2. DeepSpeed ZeRO-3 与 Checkpoint 保存机制深度解析

要理解为什么保存 checkpoint 后会 OOM,我们必须先深入理解 DeepSpeed ZeRO-3 到底是如何工作的,以及 LLaMA-Factory 在保存 checkpoint 时做了什么。

2.1 ZeRO-3 内存优化原理再回顾

ZeRO-3 是 DeepSpeed 内存优化策略的终极形态,它的核心思想是“极致分片”。它不仅像 ZeRO-2 那样分片优化器状态和梯度,还把模型参数本身也分片存储在各个 GPU 上。这意味着,在任何一个时刻,单个 GPU 上只保存了整个模型参数的一个子集。

  • 前向传播:当需要某一层的参数时,该参数所在的 GPU(所有者)会将其广播给所有其他需要该参数进行计算的 GPU。
  • 后向传播:计算得到的梯度同样被聚合到参数所有者 GPU 上,用于更新优化器状态。
  • 状态聚合:只有在需要执行诸如model.state_dict()或保存 checkpoint 这类操作时,ZeRO-3 引擎才会触发一个“收集”操作,将分散在所有 GPU 上的参数分片收集起来,在 CPU 内存或某个指定的 GPU 上拼接成完整的模型状态。

这种设计带来了巨大的内存节省,使得在有限显存的机器上训练超大模型成为可能。但代价是增加了通信开销,并且让模型状态的“完整视图”变得不再是常态,而是一个需要显式触发的临时状态。

2.2 Checkpoint 保存触发了什么?

当我们调用trainer.save_model()engine.save_checkpoint()时,为了生成一个可以独立加载、包含完整模型参数的.bin.safetensors文件,框架必须获取完整的模型参数。在 ZeRO-3 下,这个过程大致如下:

  1. 触发收集:DeepSpeed 引擎收到保存指令,开始协调所有进程。
  2. 聚合参数:每个 GPU 将其持有的参数分片发送到主进程(通常是 rank 0)。主进程在CPU 内存中将这些分片拼接成完整的参数张量。
  3. 构建状态字典:基于聚合后的完整参数,构建出state_dict
  4. 序列化保存:将state_dict序列化并写入磁盘。
  5. 释放完整状态关键步骤:为了节省主进程的 CPU 内存,在保存完成后,这个在 CPU 上聚合的完整参数副本通常会被释放。各 GPU 上依然只保留自己的参数分片。

问题就出在第5步之后,以及接下来训练步骤的衔接上。

2.3 OOM 的根本原因:状态恢复与分片重建的漏洞

保存 checkpoint 本身可能不会直接导致 OOM(除非你的 CPU 内存也严重不足)。真正的杀手是保存 checkpoint 之后,下一个训练 step 开始前的准备工作

在理想情况下,保存完成后,训练应无缝恢复到之前的分布式状态。但某些情况下,这个恢复过程可能出现偏差:

  1. 优化器状态未正确同步:保存 checkpoint 时,优化器状态可能也被收集和保存。但在恢复训练时,如果优化器状态没有严格按照 ZeRO-3 的分片方式重新分发到各个 GPU,就可能导致某个 GPU 试图加载远超其应有份额的优化器状态,瞬间爆显存。
  2. 模型参数分片缓存失效:DeepSpeed 为了性能,会缓存一些远程参数。保存 checkpoint 的聚合操作可能会干扰或清空这些缓存。当下一步前向传播开始时,系统可能需要重新从其他 GPU 拉取大量参数,如果这些通信和临时缓冲没有管理好,就可能造成显存峰值超过限额。
  3. LLaMA-Factory 的特定工作流:LLaMA-Factory 可能在其Trainersave_model方法中,除了调用 DeepSpeed 的保存,还进行了一些额外的操作,比如尝试将模型切换为评估模式、或者执行一次额外的模型前向/后向用于验证等。这些操作在 ZeRO-3 环境下,如果没有充分考虑分布式状态,极易引发混乱。
  4. deepspeed.zero.Init()上下文管理问题:模型是在deepspeed.zero.Init()上下文内初始化的,这确保了参数被正确分片。但如果在 checkpoint 保存/加载循环中,有代码意外在上下文外创建了新的张量或子模块,这个新对象就不会被 ZeRO-3 管理,成为一个完整的、未分片的“巨无霸”,直接导致 OOM。

我的经验是,最常见的原因集中在优化器状态的重建框架特定保存钩子的副作用上。接下来,我们进入实战排查环节。

3. 系统性诊断与问题定位实操

当 OOM 发生时,不要盲目增加--per_device_train_batch_size或换用更大的 GPU。首先需要进行系统化诊断,定位显存是在哪个环节被消耗的。

3.1 利用工具监控显存变化

你需要亲眼看到显存是如何涨上去的。这里推荐两个方法:

  • 使用nvidia-smi循环监控:在一个单独的终端运行以下命令,观察显存变化。

    watch -n 0.1 nvidia-smi

    在训练脚本即将保存 checkpoint 前(可以通过日志判断),密切观察所有 GPU 的显存使用量。注意是“即将保存”和“保存后第一步计算”这两个时间点。

  • 集成torch.cuda.memory跟踪:在你的训练脚本中插入内存记录点。这更精确。

    import torch def print_memory_usage(prefix): allocated = torch.cuda.memory_allocated() / 1024**3 reserved = torch.cuda.memory_reserved() / 1024**3 print(f"[{prefix}] Allocated: {allocated:.2f} GB, Reserved: {reserved:.2f} GB") # 在训练循环中 for step, batch in enumerate(train_dataloader): if step % save_steps == 0: print_memory_usage(f"Before save at step {step}") trainer.save_model(output_dir) print_memory_usage(f"After save at step {step}") # 训练步骤... loss = model(**batch).loss loss.backward() optimizer.step() optimizer.zero_grad() print_memory_usage(f"After step {step} training")

    通过对比Before saveAfter saveAfter step ... training的数值,你可以清晰看出是保存动作本身导致了显存增加,还是保存后的第一个训练 step 导致的。

3.2 检查 DeepSpeed 配置文件

你的ds_config.json是罪魁祸首的首要怀疑对象。请仔细检查以下配置项:

{ "zero_optimization": { "stage": 3, "offload_optimizer": { "device": "cpu", // 如果为“cpu”,检查是否配置正确 "pin_memory": true // pin_memory 可以提升速度,但可能增加CPU内存压力 }, "offload_param": { "device": "cpu", // 参数卸载到CPU,对缓解显存问题至关重要 "pin_memory": true }, "overlap_comm": true, // 重叠通信,一般建议开启 "contiguous_gradients": true, "sub_group_size": 1e9, "reduce_bucket_size": "auto", "stage3_prefetch_bucket_size": "auto", "stage3_param_persistence_threshold": "auto", // 关键参数! "stage3_max_live_parameters": 1e9, "stage3_max_reuse_distance": 1e9, "stage3_gather_16bit_weights_on_model_save": true // 关键参数! }, "train_batch_size": "auto", "train_micro_batch_size_per_gpu": "auto", "gradient_accumulation_steps": "auto", "fp16": { "enabled": true }, "bf16": { "enabled": false } }

需要敲黑板的两个关键配置:

  1. stage3_gather_16bit_weights_on_model_save: 这个参数默认为true。它的含义是在保存模型时,将分布在各个 GPU 上的 16 位(fp16/bf16)模型参数收集到 CPU 上,并拼接成完整的 fp16 权重进行保存。这是必须的,否则你保存的 checkpoint 将不完整。问题不在于它本身,而在于收集行为带来的副作用。确保它为true,不要关闭它。

  2. stage3_param_persistence_threshold: 这个参数控制哪些参数会“持久化”在 GPU 上,而不是在使用后立即释放。默认值“auto”通常是合理的。但如果它被设置得异常大(例如 1e9),可能会导致 DeepSpeed 试图在 GPU 上缓存比预期更多的参数,在保存 checkpoint 后的状态重建时引发混乱。建议先保持为“auto”

3.3 审查 LLaMA-Factory 的保存逻辑

查看 LLaMA-Factory 中Trainer类的save_model方法(通常位于src/llamafactory/train/trainer.py或类似路径)。你需要关注:

  • 在调用self.engine.save_checkpoint前后,是否有额外的model.eval()model.train()切换?
  • 是否有为了计算验证损失而进行的额外前向传播 (model(**batch))?
  • 保存的路径是否干净?是否尝试同时保存多个副本导致临时文件堆积?

一个常见的隐患是:为了在保存时计算一些指标(如 perplexity),代码可能无意中在torch.no_grad()上下文之外执行了前向传播,这会导致梯度计算图的构建和中间变量的保留,白白消耗大量显存。

4. 解决方案与优化策略

根据诊断结果,你可以尝试以下一种或多种组合策略。

4.1 调整 DeepSpeed 配置(首选)

修改你的ds_config.json,尝试以下组合:

{ "zero_optimization": { "stage": 3, "offload_optimizer": { "device": "cpu" }, "offload_param": { "device": "cpu" }, "stage3_gather_16bit_weights_on_model_save": true, "stage3_param_persistence_threshold": 1e5, // 显式设置为一个较小的值,例如10万 "reduce_bucket_size": 5e8, // 明确设置通信桶大小,避免“auto”的不确定性 "stage3_prefetch_bucket_size": 5e7 }, "aio": { "enabled": true, "block_size": 1048576, "queue_depth": 8, "thread_count": 1, "single_submit": false, "overlap_events": true } }

解释与操作意图:

  • stage3_param_persistence_threshold: 设置为一个具体的、较小的值(如 1e5),这可以迫使 DeepSpeed 更积极地释放不再需要的参数缓存,可能在状态重建时提供更干净的显存环境。
  • 明确设置reduce_bucket_sizestage3_prefetch_bucket_size:有时“auto”估算的值在特定模型和硬件上可能不是最优的,明确设置可以消除一个变量。
  • 启用aio(异步IO):这可以加速 checkpoint 从 CPU 内存写入磁盘的速度,可能缩短“完整参数集驻留CPU内存”的时间窗口,间接降低风险。

4.2 修改保存策略与工作流

如果调整配置无效,可能需要修改训练脚本的保存逻辑。

  • 策略一:保存后立即进行垃圾回收并清空 CUDA 缓存

    import gc import torch # 在保存 checkpoint 之后,下一个训练 step 之前 trainer.save_model(output_dir) gc.collect() # 强制进行Python垃圾回收 torch.cuda.empty_cache() # 清空 PyTorch 的 CUDA 缓存 # 注意:empty_cache() 不会释放由张量持有的显存,但可以释放一些缓存的内存分配器持有的内存。
  • 策略二:将保存 checkpoint 与训练 step 分离考虑在达到保存间隔时,不是立即保存,而是设置一个标志位,在完成当前梯度累积的多个 micro-batch 之后,下一个 step 开始之前进行保存。这可以确保保存操作发生在一个相对“静止”的状态,而不是紧挨着繁重的计算步骤。

    save_pending = False for step, batch in enumerate(train_dataloader): if step % save_steps == 0: save_pending = True # 正常的训练步骤... loss = model(**batch).loss loss.backward() if (step + 1) % gradient_accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 在参数更新后、下一个累积循环开始前保存 if save_pending: trainer.save_model(output_dir) gc.collect() torch.cuda.empty_cache() save_pending = False
  • 策略三:使用 DeepSpeed 的异步 checkpoint 保存DeepSpeed 支持异步保存 checkpoint,这可以将保存操作放到后台线程,不影响训练主线程。查看 LLaMA-Factory 是否支持或可以集成engine.save_checkpoint(save_dir, tag, client_state={}, save_async=True)

4.3 终极备选方案:切换至 ZeRO-2 或使用 CPU Offload

如果以上所有方法都失败了,而你的模型尺寸只是略微超出 GPU 显存,可以考虑:

  1. 降级使用 ZeRO-2:将ds_config.json中的stage改为2。ZeRO-2 不分片模型参数,因此没有参数聚合/分发的开销,checkpoint 保存和恢复的逻辑更简单,通常不会出现此类问题。代价是你能训练的模型最大尺寸会变小。
  2. 启用更激进的 CPU Offload:确保offload_optimizeroffload_param都已启用并指向“cpu”。这会将优化器状态和模型参数都卸载到 CPU 内存,最大程度节省显存。虽然训练速度会下降,但稳定性最高。
  3. 使用activation_checkpointing(梯度检查点):在模型配置中启用梯度检查点,用计算时间换显存空间。这可以为保存 checkpoint 时的临时状态腾出更多缓冲区。
    { "zero_optimization": { "stage": 3 }, "activation_checkpointing": { "partition_activations": true, "contiguous_memory_optimization": true, "cpu_checkpointing": true } }

5. 常见问题排查清单与实战记录

这里汇总了我遇到和从社区了解到的一些典型场景及解决思路,你可以像查手册一样对照:

问题现象可能原因排查步骤与解决方案
保存 checkpoint 瞬间 OOMCPU 内存不足,无法容纳聚合的完整参数。1. 监控htopfree -h查看 CPU 内存使用。
2. 尝试在保存前执行gc.collect()
3. 考虑增加机器 CPU 内存或使用stage3_gather_fp16_weights_on_model_save: false(不推荐,会保存分片checkpoint)。
保存后第一个训练 step 的 forward 中 OOM参数分片缓存或通信缓冲区在恢复时出错。1. 检查stage3_param_persistence_threshold,尝试调小。
2. 在保存后、训练前插入torch.cuda.empty_cache()
3. 确保没有在deepspeed.zero.Init()上下文外创建新模块。
保存后第一个训练 step 的 backward 或 optimizer.step() 中 OOM优化器状态未正确重新分片。1.这是最常见原因!确保 DeepSpeed 版本与 PyTorch、Transformers 版本兼容。
2. 在save_checkpoint后,尝试显式调用engine.load_checkpoint(指向刚保存的路径) 来强制重新加载并分片状态。这听起来奇怪,但有时能重置引擎内部状态。
3. 降级到 ZeRO-2 测试是否问题消失,以确认是 ZeRO-3 特有问题。
只有特定模型或特定大小才会出现LLaMA-Factory 的某些模型前/后处理钩子与 ZeRO-3 不兼容。1. 在 LLaMA-Factory 的 issue 中搜索你的模型名 + “zero3” 或 “oom”。
2. 尝试使用 LLaMA-Factory 的--stage sft而不是--stage pt(如果适用),因为预训练通常负载更重。
3. 尝试一个更小的模型(如 3B)测试工作流是否正确。
错误信息中包含CUDA out of memory. Tried to allocate ...但后面跟的尺寸异常大(如几十GB)几乎可以肯定是 ZeRO-3 状态混乱,某个 GPU 试图分配完整模型参数。1. 立即检查代码,确保所有模型相关操作都在deepspeed.zero.Init()上下文内初始化。
2. 检查是否在训练循环中意外调用了model.cpu()model.to(device),这可能会破坏分片状态。
3. 使用deepspeed.runtime.zero.parameter_offload.ZeroParamStatus工具进行调试(高级用法)。

我的实操心得:在我自己的案例中,最终解决问题的方法是组合策略。首先,我明确设置了stage3_param_persistence_threshold为一个较小的数值(1e5)。其次,我在 LLaMA-Factory 的save_model调用后,紧接着添加了显式的gc.collect()torch.cuda.empty_cache()。最后,也是我认为最关键的一步,我确保了我的训练脚本在启动时,所有模型相关的定义都严格包裹在with deepspeed.zero.Init():语句块内,防止了任何“漏网之鱼”的全参数张量被创建。调整之后,保存 checkpoint 变得顺滑,再也没有出现后续 step OOM 的情况。

这个问题的本质是分布式训练中状态管理的复杂性。ZeRO-3 带来了极致的显存效率,但也将模型的状态从“静态完整”变成了“动态分片”。任何需要触及“完整状态”的操作(如保存),都成为需要精心处理的临界区。希望这份详细的拆解和实战指南,能帮你顺利跨过这个深水区,让你的大模型训练之旅更加平稳。