DeepSpeed ZeRO-3推理阶段显存溢出与offload配置冲突排查

文章导读
推理阶段遇到 DeepSpeed ZeRO-3 显存溢出,常见原因不是单个参数太大,而是把训练用的 zero_optimization 配置原样搬到了推理脚本里:stage=3 会把参数分片,每次前向都要通过通信把整层参数收集回来;同时 offload_optimizer 又在推理阶段分配了根本用不到的状态。遇到 OOM,先别急着调 batch size,先确认当前推理加载过程到底启用了哪个 st
📋 目录
  1. A 先判断:推理是否真的需要 ZeRO-3
  2. B 拆开训练配置与推理配置
  3. C 用最小动作定位显存峰值
  4. D 两条干净的退出路径
A A

推理阶段遇到 DeepSpeed ZeRO-3 显存溢出,常见原因不是单个参数太大,而是把训练用的 zero_optimization 配置原样搬到了推理脚本里:stage=3 会把参数分片,每次前向都要通过通信把整层参数收集回来;同时 offload_optimizer 又在推理阶段分配了根本用不到的状态。遇到 OOM,先别急着调 batch size,先确认当前推理加载过程到底启用了哪个 stage、带了哪些 offload 项。

DeepSpeed ZeRO-3 推理显存溢出,多数是训练配置被直接复用:stage=3 的前向收集开销加 offload_optimizer 无用状态。排查顺序:确认实际生效的 ZeRO stage、拆分训练/推理配置、只保留 offload_param、再观察显存峰值。如果单卡推理或权重可合并,建议直接跳过 ZeRO-3,用合并后的权重加载,避免入坑。

先判断:推理是否真的需要 ZeRO-3

ZeRO-3 主要解决训练时优化器状态和参数分片,推理时参数固定,通信成本会变成纯开销。如果仅用单卡,ZeRO-3 无法降低单卡显存;如果多卡仅做数据并行,每卡还是需要完整参数。真正需要 ZeRO-3 的是单卡放不下模型,且没有做张量并行,只能靠多卡把参数分片放下来。这时必须接受每步前向的通信代价。

场景建议
单卡加载能放下不使用 ZeRO-3,直接普通加载
多卡数据并行普通加载,每卡复制完整权重
多卡且单卡放不下优先尝试张量并行,再考虑 ZeRO-3
单卡也放不下只剩 CPU offload 或 ZeRO-3+offload

拆开训练配置与推理配置

offload 配置冲突往往指同一份配置文件里同时出现了 offload_optimizer 和 offload_param。训练时需要 optimizer offload,推理时 optimizer 根本不存在,保留 offload_optimizer 会让 DeepSpeed 初始化一个空的优化器状态,增加 CPU 内存和初始化开销。推理配置建议只保留参数 offload:

DeepSpeed ZeRO-3推理阶段显存溢出与offload配置冲突排查
{
  "zero_optimization": {
    "stage": 3,
    "offload_param": {
      "device": "cpu",
      "pin_memory": true
    }
  }
}

这个配置仅用于推理场景,不包含 optimizer 和 scheduler;如果模型权重已经是 ZeRO-3 分片格式,最好由 DeepSpeed 自己加载,不要混用普通 torch.load 和 init_inference。pin_memory 是否开启需结合宿主内存和实际吞吐测试,不是默认越开越好。

用最小动作定位显存峰值

建议直接做两组对照:一组带 offload_param,一组临时移除 offload_param。比较相同输入下 max_memory_allocated 的差异。下面是一个通用的观测骨架,参数名以你所用的 DeepSpeed 版本为准:

import torch
import deepspeed

model = ...  # 你的模型对象
ds_model = deepspeed.init_inference(
    model,
    mp_size=1,
    dtype=torch.float16,
    replace_with_kernel_inject=True,
)

for step in range(3):
    torch.cuda.reset_peak_memory_stats()
    _ = ds_model.generate(...)  # 使用与线上一致的输入长度
    print(torch.cuda.max_memory_allocated())

如果关掉 offload_param 后峰值反而下降,说明主要开销来自前向时的参数收集;如果峰值变化不大且仍 OOM,则要转到激活值或输出缓冲区,需要缩短生成长度、减小 batch size,或改用分块推理。如果开了 offload_param 后出现 CPU 内存暴涨或速度骤降,这不是显存溢出,而是 CPU 与 GPU 之间搬运瓶颈,需要调整 offload 边界。

DeepSpeed ZeRO-3推理阶段显存溢出与offload配置冲突排查

两条干净的退出路径

路径一:合并分片权重后普通加载

如果检查点目录里存在 zero_to_fp32.py 脚本,可以先把 ZeRO-3 分片权重合并成单份权重,再用普通方式加载推理,彻底绕开 ZeRO-3 的前向收集逻辑。

python zero_to_fp32.py \
  `--checkpoint`_dir /path/to/zero3_checkpoint \
  `--output`_file consolidated.bin

合并后的权重不一定能直接用原来的 config 加载,需要确认模型的 state_dict 结构和transformers/torch 加载接口匹配。合并过程需要足够 CPU 内存,建议在内存充足的机器上执行。

DeepSpeed ZeRO-3推理阶段显存溢出与offload配置冲突排查

路径二:张量并行与 ZeRO-3 二选一

多卡推理时,如果已经通过 deepspeed.init_inference 设置了 mp_size 大于 1,就把传给推理引擎的 ds_config 中 zero_optimization 移除,避免同时启用 ZeRO-3。两者同时启用会同时发生参数分片通信和张量并行通信,显存峰值通常不降反升。调整后需要跑一次整模型生成,对比输出 token 是否与普通加载一致。

这些处理都把止血和正式方案分开:CPU offload 和 checkpoint 合并不是为了提升效率,而是为了在可接受的速度损失下把显存覆盖住;如果模型能够单卡放下,优先普通加载。