前言
动机:近期笔者完成了一项较为理论的工作,并计划将其延申出一项实际应用——Mixture of Experts (MoE)。因此近期决定较为系统的学习 MoE。幸运的是,苏神正好有一个关于 MoE 的系列,因此笔者在这里整理一份个人笔记。
继续进行 MoE 系列的学习。符号定义见MoE系列v1和MoE系列v2。
优化问题
回顾在 MoE 系列-v2 中提出的优化问题: \(\begin{equation} \max_{x_{i,j}\in\{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \quad\text{s.t.}\quad \sum_j x_{i,j} = k,\quad \sum_i x_{i,j} = \frac{mk}{n} \end{equation}\) 这里的优化条件 $\sum_j x_{i,j}=k$ 是否必要?其含义是,保证每个 token 可以被分配到 $k$ 个 Experts。但其实,我们关注的仅是,每个 Expert 被激活相同次,即所谓的 “负载均衡”,而这个被第二个优化条件 $\sum_i x_{i,j} = \frac{mk}{n}$ 所保证。因此,接下来我们应该考虑的是简化之后的问题: \(\begin{equation} \max_{x_{i,j}\in\{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \quad\text{s.t.}\quad \sum_i x_{i,j} = \frac{mk}{n} \end{equation}\) 接下来的分析同MoE系列v2一致,因此此处的推导相对较为省略。
我们首先考虑上述问题的松弛版本: \(\begin{equation} \max_{x_{i,j}\in[0,1]} \sum_{i,j} x_{i,j}s_{i,j} \quad\text{s.t.}\quad \sum_i x_{i,j} = \frac{mk}{n} \end{equation}\) 写出其拉格朗日函数可得: \(\begin{equation} \max_{x_{i,j}\in[0,1]}\min_{\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) \end{equation}\) 同前文,这里交换 $\max$ 和 $\min$ 算子顺序可得: \(\begin{equation} \min_{\beta_j}\max_{x_{i,j}\in[0,1]} \sum_{i,j} x_{i,j}(s_{i,j} - \beta_j) + \frac{mk}{n} \sum_j \beta_j\label{eq:relax-min-max} \end{equation}\) 这里同理,不难看出 $\max$ 的答案: \(\begin{equation} \left\{\begin{aligned}&\,x_{i,j}^* = 1, &\, s_{i,j} - \beta_j > 0 \\ &\,x_{i,j}^* = 0, &\, s_{i,j} - \beta_j < 0 \\ &\,x_{i,j}^* \in [0,1], &\, s_{i,j} - \beta_j = 0 \end{aligned}\right. \end{equation}\) 将 $\max$ 的答案代回 Eq. ($\ref{eq:relax-min-max}$) 可得: \(\begin{equation} \min_{\beta_j} \sum_{i,j} \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \sum_j \beta_j\label{eq:beta-obj} \end{equation}\)
Quantile Balancing (QB) 解法
上述问题等价于: \(\begin{align} & \min_{\beta_j} \sum_{i,j} \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \sum_j \beta_j \\ \iff & \min_{\beta_1, \beta_2, ..., \beta_m} \sum_{j=1}^m \left( \sum_{i=1}^n \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \beta_j \right) \end{align}\) 此时由于 $\beta_j$ 是相互独立的,因此优化不同 $\beta_j$ 的和,其实就可以视作优化 $m$ 个子问题。因此接下来只考虑单独一个子优化问题形式(省略掉下标 $j$): \(\begin{equation} \min_{\beta} \frac{mk}{n}\beta + \sum_i \max(0, s_i - \beta) \label{eq:optimization_problem} \end{equation}\) 我们假设 $s_{\sigma_1}\geq s_{\sigma_2} \geq \cdots \geq s_{\sigma_m}$ 且 $s_{\sigma_l}\geq\beta\geq s_{\sigma_{l+1}}$,那么有: \(\begin{align} \frac{mk}{n}\beta + \sum_{i=1}^{m} \max(0, s_i - \beta) = \frac{mk}{n}\beta + \sum_{i=1}^l (s_{\sigma_i} - \beta) \\ = \left\{\begin{aligned} &\,\sum_{i=1}^{mk/n} s_{\sigma_i} + \sum_{i=mk/n+1}^l \underbrace{(s_{\sigma_i} - \beta)}_{\geq 0},&\, l \geq mk/n \\ &\,\sum_{i=1}^{mk/n} s_{\sigma_i} - \sum_{\hphantom{ab}i=l+1\hphantom{ab}}^{mk/n} \underbrace{(s_{\sigma_i} - \beta)}_{\leq 0},&\, l \leq mk/n \\ \end{aligned}\right. \end{align}\) 此时不难发现,无论 $l \geq mk/n$ 还是 $l \leq mk/n$,都会导致 Eq.($\ref{eq:optimization_problem}$) 中的目标函数增大。也即是:目标函数最小值在 $l=mk/n$ 处取得。再加上假设 \(s_{\sigma_l}\geq\beta\geq s_{\sigma_{l+1}}\),自然限制住了 $\beta^\star$ 的位置:$β^\star$ 位于 $s_i$ 的第 $mk/n$ 大和第 $mk/n+1$ 大元素之间,习惯起见,我们取第 $mk/n+1$ 大元素。恢复下标 $j$,那么对于给定 $j$,$β^\star_j$ 是将 $s_{i,j}$ 从大到小排列后第 $mk/n+1$ 个元素,或者同样叫做 “$1−k/n$ 分位数”。
综上我们就可以得到下述的算法了: \(\begin{array}{|l|} \hline \text{Quantile Balancing (QB) 动态激活版} \\[4pt] \hline \text{输入: 打分矩阵 }\boldsymbol{s}\in\mathbb{R}^{m\times n}\text{, 上一步 }\boldsymbol{\beta}\in\mathbb{R}^n\text{, 衰减率 }\lambda \\ \text{输出: 分配方案 }\boldsymbol{x}\in\{0,1\}^{m\times n}\text{, 新的 }\boldsymbol{\beta}\in\mathbb{R}^n \\[4pt] \hline \begin{array}{ll} 1: & x_{i,j}=1 \text{ if } s_{i,j} - \beta_j > 0 \text{ else } 0 \\ 2: & \boldsymbol{\beta} \leftarrow \lambda\boldsymbol{\beta} + (1-\lambda)\mathop{\text{desc-sort}}(\boldsymbol{s}, \text{axis=0})_{[mk/n:mk/n+1]} \\ 3: & \text{Output } \boldsymbol{x},\boldsymbol{\beta} \end{array} \\ \hline \end{array}\)
SGD 解法
我们可以使用比较简单粗暴的方式来求解上述优化问题 (Eq. ($\ref{eq:relax-min-max}$)): \(\begin{equation} \min_{\beta_j} \underbrace{\sum_{i,j} \max(0, s_{i,j} - \beta_j) + \frac{mk}{n} \sum_j \beta_j}_{\text{记为}\ell} \end{equation}\) 计算可以得到其梯度为: \(\begin{equation} \frac{\partial\ell}{\partial\beta_j} = \frac{mk}{n} - \sum_{i=1}^m \chi(s_{i,j} - \beta_j > 0) \end{equation}\) 其中 $\chi(\text{True})=1,\chi(\text{False})=0$。有了梯度后,就可以进行梯度下降了,我们依旧考虑 SignSGD: \(\begin{equation} \beta_j \leftarrow \beta_j - \gamma\mathop{\text{sign}}\left(\frac{\partial\ell}{\partial\beta_j}\right) \end{equation}\) 因此可以利用上述的 $\beta$ 更新规则来替换掉 QB 算法中对 $\beta$ 的更新。
Reference
[1] MoE 系列笔记 v1
[2] MoE系列笔记 v2
[3] MoE环游记:7、动态激活极简解