
如果你的工作里同时出现 DeepSpeed ZeRO-3 和 MoE 训练这两件事大概率你已经开始在一个参数爆炸但计算稀疏的模型上折腾了。我见过不少同学第一次接触 MoE第一反应是“参数量都这么大了得堆多少卡才够”结果开起显存监控一看占用反而比同体量稠密模型更“温和”。原因是 MoE 的稀疏激活——不是每个 token 都会命中所有参数但前提是你得把 ZeRO-3 和专家并行安排好否则温和只是纸面上的。这篇文章不打算做概念搬运我想直接站在一个训练工程师的角度把显存账本、切分原理、配置细节和踩坑经验一条条写清楚最后给出一份可以直接起训练的最小配置。不管你是刚入门大模型训练还是已经跑过 ZeRO-2、正准备把模型切成 MoE这篇都值得当一份按图索骥的参考。1. 先把显存账算明白MoE 模型的 GPU 开销到底在哪1.1 一张卡上同时住着四类“显存房客”我在排查 OOM 的时候第一件事从来不猜而是把显存里的住户列一遍。训练一个普通的稠密模型GPU 显存主要被四类东西占着模型参数、梯度、优化器状态、激活值。前三个很好理解最后一个激活值容易被新手忽略——前向传播过程中每一层的中间张量都要留着做反向传播序列越长、batch 越大这个数越吓人。混合精度训练下一个 FP16 参数占 2 字节对应的梯度占 2 字节FP32 优化器状态通常包含模型参数副本、一阶动量、二阶动量每个参数再占大约 12 字节。加在一起就是 16 字节/参数。拿一个 7B 模型粗算光参数、梯度、优化器状态就要 112GB单张 80G 的 A100 根本塞不下。这就是 ZeRO 要解决的问题把这些“常量”按维度切到多张卡上。很多人会问既然 MoE 参数量更大为什么显存还“温和”因为 MoE 的核心是“总参数量大、单 token 计算量小”。比如一个 7B 稠密模型换成 8 个专家的 MoE每个 token 只激活其中 Top-K 个专家激活值不会等比例放大。但参数和优化器状态依然存在而且更大。所以在 MoE 场景下ZeRO 不是可选项是必选项。1.2 MoE 的“大头”不在注意力而在那片专家层MoE 最常见的做法是把 Transformer 里的 FFN 层替换成一组并行的 FFN 专家网络。注意力部分还是稠密计算每个 token 都过专家部分则是稀疏激活每个 token 由路由网络挑选 Top-K 个专家来算。假设你原来 FFN 的 hidden size 是 4096中间层是 11008现在换成 8 个同样大小的专家那么这一层总参数翻了接近 8 倍。模型参数会被这个“专家阵列”迅速撑大。Switch Transformer 那种 top-1 路由每个 token 只过 1 个专家计算量看似没涨多少但显存里的参数已经实打实涨上去了。所以讨论 MoE 显存时别再把注意力参数当成主要矛盾。如果你用一张表格看每个模块的参数量专家层几乎总是排第一。要训练这样的模型单卡全量加载基本不现实必须考虑把不同专家放到不同卡上或者让 ZeRO-3 来切分。1.3 反直觉结论并非所有参数都必须在每一块 GPU 上常驻热词里有一个问题问得特别好MoE 架构要全部参数进显存吗答案是不一定。推理时可以通过 offload 把不活跃的专家暂时放出显存训练时同样可以做到“按需加载”。这里有两套机制在起作用一套是 ZeRO-3 的通用参数分片另一套是专家并行。ZeRO-3 把参数、梯度、优化器状态全部切到不同 rank 上每个 rank 只常驻自己分到的那一份。计算某一层之前通过通信把完整参数临时拼回来算完再释放掉。MoE 专家层则更特殊——不同专家本来就可以分布在不同 GPU 上token 去哪张卡、哪张卡就加载对应的专家权重。两者结合之后一块卡上既不用放全部通用参数也不用放全部专家参数。这个结论听着反直觉但它正是 MoE 能在大规模下跑起来的原因。接下来的章节我会拆开讲这两种切分各自怎么实现、怎么配合。2. DeepSpeed ZeRO 的三代演进从“只切优化器状态”到“连参数都按层切”2.1 ZeRO-1/2/3 到底切了什么DeepSpeed 官方经常用一张很形象的图讲 ZeRO把训练状态里的“冗余”按三个维度逐步去除。我直接用一张表把区别列清楚维度ZeRO-1ZeRO-2ZeRO-3模型参数每卡全量每卡全量每卡只保存分片梯度全量聚合后再分片每卡只保留自己负责的分片每卡只保留自己负责的分片优化器状态每卡分片每卡分片每卡分片典型通信模式梯度 AllReduce梯度 ReduceScatter前向 AllGather 后向 ReduceScatter显存下降幅度4 倍左右8 倍左右跟卡数近似线性ZeRO-1 只是把 Adam 的状态切开梯度还是每轮全量 AllReduce通信成本相对固定。ZeRO-2 把梯度也切了经典做法是 ReduceScatter 之后每卡只更新自己那部分参数对应的优化器状态。ZeRO-3 更进一步参数本身也按层切分。每一层计算前要用一次 AllGather 把这一层参数广播到所有参与计算的 GPU 上计算完再释放掉不属于自己的部分。这个演进本质上是拿通信换显存。ZeRO-3 的通信量比 ZeRO-2 多了不少但换来的是可以训练远超单卡内存的模型。2.2 ZeRO-3 怎么做到“用时重建”ZeRO-3 不是把所有层都常驻显存而是“按需 AllGather、用完即丢”。以前向过程为例进入第 i 层前先对第 i 层参数做一次 AllGather让当前计算设备拿到完整层参数用完整参数计算前向输出计算结束后立刻释放除了本 rank 分片之外的其他参数副本进入下一层时重复这一过程。反向传播同样需要 AllGather 得到完整参数才能算梯度算完梯度之后再做 ReduceScatter让每个 rank 只保留自己负责的那份梯度分片。下一轮迭代时每个 rank 用自己手里的梯度分片去更新自己手里的参数分片。所以 ZeRO-3 中的“参数”不是随时都在显存里而是马上要算哪一层才把哪一层拼出来。这里有个很容易踩的坑如果你在外部代码里直接遍历 model.parameters() 拿到的是一个“分片视图”想当然地拿去算完整 loss梯度满天飞最后结果大概率是错的。正确做法是让 DeepSpeed 的 engine 接管整个前向/后向流程。2.3 当 ZeRO-3 遇上专家并行两种切分并不冲突MoE 训练里还有一个常见的并行维度专家并行。它的思路很朴素——把 E 个专家分散到 E 个 GPU 上每个 GPU 保存一部分专家路由网络决定某个 token 应该去哪张卡。你可能会疑惑ZeRO-3 已经把参数切分了为什么还要单独搞专家并行原因在于 MoE 的路由是一种“稀疏随机访问”。如果一个 token 只激活 2 个专家而这两个专家恰好不在本地那么把整层所有专家的参数 AllGather 过来就是浪费。所以更聪明的做法是专家参数不走全量 AllGather而是让 token 本身通过 All-to-All 通信去找对应的专家。DeepSpeed 的 MoE 实现在这一层做了融合通用层参数继续走 ZeRO-3 的分片逻辑专家层参数按专家维度切分并暴露一个 ep_size 参数控制单个专家会被复制到多少张卡上。两个机制不冲突因为它们切的对象和使用方式不一样。理解到这一层你再看 DeepSpeed 的配置就不会觉得“moe 那一块到底是干嘛的”。3. MoE 训练的命门负载均衡、专家容量与辅助 Loss3.1 Top-K 路由不是所有专家都要被激活MoE 层通常由两个部分组成一个轻量门控网络外加一组专家 FFN。门控网络对每个 token 生成一个在 E 个专家上的概率分布然后取 Top-K 个大。Top-1 最省通信因为一个 token 只需要去一张卡Top-2 在精度和负载均衡上更稳但 token 要等到两个专家计算完之后再做一次 All-to-All 合并。实际项目中我看到的大多数模型都选 Top-2少量用 Top-1 做极致吞吐。门控网络本身参数很少但它要承担两个任务一是把 token 分到正确的专家二是让专家利用尽量均衡。第一个任务靠正常的梯度回传第二个任务就要靠辅助损失因为路由选择里的 hard Top-K 对门控参数的梯度几乎不提供信号。3.2 负载均衡 Loss 公式拆开看负载均衡最经典的做法是 Switch Transformer 里提出的辅助损失。假设有 E 个专家一个 batch 里有 T 个 token定义两个量f_i被路由到第 i 个专家的 token 数占总 token 数的比例P_i门控网络分配给第 i 个专家的平均概率。辅助损失可以写成loss_aux α * E * Σ_{i1}^{E} f_i * P_i乘 E 是为了让 loss 的量级不随专家数变化。直观理解如果所有 token 都挤到某个专家上f_i 会很高同时门控对那个专家的平均概率 P_i 也会很高乘积加总后就会变大优化器就会去压制这种倾向。理想情况是每个专家分到 1/E 的 token此时 f_i 和 P_i 都接近 1/E各项乘积的均值是 1/E²乘以 E 后就是 1/E整体很小且稳定。这个公式看着简单但工程实现里有几个细节经常被忽略。比如 f_i 必须用 stop-gradient 的方式统计不能让路由比例这个本身不可导的数也来反向传播再比如 P_i 要从 softmax 前的 logits 上计算并且最好在 batch 内部统计而不是跨 batch 统计。3.3 专家容量与 Token 丢弃这是延迟和质量的跷跷板路由做完之后还有一个工程概念专家容量。每个专家不可能无限接收 token为了控制计算量和通信延迟通常设定容量expert_capacity capacity_factor * (T / E)capacity_factor 一般取 1.0 到 1.5。超过容量以后多出来的 token 要么被丢弃要么走残差连接直接跳过专家。丢弃会让训练质量受损尤其在训练早期门控还没学好的时候但如果不丢弃就需要额外 buffer 缓存这些 token计算等待时间会拉长。DeepSpeed 配置里常见两个字段min_capacity 和 drop_tokens。min_capacity 保证每个专家至少处理几个 token防止某些专家在 batch 边缘彻底空闲drop_tokens 控制超出容量时是否丢弃。我把 drop_tokens 设成 false 时训练 loss 更稳但吞吐明显下降设成 true 以后速度快了可 MoE 层的 aux loss 偶尔会震荡。这个 trade-off 没有标准答案得看你的模型对 token 覆盖率的敏感程度。4. DeepSpeed 配置实战把 ZeRO-3 和 MoE 跑在同一份训练脚本里4.1 环境准备与安装 DeepSpeed 的常见错误“安装 deepspeed 包一直报错”这个热词我太有共鸣了。绝大多数报错不是 DeepSpeed 本身的问题而是环境里 CUDA 和 PyTorch 的真实版本没对上。安装前先确认三件事nvcc 和 PyTorch 的 CUDA 版本一致编译器 ninja、gcc 可用不要在只有 CPU 的 Python 环境里强行装 GPU 版。最省事的做法是装预编译版本让 DeepSpeed 不现场编译 CUDA 算子。如果还是需要从源码编译我通常会加这些参数pip install deepspeed --no-build-isolation编译过程中如果报 CUDA_HOME 找不到就显式指一下export CUDA_HOME/usr/local/cuda export LD_LIBRARY_PATH/usr/local/cuda/lib64:$LD_LIBRARY_PATH还有一个非常隐蔽的坑机器上同时装了一套系统 CUDA 和一套 PyTorch 自带的 CUDA runtime版本不一致会导致 JIT 算子编译出来以后运行时直接崩。遇到这种问题先torch.version.cuda看 PyTorch 期望的版本再nvcc --version看系统版本尽量让两者一致。4.2 一份可落地的最小配置下面这份 JSON 是我常用的 ZeRO-3 MoE 训练配置骨架你可以直接抄走再调整{ train_batch_size: 256, gradient_accumulation_steps: 8, fp16: { enabled: true, initial_scale_power: 10 }, zero_optimization: { stage: 3, contiguous_gradients: true, reduce_bucket_size: 5e8, stage3_prefetch_bucket_size: 5e7, stage3_param_persistence_threshold: 1e6, stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9 }, moe: { ep_size: 4, num_experts: 8, top_k: 2, min_capacity: 4, drop_tokens: false } }逐个拆开说训练必须的部分。train_batch_size是全局 batch size不是单卡 batch size。如果你有 16 张卡梯度累积 8 步那么单卡 micro batch 就是 256 / 16 / 8 2。这个换算关系搞错batch 大小会跟你预期差一个数量级学习率也跟着失效。zero_optimization里的几个stage3参数都是控制“何时预取参数”“保留多少活跃参数”的缓冲阀门。stage3_max_live_parameters和stage3_max_reuse_distance调得越大参数在显存里存活越久省通信但费显存调小则反过来。MoE 场景我倾向把它们调得偏小因为专家层本来就靠 All-to-All 通信没必要让 ZeRO 也把专家参数缓存太多。moe块是 DeepSpeed 支持 MoE 层的核心开关。ep_size表示每个专家被切成几份、复制到几张卡上。ep_size 越大同一个专家的副本越多token 落在本卡的概率越高通信压力越小但每卡显存越大。ep_size 越小越省显存但 All-to-All 通信越密集。4.3 训练脚本里需要包一层什么逻辑如果你直接用自定义的 MoE 层光把它塞给deepspeed.initialize是不够的。DeepSpeed 官方推荐使用它实现好的 MoE 层或者在自定义层里手动接入专家并行和路由通信。下面是一个最小可用的 Python 骨架import torch import deepspeed from deepspeed.moe.layer import MoE class MyTransformerBlock(torch.nn.Module): def __init__(self, hidden_size, num_experts, ep_size): super().__init__() self.attention ... self.moe MoE( hidden_sizehidden_size, expertExpertFFN, num_expertsnum_experts, ep_sizeep_size, use_residualTrue ) model MyMoEModel(...) model_engine, optimizer, _, _ deepspeed.initialize( modelmodel, model_parametersmodel.parameters(), configds_config.json ) for batch in dataloader: loss model_engine(batch) # MoE aux loss 需要自己额外加进来 total_loss loss model_engine.moe_loss() * 0.01 model_engine.backward(total_loss) model_engine.step()几点注意model_engine.moe_loss()返回的是多个专家层的辅助损失累加值具体接口名称可能随 DeepSpeed 版本变化以你实际安装版本里的 docstring 为准训练循环不要手动调用loss.backward()用model_engine.backward()否则 ZeRO-3 的梯度分片会被绕过去如果你用 HuggingFace Trainer需要开启DeepSpeedPlugin并配置相同的zero_optimization否则配置加载顺序会互相覆盖。4.4 三个我真正踩过、并且查了很久的坑第一个坑是自定义 MoE 层和 ZeRO-3 互不兼容。我最开始把一个手写的专家 FFN 直接包在普通 Module 里param 量倒是挺大但所有参数都被 ZeRO-3 当成普通参数分片而路由逻辑根本没有把 token 发到对应专家所在的卡上。结果就是每张卡都在等一个不存在的本地参数训练卡死。后来换了 DeepSpeed 的 MoE 层让专家参数走专家并行路径才通。第二个坑是 drop_tokens 引起的震荡。我为了提升吞吐把drop_tokens开了训练 loss 前 500 步降得飞快到某个拐点后突然开始抖动验证集 perplexity 比不开时高了快 0.2。后来查代码发现是因为负载均衡 loss 权重太小门控在 torch 的极端路由偏好下疯狂丢 token模型一直在学“漏掉的信息”。最后把 capacity_factor 调到 1.25并且前 500 步冻结路由 loss 权重才稳下来。第三个坑是通信死锁。我试着在多机场景下写自己的 All-to-All 逻辑只调了 send 不调 recv结果 16 张卡全部挂起进程不报错也不退出就卡在 NCCL 的等待队列里。排查花了一天多最后把通信原语统一成 DeepSpeed 封装好的接口或者显式按 rank 顺序 pairing才解决。如果自己也遇到“进程活着但 loss 不动”先看 NCCL 日志别急着怀疑模型。5. 性能调优的几条“反直觉”经验别一上来就猛堆卡5.1 All-to-All 通信才是大瓶颈说到 ZeRO-3很多人第一反应是参数 AllGather 会拖慢速度。但 MoE 模型跑到大规模以后真正的瓶颈通常是专家路由产生的 All-to-All 通信。token 在每层路由后都要去对应专家所在的 rank这个通信是点对点交叉的节点内延迟还好跨节点就很容易成为长尾瓶颈。调这个问题的思路主要有两个一是把ep_size调大增加专家副本数让更多 token 命中本卡专家减少跨卡交换二是调整数据分布尽量让同一批 token 路由落在相近的专家上。但两者都会增加显存或者影响负载均衡所以每次调完都要同时盯三个指标吞吐、aux loss、token drop 率。5.2 显存换计算的组合拳ZeRO-3 不是唯一的省显存手段它要和别的技术叠加。我最常用的组合是开启梯度检查点activation checkpointing激活值不缓存全部反向时重算能省下大头激活优化器状态用zero_optimization.offload_optimizer放到 CPU但要注意 CPU 与 GPU 的搬运频率batch size 太小时 CPU 反而成瓶颈混合精度训练必备如果发现某些 loss 出现 inf别急着关 AMP先查梯度裁剪和初始 scale。这套组合拳打下来MoE 的总参数可能比稠密模型大了 5 倍但单卡显存反而能压在可接受范围内。不过代价是训练效率大概率比同等稠密模型低因为通信和重算都在偷时间。5.3 从小规模到上百卡的扩展注意点我的习惯永远是先在 4 到 8 卡上把数据流彻底跑通再上大规模。小规模时重点验证三件事MoE 的 aux loss 是否在健康范围、token drop 率是否小于 1%、全卡 loss 曲线是否一致。这三项过了再考虑把 ep_size、global batch、梯度累积同步放大。上大规模后很多问题会变样。比如几十卡时 All-to-All 的延迟不明显上百卡时同一个 token 可能要跨两台 NUMA 节点甚至跨交换机延迟立刻放大。此时优先把专家复制策略改成“局部优先”让同一机架内的卡共享专家副本避免一个路由请求跑到远端节点。还有一点很反直觉很多人以为卡越多显存瓶颈越小就越可以肆无忌惮地加 batch。实际上 MoE 的负载均衡是基于 batch 内 token 分布的batch 太小门控看到的数据不够多样容易在几步内偏科batch 太大每个专家的容量压力又变大。所以扩展卡数的同时要同步调整 capacity_factor 和 aux loss 的权重不是把所有配置原样搬过去就行。我自己每次搭 MoE 训练都会把三句话贴在终端上先看 moe loss再看 drop rate最后才看主 loss。路由不均衡导致的那些问题常常在主 loss 上早就被优化器“掩盖”了等发现时已经跑废了好几轮实验。先把这三条曲线盯到稳定再谈调参和加速——这个顺序不管换多少版本、多少卡都适用。