Architectures · Decoder-only
- compared
- Mixtral 8x7B vs GPT-2 small
- layers
- 32 GPT-2 12
- dmodel
- 4,096 GPT-2 768
- MLP
- 8 experts · top 2 GPT-2 1 dense
- dff
- 14,336 each GPT-2 3,072
- params
- 46.7B · 12.9B active GPT-2 124M
- context
- 32,768 GPT-2 1,024
All steps
GPT-2 vs Mixtral
Mixtral’s block is LLaMA’s (RMSNorm, RoPE, grouped-query attention) except for the MLP: each layer has 8 expert MLPs and a small router that sends each token to 2 of them. Click a label to jump to that part.
8 experts per layer · top 2 · 32 layersThe router
The router is one small matrix: a token’s vector times Wg gives one score per expert. The two highest scores win, and a softmax over just those two sets how much each winner counts. Hover the cells.
g = softmax(top-2(x · Wg))Two experts per token
Each token runs through its two experts only, and the two outputs are added with the router’s weights. Different tokens pick different experts, so some experts get more of the batch than others, and some get none.
y = ga · Ea(x) + gb · Eb(x)Stored vs used
All 8 experts must sit in memory: 46.7B parameters. But each token runs only 2 of them, 12.9B parameters, so it costs about as much compute as a 13B dense model.
46.7B stored · 12.9B used per tokenLoad balancing
Left alone, a router keeps choosing the experts that are already good, and the others stop learning. MoE training adds a small loss that is lowest when both the tokens and the router’s probability are spread evenly over the experts.
Laux = N · Σ fi · Pi
Code
# 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()