架构演进 · 仅解码器

Mixtral

Mixtral 8x7B · 用混合专家取代 MLP

对比
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

全部步骤

  1. GPT-2 vs Mixtral

    Mixtral 的块除 MLP 以外都和 LLaMA 一样(RMSNorm、RoPE、分组查询注意力):每层有 8 个专家 MLP 和一个小路由器,把每个词元送给其中 2 个。点击标签跳到对应部分。

    每层 8 个专家 · top 2 · 32 层

  2. 路由器

    路由器就是一个小矩阵:词元的向量乘以 Wg,为每个专家得到一个分数。得分最高的两个胜出,再只对这两个做 softmax,决定每个胜者占多少权重。把鼠标停在格子上看看。

    g = softmax(top-2(x · Wg))

  3. 每个词元两个专家

    每个词元只经过它的两个专家,两个输出按路由器的权重相加。不同的词元选择不同的专家,所以有些专家分到的批次比别的多,有些一个也分不到。

    y = ga · Ea(x) + gb · Eb(x)

  4. 存储 vs 使用

    全部 8 个专家都必须放在内存里:46.7B 参数。但每个词元只跑其中 2 个,12.9B 参数,所以它的计算量大约相当于一个 13B 的稠密模型。

    存储 46.7B · 每词元使用 12.9B

  5. 负载均衡

    放任不管的话,路由器会一直选择已经表现好的专家,其他专家就停止学习了。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()

延伸阅读