专家混合模型(Mixture of Experts,简称MoE)已经成为降低大规模语言模型推理成本的重要路径。与稠密模型每次前向都唤醒全部参数不同,MoE把网络拆成多个前馈专家,并引入门控模块在推理时动态挑选其中一小部分参与计算。这种稀疏激活思路让模型总参数量可以非常大,而单条样本的实际计算量却维持在较低水平。Switch Transformer和Mixtral是两种具有代表性的实现,它们在路由粒度、专家数量和负载均衡机制上采取了不同策略,直接影响推理效率与部署复杂度。

Switch Transformer的单一专家路由机制
Switch Transformer的核心创新是把门控网络简化到极致:对于每个输入词元,门控只输出一个概率最高的专家编号,也就是所谓的Top-1路由。这种做法相比此前的Top-k多专家混合,能将路由计算和专家调用的开销降到最低,因为每一个词元在任意一层都只会被发送到一个专家前馈网络。从推理视角看,这意味着显存中虽然驻留了全部专家参数,但真正发生矩阵乘法的只有被选中的那一份,计算量随专家总数近似线性下降。
在代码层面,Switch Transformer的路由可以用如下简化逻辑表达。门控先对隐状态做线性投影得到logits,再取argmax得到专家索引,最后只把对应词元送入该专家。下面的示例展示了不考虑并行通信时的单卡推理路由:
import torch
def switch_route(hidden, gate_weight, expert_list):
# hidden: [batch, seq_len, dim]
logits = torch.matmul(hidden, gate_weight) # [batch, seq_len, num_experts]
expert_idx = torch.argmax(logits, dim=-1) # Top-1 选择
outputs = torch.zeros_like(hidden)
for i, expert in enumerate(expert_list):
mask = (expert_idx == i)
if mask.any():
outputs[mask] = expert(hidden[mask])
return outputs
这种设计的劣势同样明显。由于只选一个专家,一旦某些专家被频繁选中,就会出现严重的负载倾斜,部分专家过载而其他专家闲置,导致设备利用率低下。Switch Transformer在训练时引入了负载均衡损失,但在纯推理服务中,如果输入分布偏移,依然可能出现热点。因此在工程落地时,往往需要对专家放置做手工调度,或采用容量因子限制单个专家处理的词元数量。
Mixtral的稀疏八选二路由策略
Mixtral采用了更为灵活的Top-2路由:在每层设置八个专家,每个词元选择得分最高的两个专家,并按归一化权重进行加权融合。相比Switch Transformer的强制单专家,Mixtral用略微增加的计算换来了更稳定的表达能力和更均衡的负载。推理时,一个词元会触发两份专家计算,但总体仍是八分之二的稀疏度,实际浮点运算量远小于稠密模型。
Mixtral的门控实现更强调数值稳定与融合权重。它先对八个专家的logits做softmax,取前两名,再把这两名的权重重新归一化,保证输出尺度一致。下面给出推理路由的参考实现:
import torch
def mixtral_route(hidden, gate_weight, expert_list, k=2):
logits = torch.matmul(hidden, gate_weight)
probs = torch.softmax(logits, dim=-1)
top_val, top_idx = torch.topk(probs, k, dim=-1)
top_val = top_val / top_val.sum(dim=-1, keepdim=True)
out = torch.zeros_like(hidden)
for j in range(k):
idx = top_idx[..., j]
w = top_val[..., j].unsqueeze(-1)
for i, expert in enumerate(expert_list):
mask = (idx == i)
if mask.any():
out[mask] += w[mask] * expert(hidden[mask])
return out
从部署角度看,Mixtral的Top-2结构对并行更友好。由于每个词元固定落在两个专家上,调度器可以预先把专家按卡分组,减少跨设备传输。与此同时,八选二让单一专家的容量压力被分散,推理批次中的长尾词元不会把某个专家打满。实际测试中,相同参数量级下Mixtral的吞吐表现通常优于严格单专家的Switch风格模型,尤其在多轮对话这种输入长度波动大的场景。
推理系统中的专家并行与通信开销
无论是Switch Transformer还是Mixtral,当专家数量超过单卡显存或算力时,都必须引入专家并行。典型做法是把不同专家分布到多张加速卡,门控在每层把词元通过网络发送给目标卡,计算完再送回。此时推理延迟不再只由计算量决定,还受到全互联带宽和调度策略制约。如果路由散列度高,通信量会明显上升,抵消稀疏激活带来的收益。
在Switch Transformer的极端单专家场景下,由于每个词元只去一个专家,理论上通信模式更简单,但若出现热点专家集中在一张卡,该卡会成为瓶颈。Mixtral因为Top-2,词元至少涉及两张卡,通信图更均匀,却也意味着每步要多一次专家间数据交换。工程上常用容量缓冲和丢弃策略来控制最坏情况,例如当某专家接收词元超过容量上限时,多余部分直接走残差,避免阻塞整个批次。
下面是一个简化的专家并行发送伪代码,展示如何按专家索引做跨卡分发:
def dispatch_to_experts(hidden, expert_idx, num_experts, rank_table):
# rank_table: 专家到卡号的映射
buckets = [[] for _ in range(num_experts)]
for i, e in enumerate(expert_idx.flatten()):
buckets[e].append(i)
sent = {}
for e, ids in enumerate(buckets):
dst = rank_table[e]
sent[dst] = hidden[ids]
# 实际系统调用集合通信发送 sent[dst]
return sent
总结来看,Switch Transformer用最简路由换极致稀疏,适合可控输入分布的离线批量推理;Mixtral以稍高计算成本换取鲁棒性与高吞吐,更适配线服务。理解二者在路由粒度、负载均衡和并行通信上的差异,是做MoE推理优化的第一步。
MoEinferenceSwitch_Transformer修改时间:2026-08-14 08:45:35