架构演进 · 仅解码器
- 对比
- Mixtral 8x7B vs GPT-2 small
- 层
- 32 GPT-2 12
- dmodel
- 4,096 GPT-2 768
- MLP
- 8 个专家 · top 2 GPT-2 1 个稠密
- dff
- 每个 14,336 GPT-2 3,072
- 参数
- 46.7B · 激活 12.9B GPT-2 124M
- 上下文
- 32,768 GPT-2 1,024
全部步骤
GPT-2 vs Mixtral
Mixtral 的块除 MLP 以外都和 LLaMA 一样(RMSNorm、RoPE、分组查询注意力):每层有 8 个专家 MLP 和一个小路由器,把每个词元送给其中 2 个。点击标签跳到对应部分。
每层 8 个专家 · top 2 · 32 层路由器
路由器就是一个小矩阵:词元的向量乘以 Wg,为每个专家得到一个分数。得分最高的两个胜出,再只对这两个做 softmax,决定每个胜者占多少权重。把鼠标停在格子上看看。
g = softmax(top-2(x · Wg))每个词元两个专家
每个词元只经过它的两个专家,两个输出按路由器的权重相加。不同的词元选择不同的专家,所以有些专家分到的批次比别的多,有些一个也分不到。
y = ga · Ea(x) + gb · Eb(x)存储 vs 使用
全部 8 个专家都必须放在内存里:46.7B 参数。但每个词元只跑其中 2 个,12.9B 参数,所以它的计算量大约相当于一个 13B 的稠密模型。
存储 46.7B · 每词元使用 12.9B负载均衡
放任不管的话,路由器会一直选择已经表现好的专家,其他专家就停止学习了。MoE 训练会加一个小损失,当词元和路由器概率都在专家间均匀分布时,这个损失最低。
Laux = N · Σ fi · Pi
代码
# the block is LLaMA's, with a sparse MoE layer where the MLP was
h = x + self_attn(rms_norm(x)); out = h + moe(rms_norm(h))
# moe: score the 8 experts, keep the best 2, renormalise their weights
router_logits = self.gate(x) # (tokens, 8)
weights = F.softmax(router_logits, dim=-1)
weights, experts = torch.topk(weights, 2, dim=-1)
weights /= weights.sum(dim=-1, keepdim=True) # = softmax over the top 2
y = sum(w * self.experts[e](x) for w, e in zip(weights, experts))
# each expert is a SwiGLU MLP
expert_out = self.w2(F.silu(self.w1(x)) * self.w3(x))
# training adds a balancing loss: share routed × mean probability
aux = n_experts * (f * P).sum()