稀疏的是每个 token 的 expert 选择,不是所有成本
Mixture of Experts(MoE)通常替换 Transformer block 的 dense FFN/SwiGLU 子层。模型可以拥有 个 experts,但 router 对每个 token 只激活 top- 个。它是模型结构与分布式执行机制,不是 attention、张量并行 或服务 scheduler;MoE 仍可与 TP、data parallel、pipeline parallel 和 continuous batching 组合。
把本轮所有本地 token 展平为:
router logits 与概率为:
Top-1 时,第 个 token 选择:
Top- 则选择集合 ,再按实现定义归一化或直接使用 gates:
“模型总参数更多但每 token FLOPs 接近 dense FFN”只在 expert hidden size、top- 与 shared experts 等条件匹配时才成立。Router projection、dispatch/combine、load imbalance、shared experts 和通信都没有消失。
从 source ownership 到 expert ownership,再返回
设 expert parallel group 有 个 ranks、 个 experts,每 rank 拥有两个:
rank 0 owns expert 0, 1
rank 1 owns expert 2, 3
每 rank 起始有 3 个 tokens,,Top-1 路由得到:
rank 0 source: t0→e1, t1→e2, t2→e1
rank 1 source: t3→e1, t4→e2, t5→e1
Dispatch 后 rank 0 收到 [t0,t2,t3,t5] 给 expert 1,rank 1 收到 [t1,t4] 给 expert 2。Expert 输出完成后,第二次交换必须把结果送回 source rank,再按 source_id scatter 回 [t0,t1,t2] / [t3,t4,t5] 的原顺序。
INTERACTIVE EXPLAINER
Top-1 路由、All-to-All dispatch、本地 expert 与 combine
静态图固定 N=6、H=4、E=4、p=2;重点展示 ownership 的两次变化、variable split sizes 与恢复 token 顺序所需 metadata。
图中每个 expert 画成标量函数以突出通信语义,省略真实 SwiGLU 两个矩阵、shared expert、top-2 duplicate routes、padding/capacity、token packing layout、quantization、TP-inside-EP 与通信计算 overlap。箭头表示逻辑 All-to-All 数据交换,不承诺 Ring、pairwise、NVLink 或网络实现算法。
Tensor shape 与 rank 所有权
| 阶段 | rank 本地对象 | shape / dtype | 所有权状态 |
|---|---|---|---|
| router input | X_local | [N_local,H], BF16/FP16 常见 | source rank;来自 attention/residual 路径 |
| router logits | Z_local | [N_local,E], 常以 FP32 softmax/归约 | source rank;临时 |
| route metadata | expert_id, gate, source_id | [kN_local] | 按 destination rank pack;必须能恢复 source/order |
| dispatch payload | token hidden | [N_send,H] | 从 source rank 转到 expert-owning rank |
| expert batch | 每 expert ragged rows | [n_e,H] | expert owner; 不同,常 pack 成 grouped GEMM 输入 |
| expert output | selected routes | [N_recv,H] | expert owner,准备 combine |
| combined output | Y_local | [N_local,H] | 返回 source rank并按 token 顺序还原 |
Top-2 会让 route rows 近似从 增到 ,同一 token 的两个 routes 可能去不同 ranks;combine 还需按 gate 求和。Padding-to-capacity 实现的物理 shape 可是 [E,C,H],而 variable-size 实现保持 ragged/packed [\sum_e n_e,H]。逻辑语义相同,kernel 和通信 payload 不同。
容量、负载均衡与丢 token 边界
训练实现常为每 expert 设置容量:
是 capacity factor。若某 expert 收到 ,策略可能 drop、reroute 到次优 expert、提高容量或使用无固定容量的 dropless kernels。推理中“dropless”不代表无限资源:极端路由偏斜仍会扩大最大 expert batch、workspace、尾延迟与 rank 间等待。
上面 worked example 的 expert 1 收到 4 tokens,expert 2 收到 2,expert 0/3 收到 0;rank 0 的 expert 计算量是 rank 1 的两倍。即使总选中 routes 数仍为 6,iteration latency 往往受最慢 expert-owning rank 和 collective 同步约束。
训练时可加入 auxiliary load-balancing loss、router z-loss 或容量约束;这些目标影响推理时学到的路由分布,但不是推理前向的额外“再训练”。生产评估应直接报告真实请求上的 per-expert token histogram、最大/均值比、overflow/reroute 和跨 rank bytes。
可运行的两进程 Gloo 语义实验
实验问题:单 GPU 环境能否用两个 CPU/Gloo processes 完整验证 Top-1 routing、variable-split All-to-All、expert ownership、gate、第二次 combine 与 token 顺序恢复?
预期结果:每个 expert-owning rank 只收到自己拥有的 expert ids;结果返回 source rank 后与“不通信、直接按全局 expert id 计算”的 reference 完全一致。Gloo 只证明 collective 与布局语义,不提供 NCCL/GPU 性能数字。
import os,platform,socket
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
TOKENS_PER_RANK=3; HIDDEN=4; EXPERTS=4; EXPERTS_PER_RANK=2; WORLD_SIZE=2
def free_port():
with socket.socket() as sock:
sock.bind(("127.0.0.1",0)); return sock.getsockname()[1]
def expert_forward(token,expert_id):
return token*float(expert_id+1)+float(expert_id)
def exchange_split_sizes(send_counts):
gathered=[torch.empty_like(send_counts) for _ in range(WORLD_SIZE)]
dist.all_gather(gathered,send_counts)
return [int(gathered[source][dist.get_rank()].item()) for source in range(WORLD_SIZE)]
def worker(rank,port):
os.environ["MASTER_ADDR"]="127.0.0.1"; os.environ["MASTER_PORT"]=str(port)
dist.init_process_group("gloo",rank=rank,world_size=WORLD_SIZE)
try:
torch.manual_seed(173)
global_tokens=torch.randn(WORLD_SIZE*TOKENS_PER_RANK,HIDDEN)
router=torch.randn(HIDDEN,EXPERTS)
probs=torch.softmax(global_tokens@router,dim=-1)
global_gates,global_experts=probs.max(dim=-1)
start=rank*TOKENS_PER_RANK
local_tokens=global_tokens[start:start+TOKENS_PER_RANK]
local_experts=global_experts[start:start+TOKENS_PER_RANK]
local_gates=global_gates[start:start+TOKENS_PER_RANK]
records=[]
for i,(token,expert,gate) in enumerate(zip(local_tokens,local_experts,local_gates)):
source_id=start+i; expert_id=int(expert.item())
records.append((expert_id//EXPERTS_PER_RANK,source_id,expert_id,token,gate))
records.sort(key=lambda item:item[0])
send_counts=torch.tensor([sum(dest==peer for dest,*_ in records)
for peer in range(WORLD_SIZE)],dtype=torch.int64)
recv_counts=exchange_split_sizes(send_counts)
send_meta=torch.tensor([[source,expert] for _,source,expert,_,_ in records])
send_payload=torch.stack([token for _,_,_,token,_ in records])
send_gates=torch.stack([gate for *_,gate in records])
recv_meta=torch.empty((sum(recv_counts),2),dtype=torch.int64)
recv_payload=torch.empty((sum(recv_counts),HIDDEN))
recv_gates=torch.empty(sum(recv_counts))
dist.all_to_all_single(recv_meta,send_meta,recv_counts,send_counts.tolist())
dist.all_to_all_single(recv_payload,send_payload,recv_counts,send_counts.tolist())
dist.all_to_all_single(recv_gates,send_gates,recv_counts,send_counts.tolist())
processed=torch.stack([gate*expert_forward(token,int(expert.item()))
for token,(_,expert),gate in zip(recv_payload,recv_meta,recv_gates)])
return_counts=torch.tensor(recv_counts,dtype=torch.int64)
combine_counts=exchange_split_sizes(return_counts)
returned_meta=torch.empty((sum(combine_counts),2),dtype=torch.int64)
returned_payload=torch.empty((sum(combine_counts),HIDDEN))
dist.all_to_all_single(returned_meta,recv_meta,combine_counts,return_counts.tolist())
dist.all_to_all_single(returned_payload,processed,combine_counts,return_counts.tolist())
output=torch.empty_like(local_tokens)
for meta,token in zip(returned_meta,returned_payload):
output[int(meta[0].item())-start]=token
reference=torch.stack([gate*expert_forward(token,int(expert.item()))
for token,expert,gate in zip(local_tokens,local_experts,local_gates)])
torch.testing.assert_close(output,reference)
assert all(rank*EXPERTS_PER_RANK<=int(e.item())<(rank+1)*EXPERTS_PER_RANK
for e in recv_meta[:,1])
print(f"[rank {rank}] local_experts={local_experts.tolist()} "
f"send_counts={send_counts.tolist()} recv_counts={recv_counts} "
f"received_experts={recv_meta[:,1].tolist()} max_diff={(output-reference).abs().max():.3e}",flush=True)
finally:
dist.destroy_process_group()
def main():
print(f"python={platform.python_version()} torch={torch.__version__} backend=gloo world_size=2")
print("tokens_per_rank=3 hidden=4 experts=4 top_k=1")
mp.spawn(worker,args=(free_port(),),nprocs=WORLD_SIZE,join=True)
if __name__=="__main__": main()
完整文件位于 examples/moe_all_to_all.py,已用 python -W error 执行。实际路由在两个 source ranks 上都是 [1,2,1];rank 0 收到 4 个 expert-1 routes,rank 1 收到 2 个 expert-2 routes;两边 combine 后最大差均为 0。
脚本为方便每个进程用同 seed 重建同一小 tensor/reference,生产系统不会在每 rank 复制全局 tokens。它也用三个连续 All-to-All 分别发送 metadata、payload 和 gates;实际 runtime 会压缩 metadata、融合 pack/unpack、按 dtype 对齐,并尝试把通信与 grouped GEMM overlap。
计算、通信与延迟成本
Router projection 约需 FLOPs 并产生 [N,E] logits;当 很大时,router 自身也不是零成本。Expert 主 FLOPs 约随选中 routes 増长,而不是所有 combinations,但真实时间由最拥挤 expert/rank、expert GEMM shape 和 kernel batching 决定。
若每条 remote route 携带 个、每元素 bytes 的 hidden state,dispatch 与 combine 的逻辑 payload 主项约为:
前面的 2 是去 expert owner 与返回 source;还要加 expert/source ids、gates、split sizes、padding 与 alignment。Local routes 不经过网络链路,但统一 pack/unpack 仍可能处理它们。通信可与其他 experts 的计算重叠,因此总延迟不是简单 dispatch + all experts + combine 的机械和;尾部仍常受最后一个依赖完成的 route 限制。
常见失败与误解:
- 把总 expert 参数量当成每 token 实际计算量,或反过来忽略 router/shared expert;
- 只报告平均 expert tokens,不报告 max、空 experts 和 overflow;
- All-to-All 的 rank collective 顺序不一致,导致死锁;
- token pack 后丢失 source/order metadata,combine 数值正确却写回错误 token;
- Top-2 两条 route 在 combine 前没有按 gate 累加,或 gate normalization 与训练不一致;
- capacity drop 在推理中静默改变模型语义;
- 把本机 Gloo 时间当作 NCCL、多 GPU 或跨节点 All-to-All 性能;
- expert parallel group 跨慢链路,而拓扑/placement 让大量 routes 离开高速互联域。
与上下游的组合
- Rank、Process Group 与集合通信定义 collective 顺序、shape 与逻辑所有权;本页把 All-to-All 用在 ragged token dispatch/combine。
- 张量并行切一个 expert 内部的矩阵;expert parallel 把不同 experts 放到不同 ranks。两者组合时先明确 TP group 与 EP group,不能只说“用了多卡”。
- Continuous Batching改变每 iteration 的 token mix,因而也改变 router histogram、expert batch shape 和 All-to-All splits;相同 batch size 不保证相同 MoE latency。
- 模型加载必须按 expert→rank placement 只加载本 rank experts,并正确切分其 TP shards、量化 scales 与 metadata。