MoE 路由与并行代价
从 top-k router、专家容量和负载均衡损失推导 MoE,并观察稀疏激活背后的显存与通信成本。
- 建议时长
- 9–12 小时
- 难度
- 分布式系统
- 课程位置
- 11 / 12
学完这一课,你应该能:
- 能写出 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 做三组输入:
- 随机均匀 token。
- 大量相似 token,诱发路由偏斜。
- 两种语义簇,观察专家是否分工。
比较无平衡损失与加入平衡损失:
| 指标 | 必须记录 |
|---|---|
| 质量 | train/eval loss |
| 路由 | 每 expert token、entropy、drop |
| 系统 | step time、峰值内存 |
| 专业化 | 各 expert 代表样本 |
若有多卡环境,再测通信时间占比;没有多卡,就模拟 pack/unpack 并计算理论字节量。
9. 练习与答案提示
- top-k=2、N=1024、E=8,平均负载?**答:**每 expert 约 256 token。
- 稀疏激活为何不等于省权重显存?**答:**未激活专家权重仍需存储。
- 所有 expert 完全均匀是否必然最好?**答:**不一定,可能抑制有意义的专业化。
- 最慢 expert 为什么拖慢整步?**答:**后续同步需要等待其完成。
- All-to-All 传的是模型权重吗?**答:**常见 expert parallel 主要传 token 激活与返回结果。