近期笔者完成了一项较为理论的工作,并计划将其延申出一项实际应用——Mixture of Experts (MoE)。因此近期决定较为系统的学习 MoE。幸运的是,苏神正好有一个关于 MoE 的系列,因此笔者在这里整理一份个人笔记。

从几何意义出发

问题定义

跟随着苏神的思路,研究的对象选择分析最简单的 FeedForward Network (FFN): $\boldsymbol{y}=f(\boldsymbol{x} W^{(A)}) W^{(B)}$. 其中,$x\in \mathbb{R}^d$ 是行向量,$W^{(A)}\in\mathbb{R}^{d\times D}, W^{(B)}\in\mathbb{R}^{D\times d}$ 是 FFN 的两个参数矩阵。可以将 FFN 的表达式等价用分块矩阵进行改写:

\[\boldsymbol{y} = f\left( \boldsymbol{x} \left[ W_1^{(A)} \ W_2^{(A)} \ \cdots \ W_n^{(A)} \right] \right) = \sum_{i=1}^n \underbrace{f\left( \boldsymbol{x} W_i^{(A)} \right)}_{\boldsymbol{v}_i} W_i^{(B)}\]

其中 $\boldsymbol{W}i^{(A)} = \boldsymbol{W}^{(A)}{[:, (i-1)c:ic]}, \boldsymbol{W}i^{(B)} = \boldsymbol{W}^{(B)}{[(i-1)c:ic, :]}, c=D/n$。

因此我们可以将一个完整的 FFN 理解为:多个 Expert,即多个 $\boldsymbol{y}=f(\boldsymbol{x} W_i^{(A)}) W_i^{(B)}$ 的输出之和。自然,MoE 会产生一个问题:

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

模长排序

我们自然将问题转换成下述的表述:

\[\underset{\lambda_1,\lambda_2,\cdots,\lambda_n \in \{0,1\}}{\text{argmin}} \quad \left\| \sum_{i=1}^n \lambda_i \boldsymbol{v}_i - \sum_{i=1}^n \boldsymbol{v}_i \right\|^2 \quad \text{s.t.} \quad \sum_{i=1}^n \lambda_i = k\]

记 $\gamma_i = 1- \lambda_i$,则转换为:

\[\underset{\gamma_1,\gamma_2,\cdots,\gamma_n \in \{0,1\}}{\text{argmin}} \quad \left\| \sum_{i=1}^n \gamma_i \boldsymbol{v}_i \right\|^2 \quad \text{s.t.} \quad \sum_{i=1}^n \gamma_i = n - k\]

考虑简单的情形:当 $\boldsymbol{v}_i$ 两两正交时,有:

\[\left\| \sum_{i=1}^n \gamma_i \boldsymbol{v}_i \right\|^2 = \sum_{i=1}^n \gamma_i^2 \|\boldsymbol{v}_i\|^2 = \sum_{i=1}^n \gamma_i \|\boldsymbol{v}_i\|^2\]

显然,上述的最优解就是:选择模长 $|\boldsymbol{v}_i|$ 最小的 $n-k$ 个 $\gamma_i$ 值为 1;换句话说,就是挑模长 $|\boldsymbol{v}_i|$ 最大的 $k$ 个 $\lambda_i$ 值为 1。

即:挑模长最大的 $k$ 个向量的和来逼近 $n$ 个向量的和。

当 $v_i$ 不满足两两正交的条件时,我们依然用它来作为一个近似解。它的几何意义也很直观,模长越大的向量,在求和过程中越不容易被抵消,从而作用越突出。

MoE 初现

现在我们已经有了合理的策略: \(\textcolor{red}{挑模长最大的 k 个 Experts (向量)的和来逼近 n 个 Experts (向量) 的和。}\) 但是,依赖于模长的排序势必需要得到全部向量(即:$n$ 个),这样达不到减少计算量的目的。因此我们需要能够低成本计算出向量模长的方法。

苏神通过分解 Expert 的输出来巧妙地绕过计算全部向量。具体来说,分解为模长和方向 (记作 $\boldsymbol{e}_i=\frac{\boldsymbol{v}_i}{|\boldsymbol{v}_i |}$). 只有模长最大的 $k$ 个向量才会去计算对应的 $\boldsymbol{e}_i$,从而绕开不需要计算的剩余 $n-k$ 个向量,减少计算量。

那么现在问题转化为:怎么样低成本预测每个向量的模长?答案是:使用 Router,即一个低计算需求的小模型来预测。

\[\underbrace{[\rho_1, \rho_2, \cdots, \rho_n]}_{\boldsymbol{\rho}} = h\left( \boldsymbol{x} W^{(R)} \right) \quad \in \mathbb{R}_{\ge 0}^n\]

其中,$\rho_i$ 表示第 $i$ 个 Expert (向量) 的模长。$W^{(R)}\in \mathbb{R}^{d\times n}$ 是 Router 的参数矩阵。$h(\cdot)$ 则是 $\mathbb{R} \to \mathbb{R}_{\geq 0}$ 的激活函数(因为模长 $\geq 0$)

最终我们可以得到 MoE 的基本公式如下:

\[\boldsymbol{y} = \sum_{i \in \text{argtop}_k \boldsymbol{\rho}} \rho_i \boldsymbol{e}_i\]

注意上式中的 $Top-k$ 算子。当 $k=n$ 的时候,MoE 退化为原本的 FFN,此时称其为对应的 Dense 模型。

和一般 MoE 的对比

一般的 MoE 形式如下:

\[\boldsymbol{y} = \sum_{i \in \text{argtop}_k \boldsymbol{\rho}} \rho_i \boldsymbol{v}_i\]

上述分析和一般的 MoE 相比,多了个对 $\boldsymbol{v}_i$ 的 Normalization。这里是为了便于理解,不意味着是实践环节的必备流程。

DeepSeek-V3就是这样的(多了一步re-normalize),但会不会性能“更好一些”,实际上没有保证。目前的看法是大同小异,还是看其他方面的炼丹能力。

负载均衡 (Load Balance)

问题定义

MoE 的好处是稀疏激活,能够用较小参数/训练成本达到大参数量的效果。但是负载不均衡会阻碍这个好处。举个例子:你有 8 个 Expert、总参数量 100B,但如果 4 个 Expert 是死的 (Dead Expert),实际激活的参数可能只相当于 50B 的模型。

常用 Aux Loss 形式

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

将 $\boldsymbol{\rho}$ 归一化得到 $p$,并定义 $\boldsymbol{f}=[f_1, f_2, \dots, f_n]$:

\[p_i = \frac{\rho_i}{\sum_{j=1}^n \rho_j}, \quad f_i = \begin{cases} 1/k, & i \in \text{argtop}_k \boldsymbol{\rho} \\ 0, & i otin \text{argtop}_k \boldsymbol{\rho} \end{cases}\]

综上,有 Aux Loss 为:

\[\mathcal{L}_{\text{aux}} = \boldsymbol{F} \cdot \boldsymbol{\boldsymbol{P}} = \sum_{i=1}^n F_i P_i\]

其中,$\boldsymbol{F}=\mathbb{E}[\boldsymbol{f}]$ 是 Expert 当前的负载分布,而 $\boldsymbol{P}=\mathbb{E}[\boldsymbol{p}]$ 是 $F$ 的一个光滑近似分布。

一般文献定义 Aux Loss 会多乘一个 $n$,即它们的 Aux Loss 等于这里的 $n \mathcal{L}_{aux}$. 这里就是可以自由发挥定制不同的 Aux Loss 了。

苏神在这里给出了上述 Aux Loss 能够促进负载均衡的推导思路

Aux Loss 促进负载均衡的推导

定义均匀分布 $Q=(1/n, 1/n, \dots, 1/n)$。因为 $\boldsymbol{F}$ 是当前的负载分布,因此负载均衡就可以翻译为 $\boldsymbol{F} = Q$。有了这个目标,自然可以构建出对应的优化函数:

\[\mathcal{L}_{\text{aux}} = \frac{1}{2} \| \boldsymbol{F} - \boldsymbol{Q} \|^2 = \frac{1}{2} \sum_{i=1}^n \left( F_i - \frac{1}{n} \right)^2\]

这里不要忘记,$\boldsymbol{F}$ 依赖于不可导的 $Top-k$ 算子,因此它无法直接使用。这里可以考虑使用 Straight-Through Estimator (STE) 技巧。具体而言,使用可导的光滑近似分布(和原分布同增同减趋势)来替换掉不可导的。

\[\mathcal{L}_{\text{aux}} = \frac{1}{2} \| \boldsymbol{\boldsymbol{P}} + \text{sg}[\boldsymbol{F} - \boldsymbol{\boldsymbol{P}}] - \boldsymbol{Q} \|^2 = \frac{1}{2} \sum_{i=1}^n \left( P_i + \text{sg}[F_i - P_i] - \frac{1}{n} \right)^2\]

这里的 $\text{sg}[\cdot]$ 是 stop gradient 算子。性质是:前向传播时保持不变,但强制梯度为 0. 目的是:前向传播的时候依然还是 $\boldsymbol{F}-\boldsymbol{Q}$ 但是反向传播时使用可导的光滑近似分布替代,即 $\boldsymbol{P}-\boldsymbol{Q}$。

此时对其求导有:

\[\begin{aligned} \nabla_{\boldsymbol{\theta}} \mathcal{L}_{\text{aux}} &= \frac{1}{2} \nabla_{\boldsymbol{\theta}} \sum_{i=1}^n \left( P_i + \text{sg}[F_i - P_i] - 1/n \right)^2 \\ &= \sum_{i=1}^n \left( P_i + \text{sg}[F_i - P_i] - 1/n \right) \nabla_{\boldsymbol{\theta}} \left( P_i + \text{sg}[F_i - P_i] - 1/n \right) \\ &= \sum_{i=1}^n \left( F_i - 1/n \right) \nabla_{\boldsymbol{\theta}} P_i \\ &= \nabla_{\boldsymbol{\theta}} \sum_{i=1}^n \left( F_i - 1/n \right) P_i \\ &= \nabla_{\boldsymbol{\theta}} \left( \sum_{i=1}^n F_i P_i \right) \end{aligned}\]

我们发现,Eq. (11) 的梯度和 Eq. (9) 的梯度是一样的。因此,Eq. (9) 形式的 Aux Loss 具有促进负载均衡的意义。

然而,式 (9) 只有等效梯度的意义,但没有 Loss 的意义,不算一个真正的 Loss,比如当 $F=P$ 时我们可以算出式 (11) 等于 $1/n$,但实际上我们可以构造出一个不等于 $P$ 的 $F$ 让它小于 $1/n$,所以式 (9) 并不是像正常的Loss一样越小越好,最小值也不是 $F=P$ 时取到。

总结:现在常用的 Eq. (9) 是一个”梯度工具”而非”评估工具”。它的数值大小不能反映负载均衡程度,但导出的梯度方向是对的。

构建 Aux Loss 的一般思路

首先基于 $\boldsymbol{F}$ 构建符合要求的损失,然后在实现时将 $\boldsymbol{F}$ 替换成 $\boldsymbol{P}+\text{sg}[\boldsymbol{F}−\boldsymbol{P}]$.

举个例子,的最大化同样可以推动分布尽可能均匀,即:$\boldsymbol{F}_i \log \left( \boldsymbol{F} \right)$。因此我们可以构造出:

\[\mathcal{L}_{\text{aux}} = \sum_{i=1}^n \left( \boldsymbol{P}_i + \text{sg}[\boldsymbol{F}_i - \boldsymbol{P}_i] \right) \log \left( \boldsymbol{P}_i + \text{sg}[\boldsymbol{F}_i - \boldsymbol{P}_i] \right)\]

Loss-free 方案

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

苏神这里主要是介绍了 DeepSeek 的一项工作。该篇工作注意到:一个偏置项足以达到负载均衡。因此在构建 Aux Loss 的时候不再需要优化 $\boldsymbol{P}$ 而是优化偏置项 $\boldsymbol{b}$ 即可,如下所示:

MoE 的形式为:

\[\boldsymbol{y} = \sum_{i \in \text{argtop}_k \boldsymbol{\rho}} \rho_i \boldsymbol{e}_i \quad \rightarrow \quad \boldsymbol{y} = \sum_{i \in \text{argtop}_k \boldsymbol{\rho}+b} \rho_i \boldsymbol{e}_i\]

鉴于优化目标仅为 $\boldsymbol{b}$,Aux Loss 形式为:

\[\mathcal{L}_{\text{aux}} = \frac{1}{2} \| \boldsymbol{b} + \text{sg}[\boldsymbol{F} - \boldsymbol{b}] - \boldsymbol{Q} \|^2 = \frac{1}{2} \sum_{i=1}^n \left( b_i + \text{sg}[F_i - b_i] - 1/n \right)^2\]

对上述 Aux Loss 可以求得梯度:

\[\nabla_{\boldsymbol{b}} \mathcal{L}_{\text{aux}} = \frac{1}{2} \nabla_{\boldsymbol{b}} \| \boldsymbol{b} + \text{sg}[\boldsymbol{F} - \boldsymbol{b}] - \boldsymbol{Q} \|^2 = \boldsymbol{F} - \boldsymbol{Q}\]

因此可以得到,参数更新规则是:

\[\boldsymbol{b} \leftarrow \boldsymbol{b} - \gamma \cdot (\boldsymbol{F} - \boldsymbol{Q})\]

注意,原论文给出的更新规则稍有不同:

\[\boldsymbol{b} \leftarrow \boldsymbol{b} - \gamma\cdot \text{sign}(\boldsymbol{F} - \boldsymbol{Q})\]

这里不使用 Eq. (17) 更新规则,笔者认为是为了降低 $(\boldsymbol{F} - \boldsymbol{Q})$ 带来的噪声。除 Eq. (18) 之外,苏神还给出了另一种更新规则:

\[\boldsymbol{b} \leftarrow \boldsymbol{b} - \gamma \frac{\boldsymbol{F} - \boldsymbol{Q}}{\operatorname{RMS}(\boldsymbol{F} - \boldsymbol{Q})}\]

其中,$\operatorname{RMS}(\boldsymbol{F} - \boldsymbol{Q}) = \sqrt{\frac{1}{n} \sum_{i=1}^n (F_i - Q_i)^2}$.

关于该项 Loss-free 的工作,苏神还有关于实践细节以及延伸的一些思考,笔者在此处就不过多展开。

原论文在介绍Loss-Free时,并没有上述Aux Loss的推导过程,而是直接给出式 (16) 的更新规则,这也是它Loss-Free这个名字的来源。然而,从本文给出的推导可以看出,更新规则也完全可以从Aux Loss视角得到,两者是一脉相承的。

Loss-Free 的本质创新并不是没有 Aux Loss,而是隔离了 Aux Loss 和 LM Loss 的优化参数。相比之下,常规的Aux Loss方案需要全体参数来促进负载均衡,而LM Loss优化的也是全体参数,两者的优化方向可能并不完全兼容,因此想找到一个最优的平衡点相对来说就更为困难。

考虑 Token 的难易程度

苏神基于 Eq. (18) 进行改进。首先,由于 $\boldsymbol{b}$ 存在一个冗余自由度,即:对 $\boldsymbol{b}$ 加上常数,排序结果不变,$Top-k$ 算子的输出也不变。因此我们开源将 Eq. (18) 改写为:

\[b \leftarrow b - \gamma \left[ \text{sign}(F - Q) - \overline{\text{sign}(F - Q)} \right]\]

其中,$\overline{\boldsymbol{a}} \in \mathbb{R}$ 表示 $\boldsymbol{a}$ 的全体分量的均值。上述式子本质上约束了 $\boldsymbol{b}$ 的更新总和为 0.

直观来看,每个Token的难度并不一样,所以更合理的方案应该是难的Token分配更多的计算资源,简单的token分配更少的资源,这样或许能在同样有限的资源下将效果最大化。

综上动态选择的需求,苏神给出改进形式如下:

\[\boldsymbol{y} = \sum_{i \in \text{argtop}_k \boldsymbol{\rho}+b} \rho_i \boldsymbol{e}_i \quad \rightarrow \quad \boldsymbol{y} = \sum_{i \in \text{argwhere} \boldsymbol{\rho}+b>0} \rho_i \boldsymbol{e}_i\]

此时,只要满足 $\rho+b > 0$ 的 Expert 就被选中,从而实现 “动态选择 Expert” 而不是固定选择 $k$ 个 Experts 的情况。具体来说,如果给全体 $b_i$ 都加上同一个正数,那么满足 $ρ_i+b_i>0$ 的几率将会变大,选择更多的 Expert,从而总预算也会增大。

结合 Eq. (20) and (21),最终的更新形式是:

\[\begin{equation} \boldsymbol{b}\leftarrow \boldsymbol{b} - \gamma \left[\underbrace{\mathop{\text{sign}}(\boldsymbol{F} - \boldsymbol{Q}) - \overline{\mathop{\text{sign}}(\boldsymbol{F} - \boldsymbol{Q})}}_{针对负载均衡} + \underbrace{\mathop{\text{sign}}(|\tilde{\boldsymbol{F}}|- k)}_{针对专家动态选择}\right] \end{equation}\]

Tips:

一、苏神在这里将 “负载均衡” 和 “动态选择专家” 进行解耦,似乎更偏向工程直觉而非严格的数学推导。原因如下:本文已经将 routing 修改,见 Eq. (21)。此时增加常数不改变每个 token 内部的排序,但是会改变哪些 Experts 跨过 0 阈值,进而改变负载分布。下面给出一个简单的反例:

  • 考虑两个 expert,一个 token:$\rho+b = [0.1,,-0.2]$;此时排序是 expert 1 (>) expert 2,选择集合是:${1}$;所以有:$\tilde F=[1,0],F=[1,0]$。

  • 现在给所有 $b_i$ 加同一个常数 $c=0.3$:$\rho+b+c\mathbf{1}=[0.4,\,0.1]$. 排序仍然是 expert 1 (>) expert 2,完全没有变。但选择集合变成:${1,2}$. 此时有:$\tilde F=[1,1], F=[1/2,1/2]$.

因此,针对 Eq. (22) 中 “针对负载均衡” 项的 $-\mathop{\text{sign}}(\boldsymbol{F} - \boldsymbol{Q})$ 处理势必影响负载均衡。

二、笔者对于该篇博客有点云里雾里,感兴趣的读者推荐去阅读苏神原文以得到更好地理解。

Reference

[1] [MoE环游记:1、从几何意义出发 - 科学空间 Scientific Spaces](https://spaces.ac.cn/archives/10699)
[2] [MoE环游记:2、不患寡而患不均 - 科学空间 Scientific Spaces](https://spaces.ac.cn/archives/10735)
[3] [MoE环游记:3、换个思路来分配 - 科学空间 Scientific Spaces](https://spaces.ac.cn/archives/10757)

[4] https://arxiv.org/pdf/2408.15664

[5] [MoE环游记:4、难处应当多投入 - 科学空间 Scientific Spaces](https://spaces.ac.cn/archives/10815)