3.12 混合专家MoE

1.参数多、算得省:稀疏激活是怎么来的

说起来,前面几章我们讲了不少关于模型容量和计算量的内容。一层又一层的全连接、注意力、卷积,参数加起来越多,模型能记住的东西、能拟合的函数也就越复杂。可参数一多,推理的时候每一笔计算都要走一遍,算力开销跟着水涨船高。这就出现了一个挺折磨人的矛盾:我们想要一个很大的模型来装下海量的知识,又不想每次推理都把整张网从头到尾算一遍。

混合专家(Mixture of Experts,业内一般直接叫MoE)就是冲着这个矛盾来的。它的核心想法其实挺朴素:模型的总参数可以做得很大,但每一次推理只激活其中一小部分来真正参与计算。这一小部分通常叫做专家(expert),每一个专家本身就是一个子网络,多半就是一个前馈网络。输入来了,系统只挑出几个最合适的专家去算,其余的专家这一轮就歇着。

这就是稀疏激活(sparse activation)。说穿了,密集模型(dense model)是每次都把所有参数都跑一遍,MoE则是按需取用。这样一来,总参数量可以做到几千亿,可单次推理真正动用的算力,只相当于其中一小部分。小张之前推荐给我一本讲分工的经济学小书,里头说一个工厂里不一定每个工人每时每刻都在干活,关键要让合适的人去做合适的活计,MoE大概就是这个味道。

2.路由:一个小的调度员

那么问题来了,输入来了,到底该让哪几个专家去处理呢。这件事交给一个小网络来决定,叫路由(router),也有人称它为门控(gate)。路由本身很轻量,通常就是一层线性变换加一个softmax。

我们把输入记作 xxxx 就是当前这一段要处理的向量(比如一个token对应的表示)。一共有 NN 个专家,NN 表示专家的总数,常见的有 88 个、1616 个这类数字。路由先算出当前输入对每一个专家的偏好分数 gi(x)g_i(x)ii 是专家的编号,取值从 11NNgi(x)g_i(x) 的计算方式大致是这样:

gi(x)=softmax(Wgx)ig_i(x) = \text{softmax}(W_g \, x)_i

这里 WgW_g 是路由自己的参数矩阵,softmax\text{softmax} 是一个把一组数归一化成概率分布的函数(所有输出加起来等于 11),gi(x)g_i(x) 就表示输入 xx 分配给第 ii 个专家的权重。

光有分数还不够,我们还要决定具体让谁上场。最常见的是top-kk 选法,kk 就是从 NN 个专家里挑出得分最高的那几个。比方说Mixtral那篇工作,N=8N=8k=2k=2,每次推理从8个专家里挑2个最合适的来算。最后这一段输入的输出 yy 就是这被选中的几个专家的加权求和:

y=itop-kgi(x)Ei(x)y = \sum_{i \in \text{top-}k} g_i(x) \, E_i(x)

这里 Ei(x)E_i(x) 表示第 ii 个专家对输入 xx 算出来的结果,\sum 是求和符号,意思是在被选中的那 kk 个专家里,把每个专家的输出按它的权重 gi(x)g_i(x) 加起来。没被选中的专家,它们对应的 gi(x)g_i(x) 直接当作 00,对应的 Ei(x)E_i(x) 根本不会去算,这就把算力省下来了。还要补一句,前面那个softmax是对全部 NN 个专家做的,权重和为 11,可一旦只留下 kk 个、其余置 00,这 kk 个权重加起来就不再等于 11,输出 yy 会被等比缩小。实际实现里(比如Mixtral、Switch Transformer)通常只在被选中的这 kk 个专家上重新做一次softmax,让这 kk 个权重重新归一化和为 11,这样就消除了整体缩小的问题。

这个过程有点像医院分诊。一个病人进门,分诊护士先简单判断一下,再决定把他送去内科、外科还是骨科这几个对口科室,让对应的医生去处理。护士本身不做治疗,她只负责指路。路由在这里扮演的就是这位分诊护士。

3.负载均衡:不能让几个专家忙死,其余闲死

路由这么一安排,听着挺漂亮,可真训起来马上会遇到一个老大难问题,叫负载不均衡。

事情是这样的。路由在训练初期是随机初始化的,哪几个专家碰巧一开始表现好一点,路由就会更倾向于把任务派给它们。这几个专家被派得越多,练得也越多,表现就更好,路由就更愿意派给它们。这是一个典型的正反馈循环,说穿了就是富者愈富。最后很可能出现这种局面:8个专家里就这么两三个被疯狂使用,剩下的几乎从来没被选中过,等于白养着。

这件事有两个坏处。一是浪费参数,闲着的专家占着显存却不干活。二是训练效率低,忙的那几个专家成了瓶颈。我们当然希望每个专家都被合理地使用到,任务分摊开。

怎么做到呢,常用的办法是加一个辅助损失(auxiliary loss),叫负载均衡损失。它的思路是逼着路由把输入均匀地分给所有专家。设一段训练数据里一共有 TT 个输入,TT 表示token的总数,第 ii 个专家被选中的比例记作 fif_i,路由给第 ii 个专家的平均权重记作 PiP_i。那么负载均衡损失大约长这样:

Laux=αNi=1NfiPiL_{\text{aux}} = \alpha \, N \sum_{i=1}^{N} f_i \, P_i

这里 α\alpha 是一个用来调节这个损失比重的超参数,NN 还是专家总数,\sum 同样是求和符号。当每个专家都被均匀使用、均匀打分的时候,fif_iPiP_i 都接近相等,这个损失达到最小。一旦路由偏心,少数几个 fif_iPiP_i 特别大,损失就涨上去。模型在训练时为了把这个损失压低,自然就会学着把任务均匀分摊。我记得有部动漫讲一群人组队打怪,队长要是一直把任务派给最强的两个队友,剩下的人很快就会边缘化、士气低落,最后整个队伍反而更弱。负载均衡损失起的正是这种提醒作用,逼着队长把机会轮转开。

4.参数做大、算力可控:MoE在大模型里大行其道

稀疏激活的妙处,在近几年的大语言模型里被发挥得淋漓尽致。说起来,模型要装下那么多知识、要应付那么多任务,参数量小了确实吃力。可要是真把密集模型做到几千亿参数,每次推理都得跑一遍,那张显卡的成本谁也吃不消。MoE正好让这两件事各取所需:总参数量上得去,单次推理的算力又压得住。

业内几个响当当的模型都用了这套思路。Mixtral是Mistral那边的开源模型,N=8N=8k=2k=2,总参数量不小,但单次推理只激活其中约四分之一。DeepSeek系列也大量使用了MoE结构,还在路由上做了不少改进,比如细粒度专家和共享专家这些设计(细粒度说的是把专家切得更小、数量更多,共享专家指的是有些专家无论什么输入都会被激活,专门负责那些大家都需要的通用能力)。还有Jamba这种把Mamba(一种状态空间模型)和Transformer混着用的架构,里面也掺了MoE层来扩容。

为什么大家这么爱用MoE呢,大概有这么几条原因。第一,扩参数量的时候,MoE扩的是总参数,单次计算量涨得没那么快,性价比高。第二,专家之间多少学到了一点分工,有的专家更擅长某类输入,有的更擅长另一类,整个模型的表达能力更丰富。第三,推理时的算力可控,部署起来相对友好。说起来,这跟前面章节讲的容量和效率的权衡是同一脉络的延续,只不过MoE换了个更巧妙的切入点。

5.代价:路由难训、通信昂贵、显存照旧

好处说完了,至于代价,MoE从来不是免费的午餐,它带来的麻烦主要有三样。

头一样是路由不好优化。路由本质上是一个离散的决定,选还是不选某个专家,这种选择没法直接拿梯度去调。top-kk 这一挑是一个不可导的操作,得用一些技巧(比如上面那个softmax和负载均衡损失)把梯度间接传回路由。训练初期路由很容易抖来抖去,专家之间分工迟迟稳不下来,损失曲线一阵乱跳。整体训练的稳定性,比同规模的密集模型要难照看不少。

第二样是多卡训练里的通信开销。一个几千亿参数的MoE模型,单卡装不下,得把不同的专家分布到不同的显卡上,这种做法叫专家并行(expert parallelism)。可问题是,每一次前向计算,路由都要把输入送到对应专家所在的那张卡上去算,算完再把结果送回来。这一来一回的通信,在大规模训练里开销相当可观,有时候比计算本身还费时间。打个比方,这就像一个公司在多个城市有办公室,每个办公室擅长不同的业务,客户来了要先把需求发到对口的办公室,办完了再把结果汇总寄回去,光路上跑的时间就不少。小明之前跑过一个MoE的小实验,单卡上好好的,一上多卡,吞吐量没见涨多少,倒是一堆时间花在等通信上,调了好久才把布局理顺。

第三样是显存。虽然每次推理只激活几个专家,可所有专家的参数都得老老实实待在显存里,因为谁也不知道下一个输入会轮到谁。也就是说,MoE省的是计算,没省存储。显存该装多少还是得装多少,部署大MoE模型的硬件门槛依然不低。

把这三样放一起看,MoE其实是在用更高的工程复杂度,换取参数容量和单次算力之间的那个理想平衡点。它确实让大模型这条路走得更宽了,可也意味着训一个好用的MoE模型,需要的工程功夫一点不比密集模型少。

练习

Q1. MoE 的"稀疏激活"到底是什么意思,它解决了什么矛盾?

稀疏激活是说模型总参数可以做得很大,但每次推理只激活其中一小部分专家(子网络)来真正参与计算,没被选中的专家这轮歇着。它解决的矛盾是:既想要很大的模型装下海量知识,又不想每次推理都把整张网从头算一遍。于是总参数量上得去,单次推理动用的算力却只相当于其中一小部分。

Q2. Mixtral 取 N=8N=8 个专家、k=2k=2,每次推理大概激活了总参数的多少?输出 y=itop-kgi(x)Ei(x)y=\sum_{i\in\text{top-}k}g_i(x)E_i(x) 里,没被选中的专家会怎样?

8 个专家里挑 2 个,每次大约激活四分之一的参数(粗略说约四分之一)。没被选中的专家对应的权重 gi(x)g_i(x) 直接当 0 处理,对应的 Ei(x)E_i(x) 根本不会被去算,算力就这么省下来了。实际实现里通常只在被选中的这 kk 个专家上重新做一次 Softmax 让权重归一化和为 1,避免整体被等比缩小。

Q3. 有人说"MoE 既省算力又省显存",这话对吗?

只对了一半。MoE 省的是计算:每次推理只跑几个专家。但它没省存储——所有专家的参数都得老老实实待在显存里,因为谁也不知道下一个输入会轮到谁。所以部署大 MoE 模型的显存门槛依然不低。

Q4.(面试题) MoE 训练里最头疼的负载不均衡是怎么产生的?怎么解决?多卡训练时通信开销又来自哪里?

负载不均衡是个正反馈循环:训练初期路由随机初始化,哪几个专家碰巧表现好,路由就更倾向派给它们;被派得越多练得越多、表现越好,路由更愿意派给它们,最后富者愈富,少数专家忙死、其余闲死,既浪费参数又让忙的几个成了瓶颈。解决办法是加一个负载均衡的辅助损失 Laux=αNifiPiL_{\text{aux}}=\alpha N\sum_i f_iP_i,当每个专家被均匀使用、均匀打分时它最小,逼着路由把任务均匀分摊。多卡训练时为了装下几千亿参数,得把不同专家分布到不同卡上(专家并行),每次前向路由都要把输入送到对应专家所在的卡上算、算完再送回来,这一来一回的通信在大规模训练里开销相当可观,有时比计算本身还费时间。

相关标签
深度学习MoE混合专家大语言模型