最近翻训练日志的时候发现一个现象:16 个专家里,有 2 个专家吃掉了 90% 的 token 转发,GPU 利用率看着还行,可整体吞吐就是上不去,loss 到了 3000 步之后开始来回抖。这种情况你要是也遇到过,说明手头的 MoE 模型已经出现了严重的路由负载不均衡。今天想聊的“高性能计算负载均衡”,不打算泛泛谈集群任务调度那一套,而是聚焦大模型 MoE 训练与推理场景里最典型的 token 级负载均衡问题,尤其是“等开销负载均衡”的原理和可运行代码。
如果你正在折腾 MoE 架构,或者准备在大规模分布式训练里做专家并行,这篇文章应该能帮你省几天的排查时间。我会把为什么会产生负载倾斜讲透,再把 Switch Transformer 那套 auxiliary load balancing loss 的公式拆开,最后给一份可以直接在 PyTorch 里跑通的实现,附带工程中的调参与避坑经验。
1. 先搞清楚负载不均衡的病根
1.1 MoE 路由的本质:给每个 token 找“专家”
MoE 模型的核心思路是“谁行谁上”。网关网络会为每个输入 token 计算一组概率,选出 top-k 个专家,然后只让这些专家处理该 token。这样模型虽然拥有几十亿甚至上千亿参数,实际计算时只会激活一小部分参数,理论上训练和推理的算力成本比同规模稠密模型低很多。
但问题恰恰出在这个“选人”环节上。路由器是数据学出来的,不是人工设计出来的。它天然会偏好那些“好用的”专家——谁梯度有效、谁让 loss 降得快,谁就会越来越频繁地被选中。这种赢家通吃的局面一旦形成,就会出现非常极端的分配:少数专家忙到排队,多数专家闲着看戏。
我给小白读者打个比方:假设你有 8 个仓库分拣员,包裹按地址自动派单。派单规则不是平均分,而是谁近谁快谁接单。结果地址相似、时段集中的包裹全涌向同一个分拣员,其他人无事可做。你说分拣员没问题,他是按要求接单的;问题出在派单策略没有考虑整体负载。
在训练数据里,这种“包裹扎堆”的现象非常普遍。一段文本里可能大量 token 都是停用词,一个视觉 batch 里可能全是相似背景的 patch,它们的特征分布高度聚集,天然会把 router 的概率推向同一批专家。
1.2 负载不均衡到底要付出多少代价
很多人一开始觉得,专家负载不均匀不过是浪费点算力,无所谓。真上了规模之后会发现事情远没有那么简单。
算力利用率的损失是最直观的。假设有 8 个专家,其中 2 个承担了 90% 的 token,那就有 6 个专家几乎空转。空转的专家依旧占显存,却只产生极少的有效计算。与此同时,过载专家所在的计算单元忙到排队。表面上 GPU 利用率可能不低,但真正有效的计算吞吐远低于硬件上限。我在一个 200M 参数左右的测试模型上处理过类似问题,加均衡损失之前单步耗时 4 秒出头,加完均衡损失重新训练,单步大约 3.5 秒,提升接近 20%。这个数字在不同模型上会变化,但趋势是一致的。
设备热点问题在专家并行场景下更明显。过载专家所在的 GPU 上,激活值、梯度通信量都会暴涨。All-to-All 通信本来就很吃带宽,一旦通信量集中在某几张卡上,整个集群的训练速度都会被拖累。
训练稳定性的问题则更隐蔽。长时间负载不均衡会导致某些专家长期拿不到梯度,逐步变成“死专家”;而死专家变多,router 能选的参数就更少,容易陷入更严重的倾斜,形成恶性循环。我见过不少后期 loss 震荡的 MoE 训练,排查到最后都会发现专家坍缩的影子。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 等开销负载均衡的设计思路与公式
2.1 等开销到底是怎么个“等”法
先说清楚“等开销”这个词。开销指计算开销,等开销意味着让每个专家的工作量期望趋同,也就是让每个专家处理的 token 数量大致相同。注意这里说的是统计期望上的相同,不是每一步都强制硬性平均。
为什么不直接硬性平均?因为硬约束会破坏路由器的学习空间。token 本来就有自己的特性,某些专家确实更擅长处理某种 token。如果强行规定每个专家每步只能接收固定的 token 数,路由器的个性化分配能力会被直接废掉,模型质量反而会下降。
所以工业界的主流做法是“软惩罚”。在训练总损失之外,额外加一项负载均衡损失,让路由器在优化主任务的同时,被温和地推向更均匀的分配。这个方案最早被 Switch Transformer 大规模验证,现在几乎成了 MoE 训练的标配组件。所谓“等开销负载均衡”,指的就是这一类让各专家计算开销趋于均衡的正则化手段。
2.2 公式的每一项都不是拍脑袋
Switch Transformer 核心的负载均衡辅助损失可以写成:
text复制L_aux = α * N * Σ_{i=1}^{N} f_i * P_i
我逐项拆开讲:
- N 是专家数量,乘在式子前面用来控制损失尺度,防止专家越多损失越膨胀。
- f_i 是第 i 个专家在本次前向中实际接收到的 token 占比。这个量来自 top-k 选择的离散统计,本身不可导。比如总共 100 个 token,top-2 选择后专家 3 被选中了 40 次,那 f_3 就是 0.4。
- P_i 是所有 token 对第 i 个专家的平均路由概率,来自 softmax 输出的均值。这个量是可导的,是整个损失里能反传梯度的关键。
- α 是缩放系数,控制均衡诉求对主任务的影响强度,常见起始值是 0.01。
f_i 和 P_i 相乘的含义需要仔细体会一下。某专家如果既频繁被选中(f_i 大),又得到 router 很高的平均概率(P_i 大),乘积就大,对应惩罚就高。当某个专家过载时,梯度会倾向于降低 router 给它的概率;而概率降下来之后,下一轮 top-k 选择里它被选中的次数就会变少,f_i 随之下降。虽然 f_i 本身因为离散选择无法直接反传梯度,但它随着概率变化而动态调整,整体形成了一种近似闭环的均衡机制。
P_i 的计算细节经常被实现错。注意它是对所有 token 的路由概率取平均,而不是只对“被选中位置”的概率取平均。论文里的定义就是所有 token 分配给该专家的平均概率。这个细节会影响梯度方向,实现错了一步,整个正则项的行为都会跑偏。
2.3 常见变体和需要避开的实现误区
很多开源实现里会额外把损失除以 token 数量,这在 scaling 上更友好,但核心公式保持一致。还有一层容易踩坑:如果你用的是 top-2 或 top-k,专家选中次数的总和会是 k 倍 token 数,所以 f_i 累加后等于 k,而不是 1。这不影响公式有效性,但会导致辅助损失的绝对尺度比 top-1 场景更大,调 α 时需要重新标定。
如果你在一个已经收敛得很好的 MoE 模型上突然加入辅助损失,尤其是 α 还设得比较大,会发现主任务 loss 明显恶化。这是因为 router 原有的偏好被正则项强行拉扯。正常做法是从训练开始就带上辅助损失,让 router 在学习过程中同时适应均衡约束;而不是事后补救。
3. 落地代码:手写一个等开销负载均衡实现
3.1 设计一个可以直接跑通的微型 MoE 层
为了把原理讲明白,先从零写一个简化版 MoE 层。这个实现不是性能最优的,但结构完整,方便你在本地改着玩。
python复制import torch
import torch.nn as nn
import torch.nn.functional as F
class Expert(nn.Module):
"""单个专家网络:一个简单的 MLP。"""
def __init__(self, hidden_dim):
super().__init__()
self.net = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim * 4),
nn.GELU(),
nn.Linear(hidden_dim * 4, hidden_dim),
)
def forward(self, x):
return self.net(x)
class MoELayer(nn.Module):
def __init__(self, hidden_dim, num_experts=8, top_k=2):
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
self.experts = nn.ModuleList([Expert(hidden_dim) for _ in range(num_experts)])
self.router = nn.Linear(hidden_dim, num_experts)
def forward(self, x):
# x: [batch, seq_len, hidden_dim]
flat_x = x.view(-1, x.size(-1))
router_logits = self.router(flat_x) # [num_tokens, num_experts]
probs = F.softmax(router_logits, dim=-1)
top_k_probs, top_k_indices = torch.topk(probs, self.top_k, dim=-1)
out = torch.zeros_like(flat_x)
for k in range(self.top_k):
expert_idx = top_k_indices[:, k] # 每个 token 在第 k 个位置选的专家
weight = top_k_probs[:, k].unsqueeze(-1) # 对应权重
for e in range(self.num_experts):
mask = (expert_idx == e)
if mask.any():
out[mask] += weight[mask] * self.experts[e](flat_x[mask])
return out.view(x.shape), router_logits
这个循环双层的实现效率不高,生产环境不建议直接用。但作为教学 demo,它能把“top-k 选择”、“专家分别处理”、“加权求和”这几个环节分离得很清楚。真正想上生产,可以考虑把 mask 选择改成分组矩阵乘,或者直接用 Megablocks 这类高效实现。
3.2 核心辅助均衡损失实现
负载均衡损失函数不长,但细节一个都不能错。
python复制def load_balancing_loss(router_logits, top_k=2):
# router_logits: [num_tokens, num_experts]
num_tokens, num_experts = router_logits.shape
probs = F.softmax(router_logits, dim=-1)
# 用 top-k 索引构造 one-hot,统计每个专家实际接收的 token 占比
top_k_indices = torch.topk(probs, top_k, dim=-1).indices # [num_tokens, top_k]
expert_mask = F.one_hot(top_k_indices, num_classes=num_experts)
expert_mask = expert_mask.sum(dim=1) # [num_tokens, num_experts]
f_i = expert_mask.float().mean(dim=0) # 每个专家的 token 占比
# 平均路由概率:所有 token 对每个专家 softmax 概率的均值
p_i = probs.mean(dim=0) # [num_experts]
# 辅助负载均衡损失
aux_loss = num_experts * torch.dot(f_i, p_i)
return aux_loss
有几个点值得展开说明。
f_i 为什么需要 expert_mask.sum(dim=1)?因为一个 token 在 top-2 里会选两个专家,构造 one-hot 后维度是 [num_tokens, top_k, num_experts],需要先把 top_k 维拍平求和,才能得到每个专家被选中的总次数。
为什么 P_i 用 probs.mean(dim=0)?回到论文定义,P_i 是所有 token 对专家 i 的平均路由概率,不是被选中位置的条件概率。很多实现会在这一步做成 gather 后取均值,行为会有微妙偏差,实际训练中不一定会立刻爆,但会让均衡趋势变怪。
为什么从 probs 而不是 router_logits 取 top-k?因为 softmax 是单调变换,从 logits 和从 probs 取 top-k 结果一致。代码里后面反正要用 probs 计算 P_i,顺手统一一下,不用再维护两套索引。
3.3 把均衡损失集成到训练循环里
有了 MoE 层和损失函数,下一步是把它接进训练循环。这里没有依赖真实数据集,直接喂随机张量做演示。
python复制torch.manual_seed(0)
model = MoELayer(hidden_dim=64, num_experts=8, top_k=2)
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
# 用随机输入模拟一个回归任务,只是为了演示 loss 怎么叠
for step in range(30):
x = torch.randn(2, 16, 64)
target = torch.randn(2, 16, 64)
y, router_logits = model(x)
task_loss = F.mse_loss(y, target)
aux_loss = load_balancing_loss(router_logits, top_k=2)
total_loss = task_loss + 0.01 * aux_loss
optimizer.zero_grad()
total_loss.backward()
optimizer.step()
if step % 10 == 0:
print(f"step={step}, task_loss={task_loss.item():.4f}, aux_loss={aux_loss.item():.4f}")
日常训练里你大概率不会用随机输入,但叠加逻辑是一样的:任务损失加负载均衡损失,乘上权重 α,一起反传。
我建议训练脚本里把 aux_loss 单独拎出来打日志,不要只看 total_loss。因为 task_loss 通常比 aux_loss 大一个量级,total_loss 的变化无法反映均衡约束在发生什么。单独记录后,你能观察到 aux_loss 在训练早期明显下降,随后在一个平台区域浮动,这是正常现象。
3.4 怎么验证负载均衡是否生效
很多读者跑完代码想知道“正则到底起没起作用”。方法很简单:人为制造一组不均衡的 router logits 和一组相对均衡的 logits,对比辅助损失的输出。
python复制torch.manual_seed(42)
# 模拟不均衡路由:前两个专家的 logits 被加偏置,softmax 概率会明显偏高
bad_logits = torch.randn(1024, 8)
bad_logits[:, 0] += 6
bad_logits[:, 1] += 4
# 模拟相对均衡路由:logits 本身很小,softmax 后接近均匀分布
good_logits = torch.randn(1024, 8) * 0.1
bad_loss = load_balancing_loss(bad_logits, top_k=2)
good_loss = load_balancing_loss(good_logits, top_k=2)
print(f"bad_logits aux_loss: {bad_loss.item():.4f}")
print(f"good_logits aux_loss: {good_loss.item():.4f}")
跑一次你会发现 bad case 的损失明显高于 good case,说明损失确实在惩罚不均衡分配。
除了直接看 loss 数值,每次训练迭代后还可以打印各专家的 token 数量分布:
python复制def expert_load_count(router_logits, top_k=2):
probs = F.softmax(router_logits, dim=-1)
indices = torch.topk(probs, top_k, dim=-1).indices
mask = F.one_hot(indices, num_classes=router_logits.size(-1)).sum(dim=1)
return mask.float().sum(dim=0).long().tolist()
print(expert_load_count(router_logits))
理论上,一个已经训练稳定且负载均衡的 MoE 模型里,每个专家的 token 数量不会完全一致,但分布会比较接近。像开头那种 2 个专家吃掉 90% token 的局面,就不太应该出现。
4. 工程实践中的坑与经验
4.1 α 怎么调:不是越大越好
α 是最让人头疼的超参。很多人习惯性抄 0.01,但它不是万能解。α 太小,均衡效果不足;α 太大,路由器的个性化选择会被牺牲,模型质量明显下降。
我整理了一个经验参考表,按场景粗略划分:
| α 取值区间 | 行为表现 | 适用场景 |
|---|---|---|
| 0.0001 ~ 0.001 | 几乎不干预均衡 | 路由器偏好很重要,且实验规模小,不担心坍缩 |
| 0.01 | 均衡效果明显,对模型质量影响可控 | 多数 MoE 训练的默认起点 |
| 0.05 ~ 0.1 | 均衡非常显著,模型质量可能出现小幅下降 | 严重不均衡或专家坍缩风险高时使用 |
| 大于 0.1 | 路由个性化基本被压制 | 除非专门研究负载均衡本身,否则不建议 |
注意一点:top-k 的数量会影响辅助损失的绝对尺度。前面说过,top-2 时所有专家被选中次数的总和是 2 倍 token 数,所以 aux_loss 比 top-1 场景大不少。你从 top-1 切到 top-2 时,如果发现均衡效果过强,可以把 α 减半甚至降到原来的四分之一。
4.2 损失尺度不一致造成的假象
有段时间我调试一个 MoE 模型,训练后期发现专家分布还是很差,但 aux_loss 看起来并不大。后来单独检查才发现,任务 loss 已经降到 0.02 量级,而辅助损失还稳定在 1.8 左右。乘上 α=0.01 之后,它在总损失里的占比其实超过任务损失,但数值上仍然“看起来很小”,实际影响一点都不小。
这提醒两件事。第一,不要只盯着 total_loss,把 task_loss 和 aux_loss 分开可视化;第二,如果训练后期发现任务 loss 异常波动,尝试把 α 降一个数量级或者动态衰减,看是否恢复。有时候模型质量上不去,不是模型结构问题,而是正则项在后期“喧宾夺主”。
4.3 分布式训练下的统计正确性
真正的千亿级 MoE 训练一定是多机多卡。这时统计 f_i 和 P_i 不再像单机 demo 那么简单。每张卡只看到自己的 local batch,如果不跨设备同步就计算辅助损失,等于用局部分布代替全局分布,结果会有明显偏差。
正确做法是:先在本卡统计每个专家被选中的次数 local_counts,然后对 local_counts 做 all-reduce 求和,得到全局计数;同时把本卡 token 数也做 all-reduce,最后用全局计数除以全局 token 数得到 f_i。P_i 的处理类似,先求每个专家的概率总和,除以全局 token 总数,而不是简单做 MEAN 聚合(除非确认每卡 batch 完全一致)。
写成示意代码大致是:
python复制import torch.distributed as dist
def compute_global_f_and_p(local_counts, local_prob_sum, local_num_tokens, top_k):
dist.all_reduce(local_counts, op=dist.ReduceOp.SUM)
dist.all_reduce(local_prob_sum, op=dist.ReduceOp.SUM)
global_tokens = torch.tensor(local_num_tokens, device=local_counts.device)
dist.all_reduce(global_tokens, op=dist.ReduceOp.SUM)
f_i = local_counts.float() / (global_tokens * top_k)
p_i = local_prob_sum / global_tokens
return f_i, p_i
local_prob_sum 是本地所有 token 对每个专家概率的和。这样算出来的 f_i 和 p_i 才是全局视角的。很多分布式框架内部已经封装了类似逻辑,但如果你自己写训练脚本,这个坑几乎一定会踩到。
4.4 专家坍缩与“死专家”现象
均衡损失能缓解负载倾斜,但它并不能根治“死专家”。当一个专家从训练早期就完全不被选中时,它的梯度长期为零,后续也很难再被激活。辅助损失只是惩罚过度集中,并没有强制每个专家都被启用的机制。
我在实践中一般额外做三件事:
- 定期统计每个专家的 token 接收量,对连续大量步数都为 0 的专家做标记。
- 给 router bias 或者 logits 加一点温度扰动,让冷门专家有概率被“翻牌子”。
- 配合 z-loss 使用。z-loss 是约束 router logits 尺度的一项正则,能有效减少 router 输出的方差,防止 logits 无限膨胀导致 softmax 概率走向极端。z-loss 和负载均衡损失不冲突,可以叠在一起用。
5. 等开销之外的进阶负载均衡思路
5.1 无辅助损失的动态 bias 方案
等开销负载均衡损失的优点是简单、通用,但它的代价是引入了额外超参 α,而且无论如何都会对主任务损失造成一点点干扰。为了彻底避免这种干扰,业界开始出现“无辅助损失”的负载均衡方案。
思路很直接:给每个专家维护一个动态 bias,加到路由 logits 上。负载低的专家 bias 逐渐提高,负载高的专家 bias 逐渐降低,router 在 top-k 选择时被这份偏移量“推”向冷门专家。偏差不参与主任务损失的反传,而是单独根据专家负载状态去更新。
和辅助损失的差别在于:辅助损失是在优化目标里做软约束,动态 bias 是在路由决策前做在线偏置。前者更容易实现,后者对主任务的影响更干净,但工程复杂度更高,尤其是在多机场景下,bias 的同步与更新节奏需要额外设计。
5.2 不同场景下怎么选方案
中小规模实验、快速验证模型结构的时候,我建议直接用辅助负载均衡损失,省心省事,一个函数搞定。大规模生产训练,如果对主任务指标极度敏感,不希望任何多余的正则干扰 router 的学习,可以参考无辅助损失思路。
在线推理场景的负载均衡又是另一套逻辑。推理时 token 到达模式动态变化,不存在“训练步”的全局同步窗口,更像是一个在线调度问题。这时候辅助损失就不太适合了,需要的是在推理框架层面做动态队列调度。这已经超出训练正则化的范畴,但值得对生产系统感兴趣的读者留意。
最后想分享的一点体会
我自己调试 MoE 模型时踩过最深的坑,是只盯着 total_loss 看,结果训练很久才发现专家分布已经歪得不成样子。后来把辅助损失和专家 token 统计量单独打日志、单独可视化,才真正看清模型里发生了什么。
如果你也准备引入等开销负载均衡,我建议第一步不要急着调 α,而是先把专家负载分布打印出来,搞清楚你的模型到底有多不均衡。拿了基线,再上正则项,前后对比才有意义。否则加完损失,连它有没有生效都不知道。
