ARTICLE DETAIL

资讯详情

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

MoE大模型训练显存优化:DeepSpeed ZeRO-3与专家并行实战指南

MoE大模型训练显存优化:DeepSpeed ZeRO-3与专家并行实战指南 1. 先把显存这盘账算明白MoE模型为什么会让单卡直接OOM做大规模模型训练的人十有八九都经历过那个瞬间高高兴兴把配置写好deepspeed --num_gpus8 train.py一敲几秒钟后日志里甩出一行CUDA out of memory显卡直接撂挑子。如果你训练的是稠密模型这个问题靠 ZeRO-2 基本就能扛过去可一旦换成 MoEMixture of Experts架构显存爆炸的剧本会换个方式重演——明明计算量没涨多少显存却像被捅了马蜂窝一样根本兜不住。先说结论MoE 模型不是因为算力需求大而难训而是因为“总参数量”大导致哪怕只计算一小部分参数也得把全部参数放进训练系统里管理。这里的关键矛盾在于训练时系统的显存开销往往和“总参数量”挂钩而不是和“实际激活的参数数量”挂钩。一个 1B 参数的稠密模型混合精度训练下光模型状态参数、梯度、优化器状态占多少显存业界常说一个公式混合精度 Adam 优化器每 1B 参数大约需要 12~16 字节 × 参数量。拆开算模型参数FP16 存储2 字节/参数梯度FP16 存储2 字节/参数优化器状态Adam 需要维护 FP32 的 master 参数、一阶动量momentum、二阶动量variance共 12 字节/参数加一起就是 16 字节/参数。1B 参数就是 16GB轻轻松松吃掉一张 A100 的小半壁江山。这还没算激活值、通信缓冲区、临时计算中间量这些“残余状态”。MoE 模型的问题在于它的总参数量往往比同等级算力消耗的稠密模型大一个量级。比如一个 350B 参数的 MoE 模型实际每次 token 只激活 2 个专家但训练系统要管理的是 350B 参数的全部状态。用上面的算法350B × 16 字节 5.6TB。没错是 TB 级别。单卡能装下吗显然不能。所以问题的本质不是“计算不过来”而是“放不下”。这也解释了为什么业内训练 MoE 模型几乎不会用普通的 Data Parallel必须用 ZeRO-3 或者 Expert Parallelism 这类能把参数拆分到多卡的技术。接下来的内容全围绕这件事展开。提示如果你只是做小规模的 MoE 实验比如几个专家、几十亿参数单卡硬塞也不是不行但稍微一放大就会立刻撞墙。提前理解这套显存账本能帮你少走很多弯路。2. ZeRO-3 到底“零”掉了什么从 ZeRO-1 到 ZeRO-3 的分区演进ZeROZero Redundancy Optimizer是 DeepSpeed 提出的显存优化方案核心思想就一句话把数据并行中冗余存储的模型状态变成分布式存储。你可以把数据并行想象成多人各拿一本相同的教材大家都能独立算题但书只有一本的量却每人复制了一份ZeRO 则是让每人只拿教材的一章谁需要哪一章临时找别人借。2.1 三个阶段到底在分区什么ZeRO 有明确的三阶段递进关系理解它比直接背配置重要得多。我用下面的表格把三者的分区对象说清楚ZeRO 阶段分区对象显存效果需要通信吗ZeRO-1优化器状态让人均显存减少约 3/4 的优化器状态开销基本无额外通信ZeRO-2优化器状态 梯度进一步减少梯度存储训练时有梯度规约通信ZeRO-3优化器状态 梯度 模型参数参数也不再冗余存储前向/反向时需要 All-gather 参数ZeRO-1 和 ZeRO-2 的参数仍然是每张 GPU 持有完整副本所以它们的显存节省有上限真正适合大模型的是 ZeRO-3因为它连模型参数本身都做了分区。结合前面讲的 MoE 模型“总参数量巨大”的特点ZeRO-3 几乎是训练大规模 MoE 的必备选项。2.2 ZeRO-3 参数分区的完整生命周期ZeRO-3 的具体工作方式有人会误以为“每个 GPU 永远只持有自己那一份参数”这其实只说对了一半。真正的机制是训练过程中每一层需要计算时所有 GPU 通过 All-gather 把该层参数拼出来算完以后再把不属于自己的部分释放掉。听起来很折腾但确实是拿通信换显存的核心逻辑。拿一个 Transformer 层举例。假设 8 张卡、每层 1B 参数ZeRO-3 下每张卡只常驻 1/8 即 125M 参数。前向计算开始对当前层做一次 All-gather把这 1B 参数凑齐到每张卡上计算完成立刻把不属于自己的参数排除。反向传播同样需要参数所以再 All-gather 一次。所以 ZeRO-3 相对 ZeRO-2 的额外通信量大概是每个 Transformer 层前向和反向各多一次 All-gather。这个开销听起来不大但模型一深、卡一多通信放大就很明显。这也是为什么 ZeRO-3 在千卡集群上要非常关注通信重叠和桶大小设置。2.3 ZeRO-3 不是张量并行别搞混很多人第一次接触 ZeRO-3 会问这和张量并行Tensor Parallelism有什么本质区别两者都是把参数拆开但拆完以后的“计算方式”完全不同。张量并行一次前向计算就把一个算子拆到多卡上每张卡计算一个子块算完通过 All-reduce 拼最终结果。它切的是“计算本身”通信发生在算子内部非常频繁。ZeRO-3每张卡持有完整的数据并行副本计算本身不拆分只是计算前用 All-gather 把参数重新聚合到每张卡上算完再释放。它切的是“存储”通信发生在算子之间。所以 ZeRO-3 的显存效率更高但单算子内没有并行加速张量并行牺牲部分显存效率换取计算并行。理解这个区别后你就知道为什么很多大规模训练方案是“ZeRO-3 做主TP 做补充”的组合。提示如果你的 MoE 模型单层计算本身就很大比如单个专家 MLP 超过了单卡能承载的激活值上限光靠 ZeRO-3 解决不了算子级显存问题这时候才需要张量并行去拆算子。3. MoE 的算力与通信特性为什么它不是“很多个 MLP 堆在一起”很多新手以为 MoE 就是把一个大 MLP 替换成几十个独立的小 MLP然后每个 token 随机选一个算。这个理解方向对但漏掉了两个关键工程问题怎么选和选完之后怎么调度。3.1 门控路由与 Top-KMoE 层的标准结构是一个门控网络Gate加一组专家网络Expert。门控网络接收 token 的隐藏向量输出一个在所有专家上的概率分布然后取概率最高的前 K 个专家K 通常是 1 或 2业界最常见是 Top-2把 token 送过去计算。K 的选择直接影响“稀疏性”——K 越小每个 token 激活的参数量越少计算越省但路由选择错误的概率也越高。实际训练中 Top-2 比 Top-1 更主流原因是单个专家可能视角单一Top-2 提供了一定的“交叉验证”而且梯度能同时流过两个专家的输出训练更稳定。但 Top-2 也意味着需要把 token 同时送到两个专家那里通信翻倍工程上不那么便宜。3.2 专家容量与 token 丢弃MoE 层天然有一个负载不平衡问题极少数专家可能被大量 token 选中多数专家却闲着。如果不做任何限制热门专家的计算量会高到让整个训练卡住其他 GPU 都在干等。解决办法是给每个专家设置容量上限。公式大概是专家容量 capacity_factor × (当前批次 token 总数 / 专家总数)capacity_factor1.0表示每个专家刚好吃到平均数量的 token设为 1.2 或 1.5 则允许一定程度的不平衡给热门专家一点缓冲空间。问题是超过容量上限的 token 怎么办训练时绝大多数情况下直接丢弃不参与这一层的计算。这些被丢弃的 token 会让这个位置的梯度信号缺失所以容量太小会伤害最终模型质量容量太大则会让负载失衡越来越严重。3.3 负载均衡 loss最近大家都在聊的 MoE 训练辅助函数既然 token 可能失衡就必须在训练目标函数里加一个“均衡约束”这就是负载均衡 loss 的来源。它的本质是惩罚“路由分布不均匀”的行为。数学上常见的写法是aux_loss α × N × Σ f_i × P_i简单拆一下含义N是专家数f_i是实际被路由到第 i 个专家的 token 占比P_i是门控网络对所有 token 给第 i 个专家的平均概率。如果 token 分配完全均匀f_i和P_i都会趋近1/N两者相乘求和约等于1/N再乘以N结果约等于 1是个温和的常数如果某个专家被过度选中这个乘积会明显偏大loss 就升高形成梯度推动路由更均衡。写负载均衡代码时有个细节值得注意很多开源实现里会把f_i和P_i的统计限制在一个滑动窗口内而不是整个 batch。比如只统计最近 128 步的 routing 情况而不是这次 batch 的全部 token。这样做的原因是单个 batch 的局部不平衡波动很大全量统计会让辅助 loss 抖动剧烈反而干扰主任务的收敛。这个窗口长度就是配置里的window_size属于典型的“官方文档不会细讲但实际影响很大”的隐藏参数。另外负载均衡 loss 前面那个系数α也很讲究。调太大模型会牺牲表达能力去硬凑均匀分布导致最终效果变差调太小均衡约束形同虚设卡顿和 token 丢弃会让训练不稳定。我的经验是刚开始实验用 0.01 量级起步观察到明显的路由失衡再加而不是一上来就给很大的约束。4. DeepSpeed 怎么把 ZeRO-3 和 MoE 缝在一起分区策略与 All-to-All把 ZeRO-3 和 MoE 组合到一起不是简单地把两个功能叠加就完事。MoE 层的参数结构和普通 Transformer 层完全不同普通层的参数是“每层一份大家都要用”MoE 层则是一堆专家每个 token 只用其中两三个。如果一概而论地做 ZeRO-3 参数分区All-gather 会把所有专家的参数拼到每张卡上那显存照样爆炸MoE 的稀疏优势就没了。4.1 共享层与专家层的参数处理方式完全不同DeepSpeed 对 MoE 模型的 Parameter 做了区分非专家参数Attention、LayerNorm、Embedding 这些共享层走标准的 ZeRO-3 分区专家参数moe.experts下的权重走 Expert Parallel每个 GPU 持有不同的专家子集而不是全部专家的分区切片。举个例子假设模型有 32 个专家跑在 8 卡上每张卡只需持有 4 个专家的完整参数。Token 被路由到专家 1但它不在当前卡上怎么办这就是下一节要讲的跨卡通信问题。这套混合策略的关键设计原则是越训练卡数越多每卡持有的专家数越少显存压力越小同时通信压力越大。4.2 一个 token 的“跨卡旅行”从路由到 All-to-All在 ZeRO-3 MoE 的训练流程中一个 token 要经历这么几步在本地 GPU 完成 Attention 层计算得到隐藏向量门控网络决定它要去哪几个专家如果目标专家不在本地token 就要被“搬”到目标专家所在的 GPU——这就要用到 All-to-All 通信目标 GPU 上的专家完成计算再用一次 All-to-All 把结果送回原卡原卡拿到所有被激活专家的输出做加权求和继续下一层。关键区别在于ZeRO-3 的 All-gather 是“大家一起把同一份参数拼出来”属于集体通信而 MoE 的 All-to-All 是“不同 token 去往不同的卡”属于全连接通信。前者是数据并行里典型的广播-收集模式后者更接近消息传递。这种差异直接导致通信瓶颈的表现完全不同。ZeRO-3 容易遇到的是 All-gather 带宽不足导致 GPU 频繁等待MoE 则会遇到 All-to-All 延迟抖动最坏情况下某张卡被大量 token 冲击成为热点卡。DeepSpeed 的 MoE 实现专门做了 All-to-All 调度优化把大块通信拆成小块交错执行让通信和计算重叠起来而不是傻等一整批 token 全收集齐再开算。这个优化在多机多卡时特别重要跨节点的网络往返延迟远高于卡间通信不懂这个底层逻辑的话单看 profiling 结果会一头雾水。4.3 DeepSpeed 训练引擎对 MoE 的集成方式在代码层面DeepSpeed 通过参数名里的moe.experts前缀识别哪些属于专家参数并使用专门的 MoE 策略处理它们的存储、梯度和通信。这意味着你不需要手工微调每个张量放在哪张卡——只要在配置里声明专家数引擎会处理后续。背后DeepSpeed 的模型状态分为三类共享参数的 ZeRO-3 分区、专家参数的分组放置、以及优化器状态按专家持有者的本地区分。这三类更新时机不同通信先后不同错开调度才能避免带宽打架。我在实际 profiling 时看到很多不合理的配置导致通信都挤在同一时间段发生GPU 利用率明显波浪起伏就是引擎没把三类通信错开的典型症状。5. 一份能跑的配置逐字段拆给你看配置是所有深度学习框架里最容易被忽略又最容易出问题的地方。我用一份实际跑过千亿级 MoE 的 DeepSpeed JSON 配置逐字段讲解每个关键开关的作用和调试心得。5.1 ZeRO-3 关键配置字段{ zero_optimization: { stage: 3, reduce_bucket_size: 500000000, allgather_partitions: true, allgather_bucket_size: 500000000, overlap_comm: true, contiguous_gradients: true, cpu_offload: false } }stage: 3好理解不多说。allgather_partitions设为 true 表示把分区的参数在 All-gather 后组合成连续内存减少碎片allgather_bucket_size控制每次 All-gather 的数据块大小太大会让单次通信等待明显太小会让通信次数暴增。我的经验值是先给 5e8 起步多卡网络好的话往下调网络差就往上调。overlap_comm开启通信重叠如果机器网络是 NVLink,强烈建议开。cpu_offload我这边没开因为 MoE 本身参数已经通过专家并行分摊再 offload 到 CPU 反而让 All-to-All 通信和 PCIe 传输抢资源。5.2 MoE 相关配置字段详解{ moe: { type: linear, num_experts: 32, top_k: 2, capacity_factor: 1.2, eval_capacity_factor: 2.0, min_capacity: 8, drop_tokens: true, window_size: 128, linear_wall_clock_breakdown: true } }逐条说num_experts: 专家总数。注意它是全局总数不是每卡的数量。top_k: 每个 token 激活的专家数。Top-2 是常见值想省通信可以用 1想更稳定可以试 3 但通信成本明显上升。capacity_factor: 训练时的专家容量系数。默认 1.0 到 1.5 之间。我之前用 1.2 效果不错容量太低会频繁丢token。eval_capacity_factor: 评估/验证时的容量系数。评估通常 batch 更小、数据更整齐给 2.0 是常规操作避免评估时掉 token 导致指标虚高或虚低。min_capacity: 专家容量的保底值。哪怕当前 batch 只有一个 token 路由到某个专家也至少给它计算 8 个 token 位置避免小 batch 下专家完全得不到梯度。drop_tokens: 超过容量的 token 是否丢弃。训练时开 true 是主流做法某些新框架会改为补 padding token但计算代价更高。window_size: 辅助 loss 的统计窗口前面解释过。linear_wall_clock_breakdown: 是否打印 MoE 各阶段耗时明细。我强烈建议训练初期开这项可以看到前向、反向、All-to-All 通信各自的开销占比是排查瓶颈的第一手材料。5.3 完整 JSON 示例和启动方式把各部分组合起来一份可直接用于实验的配置长这样{ train_batch_size: 4096, train_micro_batch_size_per_gpu: 16, gradient_accumulation_steps: 4, zero_optimization: { stage: 3, reduce_bucket_size: 500000000, allgather_partitions: true, allgather_bucket_size: 500000000, overlap_comm: true, contiguous_gradients: true }, moe: { type: linear, num_experts: 32, top_k: 2, capacity_factor: 1.2, eval_capacity_factor: 2.0, min_capacity: 8, drop_tokens: true, window_size: 128, linear_wall_clock_breakdown: true } }启动命令和平常训练没有本质区别deepspeed --num_gpus64 train.py --deepspeed_config ds_config.json注意一个容易踩的坑如果机器是 64 卡但模型定义里专家数为 32那专家并行的粒度会和卡数不匹配DeepSpeed 的引擎会自动把专家复制到多卡上。你要做的是确认每个 expert 的 batch size 是否在你的显存预算内。简单估算总 batch size 409664 卡每卡微批次是 16每个专家平均接待 token 数约为16 × num_micro_batches / expert_groups的某种换算显存一旦超了就减微批次而不是减 batch size。提示配置里的train_micro_batch_size_per_gpu要尽量设到显卡能承载的上限。MoE 的激活值相比稠密模型小很多很多人的显存其实还有富余不敢往上加微批次太可惜。6. 实操日记从 OOM 到跑通的三个坑与一个真相如果只看配置文档很多人会觉得 ZeRO-3 MoE 就是改几行 JSON 的事。实际上手之后我遇到过三个具有代表性的坑各自方向不同一个是显存估算失准一个是通信瓶颈识别不到还有一个是收敛过程飘到怀疑人生。6.1 显存估算公式与实测对照训练前最好先做一次纸上估算。以 350B 参数、64 卡、ZeRO-3 Expert Parallel 为例总参数一旦用 FP16 存储体积 700GB。专家参数假设占全部参数 60%也就是 420GB。64 卡平均每卡约 6.5GB 专家参数。非专家参数 280GB走 ZeRO-3 分区每卡约 4.4GB。梯度、优化器状态再各自按比例分摊每卡合计大概 20~30GB。激活值、通信 buffer、临时变量每卡留 10GB 左右。A100 80GB 完全够跑A100 40GB 就非常紧可能必须开激活重计算或者降低微批次。实测下来我发现一个规律纸上估算通常比实际少 20% 左右。原因是 DeepSpeed 在 All-gather 后构建完整参数副本时会临时占用额外显存这部分开销很隐蔽。所以预算别卡太死至少留 10~15GB 的余量。6.2 通信瓶颈怎么发现和优化第一次跑大 MoE 训练时我注意到一个怪现象单卡算力明明没打满SM 利用率只有 60% 上下但整个训练速度就是上不去。当时的第一反应是调大 batch size、开重计算根本没用。后来打开linear_wall_clock_breakdown的日志才看到 All-to-All 通信时间占了前向总时间的 40%。这个案例说明MoE 训练的性能瓶颈很可能是通信而不是计算。排查路径大致是先开linear_wall_clock_breakdown看 MoE 内部各阶段的耗时占比用nvidia-smi看 GPU 利用率和显存如果 GPU 利用率低但网络收发高基本是通信等待把capacity_factor调小一点让 token 更集中通信变少但注意别让 token 丢弃过多调整allgather_bucket_size减小通信粒度让 All-to-All 能有更多机会和计算重叠。跨节点时网络拓扑的影响非常明显。8 卡单机内 NVLink 足够快通信开销几乎可忽略一旦做到多机跨节点带宽立刻成为短板。此时最有效的做法不是调容错开关而是合理规划专家分布尽量让同一节点的卡持有互补的专家集合减少跨节点搬 token 的概率。这个优化只能通过手动控制专家分组实现DeepSpeed 不会替你自动做。6.3 收敛相关的坑负载均衡系数、专家容量、token 丢失第三个坑出现在一个 32 专家、Top-2 的模型上训练 5000 步后验证 loss 不降反升而且路由分布日志显示有 2 个专家被 90% 的 token 选中。这明显是负载均衡约束不够。我当时把辅助 loss 的系数从 0.01 提到 0.03路由分布逐渐健康但系数提到 0.1 以后模型整体表达能力下降验证 loss 又变差。这个来回测试给了我一个深刻教训任何负载均衡 loss 都是人为加上的约束本质是在“让各路专家都用起来”和“让路由自由发挥”之间找平衡。系数过大、窗口过短均衡约束会反过来伤害模型质量。另外drop_tokenstrue时 token 丢弃会带来一个很隐蔽的问题被丢弃 token 的梯度是零如果某个专家经常满载那么它收到的有效梯度不稳定参数更新方差很大。表现为训练 loss 曲线锯齿明显而不是平滑下降。遇到这种情况先把capacity_factor从 1.2 调到 1.5看丢弃比例是否回到较低水平——通常问题就解决了。6.4 回应热搜里的经典问题MoE 架构要全部参数进显存吗热搜里挂着“MoE 架构要全部参数进显存吗”这个问题。答案分训练和推理两种场景必须分开讲。训练场景不需要。用 ZeRO-3 Expert Parallel每个 GPU 只持有全部参数的分区或子集运行时按需通过 All-gather 或 All-to-All 获取需要的那部分参数用完释放。这也是 MoE 虽然总参数量巨大却依然能在有限显存里训练的根本原因。推理场景理论上也要区分但实践中大多数情况要全部加载。推理时模型不再需要梯度优化器状态可以丢弃但所有专家权重仍然需要保存在某个存储位置上——可以是显存也可以是 CPU 内存。关键是只有被路由激活的专家才需要进入计算所以如果推理引擎支持 CPU GPU 混合部署没被激活的专家可以放在 CPU 内存里显存需求就大幅下降。而多数现成推理框架的默认行为仍然是“加载全部专家到显存”因为这样路由后计算无需跨设备访问延迟最低。如果你正在做模型部署选型我的建议是显存够就全量加载简单省事显存抠得紧就选择支持将专家按需加载到显存的推理引擎比如带 CPU offload 的 MoE 推理方案。用 CPU 放专家权重的代价是推理延迟会偶发升高——某个 token 恰好路由到一个还没显存中加载的专家时必须从 CPU 调权重。最后再分享一个小技巧调试 MoE 训练时很多人一上来就盯 loss 曲线我却习惯先把moe配置段里的linear_wall_clock_breakdown打开观察每个专家在不同 batch 下的 token 分布。这个日志的价值在于它既能帮你发现通信瓶颈也能直观看到负载均衡 loss 有没有起作用——如果专家负载始终极度偏斜说明你的capacity_factor或平衡系数设置不合理。配置调优的顺序永远是先确认数据分布是否健康再去看训练速度。这个习惯帮我少走了很多弯路。
返回列表