Hoppa till innehållet
AI-grafen
F· AI engineeringtransformerarkitektur· ca 90 min· volatil — kontrolleras ofta· verifierad 2026-09-21

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 modellMoE
Totala parametrar8 B47 B
Aktiva per token8 B13 B
Beräkning per token1×~1,6×
Minne16 GB94 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-kk (oftast k=2k = 2) väljs, och utdatan är den viktade summan:

y=∑i∈top-kgi⋅Ei(x),g=softmax(Wrx)y = \sum_{i \in \text{top-}k} g_i \cdot E_i(x), \qquad g = \mathrm{softmax}(W_r x)

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:

Laux=α⋅N∑i=1Nfi⋅Pi\mathcal{L}_{\text{aux}} = \alpha \cdot N \sum_{i=1}^{N} f_i \cdot P_i

där fif_i är andelen tokens som routades till expert ii och PiP_i är den genomsnittliga routersannolikheten för den. Termen minimeras när båda är jämnt fördelade. Typiskt α≈0,01\alpha \approx 0{,}01.

Expertkapacitet. Varje expert har en gräns för hur många tokens den tar emot per batch:

C=tokens per batchantal experter×kapacitetsfaktorC = \frac{\text{tokens per batch}}{\text{antal experter}} \times \text{kapacitetsfaktor}

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:

ProblemVarför
Kommunikationexperterna ligger på olika GPU:er → all-to-all-utbyte per lager
Instabil träningroutern kan oscillera; z-loss och lägre lärhastighet på routern hjälper
FinjusteringMoE-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

Alla källor och licenser