ARTICLE DETAIL

资讯详情

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

DeepSpeed ZeRO-3与MoE训练实战:显存计算、专家并行与负载均衡调优

DeepSpeed ZeRO-3与MoE训练实战:显存计算、专家并行与负载均衡调优 印象里三年前第一次用 ZeRO-3 跑千亿参数稠密模型时单卡显存从 32G 一路压到 8G当时觉得这玩意已经是极限了。直到后来接触 MoE混合专家训练才发现稠密模型的显存焦虑只是热身MoE 的工程复杂度完全是另一个量级。这篇文章想把 DeepSpeed ZeRO-3 和 MoE 训练的核心问题一次讲透包括显存到底怎么算、专家参数要不要全部进显存、负载均衡怎么做、还有那些文档里不会写的坑。我默认读这篇文章的你已经跑过一些大模型训练至少理解数据并行和模型并行的基本概念。如果纯新手也没关系我会尽量用生活化的方式把每个概念拆开讲但建议先了解一下 Transformer 的基本结构再往下看。1. 先搞清楚 ZeRO-3 到底优化了什么聊 MoE 之前必须先把 ZeRO-3 本身的原理掰开揉碎。很多人对 ZeRO 的理解停留在显存不够就开 ZeRO但 ZeRO 不是一刀切的开关不同级别的 ZeRO 解决的显存痛点完全不同。1.1 模型训练显存的去向一次正常的混合精度训练中显存主要被四样东西消耗模型参数本身、梯度、优化器状态比如 Adam 的一阶矩和二阶矩、以及激活值。以常见的 7B 参数模型跑 Adam AMP 为例参数量 14GB半精度梯度 14GBAdam 状态是 FP32 参数副本加两个动量项这部分要 84GB 以上。也就是说真正让显存爆炸的从来不是模型参数而是优化器状态。传统的数据并行DDP为了保证每张卡都在训练同一个参数的同一份副本每张卡要完整保存上述所有状态。ZeRO 的核心思路是这些状态本来就是冗余的干脆大家分担着存。ZeRO-1 只切优化器状态ZeRO-2 顺带切了梯度而 ZeRO-3 更进一步把模型参数也切成碎片分布到所有显卡上。这就是为什么 ZeRO-3 训练时每张卡见到的都是不完整的参数切片每次前向和反向都需要先集合碎片用完再释放。1.2 ZeRO-3 和 ZeRO-1/2 的本质区别ZeRO-1/2 的训练流程跟普通 DDP 很像参数每张卡都有副本通信发生在梯度同步阶段。ZeRO-3 则是把状态全部打散通信发生在每次前向/反向执行算子的间隙。我自己的理解是ZeRO-1/2 属于数据并行上的显存优化ZeRO-3 则是一个原生的参数分片系统。因为参数每层都在不同卡上所有需要完整参数的操作比如 Transformer 某一层的输入计算都必须先走一次 all-gather 拿到完整参数念一遍算完再丢掉。这带来两个直接后果通信量大幅增加。ZeRO-2 每个训练步的通信量大约等于两倍模型参数量ZeRO-3 则要翻好几倍每层前向一次 all-gather、反向一次 all-gather为了算梯度。显存拐点出现。ZeRO-3 的显存占用并不随模型增大线性下降而是趋近于一个阈值。如果你的显存恰好卡在临界值ZeRO-3 能解决如果你的模型已经大到单卡塞不下完整一层ZeRO-3 再配合 checkpointing 才谈得上有意义。1.3 CPU offload 是什么什么时候才需要ZeRO-3 常跟 CPU offload 绑定出现但我要泼一盆冷水能把参数留在 GPU 就别往 CPU 搬。CPU offload 本质是用 PCIe 带宽换显存把优化器状态甚至参数放在内存里每步训练通过 CPU 和 GPU 之间反复搬运。带宽瓶颈非常致命一开 offload 训练速度经常直接砍半以上。我的建议是如果模型只超出显存一点点优先试 ZeRO-3 本身如果超出好几倍再考虑 offload 优化器状态stage3_offload。参数 offload 除非万不得已不要开尤其是 MoE 场景专家参数本来就走 all-to-all 通信再加一层 CPU 搬运会乱成一锅粥。2. MoE 架构的核心与工程痛点MoE 火起来不只是因为它省算力而是它打破了模型越大每步计算越贵的固定关系。这里需要先厘清一个问题MoE 到底是什么以及那位著名热词MoE 架构要全部参数进显存吗到底在问什么。2.1 MoE 的本质计算量不随参数规模线性增长稠密 Transformer 的每一层 FFN 对所有 token 都做全部计算。MoE 的做法是把 FFN 替换成若干个并列的专家网络每个 token 只激活其中少数几个。总参数量很大但单 token 的计算量只跟激活的专家数有关。一个直观的比喻稠密模型像是一家餐厅每位顾客进来之后所有厨师都要做一道菜MoE 则是几十个厨师站在那顾客只选其中两三个下单其他厨师干站着。所以 MoE 模型的实际计算密度非常低通常专家数几十个甚至上百个但每个 token 只激活 2~8 个专家算力开销可控。这也是为什么同等算力预算下MoE 能堆出远超稠密模型的参数规模。2.2 专家并行与三种经典并行方式的协同训练 MoE 时常见做法是把专家分布在多张卡上称为专家并行Expert Parallelism。最经典的布局是每个专家完整保存在一张卡上token 按照路由结果被发送到对应的专家所在卡。这里面就引入了标准 Transformer 训练中不存在的通信模式all-to-all。各位如果熟悉数据并行、张量并行、流水线并行可以这样理解 MoE 的并行维度数据并行不同卡处理不同的数据通信量小。张量并行层的参数被切开通信发生在每层内部量大且频繁。流水线并行模型按层切开通信发生在层边界。专家并行专家分布在卡上专家之间是纵向的、动态的通信token 是流动的不是层的输出流动。专家并行的特点是通信量跟被路由的 token 数 × 隐藏层大小 × 专家数有关这跟传统并行完全不是一个量级也是最难调优的部分。2.3 回到那个热词MoE 参数到底要不要全部进显存MoE 架构要全部参数进显存吗这个问题答案是训练时只有一小部分参数必须在显存中但优化器、梯度和参数切片的存放策略决定了你是否需要全部进显存。如果不做任何优化把 MoE 当稠密模型再次填进显存那当然是要全部参数进的——可能光专家就几千亿参数。但是因为 ZeRO-3 和专家并行会把参数、梯度、优化器状态分散到不同卡上显存反而从全部进单卡变成全部进集群。实际训练中有几个典型场景需要具体讨论纯数据并行 MoE每张卡都要驻留全部专家参数。这种方案显存占用极差但工程是最简单的。如果模型专家部分不算太大小规模实验可以用这种方案。专家并行 ZeRO-3 非专家参数分片非专家参数attention、embedding走 ZeRO-3 分片专家参数按专家并行放在指定卡上。这是目前主流大厂训练的标配。专家参数 ZeRO-3 分片EP 和 ZeRO-3 混用两者可以结合但要区分在哪一层分片。如果专家也分片那么每次路由前都做 all-gather 收集整个专家参数进显存再算这几乎失去了专家并行的低通信优势。所以我的经验是训练 MoE 时不要指望 ZeRO-3 把一切都分片。专家的核心优势就是局部计算、局部存储ZeRO-3 把它全局化之后反而可能因为通信把所有收益都吃回去。3. DeepSpeed 训练 MoE 的完整配置与实操前面理论讲了不少这部分直接给配置、给命令、给代码。我们以 DeepSpeed 官方兼容的 MoE 接入方式来拆解。需要提前说明的是HuggingFace Transformers 里很多 MoE 模型比如 Mixtral已经兼容 DeepSpeed但在公开框架中做大规模专家并行训练还是得依赖 Megatron-DeepSpeed 这个组合。3.1 目标场景与硬件假设假设我手上有一个瘦身版的 MoE 模型总共 16 个专家模型包括 2 层 Transformer专家用两层 MLP 实现embedding 维度 1024隐藏层 1024。硬件假设 8 张 A100 80G。这个配置比起真实大模型已经缩水很多但足够把 ZeRO-3 专家并行的核心逻辑讲清楚也方便验证配置是否生效。3.2 最小可用的 ZeRO-3 显存估算方法训练前先做显存估算比直接跑靠谱得多。估算的核心公式是总显存占用 ≈ 参数总量 × 存储因子 激活值。不同优化策略下存储因子不同。一张卡上的参数与状态占用未考虑分片时 参数量 × 2字节FP16 参数量 × 4字节FP32 参数量 × 8字节Adam 动量和方差 参数量 × 14字节。以稠密 7B 模型为例单卡需要大约14GB × 7 98GB这个量显然超过 80G不做 ZeRO 肯定跑不了。如果用 ZeRO-3 切成 8 份单卡约98GB / 8 ≈ 12.25GB这就从容多了。对于 MoE 模型总参数可能 100B但实际显存需求不能简单按总参数量算因为大量专家参数并不在单卡内存。按专家并行方式单卡显存约等于非专家参数 × 14字节 / 卡数 本卡驻留的专家参数量 × 14字节 激活值。这个公式是判断能不能跑的最快路径。跑之前先算一算不要等人报 OOM 再救火。3.3 一份可跑通的 DeepSpeed 配置以下是适合小规模 MoE 训练验证的配置模型层数故意设少一点方便观察日志输出。{ train_batch_size: 16, gradient_accumulation_steps: 1, train_micro_batch_size_per_gpu: 2, fp16: { enabled: true, loss_scale: 0, loss_scale_window: 1000, initial_scale_power: 12 }, zero_optimization: { stage: 3, reduce_bucket_size: 5e8, stage3_prefetch_bucket_size: 5e8, stage3_param_persistence_threshold: 1e6, stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9, overlap_comm: true, contiguous_gradients: true }, communication_data_type: fp16, wall_clock_breakdown: false }stage3_max_live_parameters和stage3_max_reuse_distance这两个参数决定 ZeRO-3 在参数分片后单卡上最多保留多少参数不释放。如果调得太小参数反复 all-gather 和释放通信会爆炸调得太大显存驻留量高但通信频率低。我的首选直接给 1e9对大多数场景是折中的。overlap_comm默认可以打开它让通信和计算重叠对 ZeRO-3 的收益很明显几乎不用犹豫。3.4 专家并行部分的关键代码与集成假设我有一个 MoE 层需要把专家分布到不同卡上。用 DeepSpeed 的方式可以不手写通信调用moe相关的分布式 API。下面是一个简化版伪代码重点看整体流程。import torch import deepspeed from deepspeed.moe.layer import MoE class MoETransformerLayer(torch.nn.Module): def __init__(self, hidden_size, num_experts): super().__init__() self.attention MultiHeadAttention(hidden_size) # DeepSpeed 替我们管理专家分布 self.moe MoE( hidden_sizehidden_size, expertExpertMLP(hidden_size), num_expertsnum_experts, ep_size1, # 专家并行度 use_residualTrue, capacity_factor1.0, ) def forward(self, x): x self.attention(x) x self.moe(x) return x def train(): model MoETransformerLayer(hidden_size1024, num_experts16) engine, optimizer, _, _ deepspeed.initialize( modelmodel, model_parametersmodel.parameters(), configds_config.json ) # 正常训练循环ep_size表示每个专家被复制到几张卡上一般是 1也就是每个专家只驻留一张卡。如果你卡少比如只有 8 张但希望专家数 64ep_size1会导致一张卡要驻留多个专家需要检查显存。capacity_factor是 token 容量系数控制最多允许多少 token 被路由到某个专家。小于 1 会丢弃 token大于 1 会存储更多 token。训练时通常设置为 1~1.25推理时设置为 1。这个参数跟负载均衡直接相关后面详谈。3.5 从 8 卡扩展时的通信组配置当你从单机 8 卡扩展到多机时专家并行要特别注意通信组设置。DeepSpeed 的 MoE 涉及两个通信域数据并行域和专家并行域。数据并行域处理非专家部分专家并行域处理专家间的 all-to-all。如果ep_size 1专家组跨多机时 all-to-all 会走跨机网络带宽瞬间成为瓶颈。我亲测过在 10GbE 网络环境下跑 64 专家网络直接打满训练速度掉了 60% 以上。建议是单机 8 卡能跑完就别上多机。多机时专家尽量放在同一台机器内即ep_size不超过单机卡数。使用高速互联如 InfiniBand 或者高配 RDMA否则就是话费很大的工作。4. MoE 训练的独特难点负载均衡与通信优化训练 MoE 的难点跟稠密模型完全不同。稠密模型最主要的问题是全局梯度同步和显存MoE 还要多面对两个只有在 MoE 里才会出现的问题路由分布不平衡、token 流动引起的通信风暴。4.1 负载均衡为什么会失衡以及负载均衡损失怎么写路由器Router本质上是一个线性层每个 token 经它计算后被分配到某个或某几个专家。训练初期路由器可能被几个专家带偏导致 80% 的 token 都涌向同一个专家其他专家训练严重不足整体效率和显存利用率崩掉。解决的标准手段是负载均衡损失Load Balance Loss在路由器输出时施加正则约束每个专家接收的 token 期望尽量均匀。一个经典做法是每个专家接收的 token 比例接近 1/N专家被分配到的概率也要均匀。def load_balance_loss(router_probs, expert_indices, num_experts): # router_probs: [batch_size, seq_len, num_experts] # expert_indices: [batch_size, seq_len] tokens_per_expert torch.bincount(expert_indices.flatten(), minlengthnum_experts).float() prob_per_expert router_probs.sum(dim(0, 1)) / router_probs.sum() num_tokens router_probs.shape[0] * router_probs.shape[1] loss num_experts * (tokens_per_expert / num_tokens * prob_per_expert).sum() return loss这个损失函数的核心思想是如果专家 i 接收的 token 比例过高同时 Router 给它分配的概率也高这些乘积累加会让 loss 变大反向传播时会抑制 Router 的偏置。实际应用中这个 loss 还要乘一个系数比如 0.01否则会让 Router 过于均匀损害模型质量。我的经验是只看 tokens_per_expert 的方差是不够的概率和实际分配的乘积才能真正反映负载均衡的状态。之前踩过坑只用 token 计数做均衡结果 Router 的概率分布也变得很偏训练后期效果很差。4.2 专家容量、token 丢失与辅助损失的关系容量因子capacity factor在这里就很关键了。假设一个专家最多处理 C 个 token如果路由过来的 token 超过 C超出的部分会被丢弃训练阶段会用 padding 替代推理时会丢弃。容量因子越大token 丢得越少但计算浪费越多。从训练效果来看如果丢弃了太多 token模型会学不到这些 token 的信息。但容量因子又不能设得太大否则专家计算的稀疏性优势就没了。我一般先设 1.0 开始训练观察辅助损失和每个专家的平均 token 数再逐步调整到 1.2 左右。4.3 all-to-all 通信为什么贵如何缓解MoE 的 all-to-all 通信是整个训练流程中开销最大的环节。每步训练中每个 token 都要从它的源卡被发送到目标专家所在的卡。这种通信不是简单的集合操作而是细粒度的、跨任意卡的 token 分发。缓解手段主要有几种分组 all-to-all将 token 按目标专家分组打包减少通信次数而不是每个 token 单独发送。拓扑感知的专家放置把频繁接收数据交互的专家放在相邻或同机卡上。通信与计算重叠overlap在第一个专家计算时预取下一个专家的 token 传输。DeepSpeed 的 MoE 实现已经把 many-to-many 通信封装好了但你要注意all_to_all的 buffer 大小。如果capacity_factor设得太大buffer 也跟着变大通信时间和显存占用都会增加。4.4 梯度同步的细节共享参数与专家参数的不同处理MoE 模型里有两类参数共享参数非专家参数比如 attention、layer norm这些参数在数据并行维度上可同步直接走 ZeRO-3 或 DDP 的梯度 all-reduce。专家参数只有被路由到某专家的 token 才有梯度因此每个专家的梯度只存在于专家所在的卡上。如果 ep_size1那专家的梯度根本不需要跨卡同步除非有复制。这意味着训练中的梯度同步量大幅减少非专家参数同步专家参数只在本地更新。看起来是优势但也带来了一个隐性问题如果某个专家在某个 batch 里一个 token 也没收到那这个专家的参数就完全不更新长期会僵化。这就是为什么负载均衡损失除了防止集中外还间接保证每个专家都能获得足够的训练信号。5. 我的实战配置文件与代码级调优这一部分直接给出我实际用过的、可以照抄的 DeepSpeed ZeRO-3 MoE 配置并说明几个不起眼但特别有用的调参细节。5.1 实战级 deepspeed 配置8 卡{ train_batch_size: 32, gradient_accumulation_steps: 2, train_micro_batch_size_per_gpu: 2, fp16: { enabled: true, auto_cast: true, loss_scale: 0, initial_scale_power: 16 }, zero_optimization: { stage: 3, offload_optimizer: { device: none }, reduce_bucket_size: 5e8, stage3_prefetch_bucket_size: 3e8, stage3_param_persistence_threshold: 5e5, overlap_comm: true, contiguous_gradients: true }, communication_data_type: fp16, gradient_clipping: 1.0, steps_per_print: 10, wall_clock_breakdown: true }offload_optimizer我明确设成device: none不要让 DeepSpeed 默认把优化器丢到 CPU。尤其是有 MoE 的时monitor 成本太高。5.2 路由加噪声与 Top-K 的选择训练早期均衡性差一个常用技巧是Router 加噪声Noisy Top-K Gating。在计算路由器概率前加入可学习的高斯噪声让 top-1/top-2 的选取不那么固定有助于探索和均衡。def noisy_top_k_gating(logits, noise_epsilon1e-2): if training: noise torch.randn_like(logits) * F.softplus(logits) * noise_epsilon noisy_logits logits noise else: noisy_logits logits top_k_logits, top_k_indices torch.topk(noisy_logits, k2) return top_k_logits, top_k_indices训练完成后推理阶段必须关闭噪声。我习惯在前 5% 训练步里开高噪声之后逐步衰减到 0。这样前期的均衡性会好很多后期模型质量也不会受损。5.3 观察日志与训练健康度检查MoE 训练最怕看似在训练实则一部分专家在睡觉。我每次训练都额外打印这些指标每个专家的 token 计数按 batch 统计路由概率的熵top-1 和 top-2 的重合率如果前两个选择经常一致说明 Router 有点退化drop token 数量占 token 总量的比例过高说明容量因子不够这些指标在 DeepSpeed 日志里不会直接给出需要自己写 hook 统计。但它们的价值极高越早发现问题越好等训练跑了几百步才发现专家失衡返工成本很大。6. 常见问题与故障排查最后这部分是我最想写的因为 MoE 训练踩坑的密集程度远超稠密模型。下面的问题是实际训练中被问烂或者自己踩过的整理成速查表方便大家遇到类似问题时不至于从零开始。6.1 训练速度突然暴跌症状loss 掉得挺稳但 each step 时间比预期慢 2 倍以上。思路顺序先看是不是capacity_factor调太大导致专家 buffer 海量填充通信包太大。再看 all-to-all 是否走了跨机链路。多机时如果专家分布跨机速度掉得快很正常。最后看是否开了 ZeRO-3 的overlap_comm这个通常能救命建议保持 true。6.2 一个或多个专家的参数长期不更新症状monitor 里某些专家 token count 永远是 0 或者极小。原因基本就两个负载均衡失效或者初始化太差导致 Router 一开始就锁死了某些专家。建议按照 4.4 节的负载均衡损失重新调整并把噪声系数调高。另外可以考虑用一个简单的先验估计初始化 Router 的偏置让离线的专家初始概率不至于过低。6.3 OOM 到连 ZeRO-3 都救不了症状开 ZeRO-3 还 OOM或者每张卡的显存看着没到上限但峰值直接爆。排查点是否激活了stage3_max_live_parameters驻留参数过多。激活值是否过高。即使有 ZeRO激活值依然是独立于参数体系的压力源建议开 activation checkpointing。专家参数的分布是否均匀。如果某卡负载明显偏高考虑调整专家分配或减少该卡上的驻留专家数。6.4 通信占满网络训练几乎停滞症状多机训练时网卡吞吐打满单步延迟暴增。建议按优先级把 ep_size 调整为单机卡数让专家通信尽量走 NVLink。检查通信 buffer 的分配——capacity_factor设置过大直接导致跨机通讯量增长。考虑使用 DeepSpeed 的通信日志工具看看哪个通信算子耗时最长针对性处理。6.5 开 ZeRO-3 后精度浮动变大症状收敛不稳定loss 抖动比 ZeRO-2 跟 DDP 更剧烈。ZeRO-3 中参数分片会影响数值稳定性尤其当模型是非平滑激活函数或者梯度范数较大时。我的建议是打开communication_data_type: fp16之前先实测对比不要一概而论如果精度波动明显可以回退到 fp32 通信虽然慢一些但稳定不少。7. 最后的经验之谈如果你正准备在现成代码库里加入 MoE 训练我唯一的建议是先从最小化的 expert 并行开始不要一上来就 ZeRO-3 多机 load balance loss noisy gating 全开。我见过太多人一次性把所有优化打开结果定位问题的时候无从下手。正确的节奏应该是先跑通纯数据并行 单机小 MoE确认模型逻辑正确。进而加ep_size专家并行观察 all-to-all 是否正常。再加 ZeRO-3 分片非专家参数观察显存和通信变化。最后开启负载均衡损失和噪声调容量因子。每一步都观察显存、通信、每个专家的 token 分布确认正常后再进入下一步。这个流程可以帮你避开都不知道问题出在哪的局面。实际跑 MoE 训练这几年我最大的体会是它的难点不在理论而在工程。ZeRO-3 和专家并行单独看都不复杂但两者相遇时会互相干扰一个配置不对劲整个训练效率就会肉眼可见地滑坡。把自己当成拆雷专家把每项优化分步验算是最笨但也最有效的路。
返回列表