DeepSpeed显存优化实战:混合精度训练与ZeRO详解
发布时间:2026/9/13 2:23:32 作者:尧图编辑部 阅读量:1,286

我最近在项目里踩了一个非常典型的坑用一个 7B 模型做微调单张 A100 80G 怎么塞都塞不下后来把 DeepSpeed 的混合精度训练和 ZeRO 零冗余优化器打开同一个 batch 大小下显存占用从 90 多 G 直接掉到 43G 左右。这个效果确实显著但也让我在配置上翻了不少次车。所以这篇内容不是单纯的 API 介绍而是把混合精度训练和 ZeRO 这两块真正讲透它们是解决什么问题的、配置里每一项到底在干什么、哪些参数可以无脑用、哪些必须按显存和卡数精打细算。如果你也在做大模型训练、微调或者正在被显存不足折磨这篇文章值得你花十分钟读一遍能少走不少弯路。1. 为什么要用 DeepSpeed 训大模型显存不够和通信太慢很多同学一开始接触 DeepSpeed 是因为它在 HuggingFace Trainer 里只需加一个参数就能跑起来但对它到底解决了什么问题往往一知半解。这里我先不铺概念直接带大家算一笔账你马上就能明白它的价值。1.1 单卡显存的天花板先算一笔账假设我们要训练一个参数量为 7B70 亿参数的模型。单看模型权重如果使用 FP16 存储需要 14GB 显存这看起来 80G 的卡完全够。但真正开始训练后显存中的大块头远不止“权重”这一项。在标准的 AdamW 优化器训练中每一个参数在显存里至少需要这样几份数据FP16 格式的参数副本2 字节、FP16 格式的梯度2 字节、FP32 格式的优化器状态。这个优化器状态又细分为三部分FP32 的 master weight4 字节、Adam 的一阶动量 m4 字节、Adam 的二阶动量 v4 字节。按这个公式算下来一个参数在混合精度训练中合计占用 16 字节。也就是说训练 7B 模型时仅“模型状态”就需要 7B × 16 字节 112GB 显存。这还不包括 activation 激活值、临时缓冲区、通信缓冲区等额外开销。所以单卡 80G 放不下 7B 模型的训练并不是什么玄学而是账算出来就是不够。传统的数据并行方案 DDP 的做法是每张卡都保存一份完整的模型状态每轮更新时通过 all-reduce 把梯度同步到所有卡上。这样虽然利用多卡扩大了计算能力但显存占用是线性翻倍的——8 张卡就是 8 份完整的模型状态112GB × 8再大的集群也会被显存卡死。DeepSpeed 的 ZeRO 正是针对这个“冗余存储”设计的零冗余优化器核心思路是既然每张卡最终计算出的梯度最终都会汇总并保持同步那不如把模型状态按维度切分出去每一张卡只保存一部分用通信换取显存。这也是“零冗余”这个名字的真正含义。1.2 分布式训练的三个瓶颈显存、通信、负载均衡分布式训练不是把卡插上去就能线性加速。真正训练时我们几乎都会遇到三个瓶颈。第一个是显存瓶颈这也是最直观的。模型越大、batch size 越大显存需求就越高。很多人把 activation checkpointing 和 ZeRO 混为一谈实际上它们解决的是不同层面的问题ZeRO 解决的是参数、梯度、优化器状态的冗余存储而 activation checkpointing 解决的是前向传播时中间激活值的存储。两个可以同时开但要注意它们都需要你手动在配置中分别启用。第二个是通信瓶颈。数据并行训练里每一轮迭代结束之后都要同步梯度。DDP 使用的是 all-reduce 通信原语它的通信量与模型大小和 GPU 数量直接相关。模型越大通信时间越长甚至会超过计算时间导致 GPU 大量时间在空等。DeepSpeed 的 ZeRO 在设计通信方案时尽量复用了梯度规约的过程比如用 reduce-scatter 和 all-gather 分解原有的 all-reduce让通信量和 DDP 基本持平。但这个目标只在 ZeRO Stage1 和 Stage2 成立到 Stage3 时因为前向和反向都需要重新收集参数通信量会明显增加这是后话。第三个是负载均衡瓶颈。多卡训练时如果每张卡的显存放的模型分片大小不一致很可能会出现“木桶效应”最快的那张卡要等最慢的那张卡算完才能进入下一步。这也是为什么 DeepSpeed 官方建议在配置 ZeRO-3 时开启round_robin_gradients等自动分片策略尽量让每张卡上的参数大小均匀分配。理解这三个瓶颈之后你就知道 DeepSpeed 的核心优化点在哪里了。混合精度训练主要负责降低单份数据的字节数ZeRO 则负责把整体存储“摊薄”到多台设备上两者配合才是完整的解法。如果只开其一效果会大打折扣。2. 混合精度训练用 FP16 省一半显存还能提速混合精度训练算是现在大模型训练的事实标准。简单说就是前向和反向用 FP16 计算优化器更新时用 FP32 精度做参数更新。这样既能利用 FP16 减半的显存占用又能避免数值精度不足导致的模型不收敛。2.1 FP16 为什么能省显存又为什么容易“崩”FP16 全称是半精度浮点数占用 2 字节而 FP32 占用 4 字节所以单从存储看FP16 就是标准的“省一半”。而且大多数现代 GPU 对 FP16 的矩阵运算专门做了优化算力是 FP32 的两倍甚至更多。深度学习里真正对数值精度要求高的地方是梯度更新的累积过程而矩阵乘法、卷积这类算子并不需要 32 位浮点那么高的精度所以理论上用 FP16 相当划算。但 FP16 的“坑”在于它的动态范围很窄。FP32 的指数范围大概到 1e38而 FP16 的最大值只有 65504。什么意思呢在训练过程中如果一个值算出来超过 65504就直接变成 Inf无穷大。另一个更隐蔽的问题是下溢FP16 能表示的最小归一化数值大约是 6.1e-5再小的数会变成 0。梯度在深层网络反传时往往很小很多值会落在 1e-6 甚至更小的量级FP16 下一不小心就变成 0导致参数彻底不更新loss 卡死或者发散。这里我常用一个生活化的类比FP32 就像一把精确到毫米的大标尺FP16 则像一把量程有限、最小刻度也不太够用的尺子。你量普通桌子没问题一旦遇到又大又小的极端数据就容易爆表或者测不准。所以混合精度训练不能只把模型扔到 FP16 里就完事必须配套做“损失缩放”Loss Scaling和“主权重”Master Weights来兜底。2.2 损失缩放Loss Scaling与动态缩放机制损失缩放是目前解决梯度下溢的最有效手段思路非常朴素既然梯度太小在 FP16 里会变成 0那我就在反向传播之前把 loss 乘以一个较大的缩放因子比如 2 的 16 次方即 65536这样反传得到的梯度也会被同等放大落到 FP16 可表示的范围内。更新参数前再把这个放大的梯度除以缩放因子恢复成真正的梯度值。缩放因子怎么定早期的做法是固定一个值比如 1024 或者 65536。但固定值有个问题不同模型、不同 batch size 下 loss 和梯度的量级差别很大固定缩放因子很容易导致数值溢出。所以 Apex 和 DeepSpeed 现在都使用动态损失缩放机制。简单说程序会先设定一个初始缩放因子如果在接下来一段时间内没有出现 Inf 或 NaN就认为缩放因子还有上调空间隔一段时间把它翻倍一旦检测到 Inf 或 NaN立刻把缩放因子减半并且跳过当前这个 step 的更新。这个机制在 DeepSpeed 配置里对应initial_scale_power、loss_scale_window、hysteresis这几个参数。比如我常用的配置是这样fp16: { enabled: true, initial_scale_power: 16, loss_scale_window: 1000, hysteresis: 2, min_loss_scale: 1 }initial_scale_power: 16表示初始缩放因子是 2 的 16 次方也就是 65536。loss_scale_window: 1000表示如果连续 1000 步没有出现 Inf 或 NaN就把缩放因子翻倍。hysteresis: 2是一个缓冲值意思是出现过 Inf 之后需要再连续跑 2 个 window即 2 个 1000 步没有异常才会继续增大缩放因子。2.3 Master Weights 和 FP32 副本精度兜底有了损失缩放梯度下溢的问题解决了但还有一个细节前向和反向在 FP16 下计算如果更新时直接把 FP16 的权重拿出来更新每轮迭代都会产生一点对 FP16 的舍入误差几千步累积下来权重可能和真正应该更新的方向发生较大偏离。混合精度训练因此在优化器中额外维护一份 FP32 格式的参数副本叫 master weights。每一轮优化器更新都在 FP32 副本上进行得到新的 FP32 参数后再转成 FP16 用于下一轮前向。这就是“每次更新都比纯 FP16 更准”的原因。不过在 ZeRO 机制里这份 master weight 是显存大户。前面我们算过一个参数要占 4 字节的 FP327B 模型光 master weight 就是 28GB。所以 ZeRO 通常会先切分优化器状态优先切的就是这份 FP32 数据和 Adam 动量、方差。这也是为什么 ZeRO Stage1 在配置上对显存节省最明显的原因之一。2.4 实操DeepSpeed 里如何正确开启混合精度在 DeepSpeed 配置里FP16 和 BF16 是互斥的二者不能同时设为 true。我强烈建议如果你的 GPU 是 Ampere 架构或更新A100、H100、A800 等直接优先考虑用 BF16 而不是 FP16。BF16 的指数范围和 FP32 完全一样最大最小值与 FP32 一致所以基本不存在溢出问题训练稳定性显著提升代码里也无需额外做损失缩放。它唯一的代价是尾数位比 FP16 少但实际训练中影响很小。开启 BF16 的方法是在配置里写bf16: { enabled: true }如果只能用 FP16就按上一节的方式配置fp16块。有一点必须注意不要在一份配置里同时写fp16和bf16两个块都启用否则 DeepSpeed 会启动报错。如果你从网上拷贝配置模板先检查这两个开关有没有冲突。这个坑我至少遇到三次。3. ZeRO 零冗余优化器把显存浪费降到接近零ZeRO 是 DeepSpeed 这次最值得研究的核心优化器。它并不是某个具体的优化算法而是把原有的模型状态按维度切分到多张卡上的分布式策略。理解它的关键是把“显存里到底存了什么”梳理清楚。3.1 模型状态 vs 残余状态显存被谁吃掉了ZeRO 官方将训练中的显存占用分为两类模型状态Model States和残余状态Residual States。模型状态指参数、梯度、优化器状态这三样残余状态指激活值activation、临时缓冲区temporary buffers以及不可用的显存碎片。ZeRO 的核心目标是消除“模型状态”的冗余而激活值问题需要单独用 activation checkpointing 解决。很多同学在开了 ZeRO 之后显存还是很高第一反应是“ZeRO 不管用”其实是你没开 activation checkpointing或者你的微批大小太大。激活值的显存占用与层数、序列长度、微批大小成正比。以 7B 模型、序列长度 2048 为例单层激活值常常会有几百 MB几十层堆下来轻松超过 20 到 30GB。要压显存先做激活检查点再考虑调整微批大小最后才是折腾 ZeRO 的 offload 策略。3.2 ZeRO Stage1 / Stage2 / Stage3分阶段干掉冗余ZeRO 根据切分对象的不同把优化措施分成三个级别。Stage1 只切分优化器状态。在 N 张卡上原本每张卡都保存一份完整的 12N 字节优化器状态现在每张卡只保存 12N/N 的部分。参数和梯度仍是完整保存。Stage1 的通信模式和 DDP 非常接近所以几乎不增加额外通信成本是性价比最高的第一步。Stage2 在 Stage1 的基础上把梯度也做了切分。具体做法是用 reduce-scatter 把梯度规约到各自负责的卡上每张卡只保留自己负责的那份梯度更新时再通过 all-gather 同步完整的新参数。因为梯度也是每卡完整冗余所以显存能进一步压缩。从通信量看Stage2 和 DDP 基本持平又比 Stage1 省更多显存所以它是目前大多数训练任务用的“甜点配置”。Stage3 则更进一步把参数也切分掉。每张卡只保存模型参数的一小部分在前向或反向计算到某一层时通过 all-gather 把这一层参数临时拼成完整副本参与计算算完就释放。这让 ZeRO-3 理论上可以训练无限大的模型——显存够不够只取决于你能把模型分得多细。但代价是通信量明显增大因为每一层前向和反向都要各做一次参数收集。所以 ZeRO-3 通常会配合overlap_comm: true和梯度通信与计算重叠来缓解。下面这张表是我实测比较常用的估算方式以 7B 模型、64 卡为例只算模型状态不算激活方案参数FP16梯度FP16优化器状态FP32每卡约需显存DDP14GB14GB84GB112GBZeRO Stage114GB14GB84GB / 64 ≈ 1.3GB约 29.3GBZeRO Stage214GB14GB / 64 ≈ 0.2GB84GB / 64 ≈ 1.3GB约 15.5GBZeRO Stage314GB / 64 ≈ 0.2GB0.2GB1.3GB约 1.7GB这里要注意表格里的 ZeRO-3 每卡看起来很低但实际运行时因为要生成临时完整参数副本峰值显存会比理论值高出一截而且残余状态也没算进去。所以不要真以为 64 卡能训练无限大的模型仍然要留足激活值和通信缓冲区的余量。3.3 ZeRO-Offload把负担甩给 CPU 和 NVMe当显存还是不够时ZeRO 还可以把优化器状态或者参数进一步卸载到 CPU 内存甚至 NVMe 硬盘上。最常用的是offload_optimizer在 Stage2 和 Stage3 下都能用。原理是把 Adam 的动量和方差放到 CPU 内存GPU 只保留 FP32 master weight 和计算需要的临时数据。CPU offload 的代价是 CPU 和 GPU 之间的 PCIe 通信会成为新瓶颈。我的经验是如果单机卡足够优先增加 GPU 卡数而不是开 offload只有在单卡显存实在放不下、又必须训练大模型时再开。开 offload 的时候建议把pin_memory: true打开能减少数据传输时的内存拷贝开销。NVMe offload 和 CPU offload 的思路一致适合 CPU 内存也不够的场景但速度更慢配置也更复杂需要指定nvme_path和 AIO 相关参数一般普通开发者极少用到。3.4 配置 ZeRO 的完整 JSON 模板先给一个可以直接参考的 ZeRO Stage2 配置模板再逐项解释{ train_batch_size: 32, train_micro_batch_size_per_gpu: 1, gradient_accumulation_steps: 8, fp16: { enabled: true, initial_scale_power: 16, loss_scale_window: 1000, hysteresis: 2, min_loss_scale: 1 }, zero_optimization: { stage: 2, allgather_partitions: true, reduce_scatter: true, overlap_comm: true, contiguous_gradients: true }, optimizer: { type: AdamW, params: { lr: 1e-5, weight_decay: 0.01 } } }train_batch_size是所有 GPU 上的总 batch size。train_micro_batch_size_per_gpu是每张卡每步实际计算的 batch size。gradient_accumulation_steps是梯度累积步数DeepSpeed 会自动用公式“总 batch micro batch × GPU 数 × 累积步数”去校验如果三者的乘积对不上启动时会直接报错。allgather_partitions和reduce_scatter都建议保持 true这是 Stage2 利用 gather 和 scatter 替代 all-reduce 的关键。overlap_comm: true会让梯度通信与反向计算重叠能显著减少通信等待。contiguous_gradients: true是将梯度存储到连续内存块里以减少通信碎片开启后对性能更友好。4. 实操配置详解从单卡到多卡这一节我想直接从实际改造训练脚本的视角讲清楚 DeepSpeed 的训练循环长什么样以及启动命令怎么组织。这里以 PyTorch 原生训练脚本为例因为它比直接用 Trainer 封装更能看清底层原理。4.1 安装与环境检查安装 DeepSpeed 并不复杂建议在一个独立的 Python 环境里执行pip install deepspeed装完以后先跑一下ds_report它会直接列出当前环境里的 PyTorch 版本、CUDA 版本、GPU 架构、已安装的通信库等核心信息。我见过很多启动失败最后都能归因到环境和 CUDA 版本不匹配上。DeepSpeed 版本建议选 0.9.5 以上功能更稳定对 bf16、NVMe offload 的支持也更完善。如果使用 HuggingFace Trainer还需要确保transformers和accelerate版本兼容建议都更新到较新版本。4.2 改造训练脚本关键就三处原生 PyTorch 训练脚本改成 DeepSpeed 其实不需要动主干。以最基础的分类模型训练为例原来的训练循环长这样optimizer torch.optim.AdamW(model.parameters(), lr1e-5) for batch in dataloader: loss model(batch) loss.backward() optimizer.step() optimizer.zero_grad()改成 DeepSpeed 后大致是下面这个样子import deepspeed from deepspeed.ops.adam import FusedAdam model_engine, optimizer, _, _ deepspeed.initialize( argsargs, modelmodel, model_parametersmodel.parameters() ) for batch in dataloader: loss model_engine(batch) model_engine.backward(loss) model_engine.step()这里有几个必须理解的点第一deepspeed.initialize会返回一个model_engine它本身是DeepSpeedEngine类型的对象你可以直接像调用普通模型一样调用它。第二optimizer不再需要自己创建如果配置文件中写了optimizer字段DeepSpeed 会按配置自动创建如果没有也可以从返回参数里拿到这个optimizer再用但一般没有必要。第三backward和step是 DeepSpeed 封装的方法它会自动处理混合精度下的 loss scale 缩放、梯度裁剪、梯度累积、ZeRO 分片同步等一系列底层操作。不要在model_engine.backward(loss)前后再手动调用torch.cuda.amp.GradScaler()之类的模块否则很可能出现双重缩放实际训练时 loss 要么不变、要么直接 Nan。如果你用 HuggingFace Trainer更简单只需要在TrainingArguments里传入配置路径training_args TrainingArguments( output_dir./output, deepspeedds_config.json, per_device_train_batch_size1, gradient_accumulation_steps8, )这期间 Trainer 会自动调用 DeepSpeed 的 initialize你基本不需要碰底层代码。4.3 关键参数怎么定batch、微批、累积步数batch size 的配置是我见过最容易理解错的地方。DeepSpeed 的总 batch size 由公式决定总 batch size 微批大小 × GPU 数量 × 梯度累积步数很多刚从单卡转向多卡的人会犯一个错误在配置里同时写了train_batch_size: 32和train_micro_batch_size_per_gpu: 1又在启动时用了 4 张卡结果 DeepSpeed 启动时报错说三者不匹配。实际上你可以只指定前两个值DeepSpeed 会自动帮你算出梯度累积步数或者只指定微批大小和梯度累积步数DeepSpeed 自动帮你算总 batch size。注意梯度累积步数越大每一步参数更新的间隔越长训练速度会变慢但显存占用几乎不变。所以当你碰到显存不够但还希望保持总 batch size 不变时优先减小微批大小并调大累积步数。比如从每卡 8 减到每卡 2累积步数从 1 增到 4总 batch 不变但显存可能只有原来的四分之一。多卡启动命令也很直接。单机 8 卡可以这样跑deepspeed --num_gpus 8 train.py --deepspeed_config ds_config.json多机场景需要准备一个 hostfile比如worker-1 slots8 worker-2 slots8然后执行deepspeed --hostfile hostfile --num_gpus 8 --num_nodes 2 train.py --deepspeed_config ds_config.jsonnum_gpus和num_nodes如果不写DeepSpeed 会尝试自动从宿主机探测但我建议显式指定避免在未知环境里老探测错。5. 常见问题排查与避坑实录这部分我从实际踩坑经验里挑几个高频问题整理成清单。它们不是网上搜不到而是很多人搜到了也不知道怎么往自己项目里套。我在这里直接给出判断思路和操作建议。5.1 开启混合精度后 Loss 变成 NaN 或 Inf这是混合精度训练里最高频的问题。如果你用的是 FP16先检查 FP16 的loss_scale_window和initial_scale_power。如果 loss 在训练前期还好中后期突然变成 NaN多半是动态 scale 增长太快或者学习率偏大导致梯度过冲。可以先尝试把initial_scale_power降到 8即初始缩放因子 256并增大loss_scale_window到 2000给模型更长的“稳定窗口”。如果用了 FP16 始终救不回来就切换到 BF16BF16 由于指数范围和 FP32 一致几乎不会因为溢出导致 NaN这是最省事的处理方式。还有一个小细节如果训练代码里手动做了手动混合精度处理比如自己写了torch.cuda.amp.autocast()但配置里又开启了 DeepSpeed 的fp16两套机制叠加NaN 概率会显著上升。用 DeepSpeed 时建议把手动 autocast 去掉让它全权接管。5.2 ZeRO-3 保存的模型权重不完整或加载报错ZeRO-3 下模型参数本身被分布到多张卡上每张卡的 model 对象只拥有自己分片的那部分参数。如果直接调用save_pretrained或torch.save(model.state_dict())很可能只保存了当前 rank 的分片或者出现大量missing keys导致无法直接用。这里我常用的办法是先用聚合接口把所有分片收集回第一张卡再保存if model_engine.zero_optimization_partition_weights(): model_engine.save_checkpoint( save_dircheckpoint, client_state{step: step}, save_latestTrue, exclude_frozen_parametersTrue )如果用的是 HuggingFace Trainer还可以设置--save_full_training_state true让它专门处理 ZeRO-3 的 checkpoint 聚合。需要注意训练中途的 checkpoint 和最终可加载的模型权重是两个概念如果为了继续训练需要用load_checkpoint恢复全部状态如果只是想要能直接部署的权重要额外做参数收集。5.3 训练速度不升反降比 DDP 还慢这不是 DeepSpeed 不行通常是两种原因。一是模型规模不够大却硬开了 ZeRO-3。ZeRO-3 的通信开销非常大如果你的模型只有 1B 或者更小显存不是瓶颈ZeRO-3 只会增加没必要的通信负担这种情况下 DDP 都可能更快。二是 offload 开得太激进。CPU offload 会把优化器状态通过 PCIe 搬到内存更新时再搬回来如果小 batch 每步都频繁更新通信延迟会盖过计算时间。我建议先用 Stage2 跑通观察效率和显存再决定要不要往 Stage3 和 offload 方向调。此外确保在配置文件里设置了overlap_comm: true它能办到在反向计算的同时就开始通信有效掩盖一部分延迟。如果仍然慢可以尝试增大gradient_accumulation_steps让梯度通信的频率降低总通信量不变但每轮通信的等待被更多计算步骤覆盖。5.4 开启 ZeRO-2 后激活值仍然爆显存这种情况最常出现在长序列任务上。ZeRO 解决的是参数与优化器状态的冗余激活值依然占着大量显存。你把批次微调到 1 以后如果还 OOM那就必须开 activation checkpointing。HuggingFace 的TrainingArguments里直接加--gradient_checkpointing true或者在原生脚本里自行配置。要注意 activation checkpointing 的本质是用额外的前向计算换取存储空间所以会略微降低训练吞吐但通常能节省 60% 以上的激活显存性价比很高。另一个容易被忽略的是微调阶段存在的旧显存碎片问题。长序列训练时显存分配和释放频繁产生的碎片会让可用显存看起来非常小但实际空闲空间并不少。这时候可以在启动前设置PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:64强制 PyTorch 采用较小的分配粒度碎片问题会明显缓解。最后再分享一个我个人的习惯拿到任何一份 DeepSpeed 配置先不做大改动找一个小数据集和一个极小的 batch size 跑通再把 batch size 逐步往上加。配置不是越新越好stage 也不是越高越好。对大多数 1B 到 13B 的微调任务混合精度训练加 ZeRO Stage2 已经足够真想练几十 B 级别的模型再去考虑 Stage3 加 offload。先看清自己项目的瓶颈在显存还是通信再决定从配置里哪一项下手比你照抄网上的“终极配置”靠谱得多。