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

📅 2026/8/12 19:51:13
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 下这个过程大致如下触发收集DeepSpeed 引擎收到保存指令开始协调所有进程。聚合参数每个 GPU 将其持有的参数分片发送到主进程通常是 rank 0。主进程在CPU 内存中将这些分片拼接成完整的参数张量。构建状态字典基于聚合后的完整参数构建出state_dict。序列化保存将state_dict序列化并写入磁盘。释放完整状态关键步骤为了节省主进程的 CPU 内存在保存完成后这个在 CPU 上聚合的完整参数副本通常会被释放。各 GPU 上依然只保留自己的参数分片。问题就出在第5步之后以及接下来训练步骤的衔接上。2.3 OOM 的根本原因状态恢复与分片重建的漏洞保存 checkpoint 本身可能不会直接导致 OOM除非你的 CPU 内存也严重不足。真正的杀手是保存 checkpoint 之后下一个训练 step 开始前的准备工作。在理想情况下保存完成后训练应无缝恢复到之前的分布式状态。但某些情况下这个恢复过程可能出现偏差优化器状态未正确同步保存 checkpoint 时优化器状态可能也被收集和保存。但在恢复训练时如果优化器状态没有严格按照 ZeRO-3 的分片方式重新分发到各个 GPU就可能导致某个 GPU 试图加载远超其应有份额的优化器状态瞬间爆显存。模型参数分片缓存失效DeepSpeed 为了性能会缓存一些远程参数。保存 checkpoint 的聚合操作可能会干扰或清空这些缓存。当下一步前向传播开始时系统可能需要重新从其他 GPU 拉取大量参数如果这些通信和临时缓冲没有管理好就可能造成显存峰值超过限额。LLaMA-Factory 的特定工作流LLaMA-Factory 可能在其Trainer的save_model方法中除了调用 DeepSpeed 的保存还进行了一些额外的操作比如尝试将模型切换为评估模式、或者执行一次额外的模型前向/后向用于验证等。这些操作在 ZeRO-3 环境下如果没有充分考虑分布式状态极易引发混乱。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(fBefore save at step {step}) trainer.save_model(output_dir) print_memory_usage(fAfter save at step {step}) # 训练步骤... loss model(**batch).loss loss.backward() optimizer.step() optimizer.zero_grad() print_memory_usage(fAfter step {step} training)通过对比Before save、After save和After 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 } }需要敲黑板的两个关键配置stage3_gather_16bit_weights_on_model_save: 这个参数默认为true。它的含义是在保存模型时将分布在各个 GPU 上的 16 位fp16/bf16模型参数收集到 CPU 上并拼接成完整的 fp16 权重进行保存。这是必须的否则你保存的 checkpoint 将不完整。问题不在于它本身而在于收集行为带来的副作用。确保它为true不要关闭它。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_size和stage3_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_asyncTrue)。4.3 终极备选方案切换至 ZeRO-2 或使用 CPU Offload如果以上所有方法都失败了而你的模型尺寸只是略微超出 GPU 显存可以考虑降级使用 ZeRO-2将ds_config.json中的stage改为2。ZeRO-2 不分片模型参数因此没有参数聚合/分发的开销checkpoint 保存和恢复的逻辑更简单通常不会出现此类问题。代价是你能训练的模型最大尺寸会变小。启用更激进的 CPU Offload确保offload_optimizer和offload_param都已启用并指向“cpu”。这会将优化器状态和模型参数都卸载到 CPU 内存最大程度节省显存。虽然训练速度会下降但稳定性最高。使用activation_checkpointing(梯度检查点)在模型配置中启用梯度检查点用计算时间换显存空间。这可以为保存 checkpoint 时的临时状态腾出更多缓冲区。{ zero_optimization: { stage: 3 }, activation_checkpointing: { partition_activations: true, contiguous_memory_optimization: true, cpu_checkpointing: true } }5. 常见问题排查清单与实战记录这里汇总了我遇到和从社区了解到的一些典型场景及解决思路你可以像查手册一样对照问题现象可能原因排查步骤与解决方案保存 checkpoint 瞬间 OOMCPU 内存不足无法容纳聚合的完整参数。1. 监控htop或free -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 带来了极致的显存效率但也将模型的状态从“静态完整”变成了“动态分片”。任何需要触及“完整状态”的操作如保存都成为需要精心处理的临界区。希望这份详细的拆解和实战指南能帮你顺利跨过这个深水区让你的大模型训练之旅更加平稳。