MoE这几年几乎成了大模型规模的代名词。做AI Infra的朋友应该都有体会,从GShard到Switch Transformer再到Mixture of Experts遍地开花,MoE架构凭借"参数多但算得少"的特性,让千亿万亿参数量不再是天文数字。但这玩意儿跟Dense模型完全不是一套玩法,路由计算怎么搞、专家并行怎么切、显存估算怎么做、负载均衡怎么控,每一个环节抠不细,训练的时候就等着哭吧。这篇文章我从Infra的视角把MoE的计算原理和全流程拆开揉碎讲清楚,也顺带解答一个大家最爱问的问题——MoE的权重到底要不要全部塞进显存。
我会把路由(Gate)机制、TopK选择、张量分发与合并(Dispatch/Combine)、负载均衡损失、容量因子、以及推理阶段的显存和调度考量全部串起来。不绕弯子,直接上干货。
1. 从Dense到Sparse:MoE到底改了什么
1.1 一个类比:从全科医生到专家门诊
理解MoE,最直观的方式是把一个传统的Dense大模型想象成一个全科医生。你不管生什么病,都挂同一个号,这个医生什么病都看,代价是他脑子里得装下所有科室的知识,每个科室的细节也做不到极致。
MoE的结构则更像一个专家门诊体系。外面坐着一个分诊台(Gate),里面是几十上百个专科诊室(Experts)。病人(Token)进来之后,分诊台看一眼症状,把病人分配给最对症的两三个诊室。病人真正接受诊疗的只是这两个诊室的医生,其他医生照样坐在诊室里,工资照发(参数占显存),但不开工(不参与计算)。
这背后对应MoE最核心的工程价值:参数量决定模型容量,激活参数量决定计算开销。Dense模型几乎每层每个参数都要参与计算,MoE模型则通过路由把计算量集中在少数专家上。这就是为什么大家都在冲参数量,但你实际算力吃紧的时候,MoE是唯一能让你"容量大、算得动"的架构。
1.2 MoE每个Token的口算流程
从数据流角度看,一个Token经过MoE层,实际只做四件事:
- 经过Gate网络(通常是一个线性层),得到它在所有Expert上的路由打分。
- 取TopK(通常是Top2),选出被激活的2个专家。
- 把Token的隐层向量复制分发(Dispatch)到被选中的2个专家上,各专家独立做FFN计算。
- 将2个专家的输出按照路由权重加权求和(Combine),得到该Token的最终输出。
整个过程对每个Token来说就像一次点到点的快递:分单(Gate打分)→ 分拣(TopK选择) → 派送(Dispatch) → 签收合并(Combine)。路由机制本身是一个深度学习的网络,打分不是硬编码规则,而是训练出来的。
看起来不复杂,但一旦放到分布式并行环境下,事情就没那么简单了。后文展开细讲。
2. 路由与TopK选择的计算细节
2.1 Gate网络打分的数学过程
Gate网络在MoE层中的位置是在Attention输出之后。假设我们拿到了一个Token隐层向量 (x)(维度是 (d_{model})),Gate会把它投影到所有专家上,产出路由分数 (score)。常见的公式如下:
[ score = \text{Softmax}(\text{TopK}(W_g \cdot x)) ]
其中 (W_g) 是一个形状为 (d_{model} \times N) 的矩阵,(N) 是专家数量。这里两个关键点是Softmax和TopK的执行顺序。实际工程中,先对 (W_g \cdot x) 做TopK,选出了 (K) 个专家的logits之后,再对这K个logits做归一化,得到路由权重。
这个顺序很重要,原因在于:如果先对全部N个专家的logits做Softmax,数值会被未选中的专家的logits拉低,而且Softmax整个N维的计算量在专家数量很大的时候也是额外开销;先TopK再Softmax既省算力,又能保证路由权重是在"候选集内部"进行归一化,语义更干净。
2.2 TopK选择中的排序与切分边界
实际代码里,最早的TopK实现在PyTorch中用torch.topk函数直接取。后来大家逐渐发现自带的topk在大Tensor上性能不够,因为它的排序算法在小K需求下不够"吝啬"。工程上常用近似策略:只计算TopK个元素,不去做全量排序。做法是先把Logits和阈值比较,用torch.where掩盖掉所有低于可接受阈值的项,或者直接用torch.functional.topk然后接受它的排序开销——实测在专家数小于512时,torch.topk的随机量基本可以忽略。
路由权重的数值形态也很值得注意。在推理阶段,K个专家的权重会直接作为Combine阶段的加权系数;在训练阶段,权重还会参与反向传播,因此Gate网络可以学习到更合理的打分策略。
2.3 路由中的噪声与训练稳定性
Switch Transformer那篇文章里有一个重要细节:训练时在路由打分中加入噪声。原因在于纯Greedy的路由策略会让Gate网络陷入"强者愈强"的状态——某些专家一旦前期被打高,后期就会收到更多训练信号,优势持续叠加。加入高斯噪声后,打分会产生合理的波动,让处于边缘的专家也有机会收到Token,避免训练过程彻底锁死在少数专家上。
具体实现上,噪声不是随便加的,而是用torch.randn_like(logits)再乘以一个可学习的标准差(由Gate的另一个头产生),再加到原始logits上。这样噪声幅度是学习的,而不是拍脑袋定死的,能自适应地平衡探索和利用。
3. Dispatch与Combine的工程技术
3.1 为什么Dispatch是MoE性能的关键
很多初学MoE的人不理解"分发"这类操作有什么好讲的,一个torch.gather不就解决了吗?问题出在:专家的计算是独立的,但Token的流动是批量的、并发的。常规Dense计算中,整个batch是一个规整的矩阵乘法,GPU把所有线程喂给同一个算子,利用率极高。
MoE中每个Token要去不同的专家,这意味着GPU必须处理不规则的数据流。如果用朴素的for循环逐个专家跑,极度浪费算力;用gather把Token按专家分桶吧,又得引入大量内存拷贝和显存碎片。工程上Dispatch的语义就是:把输入张量从"按Token排列"重排为"按专家排列",让每个专家拿到属于它的连续Token块,从而让处理每个专家的GEMM也能高效执行。
3.2 Combine阶段的加权求和与动态形状
Combine是Dispatch的逆操作。每个专家计算完输出后,需要把输出按原本Token的顺序归还,并乘上该Token对应该专家的路由权重。真正的实现里,Combine通常需要维护一张索引表:dispatch_mask,记录每个Token去了哪些专家、顺序是什么,这样Combine时才能精准还原。
这里有个常见的误解:Combine是否等于scatter_add?由于每个Token通常被多个专家处理,而且最终输出是多个专家输出的加权求和,因此确实可以写成:
[ \text{output}[t] = \sum_{e \in \text{selected}(t)} \text{weight}{t,e} \cdot \text{expert_output}{e} ]
翻译成工程操作就是:为每个Token维护一个加权的输出缓冲区,然后循环向里累加。但Cycle的起始顺序必须跟Dispatch完全一致,否则数据就乱了。初学阶段建议直接用官方实现来理解,不要自己去发明分发算法。
3.3 Capacity Factor:提前防御性能雪崩
上面讲过,每个专家处理Token的数量是动态的,极端情况下可能会有大量Token挤向同一个专家。如果不加限制,这个专家所在设备就会变成热点,其他专家闲置,整个Step的算力极具浪费,甚至单卡显存直接撑爆。
因此MoE的实现中引入了一个**容量因子(Capacity Factor)**概念:
[ \text{capacity} = \lceil \frac{T}{N} \times \text{capacity_factor} \rceil ]
其中 (T) 是Token总量,(N) 是专家数量。它定义了每个专家最多能处理的Token上限。capacity_factor设为1时,理论上每个专家按照均匀分布恰好分到均等的Token,但实际路由必然不均匀,所以通常设为1.25到2之间。
超过额度的Token并不会进入专家计算,而是被直接丢弃(Drop),经过残差连接直接传输到下一层。Capacity Factor越小,算力越稳定,但被丢弃的Token越多,模型质量下降越明显。这个参数调起来很Fun,属于典型的"算力与效果"天平。
4. MoE参数都需要进显存吗
4.1 MoE训练时的显存构成分解
聊到热点问题——"MoE架构要全部参数进显存吗?"先说结论:理论上训练阶段权重、梯度和优化器状态必须能在设备上寻址,但可以通过张量并行、专家并行或参数卸载,把"同时"存活于单卡的参数量降下来。推理阶段更现实,模型权重要么整机装下,要么动态换入换出。
训练时的显存构成主要包括四块:
- 模型权重(Weight):所有专家的权重矩阵
- 梯度(Grad)
- 优化器状态(Adam里的Momentum和Variance)
- 激活值(Activation)
前三者跟参数量成正比,个体量级的估算方式其实很简单:
[ \text{显存字节数} = \text{参数量} \times \text{每参数字节数} \times \text{状态份数} ]
使用Adam、FP16混合精度训练时,一份权重FP16占2字节,FP32的主权重占4字节,动量和方差各占4字节,梯度FP16占2字节。也就是一个参数在训练时大约需要 (2+4+4+2 = 12) 字节。你可能看到有的文章说16字节,那是把梯度也算成FP32了;在主流AMP场景下,12字节是一个更贴近现实的估算值。
所以,如果在单卡上放一个100亿参数MoE模型,参数量100亿×12字节 ≈ 120GB,绝大多数单卡目前都放不下。
4.2 Expert并行:把专家分布到不同设备
那怎么办?最自然的解法就是让不同专家放在不同GPU上。因为MoE本身是Sparse的,专家之间不需要参与同一个矩阵运算,天然适合设备并行。
具体来说,MoE的Expert层会做Expert Parallel切分:模型的其他层(例如Attention层)复制到所有设备(或用Tensor并行切分),而专家们均匀分配到各个设备上。每个Token的Hidden状态先经路由选择,再去对应设备上找专家计算,再把结果发回去Combine。
这套方案背后最核心的代价是通信。每个Token的所有关键信息都要走一次设备间的点对点通信,Token总量大的场景下All-to-All通信会成为瓶颈。训练超大模型时,一个MoE层的一轮迭代可能要传输几十GB的数据,网络不改成高速互联基本动不了。
所以"要不要全部参数进显存"这个问题的工程翻译其实是:在Expert Parallel下,单卡只进单卡负责的专家,整个模型依然完整分布在集群中;但如果你只有一张卡且想跑全量模型,那就必须引入检查点卸载(Offload)和动态重载,那才是真正的极端场景。
4.3 推理阶段的显存与带宽约束
推理阶段相比训练,少了梯度和优化器状态,显存压力显著下降,但推理的瓶颈往往不在显存而在带宽和时延。假设你在单机上部署MoE,权重全放显存,Feed一个Token需要激活2个专家的全部FFN权重参与计算。如果用的是大专家(比如每个专家约70亿参数),激活参数量也很大,计算延迟自然高。
针对推理还有两个常见优化:
- 专家权重按需加载(Offload):当显存不够时,只保留Gate层的权重在显存,专家权重放宿主内存或SSD,用的时候动态换入。代价是每次都要走PCIe或NVMe带宽,延迟感人。适合吞吐优先而不怎么要求低延迟的离线批量推理。
- 量化压缩:把专家权重从FP16压到INT8或INT4,显存减半甚至减到四分之一,带宽占用也同步降低。对MoE来说量化专家权重的损失往往小于量化Attention部分,因为专家输出还会被路由权重加权组合,误差在一定程度上被平均掉了。
我的实际建议是:推理阶段尽量把参数都放显存,但不要盲目追求单卡塞满,而是按带宽规划和时延预算来定。超过单卡显存容量,优先考虑多卡专家并行而不是offload,除非你的业务根本不Care首Token延迟。
5. 负载均衡的工程实现与调参心得
5.1 负载均衡损失的计算逻辑
MoE如果不加干预,路由分布经常会偏向少数专家。这就是"负载均衡"这个热搜词背后大家关心的问题:如何让每个专家尽可能收到差不多的Token量。
业界最普及的方案是给损失函数加一个辅助负载均衡损失项(Load Balance Loss),它衡量的是路由分布的均匀程度。Switch Transformer给出的经典版本公式如下:
[ \mathcal{L}{\text{balance}} = \alpha \cdot N \cdot \sum{i=1}^{N} f_i \cdot P_i ]
在上面的公式中,(f_i) 表示第 (i) 个专家接收到的Token数占总Token数的比例,(P_i) 表示所有Token分配给第 (i) 个专家的平均概率(由Gate的Softmax结果统计得到)。如果两个分布都是均匀的,乘积累加后接近1;如果路由高度集中,值就远大于1。乘上专家数N,是为了让损失尺度跟专家规模解耦。
(\alpha) 是平衡系数,实践中通常在 (0.01) 到 (0.1) 之间调整。调参时注意:(\alpha)太小,负载均衡约等于没约束,热点专家依然会出现;(\alpha)太大,负载是均匀了,但路由选择的质量会被破坏,模型效果反而变差。这跟之前讲的Capacity Factor一样,是一对需要配对调节的参数。
5.2 负载均衡损失的计算代码
下面这段代码是经典实现,数据形状我都标注清楚,方便你对号入座。假设router_logits是Gate网络输出的路由打分(shape为[T, N]),target_experts是每个Token被选中的专家ID(shape为[T, K]):
import torch def load_balance_loss(router_logits, target_experts, num_experts, alpha=0.01): """ router_logits: [T, N] 每个Token对所有专家的打分 target_experts: [T, K] 每个Token实际选中的专家ID """ T, N = router_logits.shape # 1. 计算每个专家实际接收的Token比例 f_i # 统计每个专家被选中的次数 counts = torch.zeros(N, dtype=router_logits.dtype, device=router_logits.device) counts.scatter_add_(0, target_experts.view(-1), torch.ones_like(target_experts.view(-1), dtype=router_logits.dtype)) f_i = counts / (T * target_experts.shape[1]) # 2. 计算所有Token分配给每个专家的平均概率 P_i # 对router_logits做Softmax得到概率,然后取均值 probs = torch.softmax(router_logits, dim=-1) p_i = probs.mean(dim=0) # 3. 计算负载均衡损失 loss = N * torch.sum(f_i * p_i) * alpha return loss这个代码的思路非常直观:一个专家既接收了很多Token(f_i大),又被Gate经常高概率选中(p_i大),就会显著推高损失,反向传播时优化器就会去压制Gate网络对热门专家的偏好。
5.3 工程中经常搭配的Router Z-Loss
在实际训练中,我发现只加load_balance_loss还不够。Gate网络的打分绝对值经常会出现振幅失控的问题——Logits随着训练越变越大,Softmax的分布就越来越尖锐,最后几乎变成One-Hot,对于选中专家的权重稳定性和反向传播都不友好。
解决方法就是Router Z-Loss。它把Logits的平方和压低,公式非常简单:
[ \mathcal{L}{z} = \beta \cdot \frac{1}{T} \sum{t=1}^{T} \left( \log \sum_{i=1}^{N} \exp(logits_{t,i}) \right)^2 ]
实现代码如下:
def router_z_loss(router_logits, beta=1e-3): logsumexp = torch.logsumexp(router_logits, dim=-1) loss = torch.mean(logsumexp ** 2) * beta return loss不过这股Logits膨胀的问题是真实存在且高发的,Z-Loss带来的效果非常稳。实际训练日志里,添加Z-Loss后先看Gate的logits标准差曲线,能观察到它从几千掉到几十的范围内,模型效果也会更稳定。顺便说一句,Z-Loss只用对路由打分本身做约束,不需要跟标签做交互,所以实现简单,代价几乎为零。
而且在分布式训练曲线出现大幅抖动时,先别急着调学习率,检查一下Router打分是否已经爆炸是个常被忽略的止损点。
6. 异步训练中的MoE负载不均衡排查实录
6.1 我踩过的一个典型热点问题
有一次我在训练一个8专家版本的MoE模型,发现训练曲线一切正常,但集群监控里总有几个GPU利用率拉满,另几个GPU只有30%出头。一开始我还以为网络通信有问题,后来把每个专家收到的Token数打出来,发现Expert 3收到了总Token量的41%,Expert 7只收到2%。典型的负载不均衡。
排查步骤我按顺序做了:
- 先把
load_balance_loss的alpha从0.01提到0.1——效果有限,专家分布虽然有改善,但热点依然存在。 - 查看Gate logits的分布,发现Logits数值越来越大且方差爆炸——判断Gate网络进入了"高置信度"模式,也就是前文说的Logits膨胀问题。
- 加入Z-Loss,把beta设为0.001,重新训练,观察到Load Balance Loss的曲线从2.0附近掉到1.1左右,各专家Token占比趋于均匀。
这个经历让我彻底改变了对MoE训练稳定性的认知:负载不均衡不只是推个Loss就算完,Gate的数值健康度更需要监控。否则你调再大alpha也只是把Frank专家打压一轮,下一轮照样换一个人当老大。
6.2 专家网络"废掉"与Token Drop的表象与真相
还有一种常见问题是某个专家几乎收不到Token。这时候target_experts的count统计接近0,该专家对应的权重长期得不到有效梯度更新,基本就"废掉"了。如果专家直接弃用,容量倒是省了,但参数的浪费也意味着模型容量缩水,跟当初上MoE的初衷背道而驰。
我见过把capacity_factor从1.0提到1.5后,原本被Drop的边缘Token能正常进入专家计算,这些边缘Token恰好又是那些冷门专家唯一的学习信号。所以先用大Capacity Factor保住整个Expert不荒废,同时用负载均衡Loss拉回热门Token,才是正确的组合拳,顺序反了效果完全不同。
在排查的时候,务必把每个Step的"每个专家的实际Token数"和"Drop的Token数"同时打出来看。Token数暴露是否均衡,Drop数暴露容量是否够用,这两者结合判断是"负载问题"还是"容量问题",少走很多弯路。
6.3 从训练到推理:MoE负载不均衡的另类体现
训练阶段的负载不均衡大家容易察觉,推理阶段其实同样需要重视。线上业务流量不是均匀的,MoE路由对不同的上下文可能呈现完全不同的路由分布。某个专家在闲聊类query上接收大量Token,另一个专家在技术类query上表现更活跃。
如果按照离线评测时的均匀分布做专家切分和显存规划,上线后很可能出现单卡过载。所以我一般在模型部署之前会先跑一批代表性业务数据的路由分布分析,用统计结果指导专家放置。方法很土,但确实能预防线上P0事故。
7. MoE的通信瓶颈与分布式优化方向
7.1 All-to-All通信的本质
前面反复提及Token在专家间分发,在分布式环境下这个操作对应的是All-to-All通信。具体场景是这样的:设备0上有Token组A_0,需要发给设备1的专家;设备1上有Token组A_1,需要发给设备0的专家。两者同时互发,就构成All-to-All。
这种通信模式最大的痛点是流量随Token数量线性增长。假设切分了16张卡,每个Token都要被发送到选中的专家所在的设备,每层MoE都可能产生与Batch大小成正比的通信量。网络带宽不够时,通信时间会直接超过计算时间,让GPU计算单元处于闲置等待状态。
Infra层的优化方向一般有两个:
- 减少通信次数:把多个MoE层的通信合并为一个大的All-to-All包,减少握手和框架开销。
- 通信计算重叠:利用Pipeline思想,让第N层的通信和第N+1层的计算重叠进行。用NCCL的Async接口很容易实现,但复杂度不低。
7.2 专家选择的本地性优化
另一个被忽略的调度优化方向是"路由亲缘性"。如果模型在16卡上跑专家并行,理论上一个Token可能被打到任意一个专家上。但在训练数据自然分布下,相邻位置的Token往往语义相关,路由目标会有一定聚类特征。我们可以利用缓存机制,让Gate网络输出的专家ID尽量命中当前设备已有的专家,避免跨设备路由。
不过这种优化在目前通用的实现里基本没做,因为收益不稳定。在小Batch训练中可能有些微效果,但大Batch下基本没有规律可循。我这里只提一句,免得有人拿这个方向去瞎优化浪费感情。
7.3 动态专家并行与专家容量规划
工程上还有个比较现实的问题:专家并行下,每个设备的负载取决于被分发到该设备的Token数量。而Token数量是动态的,所以各设备的显存使用也是动态的。极端情况下,一个设备因为热点专家收到过多Token,可能临时出现Out Of Memory。
最优解当然是让GPU粒度足够细,专家规模足够平均,但现实是专家数由模型结构定死,设备数由集群定死,能优化的只有Capacity Factor和Token的分配策略。务实的心态是把MoE的分布式调度看作一个流量调度问题:保守设定Capacity,让每个设备有20%-30%的弹性空间,既保吞吐也保安全。
我在做集群级容量规划时,用的粗估公式如下:
超发比例 = 1 + {\text{平均每个专家负载波动率}} + {\text{Drop率}} \times \text{buffer系数}
经验值通常是在理论均值基础上多加40%到60%的容量冗余,确保负载均衡失效时不会直接OOM。
8. MoE调参和实验记录的真实经验
8.1 训练初期的抖动不容忽视
MoE训练在最初几千步经常出现Loss大幅波动,比Dense模型更明显。原因在于Gate网络还在"试错"阶段,路由策略不稳定,每个专家的训练信号差异巨大。这种前期的随机性有时候会让冷启动专家彻底学不到东西。
我的调参做法是:
- 前2000步把
load_balance_loss的alpha调大一个数量级,让路由快速均匀化。 - 2000步之后再把alpha降下来,让专家各自去专注差异化。
这个方法有论文支持(ST-MoE的StableMoE设计思路),实测能让最终效果提升1-2个点。
8.2 测量路由分布的工具化
对于Infra工程师来说,最直观的监控不是看Loss,而是直接盯着路由的分布直方图。我们团队的做法是在训练框架里加一个Hook,每500步打印一次所有专家收到的Token数量、平均路由权重、以及Token Drop率。你把这些数值拉成曲线,可以非常直观地发现"某个专家的利用率跌到0"或者是"Gate的分给Target专家的比例是99%"之类的问题。
8.3 从8专家到64专家的扩展陷阱
最后我再多说两句专家的扩展性。Small模型上三层专家8个,一切正常;当你把专家翻到64个时,会遇到几个隐藏问题:
- Gate网络的参数量占比虽然极低,但是输入的Token数量很大,logits存储和计算也要占显存。
- 专家变多之后,TopK选取的候选数量不变,K=2和K=1的效果差异会变大,需要同步调整。
- 专家并行时,All-to-All的通信矩阵从N×N膨胀到64×64,调度开销剧增,通信占比可能超过计算。
这些问题没有一个好用的银弹,只能做实验前就规划好。搞Infra的就得有这种觉悟,模型结构上一个参数的改动,背后都是一场分布式计算的重新排兵布阵。
9. 踩坑记录与工程自查清单
整套MoE的工程模块跑通之后,我总结过一份自查清单,每次新任务或新集群都会对照着检查一遍。这里分享出来,希望能帮读者少交学费:
9.1 显存规划
- 是否确认了训练精度(FP16/BF16/FP32)对应的单参数字节数?
- 是否把权重、梯度、优化器状态都计入了显存?
- 是否给激活值和临时通信Buffer留了20%以上的冗余?
- 专家并行的切分是否考虑到了容量因子带来的额外Buffer?
9.2 分布式通信
- All-to-All通信的带宽是否经过压测验证?
- 各设备的网络拓扑是否支持全互联?树状拓扑在All-to-All场景下非常吃亏。
- 梯度AllReduce和Expert的Token通信是否互相挤占带宽?建议分优先级。
9.3 路由健康度
- 是否监控了每个专家的Token接收量?
- 是否监控了Gate Logits的数值范围?
- Load Balance Loss的曲线是否在合理区间?
- Token Drop率是否低于1%?高于5%就需要调整Capacity Factor。
9.4 训练稳定性
- 是否加入了Router Z-Loss?
- 是否设置了前2000步的高alpha预热期?
- 是否周期性保存了Gate网络的checkpoint?路由发散的情况下这是回滚的关键。
上面这些条,每一条背后都是我或者同行实打实填过的坑。MoE给了你模型容量的自由,也附赠了分布的复杂度。搞AI Infra的,既要懂模型优化,也要懂集群调度,这两只脚缺一只都站不稳。