☰
DeepSpeed ZeRO-3与MoE组合实战:大模型训练显存优化与稀疏激活解析
2026/10/1 1:07:36 网站建设 项目流程

半夜盯着nvidia-smi发呆的场景,跑大模型的人多少都经历过:一块 80G 的 A100,加载 7B 参数做全量微调,优化器状态一开直接 OOM;换成 8 卡分布式训练,模型反而跑起来了,但通信开销把训练效率拖下去一大截。这正是 DeepSpeed ZeRO-3 与 MoE 训练要解决的两个核心问题。简单说,ZeRO-3 负责把模型参数、梯度、优化器状态全部拆开摊到多张卡上,MoE 则让模型在参数总量膨胀的同时不把计算量同步抬上天。把这两件事放在同一篇文章里讲,是因为今天训练大规模稀疏模型时,它们几乎总是成对出现。

这篇文章会从显存账本讲起,把 ZeRO-3 的切分机制、MoE 的路由原理、两者怎么配合落地,以及我在实际训练中踩过的坑一次说清楚。适合正在做多卡分布式训练、或者准备上手稀疏大模型的同学参考。不涉及数学推导,全部用算账和实操视角来讲。

1. 显存容量矛盾:为什么单卡永远装不下大模型

1.1 一个 7B 模型实际要吃掉多少显存

很多人对显存占用的理解还停留在“模型权重多大就要多大显存”,实际上全量训练时的显存需求比这大得多。以常见的 7B 参数模型为例,现在主流做法是混合精度训练,模型权重复制一份 BF16 用作前向反向计算,同时保留一份 FP32 的“主权重”,优化器用 Adam,还得维护一阶动量 m 和二阶动量 v。逐项算下来:

  • 模型权重(BF16):2 字节 × 7B = 14GB
  • 梯度(BF16):2 字节 × 7B = 14GB
  • FP32 主权重:4 字节 × 7B = 28GB
  • Adam 的 m 和 v:各 4 字节 × 7B,合计 56GB

四项加起来,单卡全量训练 7B 模型大约需要 112GB 显存。这就是为什么单张 80G 的 A100 在真正全量微调时一定会 OOM。顺着这个公式继续算,13B 模型需要超过 200GB,68B 模型超过 1TB。显存墙不是靠换一张更大的卡能解决的——你买到 1000GB 显存的服务器之前,必须先学会怎么把显存用高效。

我见过不少同学最先想到的解法是梯度累积、混合精度、或者干脆用 LoRA 只训练低秩 adapter。这些方案都能解决问题,但各有边界:LoRA 只适合微调场景,预训练和全量微调仍然要面对全参数优化;梯度累积只是把一次大步拆成小步,并没有减少单步所需的总显存。这时候就需要分布式训练里真正管用的手法——把状态切出去。

1.2 从 ZeRO-1 到 ZeRO-3 的演进逻辑

ZeRO 的核心思想其实特别朴素:一个分布式训练系统里有三种主要状态——模型参数、梯度、优化器状态。传统 Data Parallel(DDP)模式下,每张卡都保存一份完整副本,数据并行只切数据不切模型,显存浪费严重。

ZeRO-1 只切优化器状态,每张卡只负责一部分参数的 m、v 和主权重,通过通信让每张卡轮流得到全部状态完成更新,显存瞬间下降一半以上。ZeRO-2 在 ZeRO-1 基础上把梯度也切成多份,反向传播后每张卡只保存自己负责的那一段梯度。ZeRO-3 更进一步,把模型参数本身也切成碎片,前向计算需要哪层参数,就通过通信把那层参数临时“拼”出来,用完立刻丢弃。

从 ZeRO-1 到 ZeRO-3,显存占用一路走低,但通信量一直在涨。三者对比下来:

ZeRO 阶段切分内容单卡显存占用(7B 模型)通信开销变化
DDP不切约 112GB基准
ZeRO-1优化器状态约 70GB略增
ZeRO-2梯度 + 优化器状态约 56GB明显增加
ZeRO-3参数 + 梯度 + 优化器状态约 28GB大幅增加

从数字上能看出:ZeRO-3 对显存极度友好,代价是通信量成为主要瓶颈。实际训练中通常用通信重叠、梯度分桶、参数预取这些技巧把这些代价“藏”起来,而这些技巧的具体效果,和你的模型结构、卡间网络都有直接关系。

2. ZeRO-3 机制拆解:参数分片、动态全收集、梯度规约

2.1 切分对象与切分粒度

ZeRO-3 切的最小组件不是整个模型层,而是每个参数张量。一个 7B 模型有几百个参数张量,每个张量会被按 rank 数量切成 N 份,每个 rank 只持有原本那个张量的 1/N。这种细颗粒度的好处是负载均衡:模型里不同的层参数量差异不大,但 embedding 层和 classifier 头经常特别大,按张量切分更容易保持各卡之间的显存压力均匀。

这里有一个容易误解的点:ZeRO-3 并不会把所有参数均匀塞到每个 rank 上之后就完事。它需要在 forward 过程中按需触发 all-gather 操作,把当前计算层所需的参数重新组合成完整版本。某个参数张量如果位于第 12 层,那么在第 12 层计算之前,所有 rank 都要通信一次,把自己手里的 1/N 碎片拿出来拼成一个完整参数,供这一层的矩阵乘法使用。计算完之后,完整的副本又会被抛弃,各 rank 重新只保留自己的那 1/N 碎片。

生活化理解:这就像一组人合力读一本分成很多册的辞典,谁先查某个词条,就把包含这个词条的那一册借出来,其他册还给书架。只借当下要用的那一册,而不是把所有册都摊在桌上。代价是借还书(通信)频繁,好处是桌面(显存)上永远只有一本书占地方。

2.2 forward 和 backward 中的两类关键通信

ZeRO-3 在 forward 阶段的主要通信模式是 all-gather。每个参数张量在被使用前都要做一次全收集,让每个 rank 都拿到完整参数参与计算。到 backward 阶段,每个 rank 对自己拿到的完整参数算出一份梯度,但这部分梯度是“重复的、不完整的”——毕竟每个 rank 只负责该张量的一部分分片的更新。因此反向传播结束时要做一个 reduce-scatter:把各 rank 手上重复的梯度按分片位置求和后切回自己的那一段,最终每张卡只保存它负责的那一段梯度。

这就产生了一个实操经验:如果模型非常大,层数很多,all-gather 和 reduce-scatter 的次数会呈线性增长。默认情况下,每次参数收集都是同步的,也就是说 GPU 在等通信返回时经常处于空闲状态。解决办法是开通信重叠(overlap_comm)和参数预取(prefetch_bucket_size):在计算上一层的间隙,提前发起下一层参数的收集请求,让通信和计算并行。实测下来,模型在 64 卡规模下开不开 overlap,吞吐差距经常在 20% 以上。

另外,contiguous_gradients这个配置也很重要。它会把多个小梯度的 reduce-scatter 合并成一次大块通信,避免零碎的通信请求把通信链路打满。如果库里默认是关闭状态,我建议在 stage 3 下显式打开。

2.3 一份可直接使用的 ZeRO-3 配置示例

实际使用 DeepSpeed 时,stage 3 的配置一般在 JSON 里完成。下面这份配置我在训练 6B 稠密模型时验证过,可以直接抄:

{ "train_batch_size": 32, "gradient_accumulation_steps": 4, "optimizer": { "type": "AdamW", "params": { "lr": 3e-4, "betas": [0.9, 0.999], "eps": 1e-8, "weight_decay": 0.01 } }, "zero_optimization": { "stage": 3, "contiguous_gradients": true, "reduce_bucket_size": 5e7, "reduce_scatter": true, "overlap_comm": true, "prefetch_bucket_size": 5e7, "offload_optimizer": { "device": "cpu", "pin_memory": true } }, "fp16": { "enabled": true, "loss_scale": 0, "initial_scale_power": 16 }, "steps_per_print": 100 }

重点解释几个参数。reduce_bucket_size和prefetch_bucket_size控制每次通信携带的梯度/参数字节数,单位是字节数还是元素数要看你用的 DeepSpeed 版本,默认 5e7 左右是安全值。调太大,单个 bucket 可能触及通信缓冲上限;调太小,通信次数太多反而慢。

offload_optimizer把 Adam 状态放到 CPU 内存,这能让单卡显存再降一截,但 CPU 内存占用直接上涨,且优化器步骤变慢。建议先在纯 GPU 模式下跑通,确认性能瓶颈后再决定是否开 offload。注意如果开了offload_optimizer又把pin_memory设为 true,主机会锁定一部分内存用于 DMA 传输,这需要你在机器上预留足够物理内存,否则系统会开始 swap,速度断崖式下跌。

启动命令也很简单:

deepspeed --num_gpus 8 train.py --deepspeed ds_config.json

在train.py里只需要完成模型初始化后调用deepspeed.initialize,框架会自动接管参数的切分、梯度通信和优化器步骤。第一次跑通 stage 3 的常见 bug 是模型用了自己实现的参数共享(比如 word embedding 和 output layer 绑权值),这类共享参数在一个 forward 里被多个层引用,ZeRO-3 会对同一参数反复做 all-gather,导致性能极差或者显存反而更高。遇到这种情况,建议要么拆开共享,要么用 DeepSpeed 提供的特殊处理接口,而不是自己在 forward 里直接缓存param.data。

3. MoE 稀疏激活:参数翻倍,算力不翻倍

3.1 专家路由的直觉理解

MoE(Mixture of Experts)架构最近两年热度极高,从 Switch Transformer、Mixtral 到 DeepSeek-MoE,核心思路都一样:把 Transformer 里的 Feed-Forward Network 换成一个 Router 加多个专家 FFN。每个 token 并不经过全部专家,而是由 Router 给每个专家打分,选出得分最高的 top-k 个专家做计算。

这里的关键收益和代价需要分清楚。如果模型有 64 个专家,但每个 token 只激活 2 个专家,那么推理时单个 token 的计算量近似于一个 2 专家规模的 FFN,而不是 64 个专家规模的 FFN。这就是“参数总量可以很大,但计算量相对可控”的来历。

不过要注意一个容易混淆的点:训练和推理对显存的要求不一样。推理时模型权重已经被固定,加载到显存之后,只有被激活的专家会进入计算流程;但训练时要对每个专家计算梯度、更新参数,即使一个 token 只走两个专家,所有 64 个专家的权重依然需要被存储和更新。所以训练 MoE 时,模型参数“进显存”这件事是绕不开的,差别只在于进一张卡的显存,还是分散在多张卡的显存里。

自然语言里“moe架构要全部参数进显存吗”这个问题,完整回答是:权重必须存在于显存体系里,但通过 ZeRO-3 或专家并行,参数可以分摊到多张卡上,运行期按需收集,不需要任何单卡装下全部参数。

3.2 负载均衡:MoE 训练的第一个拦路虎

MoE 训练最容易出问题的地方是 Router 崩溃。所谓“崩溃”不是指程序报错,而是 Router 快速偏好少数几个专家,导致这几个专家被大量 token 挤爆,其他专家几乎闲置。专家之间负载严重不均衡时,模型退化得非常厉害,相当于大部分计算和参数都没有被有效利用。

解决方案在工程上已经比较成熟:引入负载均衡辅助损失,也就是 aux loss。基本思路是统计每个专家分配到的 token 比例 f_i,以及 Router 对所有 token 分配给专家 i 的平均概率 P_i,然后把 N × f_i × P_i 的和作为惩罚项加进总损失。当 token 分布均匀时这个值最小;某个专家聚集大量 token 时惩罚变大,训练会自动把 Router 往平均分配方向推。

DeepSpeed 的 MoE 配置里对应的就是load_balance_loss_weight。这个权重的设置非常需要拿捏:开大了 Router 会变得过于“和稀泥”,每个专家都差不多,稀疏激活退化成均匀计算,能力上不去;开小了负载均衡失效,一会儿又出现专家过载。我自己的经验是从 0.01 起步,观察训练日志里每个 step 的专家分配统计,再微调。

除了 aux loss,实际训练里还要注意noisy_gate_policy。常见选项是Jitter,给 Router 打分加上一点带噪声的扰动,让 token 分布更随机,一定程度上防止 Router 在训练初期就锁定在偏科状态。这个机制可以理解为给专家筛选加上一点“随机选择权”,避免所有 token 都挤向同一个最初得分高的专家。

3.3 DeepSpeed-MoE 的并行方式:专家并行和 All-to-All 通信

MoE 的专家层在 DeepSpeed 中通常采用专家并行(Expert Parallelism)。做法是把 N 个专家分配到不同 GPU 上,每个 GPU 只保存一部分专家。某个 token 被 Router 选到某个专家时,它所在的 GPU 需要把 token 的 hidden state 发给目标专家所在的 GPU,计算完成后再把结果传回来。这个过程就是 all-to-all 通信。

这也是 MoE 训练和纯稠密模型训练最大的差别:稠密模型的通信主要是梯度同步和数据并行,通信模式比较规律;MoE 则会在每次前向和反向中穿插多次 all-to-all。这些通信请求是随 token 分配动态变化的,很难静态优化,因此网络带宽会成为 MoE 训练的主要瓶颈之一。

如果你的集群是跨节点多机环境,我的建议是尽量把专家并行范围控制在单个节点内。也就是说ep_size不要跨过多的机器,因为节点间网络(比如以太网)的带宽通常低于节点内 NVLink,而 all-to-all 对带宽极其敏感。跨节点之后,训练吞吐的下降往往不是来自算力,而是来自这些动态通信请求的排队时延。

4. ZeRO-3 与 MoE 组合实战:配置、通信与调参

4.1 两种省显存机制如何叠加

把 ZeRO-3 和 MoE 放在一起训练,实际上是在同一份显存预算里做了两套分配策略:对于模型中的稠密层,使用 ZeRO-3 进行全参数切分;对于 MoE 层,使用专家并行把不同专家放在不同卡上,同时利用 ZeRO-3 把专家层内部的参数再做分片处理。这样一张卡上既不会保存全部稠密层参数,也不会保存全部专家参数,任意时刻显存里主要保留的是当前计算层需要的那份完整参数。

有人可能会问:既然 MoE 层已经用专家并行分散了,为什么还要 ZeRO-3 再切一遍?原因是 MoE 层之外还有大量注意力层、层归一化、embedding 层,这些稠密部分不会因为 MoE 而自动变小。如果只靠专家并行不动这些稠密层,显存中占比很大的注意力权重依然每张卡各有一份,仍然装不下很大的模型。所以实际 DeepSpeed-MoE 的推荐组合是:稠密层走 ZeRO-3 分片,专家层走 expert parallelism,两者共用一套通信规划机制。

关于训练时的显存占用,可以借助nvidia-smi dmon或 DeepSpeed 自带的日志观察每张卡的显存曲线。如果发现某一张卡明显比其他卡多占用几个 GB,很可能是专家并行分配时的ep_size设置不当,或者某一层专家被 Router 频繁调用但物理分配不平衡导致的缓存压力。

4.2 一份 DeepSpeed-MoE 的完整配置示例

下面是 DeepSpeed-MoE 训练中比较典型的一份配置文件,注释部分是我在实际项目里反复调出来的经验值:

{ "train_batch_size": 32, "gradient_accumulation_steps": 4, "zero_optimization": { "stage": 3, "overlap_comm": true, "contiguous_gradients": true, "reduce_bucket_size": 5e7, "prefetch_bucket_size": 5e7, "offload_optimizer": { "device": "cpu", "pin_memory": true } }, "moe": { "enable": true, "num_experts": [8], "top_k": 2, "ep_size": 2, "min_capacity": 4, "drop_tokens": false, "noisy_gate_policy": "Jitter", "load_balance_loss_weight": 0.01 }, "fp16": { "enabled": true, "loss_scale": 0, "initial_scale_power": 16 } }

逐个解释 MoE 相关参数:

  • num_experts: [8]:这是一个数组,表示每一层 MoE 层的专家数量。如果模型有 12 层要替换成 MoE,且有 8 个专家的层和 16 个专家的层,就写成[8, 8, 8, 8, 16, 16, ...]这样的序列。不能简单写一个数字,DeepSpeed 会按数组长度依次对应每一层。
  • top_k: 2:每个 token 路由到的专家数量。top_k=2 是常见选择,效果比 1 稳,计算开销也只多一个专家。
  • ep_size: 2:专家并行的卡数。这里设为 2 表示每个专家被复制并分配到 2 个 GPU 上(其实是把一个专家按卡再分片),适合检测连通性和通信开销。生产环境建议先按单节点卡数设置,不要盲目跨节点。
  • min_capacity: 4:每个专家每个 batch 的最小容量。容量就是专家最多能处理的 token 数,容量太小时很多 token 溢出;太大会让每个专家计算大量本不该由它处理的 token。
  • drop_tokens: false:当 token 超出专家容量时的策略。false 表示不丢弃而是重新路由到其他专家,效率更稳,但可能造成重复计算。真正大规模训练中我建议保持 false,避免数据丢失。
  • load_balance_loss_weight: 0.01:这个权重前面讲过,是负载均衡的惩罚系数。

4.3 训练流程中的关键实操建议

我强烈建议第一次接触 DeepSpeed-MoE 的团队,不要一上来就在大模型上完整跑。先构造一个“麻雀式”的小模型:两层 Transformer、每层 4 个专家、top_k 1,数据用随机生成的假数据。这样做的目的很简单——把通信链路、Router 初始化、all-to-all 的执行路径全部跑通,确认没有在分布式环境里埋雷,再扩大规模。

小模型跑通之后,有几个点值得专门记录:

第一,训练刚开始 200 步内的 loss 曲线波动是非常正常的。MoE 的 Router 在初期相当于随机打分,专家之间负载不均衡是常态。不要因为这一步的 loss 高就立刻判定模型废了,先观察 200 到 500 步,看 aux loss 是否进入下降趋势。

第二,drop_tokens和min_capacity的配合逻辑。如果drop_tokens设为 false,那么当某个专家容量超限时 token 会被重新分配,但这也意味着你实际上创建的 token-router 映射被打破了,会有额外通信。训练中遇到吞吐异常下降时,优先查这两项。

第三,不要忽略noisy_gate_policy的影响。在训练初期噪声可以让 Router 保持“探索欲”,但训练后期如果仍然保持强噪声,会让已收敛的 Router 不断被扰动。部分框架允许动态调整或衰减噪声,如果平台不支持,建议在固定迭代步数后手动切换策略。

5. 高频问题与排查实录

5.1 常见问题速查表

我在多个训练任务里反复遇到下面这些问题,整理成一张速查表:

问题现象根因分析排查与解决
开了 stage 3 仍然 OOM通信缓冲区、梯度累积缓冲、优化器状态没有真正被分片覆盖检查 reduce_bucket_size 是否过大;确认 optimizer 是否走 DeepSpeed 接管;考虑开启 offload_optimizer
训练吞吐大幅下降通信次数过多,all-gather/ reduce-scatter 没有被隐藏打开 overlap_comm、contiguous_gradients;确认 prefetch_bucket_size 设置合理;检查模型是否在 forward 中重复使用同一参数
部分专家显存占用异常高专家并行分配不均,或 Router 严重偏科检查 ep_size 设置;用训练日志观察专家 token 分配比例;调大 load_balance_loss_weight
训练loss震荡不收敛Router 对负载分配失去控制,噪声过大降低学习率;调小 noisy_gate_policy 的扰动幅度;检查 aux loss 权重是否过高
开 offload 后系统卡死/极慢CPU 内存不足、swap 被触发关闭 pin_memory;降低 offload 比例;升级物理内存;改用 NVMe offload
all-to-all 通信消耗时间过长专家分布在跨节点,网络带宽不足缩小 ep_size 到单节点内;优先用 NVLink 连接专家通信;减少 top_k 或专家层数量

5.2 几类典型问题的现场记录

以 stage 3 下仍然 OOM 举例。很多人以为开了 ZeRO-3 就一定不会显存爆炸,实际上 ZeRO-3 只负责把参数、梯度和优化器状态切出去,并不负责限制前向计算时的激活值。激活值显存在大 batch 下照样可以吃掉几十 GB。排查方式是把train_batch_size和gradient_accumulation_steps组合调整:先设train_batch_size=1观察显存占用,如果显存依然打在 80G 边缘,明显是激活值或者通信缓冲的问题;如果降下来了,就是批量大小相关的问题。

另一个常见坑是模型 executor 或 loader 在 ZeRO-3 环境下被重复执行。比如某些库在model.generate时会对权重自动做缓存,这个缓存可能完全绕开 DeepSpeed 的分片管理,导致显存立刻上升。遇到这种情况,可以检查模型对象的子模块是否在每次 forward 里通过.to(device)移动参数,这种“手动搬运”会破坏 ZeRO-3 的动态参数收集逻辑。

MoE 训练中比较隐蔽的问题是 token 丢包。min_capacity设置过小时,大量 token 超出专家容量,drop_tokens=false的语义会把这些 token 重新分配,但有些版本的框架在 re-route 时的随机性可能导致同一个 token 被多个专家重复计算。这类问题不会让程序报错,但会让 loss 曲线出现一些奇怪的“跳跃”,最终模型精度也偏低。建议在训练脚本里定期打印每个专家的 token 计数,用统计数据而不是手感来判断是否出现丢包。

还有一点是关于负载均衡损失的监控。不要只盯总 loss,一定要单独把 aux loss 从训练日志里拉出来看曲线。如果 aux loss 从第一步起就一直非常低,可能并不是因为路由均衡,而是因为权重太低没有被学习到。反过来,如果 aux loss 一直降不下来,就要怀疑 Router 本身的结构或初始化有问题。

5.3 组合场景下的实测心得

把 ZeRO-3 和 MoE 一起上之后,最直观的感受是“显存压力瞬间释放,通信压力集中爆发”。我在一次实验里用 8 卡 H800 训练 16B MoE 模型,开 ZeRO-3 之前连参数都加载不进去;开完之后显存占用非常健康,但训练吞吐一上来就比同规模稠密模型低了不少。用性能剖析工具打点之后发现,绝大多数时间花在了 all-to-all 通信的等待上,而不是计算本身。这说明显存问题确实被解决了,但接下来要优化的是通信拓扑。

结合以上经验,再给一个比较实用的建议:MoE 模型训练时的ep_size选择,一定要同时考虑显存上限和通信拓扑。如果显存非常紧张,可以适当增大ep_size,让每个专家分到更多卡的显存,但要注意跨节点通信成本;如果显存还有余量,优先降低ep_size,减少通信频率,你会看到吞吐立刻改善。

最后再说说调参顺序。训练 MoE 的过程中,变量太多,一次只改一个参数是黄金法则。我通常的顺序是:先固定num_experts和top_k,然后调load_balance_loss_weight观察 aux loss 和专家负载曲线;再调min_capacity和drop_tokens处理 token 溢出;最后才碰noisy_gate_policy这类影响收敛动态的参数。每次都记录日志对比,不要凭感觉。我对 ZeRO-3 和 MoE 最深的体会是:这两个机制都极其依赖训练时的可观测性,日志里多打印专家分布和通信耗时,比任何理论分析都管用。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询