BotOf TechAI / IoT / Full-Stack / 植物养护知识分享
W1103 / 剖析并复现系统论文

MoE 路由与并行代价

从 top-k router、专家容量和负载均衡损失推导 MoE,并观察稀疏激活背后的显存与通信成本。

建议时长
9–12 小时
难度
分布式系统
课程位置
11 / 12
LEARNING OBJECTIVES

学完这一课,你应该能:

  • 能写出 token 到 expert 的 top-k 路由
  • 能解释 capacity factor、token drop 和路由坍塌
  • 能区分总参数、激活参数与常驻显存
  • 能测量专家利用率与 All-to-All 通信瓶颈

1. MoE 为什么稀疏

Dense FFN 对每个 token 激活同一组参数。MoE 准备多个 expert,用 router 为每个 token 选少数几个:

p(expert | x) = softmax(Wᵣx)
y = Σᵢ∈top-k pᵢ Eᵢ(x)

总参数可很大,但单 token 只激活 k 个 expert。这提高模型容量与每 token 计算的比例,但不等于“其他专家不存在”:权重仍需存储、加载或分布到设备。

2. 一个 Toy MoE

import torch
from torch import nn

class ToyMoE(nn.Module):
    def __init__(self, d_model=128, hidden=256, experts=4, top_k=2):
        super().__init__()
        self.top_k = top_k
        self.router = nn.Linear(d_model, experts, bias=False)
        self.experts = nn.ModuleList([
            nn.Sequential(
                nn.Linear(d_model, hidden),
                nn.GELU(),
                nn.Linear(hidden, d_model),
            )
            for _ in range(experts)
        ])

    def forward(self, x):
        # x: [tokens, d_model]
        probs = torch.softmax(self.router(x), dim=-1)
        weights, indices = probs.topk(self.top_k, dim=-1)
        output = torch.zeros_like(x)

        for expert_id, expert in enumerate(self.experts):
            token_pos, slot = torch.where(indices == expert_id)
            if token_pos.numel() == 0:
                continue
            expert_out = expert(x[token_pos])
            output[token_pos] += (
                weights[token_pos, slot, None] * expert_out
            )
        return output, probs, indices

这个实现清楚但不高效。真实系统会按 expert 重排 token,批量执行,再恢复顺序;跨设备时还需要通信。

3. Capacity factor 与 token drop

若总 token 数 N、专家数 E、top-k 为 k,平均每 expert 接收约 (Nk/E)。系统通常设容量:

capacity = ceil(capacity_factor × N × k / E)
  • factor 太低:热门 expert 溢出,token 被丢弃或路由到次选。
  • factor 太高:预留空间浪费,batch 不均衡。
  • 路由分布极偏:最慢 expert 决定整步时间。

实现时记录每 expert token 数、溢出数和最大/平均负载比。

4. 路由坍塌

Router 可能早期偏好少数 expert,热门 expert 获得更多训练信号,进一步变热门。缓解方法:

  • 负载均衡辅助损失
  • router noise / jitter
  • capacity 调整
  • expert bias 或无辅助损失的平衡策略
  • 更合理的初始化与 warmup

但平衡不是唯一目标。强行均匀可能让相似 token 被拆散,损害专业化。要同时看主任务质量与路由指标。

5. 负载均衡指标

设 (f_i) 是路由到 expert i 的 token 比例,(p_i) 是平均 router probability。Switch 风格辅助项鼓励两者平衡:

L_aux ∝ E × Σᵢ fᵢ pᵢ

实际还可监控:

  • expert utilization histogram
  • router entropy
  • top-1/top-2 gap
  • overflow/drop rate
  • expert specialization 样例
  • 最慢 expert 执行时间

Router entropy 高不保证负载一定平衡,也不保证专家形成有用分工。

6. Expert parallel 与 All-to-All

专家分布在不同 GPU 时,每个设备先拥有本地 token,路由后需要:

local tokens
→ pack by destination expert
→ All-to-All
→ expert compute
→ All-to-All return
→ restore token order

瓶颈可能从矩阵计算转为网络通信。吞吐受 token 分布、消息大小、拓扑、并行组合和最慢设备影响。理论激活 FLOPs 减少,不代表端到端 latency 按比例下降。

7. 总参数、激活参数与显存

报告 MoE 时分开:

  • total parameters:所有 experts + shared layers。
  • active parameters/token:本次路由实际参与计算。
  • resident parameters/device:设备上常驻多少权重。
  • communication volume:token 激活跨设备传输量。

“总参数 8×、激活参数不变”容易掩盖权重显存和通信成本。

8. 本课实验

对 Toy MoE 做三组输入:

  1. 随机均匀 token。
  2. 大量相似 token,诱发路由偏斜。
  3. 两种语义簇,观察专家是否分工。

比较无平衡损失与加入平衡损失:

指标必须记录
质量train/eval loss
路由每 expert token、entropy、drop
系统step time、峰值内存
专业化各 expert 代表样本

若有多卡环境,再测通信时间占比;没有多卡,就模拟 pack/unpack 并计算理论字节量。

9. 练习与答案提示

  1. top-k=2、N=1024、E=8,平均负载?**答:**每 expert 约 256 token。
  2. 稀疏激活为何不等于省权重显存?**答:**未激活专家权重仍需存储。
  3. 所有 expert 完全均匀是否必然最好?**答:**不一定,可能抑制有意义的专业化。
  4. 最慢 expert 为什么拖慢整步?**答:**后续同步需要等待其完成。
  5. All-to-All 传的是模型权重吗?**答:**常见 expert parallel 主要传 token 激活与返回结果。

10. 延伸阅读