Mixture of Experts
Kunna förklara routing, expertkapacitet och lastbalansering i MoE-modeller.
Förkunskaper
Intuition
En vanlig transformer kör hela MLP-blocket för varje token. En mixture of experts ersätter MLP:n med många parallella experter och låter en liten router välja ut ett fåtal per token.
┌──────────┐
token ──→ │ router │──→ väljer expert 3 och 17 av 64
└──────────┘
│
┌──────────┴─────────┐
expert 3 expert 17
└──────────┬─────────┘
viktad summa
Vinsten: modellen kan ha mycket fler parametrar utan att varje token kostar mer att beräkna.
| Tät modell | MoE | |
|---|---|---|
| Totala parametrar | 8 B | 47 B |
| Aktiva per token | 8 B | 13 B |
| Beräkning per token | 1× | ~1,6× |
| Minne | 16 GB | 94 GB |
Haken syns på sista raden: alla experter måste finnas i minnet även om bara två används. MoE byter minne mot kapacitet — det är en helt annan affär än kvantisering.
Formellt
Routern är ett linjärt lager följt av softmax över experterna. Top- (oftast ) väljs, och utdatan är den viktade summan:
Lastbalansering är det centrala problemet. Utan motverkan uppstår en självförstärkande spiral: en expert som råkar väljas oftare tränas mer, blir bättre, och väljs ännu oftare. Till slut används en handfull experter och resten är döda vikter.
Lösningen är en hjälpförlust som straffar ojämn fördelning:
där är andelen tokens som routades till expert och är den genomsnittliga routersannolikheten för den. Termen minimeras när båda är jämnt fördelade. Typiskt .
Expertkapacitet. Varje expert har en gräns för hur många tokens den tar emot per batch:
Med kapacitetsfaktor 1,25 tål varje expert 25 % över sin proportionella andel. Tokens över gränsen droppas — de passerar oförändrade via residualkopplingen. Låg faktor sparar minne men droppar mer; hög faktor slösar.
Tre praktiska problem:
| Problem | Varför |
|---|---|
| Kommunikation | experterna ligger på olika GPU:er → all-to-all-utbyte per lager |
| Instabil träning | routern kan oscillera; z-loss och lägre lärhastighet på routern hjälper |
| Finjustering | MoE-modeller överanpassar lättare; färre epoker, mer regularisering |
Varianter: delade experter som alltid är aktiva (DeepSeek), expertval i stället för tokenval (varje expert väljer sina tokens, vilket ger perfekt balans per konstruktion), och mycket finkorniga experter (hundratals små i stället för åtta stora).
När MoE lönar sig: när du har gott om minne men vill ha fler parametrar per FLOP. När minnet är begränsat är en tät modell nästan alltid rätt val.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
class MoE(nn.Module):
def __init__(self, d=512, dolt=2048, n_experter=8, k=2, kapacitetsfaktor=1.25):
super().__init__()
self.k, self.n = k, n_experter
self.kapfaktor = kapacitetsfaktor
self.router = nn.Linear(d, n_experter, bias=False)
self.experter = nn.ModuleList([
nn.Sequential(nn.Linear(d, dolt), nn.GELU(), nn.Linear(dolt, d))
for _ in range(n_experter)])
def forward(self, x):
B, T, d = x.shape
xf = x.reshape(-1, d) # (N, d)
N = xf.size(0)
logits = self.router(xf)
sannolikhet = F.softmax(logits, dim=-1) # (N, n)
vikt, val = sannolikhet.topk(self.k, dim=-1) # (N, k)
vikt = vikt / vikt.sum(dim=-1, keepdim=True) # normalisera de valda
kapacitet = int(self.kapfaktor * N * self.k / self.n)
ut = torch.zeros_like(xf)
droppade = 0
for e in range(self.n):
mask = (val == e)
index = mask.any(dim=-1).nonzero(as_tuple=True)[0]
if len(index) > kapacitet: # över kapacitet → droppa
droppade += len(index) - kapacitet
index = index[:kapacitet]
if len(index) == 0:
continue
w = (vikt * mask.float()).sum(dim=-1)[index].unsqueeze(-1)
ut[index] += w * self.experter[e](xf[index])
# Lastbalanseringsförlust: jämn fördelning av både tokens och sannolikhet
f = torch.zeros(self.n, device=x.device)
for e in range(self.n):
f[e] = (val == e).any(dim=-1).float().mean()
P = sannolikhet.mean(dim=0)
aux = self.n * (f * P).sum()
self.senaste = {"aux_loss": float(aux),
"droppad_andel": round(droppade / max(N * self.k, 1), 4),
"fordelning": [round(float(v), 3) for v in f]}
return ut.view(B, T, d), aux
# Parameterräkning: totalt mot aktivt
def parametrar(d=4096, dolt=14336, n_experter=8, k=2, lager=32):
per_expert = 3 * d * dolt
totalt = lager * n_experter * per_expert
aktivt = lager * k * per_expert
return {"totalt_md": round(totalt / 1e9, 1), "aktivt_md": round(aktivt / 1e9, 1),
"kvot": round(totalt / aktivt, 1)}
print(parametrar())
# {'totalt_md': 45.1, 'aktivt_md': 11.3, 'kvot': 4.0}
# ↑ fyra gånger fler parametrar än vad som beräknas per token
# Kollapsad routing ser ut så här
kollaps = [0.71, 0.24, 0.03, 0.01, 0.01, 0.0, 0.0, 0.0]
balanserad = [0.13, 0.12, 0.13, 0.12, 0.13, 0.12, 0.12, 0.13]
for namn, f in (("kollapsad", kollaps), ("balanserad", balanserad)):
import math
entropi = -sum(p * math.log(p) for p in f if p > 0)
print(f"{namn:<11} entropi {entropi:.3f} av max {math.log(8):.3f}")
# kollapsad entropi 0.783 av max 2.079
# balanserad entropi 2.079 av max 2.079
Att logga routingens entropi är den enklaste kollapsdetektorn: faller den under ungefär halva maxvärdet används bara en bråkdel av experterna, och hjälpförlustens vikt behöver höjas.
Behärskning innebär
- Förklarar routing och top-k-val
- Beskriver lastbalanseringsproblemet
- Vet vad MoE sparar och vad det kostar
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Switch Transformers: Scaling to Trillion Parameter Models with Simple and Efficient Sparsity — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Mixtral of Experts — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer — arXiv (öppen åtkomst; licens per artikel)