DeepSpeed ZeRO-3与MoE训练:显存优化、通信开销与配置实战
发布时间:2026/10/2 5:07:57来源:尧图网络
搞过大模型训练的人迟早会撞上两座山显存不够用参数放不下。DeepSpeed 的 ZeRO-3 是分布式训练里最常用的显存破局方案之一而 MoE 架构又是当前大模型规模竞赛里绕不开的路Mixtral、DeepSeek-V3 这些名字大家都听过。这篇文章我尽量用说人话的方式把 ZeRO-3 的原理、MoE 训练里的显存和负载问题、两者怎么配合以及 DeepSpeed 安装配置中最容易翻车的环节一次讲清楚。你不用提前懂分布式系统只要用过 PyTorch、训过模型读完就能照着调。1. 先搞明白 ZeRO-3 到底省的是哪块显存1.1 数据并行为什么浪费很多人第一次接触多卡训练用的都是DistributedDataParallel也就是数据并行。思路很简单每张卡上放一份完整的模型副本每张卡吃不同的 batch前向反向各自算完后把梯度做一次 all-reduce 同步再各自更新参数。这个方案在小模型时代非常好用但一到几十亿参数就顶不住了。原因不是“模型太大装不下”这么一句话能概括的而是一个反直觉的事实训练阶段模型参数的显存开销只是冰山一角真正把显存吃干榨净的是梯度、优化器状态这些“训练附属物”。打个比方数据并行就像每个学生都去复印一整本教材自己拿回家划重点。十个人就复印十本内容完全一样纯属浪费。但大模型训练时大家又不得不“人手一本”因为每张卡都要独立完成前向和反向计算没有完整参数就动不了手。ZeRO 的核心思路就是把这份冗余干掉。它全名叫 Zero Redundancy Optimizer翻译过来就是“零冗余优化器”。ZeRO-1 先切掉了优化器状态的冗余ZeRO-2 再切掉梯度的冗余到了 ZeRO-3模型参数本身也被切碎分到各卡上每张卡只保留一小片拼图真正做到了“教材拆成章节人手几页需要时互相借”。1.2 一个参数在训练时到底占多少字节要理解 ZeRO-3 省了多少先得学会算“一个参数在训练时占多少显存”。这里有一个非常经典的估算大家可以当成口诀背下来模型参数本身如果是 FP16 存储每个参数占 2 字节梯度反向传播时要算梯度一般也用 FP16 存储每个参数占 2 字节优化器状态这个是最大的头。以最常用的 Adam 优化器为例它需要维护一份 FP32 的 master 权重副本、一份动量momentum和一份方差variance加起来是 4 4 4 12 字节。一个参数在训练时粗略算下来要占 2 2 12 16 字节。如果用了 offload 或者不同精度数字会有浮动但量级就是这样。拿一个 7B 参数模型来算70 亿 × 16 字节 ≈ 112GB。单张 A100 80GB 根本放不下更别说还有 batch 数据和中间激活值。所以很多 7B 模型的“训练需求”其实不是算力而是纯显存。再看 ZeRO-3 怎么救场。假设用 32 张卡方案模型参数梯度优化器状态单卡权重侧显存7B 模型普通数据并行14GB14GB84GB112GB 每卡全量ZeRO-114GB14GB84/32≈2.6GB约 30.6GBZeRO-214GB14/32≈0.4GB84/32≈2.6GB约 17GBZeRO-314/32≈0.4GB14/32≈0.4GB84/32≈2.6GB约 3.5GB这张表还不含激活值和通信 buffer但分量已经很清楚了普通数据并行在 32 卡上每张仍要扛 112GB 权重相关开销而 ZeRO-3 把这块压到了 3.5GB 左右剩下的显存全部可以用来放大 batch、增大输入序列这才是它能撑起超大模型训练的根本原因。1.3 ZeRO-3 的“通信换显存”逻辑ZeRO-3 不是没有代价代价主要落在通信上。因为模型参数被切到所有卡上前向计算某一层时需要先做一次 all-gather 把这一层的参数临时凑齐反向传播时为了算梯度又要把这一层参数再凑齐一次。可以把它理解成以前你手边永远有全套教材随查随用现在教材被拆了每次查资料前都得打电话让同学把相关章节拍照发过来用完就删。省了书柜空间但电话费暴涨。为了把“电话费”降下来ZeRO-3 做了不少优化。常见的包括overlap_comm: 让通信和计算重叠前向算上一个算子时后台提前拉取下一层的参数stage3_prefetch_bucket_size与stage3_max_live_parameters: 控制预取的粒度平衡“多拉一点省通信”和“少留一点省显存”stage3_max_reuse_distance: 控制参数在显存里保留多久如果某个参数很快还会再用就暂时不释放。这些参数在后面的配置文件里都会用到。理解了“通信换显存”的底层逻辑遇到训练慢的问题时你就知道该往哪个方向查了。2. MoE 训练为什么是另一套玩法2.1 MoE 的“多而不算”是怎么做到的MoE 全称 Mixture of Experts翻译成“混合专家”。它把一个 transformer 里的 FFN 层替换成一堆专家网络每个 token 只激活其中一小部分专家。比如经典的 MoE 配置是 8 个专家、每次激活 top-2。也就是说模型结构里明明有 8 份 FFN 权重但每来一个 token门控网络router只选出 2 个最“有把握”的专家来计算其余 6 个专家原地待命。这里要分清一个概念MoE 省的其实是“计算量”或者说“FLOPs”而不是“参数量”。因为参数确实全都存着但推理时每个 token 只用一小部分参数算。所以 MoE 模型可以在总参数量达到几十亿上百亿的情况下单 token 的计算成本只相当于一个小得多的 dense 模型。这也是为什么大家叫 MoE “稀疏模型”它稀疏在“每次激活的路径”上而不是“参数文件稀疏”。文件下载下来一样是几十上百 GB。2.2 训练 MoE所有参数真的都要进显存吗这个问题非常经典也是很多人在安排训练资源时算错账的地方。分两种情况说。推理阶段确实不需要所有参数进显存。每个 token 只需要把 top-k 专家的权重装进显存进行计算用不到的专家可以留在 CPU 内存或磁盘上按需加载。训练阶段情况有所不同。虽然某个 batch 里只有部分专家被激活但反向传播时优化器要更新所有专家的参数——没被激活的专家梯度为 0但优化器状态里的动量项和方差项依然需要保留、迭代。也就是说专家参数即使没参与计算它们对应的优化器状态依然占着显存。所以在训练时MoE 模型并不会因为“每个 token 只用 2 个专家”就把显存需求砍到原来的四分之一。你仍然需要为全部专家参数准备好存储空间、梯度空间和优化器状态空间。具体怎么省靠的是下面的 ZeRO-3 和专家并行而不是指望稀疏激活。2.3 负载不均几乎所有 MoE 训练都会遇到的坑MoE 还有个逃不掉的经典问题负载不均衡。门控网络天然有“赢者通吃”的倾向训练到后面某些专家会被大量 token 选中忙得不可开交另外一些专家则长期处于“失业”状态。这不是小问题。如果少量专家被过度使用首先会导致这批专家的 on-chip 显存和算力瓶颈更严重的是欠学的专家会进一步得不到梯度信号形成恶性循环最终模型的有效表达能力下降甚至训练直接崩溃。解决办法是在损失函数里加入辅助负载均衡损失aux loss。思路是统计每个专家在当前 batch 中分到的 token 比例再和 router 输出的平均概率做点积两头越接近损失越小两头差距越大损失越大。这样每当门控想偷懒集中选择少量专家时损失函数就会拽它一把。具体代码我放到后面第 5 节那里可以直接抄。3. ZeRO-3 和 MoE 怎么组合收益最大3.1 统一分片和专家并行哪个更合理现在核心问题来了ZeRO-3 和 MoE 到底怎么配先说一种实现直接把 MoE 模型当成普通大模型所有参数包括专家参数都丢给 ZeRO-3 统一分片。这种做法在 Transformers 里跑 Mixtral 这种现成架构时非常省事ZeRO-3 会自动把每一层的参数切到所有卡上前向计算时临时 all-gather 整层参数。但这里有个隐忧MoE 的专家是稀疏激活的本来每个 token 只需要一小撮专家可 ZeRO-3 做 all-gather 时是把整层参数全部拉到当前卡上通信开销较大。换句话说统一分片方案“能用”但不是 MoE 场景下的最优解。更专业的方案叫“专家并行”Expert Parallelism。思路是MoE 层里把不同的专家显式放到不同 GPU 上例如 8 个专家分布在 8 张卡每张卡只负责一个专家。token 经过 router 选择后通过 all-to-all 通信被发送到对应专家所在的 GPU 上去计算。这样本来一次只激活 2 个专家就只需要通信那 2 个专家的数据而不是把整层参数都拉一遍。DeepSpeed 官方对 MoE 训练也更推荐“专家并行 ZeRO-1/2 组合”的方式专家层走专家并行attention 或其他 dense 层走数据并行加 ZeRO 分片。讲这句话的意思是很多人在群里问“ZeRO-3 能不能训 MoE”答案是可以但如果追求性能建议用 DeepSpeed-MoE 的 expert parallel 路线而不是硬套 ZeRO-3。3.2 用一个稀疏模型算笔显存账为了让大家有个直观感觉我们来算一笔账。假设一个稀疏模型总参数量是 40B其中包含了多个专家。训练时每个参数按 16 字节估算总显存需求大约 40 × 16 640GB。如果用 16 张 A100 80GB每张卡理论上只需要承担 640 / 16 40GB 的权重相关显存剩下 40GB 留给激活值和通信 buffer完全可行。这种情况下不管用 ZeRO-3 还是专家并行单卡都不需要装下完整的 40B 参数。反过来如果这个模型是 dense 的 40B用普通数据并行每张卡上来就要扛 640GB16 张卡每张依旧 640GB完全跑不动。这就是稀疏模型加参数分片的双重优势参数量大但通过分片大家分摊计算量又因为稀疏被控制住。两者叠加MoE 模型才能在实际集群上跑得起来。3.3 组合训练时的配置与取舍如果你决定用 DeepSpeed 训练一个 MoE 模型有几个配置层面的取舍要提前想清楚。如果直接用 HuggingFace Transformers 里的 Mixtral 等模型最稳妥的做法是把整个模型交给 ZeRO-3负载均衡损失由模型内部自己处理你只需要在 JSON 配置里把zero_optimization.stage设为 3。如果自己动手实现 MoE 层建议先做专家并行再叠加 ZeRO-1 或 ZeRO-2 来管 dense 层参数。DeepSpeed 在 MoE 场景下有一个独立的moe配置段会告诉我模型里哪些层是专家层从而走专门的通信路径具体字段建议以官方文档的 DeepSpeed-MoE 章节为准。如果既想省显存又想省事可以先用 ZeRO-3 跑通一个小规模 MoE 实验再逐步迁移到专家并行。不要一上来就追求最优结构否则调试难度会叠加。4. DeepSpeed 安装与 ZeRO-3 配置实操4.1 安装阶段最容易翻车的几个点很多人在安装 DeepSpeed 时就卡住了网上最常见的错误集中在编译环节。这里把高频坑列一遍都是我实测过的。第一先确认 PyTorch 版本和 CUDA 版本匹配。torch.cuda.version显示的 CUDA 版本最好和系统里的nvcc --version保持一致。不一致时编译会指向错误的 CUDA 头文件报一堆莫名其妙的 error。第二装一个ninja。DeepSpeed 默认走 JIT 编译算子有了 ninja 之后编译会快非常多没有它可能要编译半小时以上而且容易超时中断。可以直接pip install ninja。第三如果你只是先想跑起来不要急着编译全部算子。可以设置DS_BUILD_OPS0让 DeepSpeed 先用纯 Python 模式跑基础功能和 ZeRO-3 都能用。等到真正要用 CPU Adam、NVMe offload 这类高性能算子了再重新编译。pip install ninja setuptools wheel pip install deepspeed # 使用国内 PyPI 镜像源可以明显提速第四遇到aio_setup.h not found、libaio相关的报错时多半是没装 libaio 内核库。如果不打算使用 NVMe offload可以直接用DS_BUILD_AIO0跳过DS_BUILD_AIO0 DS_BUILD_OPS0 pip install deepspeed第五装完后可以用一行命令验证安装是否正常ds_report它会输出 DeepSpeed 版本、可用算子列表、CUDA 信息等。看到各项都 OK再去做训练后续排错会简单很多。4.2 一份可以直接跑的 ZeRO-3 配置模板下面这份配置我经常拿来当模板换成自己的路径和显存参数就能跑{ train_batch_size: 16, gradient_accumulation_steps: 4, train_micro_batch_size_per_gpu: 1, fp16: { enabled: true, auto_cast: true, initial_scale_power: 16, loss_scale_window: 1000, hysteresis: 2, min_loss_scale: 1 }, zero_optimization: { stage: 3, offload_optimizer: { device: cpu, pin_memory: true }, offload_param: { device: cpu, pin_memory: true }, overlap_comm: true, contiguous_gradients: true, reduce_bucket_size: auto, stage3_prefetch_bucket_size: auto, stage3_param_persistence_threshold: auto, stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9, gather_16bit_weights_on_model_save: true, round_robin_gradients: true }, gradient_clipping: 1.0, steps_per_print: 100, wall_clock_breakdown: false }内外两层zero_optimization里我的经验是offload_optimizer和offload_param同时打开后显存能大幅下降但 CPU 内存和 PCIe 带宽会成为瓶颈适合单机多卡或模型极大但不追求速度的场景stage3_max_live_parameters越大参数在显存里留得越久通信次数越少但显存占用越高。遇到通信瓶颈可以先调大它遇到 OOM 就先调小gather_16bit_weights_on_model_save务必设为true否则保存 checkpoint 时每张卡只存自己的分片weights 文件缺一大块后面加载模型会特别痛苦。4.3 用 Transformers Trainer 接上 ZeRO-3用 HuggingFace Transformers 训练时最简单的集成方式是直接在 Command Line 里加--deepspeeddeepspeed --num_gpus8 train.py \ --model_name_or_path gpt2 \ --per_device_train_batch_size 1 \ --gradient_accumulation_steps 4 \ --deepspeed ds_config.json \ --output_dir output或者在 Python 里通过TrainingArguments传递from transformers import TrainingArguments, Trainer training_args TrainingArguments( output_diroutput, per_device_train_batch_size1, gradient_accumulation_steps4, deepspeedds_config.json, fp16True, save_strategysteps, save_steps500, ) trainer Trainer( modelmodel, argstraining_args, train_datasettrain_dataset, ) trainer.train()启动后日志里会打印DeepSpeed info: stage3之类的信息。如果看到这个说明 ZeRO-3 已经生效。注意用了 ZeRO-3 后训练脚本里的model会被 DeepSpeed 接管你自己的 precision 设置和torch.cuda.amp相关逻辑都要让位给配置文件里的fp16段避免两边打架。5. MoE 负载均衡代码与训练调优实战5.1 负载均衡辅助损失实现自己写一个最简版 MoE 门控并加负载均衡损失其实不复杂。下面是一个基于 Switch Transformer 思路的简化实现可以当作参考。import torch def load_balancing_loss(router_probs, expert_indices, num_experts): 计算 MoE 的辅助负载均衡损失。 router_probs: [batch, seq, num_experts]每个 token 对每个专家的概率 expert_indices: [batch, seq]每个 token 选中的专家 id tokens_per_expert torch.bincount( expert_indices.flatten(), minlengthnum_experts, ).float() total_tokens tokens_per_expert.sum() fraction_per_expert tokens_per_expert / total_tokens prob_per_expert router_probs.mean(dim(0, 1)) aux_loss num_experts * (fraction_per_expert * prob_per_expert).sum() return aux_loss这个损失的直观含义是如果路由完全均匀每个专家分到的 token 比例接近1/num_expertsrouter 给每个专家的平均概率也接近1/num_experts点积结果比较小如果某个专家被大量 token 选中fraction_per_expert会变大同时 router 给它的prob_per_expert也会偏大损失就会被拉高。实际训练时一般把它乘以一个系数例如 0.01加到总损失里。系数太大会干扰主任务学习系数太小负载均衡就形同虚设。这个系数我通常会先用 0.01 起步观察每个 step 打印的专家分布统计再调。5.2 训练不稳、loss 震荡怎么排查ZeRO-3 加 MoE 的组合里常见的坏症状就是 loss 上蹿下跳或者一张卡显存高、另外几张卡闲得发慌。按下面的顺序排查比较高效。第一步看路由分布。训练时打印每个专家接收的 token 数量如果集中在少数几个专家优先调负载均衡损失系数如果还不行可以考虑在训练早期冻结路由网络先在主干结构上稳定特征表达。第二步看混合精度配置。ZeRO-3 加 FP16 对梯度缩放策略敏感hysteresis和loss_scale_window不够大时loss 容易出现尖峰。可以先把initial_scale_power调低到 8观察 loss 曲线是否平滑。第三步看 batch 是否太小。MoE 的 router 要在足够多的 token 上做统计才有意义。micro batch 太小同一个 batch 里每个专家分到的 token 数量波动很大负载均衡损失就像一个随机噪音源。尝试增大train_micro_batch_size_per_gpu或gradient_accumulation_steps让统计更稳定。5.3 显存优化与通信瓶颈的调参顺序当你遇到 OOM我的建议是按这个顺序一步步来而不是上来就堆硬件。先降低train_micro_batch_size_per_gpu这是最快见效的手段。如果还 OOM打开offload_param或offload_optimizer把参数或优化器状态挪到 CPU。如果 CPU 内存也不够再考虑 NVMe offload但速度会明显下降。如果显存没问题但训练速度上不去大概率是通信瓶颈。这时把overlap_comm打开并适当调大stage3_prefetch_bucket_size和stage3_max_live_parameters让通信和计算更重叠减少卡等数据的时间。还可以尝试reduce_bucket_size和gradient_accumulation_steps配合把梯度归约次数降下来。最后说一个我踩过的坑ZeRO-3 训练中途保存模型时要小心如果发现 checkpoint 里某个参数文件缺失或形状不对多半是gather_16bit_weights_on_model_save没开。另外在 MoE 模型里验证 checkpoint 是否完整建议额外打印每个 expert 的参数统计确认多个专家的头部信息都正常存在而不是全部被塞到同一张卡的分片里。我个人的体会是ZeRO-3 不是银弹。小规模实验该用 ZeRO-2 就用 ZeRO-2速度更快模型大到参数显存成为瓶颈再切到 ZeRO-3MoE 模型则优先研究专家并行这比无脑堆 ZeRO-3 高效得多。最后留一个小技巧训练前先用ds_report确认编译状态训练中定期看 gateway 的breakdown日志学会从通信时间占比判断瓶颈比盯着 loss 曲线瞎猜强得多。
网站建设高端定制企业官网