MoE 专题索引

MoE 引入

Scaling Law 表明,在训练数据充足的前提下,用 更多的参数更大的 FLOPS 来扩展语言模型,可以得到强得多的模型。但对经典的 Dense 模型来说,这两者其实是同一个旋钮:$N$ 个参数的模型,每个 token 前向约需 $2N$、训练约需 $6N$ 次浮点运算。参数量一涨,训练和推理成本同比例上涨,而推理是模型上线后每天都要付的账。

Mixture of Experts(MoE)把这把锁撬开:将 FFN 拆成多个 expert,每个 token 只激活其中一小部分,每 token 计算量于是变成 $6N_{\text{active}}$ 而不是 $6N_{\text{total}}$。模型因此有了两个可以独立调节的参数量——总参数量决定容量,激活参数量决定成本。

Comparing a dense and sparse expert Transformer, Credit: A review of sparse expert models in deep learning
Comparing a dense and sparse expert Transformer, Credit: A review of sparse expert models in deep learning

MoE 本身可以追溯到 1990 年代初,但与现代 Transformer 的结合才是重新点起兴趣的原因。GShard 首创了大规模分布式 MoE 训练,引入了 Expert Parallelism 和负载均衡辅助损失1。Switch Transformer 证明了 MoE 可以扩展到万亿参数并保持训练稳定2。GLaM 显示 MoE 能以 Dense 模型一小部分的训练成本达到相当的质量3。Tutel4、DeepSpeed-MoE5 等框架进一步推进了 MoE 训练系统。

稀疏度由此成为与深度、宽度并列的第三个 scaling 维度,且逐年抬升: Mixtral 8x7B 47B 总参数激活 13B(3.6 倍), DeepSeek-V3 671B 激活 37B(18 倍), Kimi K3 2.78T 激活 104B(27 倍), ERNIE 5.0 的激活率已压到 3% 以下。时间线和「参数扩张比 vs expert 稀疏度」为什么会分叉,见 主流 LLM 的 MoE 稀疏度怎么变的

Mixture of Experts

作为基于 Transformer 的 MoE 模型,主要由以下两部分组成:

  • 稀疏 MoE 层:它将 FFN 拆成多个子层,每一个子层被称为 Expert 。一般来说,这些 Expert 都是 FFN,但是也可以是更复杂的网络,甚至是 MoE 本身
  • Router:也被称为 Gating Network,这部分用于决定将哪些 token 被发送到哪些 Expert
    Switch Transformer
    Switch Transformer

MoE 的整个计算过程如下所示:

  1. Routing:decide target experts for each token

    1. 每个 token 与 router 权重相乘,得到它对全部 expert 的分数
    2. 再经 softmax(或 sigmoid)和 top-k,决定激活哪几个 expert、以及各自的 gate 权重
    3. 输出是一张 Expert Indices 和对应的概率。
  2. Dispatch Phase

    1. 按 expert id 把 token 重排 layout transform,让 「属于同一个 expert 的 token」 在显存里连成一段,后面的 Grouped GEMM 需要连续输入。
    2. top-k > 1 时每个 token 会复制 $k$ 份,token 数从 $s$ 涨到 $s\cdot k$
  3. Computation

    1. 每个 expert 对自己那一段 token 独立跑一遍 FFN。expert 之间互不看对方的输入。
  4. Combine Phase

    1. 把各 expert 的输出按原 token 顺序还原,再按 gate 权重把 $k$ 份副本加权求和,shape 塌回 $(s, h)$
  5. Gating function: decide target experts for each token

  6. Dispatch Phase

    1. Layout transformation: tokens to the same target experts are grouped in a continuous memory buffer
    2. Alltoall: dispatch tokens to their corresponding experts
  7. Expert Compute: each expert process its tokens

  8. Combine Phase

    1. Combine processed tokens batch to their GPUs
    2. Layout transform: restore tokens to their original positions

当把不同 expert 放到不同 GPU 上(Expert Parallelism)时,Permutation / Un-Permutation 里会各插入一次 all-to-all,把 token 送到持有该 expert 的设备、算完再送回来,Expert Parallelism 会进一步详细讨论。

A Mixture-of-Expert Layer, Credit: MegaBlocks Paper
A Mixture-of-Expert Layer, Credit: MegaBlocks Paper

对应于伪代码如下所示:

General MoE Training Process, Credit: HetuMoE
General MoE Training Process, Credit: HetuMoE

沿用前面 MegaBlocks 那张四阶段图的输入 the quick brown fox jumped over,取 $s=6$、$h=8$、$E=4$、$k=2$,只看单卡。图里画的是 top-1 加 capacity 丢弃,这里换成更常见的 top-2 dropless,四个阶段仍一一对应。

flowchart LR
    A["hidden_states 6×8"] --> B["router logits 6×4"]
    B --> C["topk_ids 6×2"]
    C --> D["sorted_tokens 12×8"]
    D --> E["变长四段 2,2,5,3"]
    E --> F["outs 12×8"]
    F --> G["reshape 6×2×8"]
    G --> H["输出 6×8"]

四段计算流程可以看下图,可以对应着下面的解析对着看。

MoE 层的四段计算流程,默认参数即本节的例子单独打开 ↗

Routing

MoE 里第一件事是把 token 路由到不同的 experts。一个 token 是长度为 d 的向量,乘上 router 那个 (d, E) 的矩阵,就得到它在每个 expert 上的分数:

代码实现如下。经过 Attention 之后当前的 tensor shape 为 (s, d),router 是一个 linear 层 (d, n_exp),输出 (s, n_exp),表示每个 token 路由到各个 expert 的权重,也就是上图的 Router Logits。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
class Qwen3MoeTopKRouter(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.top_k = config.num_experts_per_tok
        self.num_experts = config.num_experts
        self.norm_topk_prob = config.norm_topk_prob
        self.hidden_dim = config.hidden_size
        self.weight = nn.Parameter(torch.zeros(self.num_experts, self.hidden_dim))

    def forward(self, hidden_states):
        hidden_states = hidden_states.reshape(-1, self.hidden_dim)
        router_logits = F.linear(hidden_states, self.weight)  # (seq_len, num_experts)
        router_probs = torch.nn.functional.softmax(router_logits, dtype=torch.float, dim=-1)
        router_top_value, router_indices = torch.topk(router_probs, self.top_k, dim=-1)  # (seq_len, top_k)
        if self.norm_topk_prob:
            router_top_value /= router_top_value.sum(dim=-1, keepdim=True)
        router_top_value = router_top_value.to(router_logits.dtype)
        router_scores = router_top_value
        return router_logits, router_scores, router_indices

Dispatch

expert 计算要合成一次 Grouped GEMM,这要求同一个 expert 的 token 在显存里连续,而 routing 阶段里它们是散的。

1
2
idxs = topk_ids.view(-1).argsort()                    # (s·k,)
sorted_tokens = x[idxs // self.num_experts_per_tok]   # (s·k, h)

第一行先把 (6, 2)topk_ids 按行摊平成 (12,),这是行数从 6 涨到 12 的起点:

摊平下标 0 1 2 3 4 5 6 7 8 9 10 11
来源 token the the quick quick brown brown fox fox jumped jumped over over
目标 expert 1 2 2 3 0 3 2 0 1 2 2 3

argsort 按目标 expert 排序,得到 idxs = [4, 7, 0, 8, 1, 2, 6, 9, 10, 3, 5, 11]——注意它是摊平坐标系里的下标,第二行还要把它换回 token 下标才能取到 hidden state。换算靠的是「摊平时每个 token 连着占 k 个位置」这个规律,所以整除即可:idxs // k 得到 [2, 3, 0, 4, 0, 1, 3, 4, 5, 1, 2, 5],于是 sorted_tokens 是一个 (12, 8) 的张量:

排序后位置 0 1 2 3 4 5 6 7 8 9 10 11
token brown fox the jumped the quick fox jumped over quick brown over
归属 expert E0 E0 E1 E1 E2 E2 E2 E2 E2 E3 E3 E3

the 在位置 2 和 4 各出现一次,fox 在 1 和 6,brown 在 0 和 10——top-2 意味着每个 token 有两份副本,各自去一个 expert,这正是 (12, 8) 里 12 = s·k 的来源。idxs 这个数组之后还要用一次:combine 阶段靠 y[idxs] = outs 把结果按同一张索引表送回原位,所以它必须一直留在显存里。

Computation

每个 expert 的 token 是一段连续区间,偏移由 tokens_per_expert 的前缀和给出:

1
2
3
4
5
6
7
8
outputs, start_idx = [], 0
for i, num_tokens in enumerate(tokens_per_expert):   # [2, 2, 5, 3]
    if num_tokens == 0:
        continue
    end_idx = start_idx + num_tokens
    outputs.append(self.experts[i](sorted_tokens[start_idx:end_idx]))
    start_idx = end_idx
outs = torch.cat(outputs, dim=0)                     # (s·k, h)
expert 切片 输入 shape 中间态 输出 shape
E0 [0:2] (2, 8) (2, 16) (2, 8)
E1 [2:4] (2, 8) (2, 16) (2, 8)
E2 [4:9] (5, 8) (5, 16) (5, 8)
E3 [9:12] (3, 8) (3, 16) (3, 8)

每个 expert 就是一个普通 FFN,权重 w_in(8, 16)w_out(16, 8),所以 h 进出不变、只有行数按切片长度变化,torch.cat 之后 outs 回到 (12, 8)num_tokens == 0 时要 continue 跳过的原因也在这里:(0, 8) 的 GEMM 是个无意义的空 kernel。

这个 Python 循环只是用来讲清语义的——真实实现绝不会这样跑,因为 E 上千时它意味着上千次 kernel launch。生产代码把四段区间合成一次 Grouped GEMM,靠传入每段的长度让一个 kernel 内部分组计算,这是 MoE 计算侧最关键的优化。

Megatron 的 self.experts(dispatched_input, tokens_per_expert, permuted_probs) 之所以要把 tokens_per_expert 当参数传进去,就是为了喂给 Grouped GEMM 做分组。

Combine

moe_ep 里逆排序就是一行 y[idxs] = outs,用的正是 dispatch 阶段那张索引表,把 (12, 8) 从「按 expert 排序」还原成「按 token 摊平」的顺序。剩下的加权在 moe_ep 之外:

1
2
y = y.view(6, 2, 8) * topk_weight.unsqueeze(-1)   # (6, 2, 8) 乘上 (6, 2, 1)
out = y.sum(dim=1)                                # (6, 8) 沿 k 维塌回

the 为例,它的两份副本分别由 E1E2 算出,按 gate 权重合并:

$$ \text{out}(\texttt{the}) = 0.647 \cdot E_1(\texttt{the}) + 0.353 \cdot E_2(\texttt{the}) $$

0.647 与 0.353 是 0.55、0.30 归一化到和为 1 的结果,Mixtral 和 DeepSeek-V3 都会做这步归一化。至此 shape 完整绕回 (s, h),MoE 层对外看起来和一个普通 FFN 完全一样。

Expert Parallelism

单卡那四步还在。EP 只改一件事:不同 GPU 持有不同 expert,权重不动,token 按路由结果往返一趟。Dispatch 变成「本地 permute + all-to-all」,Combine 对称地做回来。ep_size == 1 时两个 if 整段消失,就是上一节。

摊到 EP=2:GPU0E0E1GPU1E2E3。两张卡各有 6 个 token,GPU0 上仍是 the quick brown fox jumped overGPU1 上是 v0v5。下面这张动画是一趟往返:chip 在 dispatch 跨泳道飞过去、combine 原路飞回来;切换视角可以看到同一份代码在两个 ep_rank 上跑出不同的行数。

EP=2 下的 MoE 层:expert 权重不动,token 往返一趟,中间是两次搬数据的 all-to-all单独打开 ↗

发送:一次排序,两个用途

all-to-all 要求发送方把缓冲区切成 $R$ 段,第 $j$ 段整段发给 rank $j$——发往同一张卡的 token 必须连续。expert 按 id 成块分卡(experts_per_rank = 2,所以 E0E1GPU0E2E3GPU1),于是上一节那次按 expert id 的排序,自动就等价于按目标 rank 排序。一次排序同时满足两件事:Grouped GEMM 要按 expert 连续,all-to-all 要按目标 rank 连续。

沿用上一节排好的 12 行,只多看一行「目标 rank」:

位置 0 1 2 3 4 5 6 7 8 9 10 11
token brown fox the jumped the quick fox jumped over quick brown over
目标 expert E0 E0 E1 E1 E2 E2 E2 E2 E2 E3 E3 E3
目标 rank 0 0 0 0 1 1 1 1 1 1 1 1

[0:4] 留给自己,[4:12] 整段发给 GPU1,于是 input_splits = [4, 8]。没有额外搬运,只是对同一个缓冲区换了个切法。

input_splits 本地就能算,但 output_splits 不行——不知道对方要给自己多少,而 all-to-all 必须预先知道收发长度。所以先对计数做一次 all-to-all。GPU1 的 routing 设为 v0→E0,E2v1→E0,E2v2→E0,E3v3→E1,E2v4→E1,E3v5→E2,E3,本地统计 [3, 2, 4, 3]

flowchart LR
    A0["GPU0 本地统计 2,2,5,3"] -->|"留给自己 2,2 发给 GPU1 5,3"| X(("all-to-all 交换计数"))
    A1["GPU1 本地统计 3,2,4,3"] -->|"发给 GPU0 3,2 留给自己 4,3"| X
    X --> B0["GPU0 收到 2,2,3,2 即 E0 共 5 个 E1 共 4 个"]
    X --> B1["GPU1 收到 5,3,4,3 即 E2 共 9 个 E3 共 6 个"]

GPU0input_splits = [4, 8]output_splits = [4, 5]。数据 all-to-all 之后 sorted_tokens(12, 8) 变成 (9, 8),同一时刻 GPU1(15, 8)过了这一步行数不再是 $s\cdot k$,而取决于 routing 派给本卡 expert 的数量——9 与 15 的落差就是 straggler。

接收:为什么还要再排一次

这是 EP 相对单卡真正多出来的一步。发送时按目标 rank 切好了,但接收缓冲区是按来源 rank 拼起来的:先是 GPU0 的 4 行,再是 GPU1 的 5 行。每个来源都同时给 E0E1 送了 token,expert 又交错了:

位置 0 1 2 3 4 5 6 7 8
token brown fox the jumped v0 v1 v2 v3 v4
来源 rank 0 0 0 0 1 1 1 1 1
归属 expert E0 E0 E1 E1 E0 E0 E0 E1 E1

E0 落在 0、1、4、5、6,E1 落在 2、3、7、8。Grouped GEMM 吃不了这个布局,必须按本地 expert 再排一次,得到 gatherd_idxs

位置 0 1 2 3 4 5 6 7 8
token brown fox v0 v1 v2 the jumped v3 v4
归属 expert E0 E0 E0 E0 E0 E1 E1 E1 E1

E0 连续 5 行、E1 连续 4 行。所以 EP 下的 layout transformation 实际是两次排序:发送前靠 idxs(按目标 rank,顺带按 expert),接收后靠 gatherd_idxs(按本地 expert)。两张表都要留着,Combine 会反用。

Combine:原路返回

严格对称:sorted_tokens[gatherd_idxs] = outs 撤销第二次排序,input_splitsoutput_splits 对调再走一次 all-to-all,把 (9, 8) 换回 (12, 8),最后 y[idxs] = outs 撤销第一次排序,reshape 加权求和回到 (6, 8)

整层除了搬计数的那次小 all-to-all,真正搬隐藏态的是 dispatch 与 combine 各一次,量级随 $h \times k$ 增长。这就是未优化的跨节点 all-to-all 能吃掉六成训练时间的原因,也是 DeepEP 这类通信库存在的理由。

上面这些都收在 moe_ep 的两个 if self.ep_size > 1 里;其余行就是上一节的单卡路径。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
def moe_ep(self, x, topk_ids):
        cnts = topk_ids.new_zeros((topk_ids.shape[0], self.n_routed_experts))
        cnts.scatter_(1, topk_ids, 1)
        tokens_per_expert = cnts.sum(dim=0)
        idxs = topk_ids.view(-1).argsort()
        sorted_tokens = x[idxs // self.num_experts_per_tok]
        if self.ep_size > 1:
            tokens_per_expert_group = torch.empty_like(tokens_per_expert)
            dist.all_to_all_single(tokens_per_expert_group, tokens_per_expert, group=self.ep_group)
            output_splits = tokens_per_expert_group.view(self.ep_size, -1).sum(dim=1).cpu().tolist()
            input_splits = tokens_per_expert.view(self.ep_size, -1).sum(dim=1).cpu().tolist()
            gathered_tokens = All2All.apply(sorted_tokens, output_splits, input_splits, self.ep_group)
            gatherd_idxs = idxs.new_empty(gathered_tokens.shape[0], device="cpu")
            s = 0
            for i, k in enumerate(tokens_per_expert_group.cpu()):
                gatherd_idxs[s : s + k] = i % self.experts_per_rank
                s += k
            gatherd_idxs = gatherd_idxs.to(idxs.device).argsort()
            sorted_tokens = gathered_tokens[gatherd_idxs]
            tokens_per_expert = tokens_per_expert_group.view(self.ep_size, -1).sum(dim=0)
        tokens_per_expert = tokens_per_expert.cpu().numpy()

        outputs = []
        start_idx = 0
        for i, num_tokens in enumerate(tokens_per_expert):
            if num_tokens == 0:
                continue
            end_idx = start_idx + num_tokens
            expert = self.experts[i + self.ep_rank * self.experts_per_rank]
            outputs.append(expert(sorted_tokens[start_idx:end_idx]))
            start_idx = end_idx

        outs = torch.cat(outputs, dim=0) if len(outputs) else sorted_tokens.new_empty(0)
        if self.ep_size > 1:
            sorted_tokens = torch.empty_like(outs)
            sorted_tokens[gatherd_idxs] = outs
            gathered_tokens = All2All.apply(sorted_tokens, input_splits, output_splits, self.ep_group)
            outs = gathered_tokens

        y = torch.empty_like(outs)
        y[idxs] = outs
        return y

通信成本分析

$s$ 是每卡的 token 数(local batch × seq len)、$h$ 是 hidden size、$k$ 是 top-k,再补三个:$h_{ff}$ 是单个 expert 的 FFN 中间维,$R$ 是 EP size,$b$ 是激活的字节数(BF16 为 2,FP8 为 1)。

一层的字节数。 本卡的 $s$ 个 token 摊平成 $s k$ 个 token,均衡路由下每个目标 rank 分到 $s k / R$ 个,其中发给自己的那一段是本地拷贝、不上网络。所以 dispatch 一次的出网字节是

$$ V_{\text{dispatch}} = s \cdot k \cdot h \cdot b \cdot \frac{R-1}{R} $$

combine 把结果原路送回,形状完全对称,字节数相同。于是一层前向搬两份、量级 $2skhb$

此外只有两项小开销:交换 input_splitsoutput_splits 的那次计数 all-to-all 是 $R$ 个整数,以及随 token 一起发过去的 topk_weights,都比 $h$ 维的隐藏态小三四个数量级。shared expert 不参与路由,它的通信量是 0。

代入 DeepSeek-V3 的形状($h = 7168$、$k = 8$、BF16)、每卡 $s = 4096$ 个 token,$R$ 较大时 $\frac{R-1}{R} \approx 1$:单层单向约 448 MiB,一层往返约 0.9 GiB,58 个 MoE 层前向约 51 GiB,训练一步约 150 GiB。按每卡 50 GB/s 的 IB 带宽,这是 3 秒量级的纯通信时间——它必须被藏进计算里,否则整个 step 都在等网络。

通信/计算比:$s$ 和 $k$ 都会消掉。 均衡时每张卡收到的 token 数也是 $s k$($R$ 张卡各发来 $sk/R$)。SwiGLU 的 expert 有三个 $h \times h_{ff}$ 量级的矩阵,每个 token 前向约 $6 h h_{ff}$ FLOPs,所以

$$ \frac{T_{\text{comm}}}{T_{\text{comp}}} = \frac{2 s k h b}{s k \cdot 6 h h_{ff}} \cdot \frac{F}{B} = \frac{b}{3 h_{ff}} \cdot \frac{F}{B} $$

$s$、$k$、$h$ 全部约掉,只剩下单 expert 的中间维 $h_{ff}$ 和硬件的「算力/带宽」比 $F/B$。这个式子解释了不少现象:调大 batch 不改善通信占比(两边同比例涨)、提高 top-k 也不改善、加宽 $h$ 反而略微有利(计算涨得比通信快,因为 $h_{ff}$ 通常随 $h$ 一起涨)。

代入几组配置看看量级。$F$ 取 400 TFLOPS(BF16 grouped GEMM 的实测可达值,不是峰值),H800 上 NVLink 单向约 160 GB/s、每卡 IB 约 50 GB/s;上面那个式子里省掉的 $\frac{R-1}{R}$ 折扣按 $7/8$ 代入(EP=8 时有 1/8 留在本卡,EP=64 跨 8 机时有 1/8 留在本机):

配置 $h_{ff}$ 主要带宽 $T_{\text{comm}} / T_{\text{comp}}$
Mixtral 类,EP=8 单机 14336 NVLink 160 GB/s ≈ 0.10
V3 类,EP=8 单机 2048 NVLink 160 GB/s ≈ 0.71
V3 类,EP=64 跨 8 机 2048 IB 50 GB/s ≈ 2.3
上一行 + node-limited routing(≤4 节点) 2048 IB 50 GB/s ≈ 1.1
上一行 + FP8 dispatch 2048 IB 50 GB/s ≈ 0.86

第一行和第三行差了 20 多倍,而它们之间只隔着「expert 细不细」和「EP 有没有跨出单机」这两件事。第三行的 2.3 意味着通信比计算还贵一倍多,哪怕 overlap 做到完美也仍有 $2.3/3.3 \approx 70\%$ 的时间在等网络——和 DeepEP 那边「未优化的跨节点 all-to-all 吃掉六成训练时间」是同一个量级。

这也给前面 EP 与 TP 的讨论 补上了定量的一半:ETP=$T$ 把每卡的有效中间维压到 $h_{ff}/T$,按上式通信/计算比直接乘 $T$,同时 ETP group 内还要为 expert 的输入输出额外付一次 all-gather 加一次 reduce-scatter。对 $h_{ff}$ 本来就只有 2048 的细粒度 MoE,这是两头都吃亏,所以 Parallel Folding 的结论是优先给 EP、把 ETP 留在 1。

去重:跨节点的账不是按 $k$ 算的。 上面的模型假设一个 token 的 $k$ 个副本各自独立上网,但如果其中几个 expert 落在同一台机器上,跨节点其实只需要发一份,到了对端再用 NVLink 复制给本机的几张卡。这正是 DeepEP 分层转发的立足点:设集群有 $M$ 个节点,RDMA 流量与 $k$ 无关,只与「这个 token 命中了几个节点」$\bar m \le \min(k, M)$ 有关:

$$ V_{\text{RDMA}} \approx S \cdot \bar m \cdot h \cdot b, \qquad V_{\text{NVLink}} \approx S \cdot k \cdot h \cdot b $$

DeepSeek-V3 的 node-limited routing(每个 token 最多命中 4 个节点)就是把 $\bar m$ 硬性钉在 4,让 $k=8$ 的模型只付 4 份 IB 流量——上表第四行的 2 倍改善就是这么来的。它还有一层作用:把 IB 与 NVLink 的负载比调到和两者的带宽比匹配,避免其中一段闲着。

带宽项之外还有延迟项。 按 $\alpha + \beta n$ 模型看,NCCL 的 all-to-all 展开成 $R$ 对 send/recv,消息数随 $R$ 线性增长,而每条消息的大小随 $R$ 线性缩小:上面那个例子在 $R=64$ 时每条 7.3 MB,还算健康;但 decode 阶段 $S$ 只有几十,每条就掉到几百 KB,彻底进入 latency-bound 区间,实测带宽远低于峰值。这是 DeepEP 要为 prefill 和 decode 准备两套 kernel 的原因。另外两阶段 all-to-all 的第一阶段必须把计数取回 host 才能分配输出张量,这次 device-host 同步的开销与数据量无关,小 batch 下占比同样不可忽略。

不均衡按 max 计费,不按 mean。 all-to-all 是硬同步点,所有 rank 都要到齐才能返回,所以实际耗时由最重的那条边和最慢的那张卡决定。回到 EP=2 那个例子:GPU0 收 9 行、GPU1 收 15 行,均值 12,$\max/\text{mean} = 1.25$——计算和通信双双按 25% 的溢价结算,而多出来的时间里 GPU0 只是在空转。$R$ 越大,$\max$ 与 mean 的差距越容易被少数热点 expert 拉开,所以 负载均衡 从来不只是算力利用率问题,它同时是通信问题。

综上,能动的旋钮就这么几个:

  • 减字节(FP8 dispatch;combine 是加权求和、精度敏感,一般仍走 BF16)
  • 减份数(node-limited routing 限制 $\bar m$、分层去重)
  • 换带宽(把 EP 折进 NVLink 域,见 Megatron MoE Parallel Folding
  • 藏起来(DualPipe、1F1B 里与 GEMM overlap,不减少通信量只是遮住它)
  • 以及把 max 拉回 mean(负载均衡)。

EP 与其他并行训练方式组合

EP 只负责对 MoE 层中的 Expert 进行切分:不同 GPU 保存并计算不同的 Expert,Token 再根据 Router 的结果通过 All-to-All 被分发到对应设备。EP 本身并不规定 Attention、LayerNorm 等 non-MoE 模块如何并行,也不负责切分 batch、sequence 或 hidden states 维度。

如果仅对 Expert 使用 EP,而在 EP 组内复制 non-MoE 模块及其输入,那么各 GPU 会对相同的 Token 重复执行 Attention 等稠密计算。这虽然能够分摊 Expert 的参数和计算,却没有有效并行化模型的其余部分。

因此,实际训练和推理中,不会只用 EP,通常会将 EP 与其他并行方式组合使用。

针对 MoE 的并行训练,megascale-moe 对 EP 与其他并行训练策略如何选择进行了详细的讨论:

在 Megatron 也有类似的讨论 Megatron Core Parallelism Guide

EP 与 TP 的讨论

DeepSeek V3 之前,TP 在 Dense 大模型训练中一直是最核心的并行维度之一。而 DeepSeek V3 训练阶段采用的并行方案完全没有采用 TP,Attention、shared expert 等没有切 TP:

  • PP=16
  • EP=64,跨 8 个节点
  • TP=1
  • ZeRO-1 数据并行
  • 256 个 routed experts,每个 token 激活 8 个

在 DeepSeek V3 方案中,通过依靠 PP、EP、FP8、重计算和 ZeRO-1 优化器状态分片解决显存问题,再用 DualPipe 隐藏 PP 与 EP All-to-All 通信。

对于 DeepSeek V3 方案,当时引起了不少的讨论:

Q: 对于细粒度 MoE 训练,FFN 部分的 Expert 是应该完整保存,还是用 TP 内部切碎?

  • DeepSeek-V3 的 expert 很细粒度 (1+8 in 256 expert),单个 expert 的 FFN intermediate size 相对小 2048。如果继续采用 TP=8,那么每个 expert 的第二个 GEMM 就只剩 256 列,tensor core 完全喂不饱。
  • 作为对比,Mixtral 8x7b 模型中 2 in 8 expert,expert 中间维是 14336,切成 TP=8 之后每份还有 1792。
  • 因此,对于 DeepSeek-V3 这种细粒度的 MoE,选择不使用 TP 去切 MoE 是正确的选择,否则矩阵切的太碎,降低 GEMM 效率

在 Megatron MoE Parallel Folding 实验基本支持了这一点:ETP 的通信开销显著高于 EP,细粒度 MoE 尤其明显,因此建议保持较小的模型并行度,并优先 EP 而不是 ETP6

Q: 对于细粒度 MoE 训练,如果 FFN 部分的 Expert 不采用 TP,Attention 也不切 TP 吗?

  • 对于 DeepSeek V3 训练而言,他确实是这么做的,Attention 和 MoE 都没有 TP。
  • 训练期不切 TP 的真实原因是带宽而不是算法:H800 为出口合规把 NVLink 从 900 GB/s 砍到 400 GB/s(单向实测约 160 GB/s),与 IB 形成约 4:1 的差距,而 TP 恰好是最依赖节点内带宽的那个维度。DeepSeek 自己的表述是「训练期避开 TP 是因为 NVLink 带宽受限,推理期仍可选择性使用 TP 来改善 TTFT 与 TPOT」7
  • 因此,Attention 切不切 TP 与 Expert 切不切 TP 是两个独立的决定:后者取决于 GEMM 形状,前者取决于显存压力与节点内带宽,不应该被绑在一起。V3 能在 Attention 上省掉 TP,还额外依赖 MLA 把 KV cache 压掉九成,换成普通 GQA 这笔账就不成立。

Q: MoE 训练与推理不同场景下,EP 与 TP 的不同配置

DeepSeek V3 报告,训练没有用 TP,推理部署实际上重新用了 TP:

  • Prefill:Attention TP4 + SP + DP8,MoE EP32
  • Decode:Attention TP4 + SP + DP80,MoE 使用更宽的 EP

因为推理时的约束变成了 KV Cache、延迟、请求批量和显存带宽,不再等同于训练。在 decode 阶段,EP 越大,每卡 expert 越少、腾出的显存越多、batch 能开越大,per-expert 的 GEMM 反而更胖。这是 EP320 的真正动机,和训练期省显存是两回事。

Q: MoE EP 并行时,大 EP 带来的通信问题

Megatron MoE Parallel Folding

Megatron Core MoE 针对 Attention 和 MoE 这种 mismatch 进行了详细讨论,给出了 MoE Parallel Folding6 的解决方案。

Parallel Folding 为注意力层和 MoE 层引入各自独立的并行 group,两边的并行度不再互相绑定。

  • 注意力层 在 $\text{TP} \times \text{CP} \times \text{DP} \times \text{PP}$ 上组 group,针对序列级的 dense 计算优化。
  • MoE 层 在 $\text{ETP} \times \text{EP} \times \text{EDP} \times \text{PP}$ 上组 group,其中 ETP(Expert Tensor Parallelism)和 EDP(Expert Data Parallelism)是 MoE 专用的维度。

唯一的约束:Pipeline Parallelism(PP)在两套布局里必须保持一致,才能保证梯度在模型里正确流动。

MoE Parallel Folding 通过允许 EP 在注意力并行配置的任意子 group 上「折叠」,消除了 EP $\leq$ DP 的限制。这带来四项关键优势:

  1. 打破 EP $\leq$ DP 约束:EP 现在可以通过在 TP $\times$ CP group 上折叠而超过 DP。考虑注意力配置为 TP=4、CP=2、DP=8、PP=4(共 256 张 GPU):
    • 传统方式:EP $\leq$ DP = 8,所以 EP 最大只能到 8。
    • 有了 Parallel Folding:MoE 用 ETP=1、EP=64、EDP=1(PP 同为 4)。EP 在 TP $\times$ CP $\times$ DP group 上「折叠」,实现 8 倍高的专家并行度,同时注意力层保持它自己的最优配置 TP=4、CP=2。
  2. 降低最少 GPU 需求:传统配置下 CP=8、EP=8 至少要 64 张 GPU。有了 Folding,CP 和 EP 共享同一组 GPU,只需要 8 张。
  3. 支持独立优化:注意力可以为大矩阵用高 TP,而 MoE 用 ETP=1 保住完整专家宽度和更好的 GEMM 效率。
  4. 把高带宽通信留在 NVLink 域内:CP(给注意力)和 EP(给 MoE)的 all-to-all 通信都可以留在 NVLink 连通的 GPU group 内,避开更慢的跨节点传输。

如果把 ETP1 换成 ETP2,则对应的计算流程如下所示:

其中 ETP 部分通信流程如下所示:

FSDP + EP

FSDP + EP + SP

MoE Load Balance

RL 推理与训练路由一致性

均衡问题基本解决之后,MoE 路由上真正活跃的争议已经换了题目:同一个 token 在训练和推理时会不会被路由到不同的 expert

这在预训练时不重要,但在 RL 里很致命:

  • rollout 由推理引擎产生,训练侧在另一套 kernel 和并行策略下重算 logprob,路由的微小差异会被放大成 importance ratio 的偏差。
  • MiMo-V2-Flash 的 Rollout Routing Replay(R3) 是直接的应对——把 rollout 时的路由决策存下来,训练时重放,强制两侧走同一批 expert。
  • GLM-5 在讨论 DSA 的 indexer 时也明确对照了这个思路,但指出 indexer 的 $k = 2048$ 远大于 MoE 中常用的 $k$,直接照搬的存储与通信成本不可接受。

详细讨论在 Routing ReplayTraining Inference Mismatch。值得注意的是,它和负载均衡是互相冲突的两个目标:bias 控制器要求路由随负载动态调整,而一致性要求路由可复现。K3 把 bias 在推理时冻结,本质上就是在这两者之间做的取舍。


  1. GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding, https://arxiv.org/abs/2006.16668 ↩︎

  2. Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity, https://arxiv.org/abs/2101.03961 ↩︎

  3. GLaM: Efficient Scaling of Language Models with Mixture-of-Experts, ICML 2022, https://arxiv.org/abs/2112.06905 ↩︎

  4. Tutel: Adaptive Mixture-of-Experts at Scale, MLSys 2023, https://arxiv.org/abs/2206.03382 ↩︎

  5. DeepSpeed-MoE: Advancing Mixture-of-Experts Inference and Training to Power Next-Generation AI Scale, ICML 2022, https://arxiv.org/abs/2201.05596 ↩︎

  6. MoE Parallel Folding: Heterogeneous Parallelism Mappings for Efficient Large-Scale MoE Model Training with Megatron Core, https://arxiv.org/abs/2504.14960 ↩︎ ↩︎

  7. Insights into DeepSeek-V3: Scaling Challenges and Reflections on Hardware for AI Architectures, ISCA 2025, https://arxiv.org/abs/2505.09343 ↩︎