← 返回首页
读书笔记

MoE 基础

提到2024 年末火出圈的 Deepseek 模型,很多人第一反应就是其为大模型奠定了正确的MoE (Mixture of Experts, 混合专家模型)路线。在大模型狂飙突进的今天,MoE已经从一个学术界的“古老”概念,一跃成为了 DeepSeek、GPT-4 等顶尖模型背后不可或缺的动力源泉。

站在巨人的肩膀上:为了深入了解 MoE ,我最近花了不少时间学习苏剑林MoE 相关的博文,并对一些工作进行了深入了解,受益匪浅。本篇博客是学习苏神博客和一些论文的笔记。

从几何角度看MoE的原理

我们以 FFN 为例子,来介绍大模型的的 MLP 层:

y=f(xW(A))W(B)y=f (x W^{(A)})W^{(B)}

其中 xRdx \in \mathbb{R}^d 是输入向量,W(A)Rd×DW^{(A)} \in \mathbb{R} ^ {d \times D}, W(B)RD×dW^{(B)} \in \mathbb{R} ^ {D \times d} 是参数矩阵,假设 nn 是可以整除 DD 的整数,上述公式可以等价地用分块矩阵写成:

y=f(x[W1(A) W2(A) ... Wn(A)])[W1(B)W2(B)...Wn(B)]y=f \left(x \left[ W_1^{(A)} \space W_2^{(A)} \space ... \space W_n^{(A)} \right] \right) \left[ %左括号 \begin{array}{ccc} W_1^{(B)} \\ W_2^{(B)} \\ ... \\ W_n^{(B)} \end{array} \right]

由此可见,FFN 可以等价地表示成 nn 个向量之和,每个向量表示了一个小模型 f(xWi(A))Wi(B)f(x W_i^{(A)})W_i^{(B)}的输出,每个小模型计算量相同,每个小模型就是一个 Expert。

“能否只挑 kk 个向量的和来逼近 nn 个向量的和?这样就可以将计算量降低到 n/kn/k了”

优化问题:

arg minλ1,λ2,...,λn0,1i=1nλivii=1nvi2 s.t.i=1nλi=k\argmin_{\lambda_1, \lambda_2, ..., \lambda_n \in {0,1}} \left \| \sum_{i=1}^n{\lambda_i v_i} - \sum_{i=1}^n{v_i} \right \| ^2 \space s.t. \sum_{i=1}^n\lambda_i = k

直接求解较为困难,我们可以考虑取 viv_i 两两正交时的解作为正式解: “挑模长最大的 kk 个向量”。

为了避免模长的计算,我们需要设计一个小的模型“Router”来计算模长,获得前 kk 大小模长的向量后,再计算对应的方向:

y=iargtoppkρρieiy= \sum_{i \in argtop p_k \rho} \rho_i e_i

这就是 MoE 模型的基本公式。

MoE的负载均衡

我们希望每个 Expert 都在干活,并且尽量干一样的活,避免某些 Expert 浪费算力,充分发挥训练算力的寻求。

y=iargtoppkρρieiy= \sum_{i \in argtop p_k \rho} \rho_i e_i

促进负载均衡的常规思路是添加与之相关的损失函数,我们通常称之为“Aux Loss(Auxiliary Loss)”。

// 待拓展

LOSS-Free 的引入

Aux Loss固然简单直观,但它也有一个明显的缺点——权重不好调——调低了无法促进均衡,调高了容易损害LM Loss,所以业界一直有寻找替代方案的尝试。

DeepSeek 引入了一种非常优雅的做法:

y=iargtopk(ρ+b)ρieiy=\sum_{i \in argtopk ( \rho + b ) } \rho_{i}e_{i}

其中, bb 是输入无关的向量,训练后就保持不变,它仅仅参与分配过程而非 MoE 的前向计算。

我们先定义 f=[f1,f2,...,fn]f = [f_1, f_2, ..., f_n]:

fi={1/k,iargtopkρ+b>00,iargtopkρ+b>0f_{i} = \begin{cases} 1/k, i \in {argtop}_k\rho + b\gt 0 \\ 0, i \notin {argtop}_k\rho + b\gt 0 \end{cases} \nonumber

bb 偏置下 Expert 当前的分布 F=E[f]F = \mathbb{E}[f]。其中我们的目标均匀分布定义为: Q=(1/n,1/n,...,1/n)Q=(1/n, 1/n, ..., 1/n),负载均衡就相当于最小化

Laux=12FQ2=12i=1n(Fi1/n)2\mathcal{L}_{aux} = \frac{1}{2} \left \| F-Q\right \|^2 = \frac{1}{2} \sum^n_{i = 1} \left( F_i - 1/n \right) ^2

这个目标是不可导的,我们使用 STE 解决这个问题:找一个可导且和 FF 具有相同增减趋势的量作为 FF 的光滑近似:

Laux=12b+sg[Fb]Q2=12i=1n(bi+sg[Fibi]1/n)2\mathcal{L}_{aux} = \frac{1}{2} \left \| b + sg[F - b]-Q\right \|^2 = \frac{1}{2} \sum^n_{i = 1} \left( b_i + sg[F_i - b_i] - 1/n \right) ^2

它的梯度是:

bLaux=b12i=1n(bi+sg[Fibi]1/n)2=FQ\nabla_b \mathcal{L}_{aux} = \nabla_b \frac{1}{2} \sum^n_{i = 1} \left( b_i + sg[F_i - b_i] - 1/n \right) ^2 = F-Q

因此使用 SGD 更新 bb 就是:

bbγ(FQ)b \gets b-\gamma \left( F-Q\right)

采用SGD 来更新 bb 的符号梯度下降版本就是:

bbγsign(FQ)b \gets b-\gamma sign \left( F-Q\right)

将固定专家数扩展为可变专家数

在先前介绍的MoE架构中,我们习惯了为每个 Token 分配固定数量K的专家(Top-k)。但稍微动脑筋想想就会发现:Token 与 Token 之间的难度是不平等的。 一个简单的字可能只需要 1 个专家顺手处理,而一个复杂的数学逻辑可能需要 4 个专家协作。如果我们给所有 Token 都分 k=2,那对简单的 Token 是浪费,对难的 Token 是抠门。

这部分接着上一部分进行扩展,利用 Loss-Free 方案中的偏置项(Bias),把 MoE 变成动态专家系统。

从 Top-K 到 Argwhere

让我们首先回顾 MoE 的基本形式:

y=iargtopk(ρ)ρieiy=\sum_{i \in argtopk ( \rho ) } \rho_{i}e_{i}

ρ\rho 是专家打分,几何意义上来说是预测模长,eie_{i} 是专家的输出;

为了解决负载均衡,引入偏置项 bb:

y=iargtopk(ρ+b)ρieiy=\sum_{i \in argtopk ( \rho + b ) } \rho_{i}e_{i}

上一节内容讲到,bb 是作为相对大小调节 ρ\rho 的相对大小,bb 整体有一个均值在这个过程中是不起作用的,这个冗余度将在这一节中被使用。

为了实现动态数量,我们干脆去掉 Top- kk的硬性约束,使用门槛:

y=iargtopk(ρ+b>0)ρieiy=\sum_{i \in argtopk ( \rho + b \gt 0) } \rho_{i}e_{i}

只要 预测专家得分和补偿得分的和大于 0,专家就被激活。这样每个词选中的专家数就不再固定。

数学实现

我们的训练目标依旧是找到合适的 bb ,而 bb 的优化目标有两个:

  1. 负载均衡:保证流量平均分配给每个专家。
  2. 预算控制:考虑第一个优化目标,如果不加限制,那么 bi=+b_{i} = +\infty,即选中所有专家,这也是一种“均衡”,计算量会变得特别大。所以要让平均激活数维持在 kk

定义指示函数 f=[f1,f2,...fn]:f=[f_{1}, f_{2},...f_{n}]:

fi={1,ρi+bi>00,ρi+bi0f_{i} = \begin{cases} 1, \rho_i + b_i \gt 0 \\ 0, \rho_i + b_i \le 0 \end{cases} \nonumber

记其期望值(Batch 平均值)为 F~=E[f]\tilde{F} = \mathbb{E}\left[ f \right];

负载均衡实现:让 FF 趋近于 QQ

bbγsign(FQ)b \gets b-\gamma sign \left( F-Q\right)

为了不影响整体水平,而且本身这个冗余度在负载均衡中没有影响,所以上边的更新规则可以先去均值:

bbal=sign(FQ)sign(FQ)b_{bal}=sign\left( F-Q \right) - \overline{sign\left( F-Q \right)}

这样,我们就将 bb 的均值这个自由度让出来给后续的预算控制。

此时,实际的平均激活数是 F~| \tilde{F}|,而目标是 kk。如果实际激活数大于目标,则说明 bb 太大,此时需要我们调小 bb

bbudget=sign(F~k)b_{budget}=sign( | \tilde{F} |- k )

将上述过程整理,得到最终的更新公式:

bbγ[sign(FQ)sign(FQ)+sign(F~k)]b \gets b-\gamma \left[ sign \left( F-Q\right) - \overline{sign\left( F-Q \right) } + sign( | \tilde{F} |- k ) \right]

如果我们希望总预算大,就可以让整体的 bb 变大,这样 ρi+bi>0 \rho_i + b_i \gt 0  的概率就越大,预算就越多。

显然地,如果我们希望预算只是不超过 kk,而并非希望其等于 kk,则可以让 bbudgetb_{budget} 最小为 0,这样当预算不超过 kk 时不因为预算问题对 bb 进行更新:

bbγ[sign(FQ)sign(FQ)+max(sign(F~k),0)]b \gets b-\gamma \left[ sign \left( F-Q\right) - \overline{sign\left( F-Q \right) } + max \left( sign ( | \tilde{F} |- k ) , 0 \right) \right]

简化

苏神在实验中发现,其实还可以更简单。如果我们让每个专家的“绝对平均负载”都去逼近 k/nk/n,那么均衡和预算就自动合二为一了:

定义目标 Q~=(k/n,k/n,...,k/n)\tilde{Q} = (k/n, k/n, ..., k/n),更新规则可以简化为:

bbγsign(F~Q~)b \gets b-\gamma sign \left( \tilde{F}-\tilde{Q}\right)

根据苏神的实验,这些优化过程效果上大同小异,但是最后简化后的更新过程在负载均衡和预算控制两个指标在训练前期抖动大得多。

同时,苏神还在这里提到了一个技巧:使用 RMS Norm 来替换 sign 函数,增加稳定性。在这个任务中使用 RMS Norm 替换后更新的数量级相同,同时 RMS Norm 保留了相对大小:

bbγ(F~Q~)/F~Q~RMSb \gets b-\gamma \left( \tilde{F}-\tilde{Q}\right) / \left\| \tilde{F}-\tilde{Q} \right\|_{RMS}

初始化

最后问题来到了一个“有意思但是不算十分关键的问题”: bb 的初始化。

如果我们偷懒直接把 b 全设为 0,且 Router 用的激活函数是 Sigmoid:

数学假设:为了给 b 一个完美的初始值,我们需要做点简单的统计假设:

二分法寻找“起跑线”

我们的目标是:找到一个初始的 binit,使得在训练第一步时,激活的专家数正好接近我们的预算 k:

E[Sigmoid(Logits)+binit>0]k.\mathbb{E}[Sigmoid(Logits)+b_{init}>0]≈k.

总结

从上一章对 DeepSeek 的 “Auxiliary-Loss-Free Load Balancing Strategy for Mixture-of-Experts”这篇工作的解读,再到这一部分对这一工作的探索,确实让人感觉到一些无比“优雅”且充满美感的工作。

均匀分布的反思

其实看到这里,如果我们稍微思考一下,就会发现,其实现实时间中的大部分情况都是“非均匀的”,而我们之前一直在强调“负载均衡”。可是均衡的负载在除了效率方面,一定是最好的吗?

Shared Expert

MoE 的基本形式如下:

y=iargtopk(ρ)ρieiy=\sum_{i \in argtopk ( \rho ) } \rho_{i}e_{i}

Loss-Free 将 argtopk(ρ)argtopk ( \rho ) 替换成 argtopk(ρ)+bargtopk ( \rho ) + b,在动态专家系统中,我们进一步将其推广为 argtopk(ρ)+b>0argtopk ( \rho ) + b > 0

如果负载完全均衡,是否意味每个专家都成了“全才”或“庸才”?

Shared Expert:

数学形式如下:

y=i=1sei+iargtopksρ[s:]ρi+sei+sy=\sum_{i=1}^s e_i + \sum_{i \in argtop_{k-s} \rho_{[s:]} } \rho_{i + s}e_{i + s}
  1. 压缩共性:等号后的第一项是必然选中的部分,用 DeepSeek 的话来说就是把所有 Token 都要用到的公共知识塞进共享专家;
  2. 减少冗余:让剩下的路由专家专心地去学习哪些偏门的专业知识;
  3. 几何意义:路由专家只需要学习“残差”。从几何上来看,让专家向量之间更容易满足正交假设,提高参数效率。

比例因子

实际上,前边的式子是一个理想的过程:每个Router Expert 前边有一个预测模长 ρ\rho。如果直接相加,会导致“力量失衡”。为了让两个部分在初始化部分的模长一样,引入一个比例因子 λ\lambda

y=i=1sei+λiargtopksρ[s:]ρi+sei+sy=\sum_{i=1}^s e_i + \lambda\sum_{i \in argtop_{k-s} \rho_{[s:]} } \rho_{i + s}e_{i + s}

Fine-Grained Expert

除了 Shared Expert 以外,DeepSeekMoE提出的另一个改进点是Fine-Grained Expert:在总参数量和激活参数量都不变的情况下,Expert 的粒度越细,效果往往越好。

例如, 2n2n2k2k 的系统就比 nnkk 的效果要好。原论文的说法是丰富了 Expert 的组合多样性:

(nk)(2n2k)(4n4k)...\binom{n}{k} \ll \binom{2n}{2k} \ll \binom{4n}{4k} ...

虽然总参数量没变,但是组合的可能性爆炸式增长。

最优分配

目前我们已经整理的负载均衡的具体实现一共有两种方案:Aux Loss 和 Loss-Free 方案。下边讨论第三种思路:最优分配。它将负载均衡视作等式约束下的线性规划问题。

问题描述:假设有 mm 个 Token,第 ii 个 Token 的 Rouer 对 nn 个 Expert 的打分记做:

si=(si,1,si,2,...,si,n),s_i = \left( s_{i,1}, s_{i,2}, ..., s_{i,n}\right) ,

那么共有 mnmn 个分数。目标是基于这个分数指定一个分配方案来决定每个 Token 应该激活哪些 Expert。

这在数学上是一个典型的带约束的线性规划问题

maxxi,j{0,1}i,jxi,jsi,js.t.jxi,j=k,ixi,j=mkn\max_{x_{i,j} \in \{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \\s.t.\sum_{j} x_{i,j} = k, \sum_{i} x_{i,j} =\frac{mk}{n}

上式属于整数优化问题,难以求解,考虑其松弛版本:

maxxi,ji,jxi,jsi,js.t.jxi,j=k,ixi,j=mkn\max_{x_{i,j}} \sum_{i,j} x_{i,j}s_{i,j} \\s.t.\sum_{j} x_{i,j} = k, \sum_{i} x_{i,j} =\frac{mk}{n}

其中 xi,j[0,1]x_{i,j} \in [0,1]是分配系数, si,js_{i,j} 是原始得分。这样就是一个有界区域内的一个线性规划问题。

进一步,我们考虑约束优化问题的 max-min 形式,即拉格朗日乘子法:

maxxi,j[0,1]minαi,βji,jxi,jsi,jiαi(jxi,jk)jβj(ixi,jmkn)\max_{x_{i,j} \in [0,1]} \min_{\alpha_i, \beta_j} \sum_{i,j} x_{i,j}s_{i,j} - \sum_{i} \alpha_i \left( \sum_{j} x_{i,j} - k \right) - \sum_{j} \beta_j \left( \sum_{i} x_{i,j} - \frac{mk}{n} \right)

该形式和原问题松弛版本的形式等价:如果后两项括号内的形式不满足等于 0,则在 maxmax 后的 minmin 步骤可以取负无穷。

上边的目标关于 xi,j,αi,βjx_{i,j}, \alpha_i, \beta_j 都是线性的,同时 xi,jx_{i,j}的可行域为凸集,满足 Minimax theorem 条件,交换 maxmaxminmin 的顺序后,整理一下,我们得到:

minαi,βjmaxxi,j[0,1]i,jxi,j(si,jαiβj)+iαi+jβjmkn\min_{\alpha_i, \beta_j}\max_{x_{i,j} \in [0,1]} \sum_{i,j} x_{i,j} \left( s_{i,j} - \alpha_i - \beta_j \right) + \sum_{i} \alpha_i + \sum_{j} \beta_j \frac{mk}{n}

注意到,当前优化目标的第一项我们是可以进行分析的:

也就是说,最后的优化目标和一开始非松弛版本的优化目标是一致的,二者完全等价。

xi,jx_{i,j}^* 代入优化目标,得到 xi,j(si,jαiβj)=max(0,si,jαiβj)x_{i,j}^* ( s_{i,j} - \alpha_i - \beta_j) = \max(0, s_{i,j} - \alpha_i - \beta_j),优化目标可以简化为:

minαi,βjmaxxi,j[0,1]i,jmax(0,si,jαiβj)+iαi+jβjmkn\min_{\alpha_i, \beta_j}\max_{x_{i,j} \in [0,1]} \sum_{i,j} \max(0, s_{i,j} - \alpha_i - \beta_j) + \sum_{i} \alpha_i + \sum_{j} \beta_j \frac{mk}{n}

我们使用交替最小化思路求解:先固定 αi\alpha_iβj\beta_j ,再固定 βj\beta_jαi\alpha_i,交替进行。由于 αi\alpha_iβj\beta_j 有明显的对偶性,因此这两个步骤是在求解同一个问题。

我们先来看在固定 βj\beta_j 来求 αi\alpha_i,问题等价于:

minαi,jmax(0,si,jαiβj)+kiαi\min_{\alpha} \sum_{i,j} \max(0, s_{i,j} - \alpha_i - \beta_j) + k \sum_i \alpha_i

每一项 αi\alpha_i 是单独累加起来的,所以将其划分为 mm 个独立的子优化问题:

minαkα+jmax(0,sjβjα)\min_\alpha k \alpha + \sum_j \max(0, s_j - \beta_j - \alpha)

sjβjs_j - \beta_j 按照从大到小排列成 sσ1βσ1sσ2βσ2...sσnβσns_{\sigma_1} - \beta_{\sigma_1} \ge s_{\sigma_2} - \beta_{\sigma_2} \ge ... \ge s_{\sigma_n} - \beta_{\sigma_n},第 jj 大的元素为 sσjβσjs_{\sigma_j} - \beta_{\sigma_j},假设我们已知 sσlβσlsσl+1βσl+1s_{\sigma_l} - \beta_{\sigma_l} \ge s_{\sigma_{l+1}} - \beta_{\sigma_{l+1}},那么目标函数变成:

kα+j=1l(sσjβσjα)={j=1k(sσjβσj)+j=k+1l(sσjβσjα0),lkj=1k(sσjβσj)j=l+1k(sσjβσjα0),lkk\alpha+ \sum_{j=1}^l \left ( s_{\sigma_j} - \beta_{\sigma_j} - \alpha \right ) = \begin{cases} \sum_{j=1}^k (s_{\sigma_j} - \beta_{\sigma_j}) + \sum_{j=k+1}^l \underbrace{(s_{\sigma_j} - \beta_{\sigma_j} - \alpha}_{\ge 0}), l \ge k \\ \sum_{j=1}^k (s_{\sigma_j} - \beta_{\sigma_j}) - \sum_{j=l+1}^k \underbrace{(s_{\sigma_j} - \beta_{\sigma_j} - \alpha}_{\le 0}), l \le k \end{cases}
这一部分个人理解:

梯度下降

在QB 迭代格式中, α\alpha 的计算是比较便宜的,真正昂贵的是 β\beta的计算,也就是说需要跨全体 Token 进行排序。假设给定 α\alphaβ\beta的优化目标是:

minβji,jmax(0,si,jαiβj)+mknjβj记为L\min_{\beta_j} \underbrace{ \sum_{i,j} \max(0, s_{i,j} - \alpha_i - \beta_j) + \frac{mk}{n}\sum_{j} \beta_j }_{记为\mathcal{L}}

很明显的是, l\mathscr{l}是可导的,这一步的最优解我们可以使用梯度下降:

Lβj=mkni=1mX(si,jαiβj>0)\frac{\partial{\mathcal{L}}}{\partial{\beta_j}}= \frac{mk}{n} - \sum_{i=1}^m \mathcal{X} (s_{i,j} - \alpha_i - \beta_j \gt 0)

其中, X\mathcal{X} 为示性函数,考虑 SignSGD:

βjβjγsign(Lβj)\beta_j \gets \beta_j - \gamma sign \left( \frac{\partial{\mathcal{L}}}{\partial{\beta_j}} \right)

最优分配问题下的进一步探索

在上一个章节,我们通过求解如下最优分配问题来实现负载均衡:

maxxi,j{0,1}i,jxi,jsi,js.t.jxi,j=k,ixi,j=mkn\max_{x_{i,j} \in \{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \\s.t.\sum_{j} x_{i,j} = k, \sum_{i} x_{i,j} =\frac{mk}{n}

其中,第一个约束条件表示每个Token 分配的Expert,第二个 约束条件表示每个 Expert 被激活的次数,其中我们真正需要的是后者:平均来说每个 Token 激活 kk 个 Expert 以及每个 Expert 的负载均衡。

只要保证每个专家平均分到 mk/n 个 Token,那么全局来看,每个 Token 平均激活的专家数自然就是 k。

本文考虑更加简化的问题

maxxi,j{0,1}i,jxi,jsi,js.t.ixi,j=mkn\max_{x_{i,j} \in \{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \\ s.t. \sum_{i} x_{i,j} =\frac{mk}{n}

考虑其等级 maxmin\max-\min形式:

maxxi,j[0,1]minαi,βji,jxi,jsi,jjβj(ixi,jmkn)\max_{x_{i,j} \in [0,1]} \min_{\alpha_i, \beta_j} \sum_{i,j} x_{i,j}s_{i,j} - \sum_{j} \beta_j \left( \sum_{i} x_{i,j} - \frac{mk}{n} \right)

交换 max\maxmin\min 顺序,整理得:

minβjmaxxi,j[0,1]i,jxi,j(si,jβj)+jβjmkn\min_{\beta_j}\max_{x_{i,j} \in [0,1]} \sum_{i,j} x_{i,j} \left( s_{i,j} - \beta_j \right) + \sum_{j} \beta_j \frac{mk}{n}

同样的, max\max 这一步可以先完成:

将上边的 xi,jx_{i,j}^* 带回,优化目标简化为:

minβji,jmax(0,si,jβj)+mknjβj\min_{\beta_j} \sum_{i,j} \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \sum_j \beta_j

分解为 mm 个独立的子问题:

minβimax(0,si,jβj)+mknβ\min_{\beta} \sum_{i} \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \beta

同样的,我们使用 QB 的思路去求解,也可以使用 SignSGD 的思路去求解。

Expert Choice Routing

我们表面上是在做 Token Choice(每个 Token 独立决定去哪,符合推理逻辑),但通过引入这个从对偶问题求出来的 β\beta ,我们实际上达到了 Expert Choice 的完美均衡效果。

本文由 GJJ 创作,内容来源于 Notion 数据库,随时可在 Notion 中编辑更新。 本站由 DeepSeek-v4-flash 辅助构建,项目参考 NotionNext

← 返回首页
61
文章
6
标签
3
分类
962
运行天数