ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

ZeRO详解:从显存账单到三阶段分片,彻底搞懂DeepSpeed显存优化

ZeRO详解:从显存账单到三阶段分片,彻底搞懂DeepSpeed显存优化 上周组里做技术分享我把训练 7B 模型时的显存账单拉出来很多人第一反应是算错了112GB。要知道这还只是模型参数、梯度和优化器状态三个大头还没算激活值和临时 buffer。这桩“显存去哪了”的公案追到底就绕不开微软那篇 ZeRO 论文——Zero Redundancy Optimizer。这篇论文我前后读了三遍第一遍觉得只是把数据并行的状态切开了第二遍发现三阶段的通信开销大有讲究第三遍才注意到 ZeRO-R 这些“残差状态”的优化。这篇细读笔记我打算把每一遍读到的重点都摊开来讲顺便附上我踩过的一些坑。适合刚接触分布式训练的同学也适合那些被 OOM 折磨过、想彻底搞懂 DeepSpeed 后端到底在做什么的人。1. 追根溯源训练一个模型显存是怎么被吃掉的1.1 从单个参数算起一条 16 字节的显存账单要理解 ZeRO先得搞清楚一个朴素问题混合精度 Adam 训练一个 Transformer 模型单个模型参数到底占多少显存很多人会脱口而出“2 字节因为用 fp16 存权重”。但训练不是推理权重之外还有梯度和优化器状态。在主流训练配置下一个模型参数在显存里的完整开销是这样拆的状态精度每个参数占用fp16 权重副本前向/反向用FP162 字节fp16 梯度FP162 字节fp32 master 权重真正的更新对象FP324 字节Adam momentum一阶动量FP324 字节Adam variance二阶动量FP324 字节加起来正好 16 字节。注意这里有个特别容易被忽略的点Adam 的更新发生在 fp32 上。每一步更新用的是 fp32 master weight而不是 fp16 权重。所以优化器状态不是大家以为的 8 字节而是 12 字节——4 字节 master 权重加 4 字节 momentum 加 4 字节 variance。拿 70 亿参数模型来算112GB。随便一张 A100 40GB 或 H100 80GB 都放不下这就是大模型训练显存压力的根源。1.2 数据并行为什么让成本成倍叠加既然单卡放不下最直接的办法就是数据并行——每张卡一份完整模型副本各自吃一个数据切片。这确实解决了“模型放不进单卡”的问题吗没有它只是让每张卡都能独立完成训练。代价是模型状态被复制了 N 份N 张卡就是 N 份冗余。8 卡训 7B 模型模型状态总量是 112GB × 8 896GB其中 7/8 是纯冗余。每张卡还要吃下完整的 112GB 模型状态。数据并行训练时反向传播结束后要把所有卡上的梯度做一次全局 all-reduce保证每张卡拿到的梯度是一致的然后各自独立更新参数。这套流程本身没错问题在于所有卡上的参数、梯度、优化器状态在任意时刻都是完全相同的副本。既然完全相同为什么每张卡都要完整保存一份1.3 一个容易被忽略的推导前提我在团队分享时被追问过为什么不能用 fp16 的权重直接做 Adam 更新答案是数值稳定性。fp16 只有大约 3 位有效十进制精度Adam 更新量往往比权重小几个数量级直接累加在 fp16 上会直接丢精度。所以才需要 fp32 master weight 参与更新fp16 权重只是前向/反向计算的载体。这个前提直接决定了后面 ZeRO 三阶段的显存公式怎么推导16 字节里参数占 2梯度占 2优化器状态占 12。优化器状态才是真正的大头占了 75%。2. ZeRO 想明白的一件事显存冗余是最大的浪费源2.1 核心思想从“每卡都存全部”到“全体合起来只存一份”ZeRO 的题目全称是 Zero Redundancy Optimizer关键词不是“分片”而是“redundancy”。论文的核心观察非常简单直白数据并行训练中参数、梯度、优化器状态在不同 rank 之间完全一致这些一致就是冗余。ZeRO 的思路是既然这些状态在每张卡上都是同样的内容那就别再每人存一份了。把状态按 rank 切成 N 份每张卡只保存自己负责的 1/N需要的时候通过集合通信把完整状态汇聚出来。用通信换显存把显存占用从“单卡必须放下完整模型”变成“所有卡合起来放下一个模型”。打个比方以前是一车人每人背一个同款行李袋里面装的还是完全相同的东西ZeRO 做的事情是把一整套行李拆成 N 份分给 N 个人分别背着谁需要某件物品找一个离得近的人互相广播一下。省下的是 N-1 份重复行李的空间花的是大家互相喊话的时间。2.2 “分片”和“切分”的本质差别很多人第一次接触 ZeRO 会把它和模型并行混淆尤其是到了第三阶段参数分片之后两者看起来都是在“把模型拆开”。但它们的本质完全不同。张量并行 / 流水线并行切的是计算。每张卡只负责计算模型的一部分层或一部分张量切片计算图和激活值在每个 rank 上是不完整的。ZeRO不切计算。每个 rank 仍然处理不同的 micro-batch计算图和激活值在逻辑上是完整的只是参数、梯度、优化器状态的存储被切开了。用表格对比更清晰维度数据并行 DDP张量并行ZeRO计算分布各 rank 算完整模型的不同数据各 rank 算模型的一部分各 rank 算完整模型的不同数据参数存储每 rank 一份完整副本每 rank 一份切片每 rank 一份切片显存压力单卡必须能放下完整模型单卡只需放下切片单卡只需放下切片通信频率每步一次梯度 all-reduce每层前向/反向多次通信每步 1~3 次集合通信ZeRO 的本质是“存储层面的并行”不改变计算图。这也是为什么它能直接嵌进现有的 PyTorch 训练流程改动成本远低于重写一层张量并行的模型代码。2.3 为什么会值得通信成本与显存收益的对比很多人的第一反应是把状态拆开之后训练时还要 gather 来 gather 去通信开销会不会爆炸ZeRO 的底气在于 GPU 互联带宽和显存容量之间的“剪刀差”。以 A100 为例单卡 NVLink 带宽约 600GB/s8 卡全互联。一次 7B 模型的全量梯度 all-reducefp16 梯度大约 28GB 数据单机内也就几百毫秒到一秒量级。这点时间换来的是每卡 112GB 降到几十 GB能让原来根本跑不动的模型在同样硬件上跑起来。显存是稀缺资源带宽相对没那么稀缺这笔交易在大多数场景下非常划算。3. 三个阶段逐一拆解Pos、Pg、Pr 到底做了什么ZeRO 论文把优化拆成三个阶段分别处理优化器状态分区、梯度分区、参数分区。论文里给它们起了专门的名字Posoptimizer state partition、Pggradient partition、Prparameter partition。逐级往下显存越来越省通信模式也在悄悄改变。3.1 第一阶段Pos优化器状态分片最划算的一步第一阶段只动优化器状态。为什么优化器状态可以分片因为 Adam 的 momentum 和 variance 完全由该参数对应的梯度更新而每个 rank 在数据并行中拿到的梯度是全量的——经过 all-reduce 之后大家的梯度一致。既然每个 rank 都有完整梯度优化器状态的更新就可以各算各的不需要知道其他 rank 的优化器状态。显存从 16 字节/参数降到 4 12/N 字节/参数其中参数 2 字节 梯度 2 字节不变优化器状态 12 字节被 N 个 rank 瓜分。N8 时是 5.5 字节/参数7B 模型对应每卡 38.5GB。这一步的代价几乎为零梯度 all-reduce 的通信量跟 DDP 完全一样没有增加任何通信负担。所以我在实际项目中总是推荐先开 stage 1改动最小、风险最低而且效果立竿见影。3.2 第二阶段Pg梯度也按需分片反向传播像接力赛第二阶段把梯度也分片了。阶段一的通信模式还是 DDP 那样的全局 all-reduce每张卡都要拿到完整梯度。但阶段二之后每个 rank 只需要自己负责的那 1/N 优化器状态对应的梯度切片所以全局 all-reduce 可以换成 reduce-scatter——各 rank 把自己的梯度切片 reduce 到对应的 owner 上其他切片不用给自己送。这也是 ZeRO 通信设计的巧妙点不是所有梯度都需要全部汇聚到每张卡只需要汇聚到它的“消费者”那里。梯度分片的合法性来自第一阶段优化器状态已经按 rank 分片那么梯度的唯一消费者就是对应的优化器状态切片。显存变成 2 2/N 12/N 字节/参数N8 时约 3.75 字节/参数7B 模型对应每卡 26.25GB。通信量反而比 DDP 略降因为 reduce-scatter 本身比全量 all-reduce 要省一部分数据。3.3 第三阶段Pr参数也不再每人一份彻底打破单卡上限第三阶段把最后的参数副本也分片了每个 rank 只保留 1/N 的模型参数。显存降到 16/N 字节/参数N8 时 2 字节/参数7B 模型对应每卡 14GBN16 时 7GBN32 时 3.5GB。这意味着模型规模可以随卡数线性扩展单卡显存上限不再是硬约束。代价是通信模式变了。前向计算时需要把该层权重 all-gather 到所有 rank反向计算梯度时还要再 all-gather 一次加上梯度 reduce-scatter通信量大约是 3 倍的模型大小比 DDP 高约 50%。这也是 stage 3 在通信拓扑差比如跨节点、网络带宽低时吞吐明显下降的原因。3.4 一张表看完全部阶段阶段每参数显存公式每参数显存N87B 模型每卡通信量相对模型大小DDP16 字节16 字节112GB约 2 倍Stage 1Pos4 12/N5.5 字节38.5GB约 2 倍无增加Stage 2Pg2 14/N3.75 字节26.25GB约 2 倍略降Stage 3Pr16/N2 字节14GB约 3 倍这里说的“模型大小”指的是 fp16 参数所占的字节数。通信量以每次迭代为单位估算实际开销还受 bucket 大小、通信与计算重叠程度的影响。3.5 别忘了 ZeRO-R 和 Offload论文后半部分还有 ZeRO-R处理的是“残差状态”——激活值、临时 buffer 和显存碎片。激活值分区activation partition就是把 Transformer 里巨大的中间激活也切到多卡恒定 buffer 分区消除冗余碎片整理则通过预先分配固定大小的内存块避免训练过程中显存碎片化导致 OOM。实际操作中我见过不少项目开了 stage 3 之后仍然 OOM查下来根本原因是激活值没做 checkpointing或者是显存碎片太多。所以在读这篇论文时别把注意力全放在三阶段上ZeRO-R 里的激活管理才是很多边缘情况的解药。4. 纸上谈兵结束用 DeepSpeed 实测一个 7B 模型4.1 一套能直接改的 DeepSpeed 配置理论说得再漂亮不如跑一遍。我在 8×A100 40GB 环境上用一个 7B 模型做过对比配置长这样{ train_batch_size: 8, gradient_accumulation_steps: 1, optimizer: { type: AdamW, params: { lr: 1e-5, weight_decay: 0.01 } }, zero_optimization: { stage: 2, allgather_partitions: true, reduce_scatter: true, overlap_comm: true, contiguous_gradients: true }, fp16: { enabled: true, loss_scale: 0 } }启动命令很简单deepspeed --num_gpus8 train.py --deepspeed_config ds_config.json有几个配置项值得解释一下overlap_comm: true让梯度通信和反向计算重叠这是吞吐优化的大头。contiguous_gradients: true把梯度放进连续内存块避免碎片同时能提高通信效率。loss_scale: 0让 DeepSpeed 动态调整 loss scale混合精度训练必备。4.2 实测下来看到的现象和数字我在 7B 模型上的观测大致如下不同模型和框架版本会有出入仅供量级参考配置每卡显存模型状态部分实测每卡峰值显存每步吞吐相对 DDPDDP112GB直接 OOM无法训练Stage 138.5GB约 44GB约 98%Stage 226.25GB约 33GB约 96%Stage 314GB约 22GB约 85%峰值显存比模型状态大是因为还有激活值、通信 buffer 和框架自身的开销。stage 2 到 stage 3 的吞吐下降比较明显主要来自前向/反向两次参数 all-gather。单机 NVLink 环境下能接受跨节点环境里这个差距会进一步拉大。4.3 我在实践中踩过的几个坑坑 1stage 3 下不能随便访问model.parameters()在 stage 3 下参数的存储是分片的某些 rank 上根本没有完整权重。我在写自定义 loss 时需要拿权重做约束直接访问model.parameters()拿到的只是一堆碎片。正确做法是用 DeepSpeed 提供的作用域from deepspeed.zero import GatheredParameters with GatheredParameters(model.parameters()): # 此时参数是完整的 do_something_with_full_params()坑 2checkpoint 里的权重是碎片直接 load 会炸stage 2 以上保存的 checkpoint每张卡只保存自己那份状态。想转成完整的 PyTorch 权重DeepSpeed 会生成一个zero_to_fp32.py脚本python zero_to_fp32.py --input_dir ./checkpoint --output_file ./pytorch_model.bin我之前偷懒直接拿 rank 0 的模型权重去评估结果 loss 不对模型根本没收敛查了半天才反应过来是碎片没合并。坑 3gradient checkpointing 和 stage 3 一起用时需要注意重复 gather开 gradient checkpointing 后前向过程会丢弃部分激活值反向时重算。重算过程中会再次触发参数 gather。如果代码里手动管理了参数收集和释放容易在反向时重复 gather 甚至 gather 了已被释放的切片。我遇到过一次奇怪的死锁最后是把手动管理的GatheredParameters改成直接依赖框架自动管理问题才消失。坑 4跨节点训练网络拓扑比想象中更敏感单机内 NVLink 跑 stage 3 很顺跨节点走 RoCE 后吞吐直接掉了不少。原因就是参数 all-gather 太依赖节点间带宽。我的建议是跨节点场景优先 stage 2或者给 stage 3 调大allgather_bucket_size减少通信次数。调参幅度看具体网络没有通解。坑 5DataLoader 的采样器不能被忽略开 ZeRO 之后数据并行的语义没变但 stage 3 里的广播逻辑会对数据分布敏感。我用过不均匀的采样器结果某些 rank 数据量明显不同损失震荡得厉害。后来一律用带shuffleTrue的DistributedSampler问题就消失了。这个坑跟 ZeRO 本身关系不大但容易在折腾分布式时被忽略。5. 读图纸之外的思考ZeRO 的“分片”边界在哪里5.1 ZeRO 不负责计算的划分ZeRO 虽然叫“零冗余优化器”但它不是万能的。有一个非常关键的边界它不切计算。也就是说如果模型单层都放不进一张卡ZeRO 就救不了你。前向和反向时stage 3 需要把某一层的参数 all-gather 到各个 rank 上参与计算。如果这一层的参数本身超过单卡显存gather 时就会 OOM。比如超大规模 MoE 的 expert 层或者超级大的 embedding 表都属于这种情况。这时候只能上张量并行把单层内切分到多卡或流水线并行按层切分计算。三种技术并行不悖实际大规模训练里往往是 ZeRO 张量并行 流水线并行叠加使用。5.2 什么时候用哪个 stage我的选择经验基于这些实测和踩坑我给团队的建议是场景推荐配置理由单机 4~8 卡NVLink 互联Stage 2吃满单机带宽吞吐损失小显存收益大单机 8 卡以上Stage 2 或 Stage 3卡多时 stage 3 的显存收益更有吸引力跨节点、节点间带宽一般Stage 1 或 Stage 2通信压力小吞吐稳定单卡或双卡跑超大模型Stage 3 Offload Activation Checkpointing以吞吐换显存目标只是“能跑起来”具体到每个项目还要看模型大小、batch 需求和激活值规模。我的经验是先开 stage 2 试一轮观察显存和吞吐如果显存仍紧再上 stage 3 并配合激活重算如果 stage 3 吞吐损失不可接受再考虑张量并行或 offload。没有银弹但至少可以按这个顺序排查。6. 三遍细读之后我留下的几个观念变化第一遍读 ZeRO我以为它只是一个“把 DDP 状态切开”的工程技巧。第二遍才意识到通信量和显存之间的 trade-off 设计才是精髓——用几乎不增加的通信成本换来 4 倍显存收益这是非常漂亮的权衡。第三遍再看发现真正难的是把论文里的思想落到 DeepSpeed 的配置和训练流程里那些隐藏的坑碎片 checkpoint、参数 gather、网络拓扑影响往往比论文本身更折磨人。现在我看任何分布式训练方案第一反应不再是“加卡”而是先问一句“这里面的状态冗余有没有消干净”。这一观念转变就是从 ZeRO 论文里来的。如果非要给刚接触分布式训练的朋友一个建议我的建议是先把这篇论文的 16 字节账单和三阶段公式吃透再去碰那些花哨的框架功能你会少走很多弯路。
返回列表