Hoppa till innehållet
AI-grafen
E· Universitettransformerarkitektur· ca 60 min· utvecklande· verifierad 2026-09-20

MLP-blocket: GELU, SwiGLU

Kunna förklara feed-forward-blocket och gated-varianter i Llama-stil.

Förkunskaper

Intuition

Varje transformerblock har två delar: attention (flyttar information mellan positioner) och MLP (bearbetar varje position för sig).

Klassiskt MLP-block: MLP(x)=W2 GELU(W1x)\text{MLP}(x) = W_2\,\text{GELU}(W_1 x) med en expansion på 4× i mitten: 4096 → 16384 → 4096.

MLP:n står för ungefär två tredjedelar av parametrarna i en transformer — den är där det mesta av «kunskapen» sitter, medan attention sköter dirigeringen.

Formellt

Gated Linear Units (Shazeer 2020) lägger till en grind: SwiGLU(x)=(Swish(Wgatex)⊙Wupx)Wdown\text{SwiGLU}(x) = \big(\text{Swish}(W_{\text{gate}}x)\odot W_{\text{up}}x\big)W_{\text{down}}

Där Swish(z)=z⋅σ(z)\text{Swish}(z) = z\cdot\sigma(z) (även kallad SiLU). Grinden låter nätet multiplikativt dämpa eller släppa fram varje kanal beroende på indata — en form av innehållsberoende filtrering som en vanlig MLP inte kan uttrycka lika lätt.

Tre matriser i stället för två betyder fler parametrar vid samma bredd. Därför justeras expansionsfaktorn från 4× till ungefär 83≈2,67×\tfrac{8}{3}\approx 2{,}67\times, så att parameterantalet blir jämförbart:

VariantMatriserDold dim (d = 4096)Parametrar per block
GELU-MLP216 3842⋅4096⋅16384≈1342\cdot 4096\cdot 16384 \approx 134 M
SwiGLU311 0083⋅4096⋅11008≈1353\cdot 4096\cdot 11008 \approx 135 M

Vid jämförbart parameterantal presterar SwiGLU konsekvent något bättre — därför används den i Llama, Mistral, Qwen och PaLM.

Llamas märkliga dolda dimension 11 008 är just 83⋅4096=10922\tfrac{8}{3}\cdot 4096 = 10922 avrundat uppåt till en multipel av 256 (för effektiv matrismultiplikation).

Kod

import torch, torch.nn as nn, torch.nn.functional as F

class SwiGLU(nn.Module):
    def __init__(self, d, dold=None, multipel=256):
        super().__init__()
        dold = dold or int(((8 / 3) * d + multipel - 1) // multipel * multipel)
        self.gate = nn.Linear(d, dold, bias=False)
        self.up   = nn.Linear(d, dold, bias=False)
        self.down = nn.Linear(dold, d, bias=False)
        self.dold = dold
    def forward(self, x):
        return self.down(F.silu(self.gate(x)) * self.up(x))

class GeluMLP(nn.Module):
    def __init__(self, d, dold=None):
        super().__init__()
        dold = dold or 4 * d
        self.upp = nn.Linear(d, dold, bias=False)
        self.ner = nn.Linear(dold, d, bias=False)
    def forward(self, x):
        return self.ner(F.gelu(self.upp(x)))

for namn, m in (("GELU-MLP", GeluMLP(4096)), ("SwiGLU", SwiGLU(4096))):
    print(f"{namn:9s} parametrar {sum(p.numel() for p in m.parameters()):,}")
# GELU-MLP  parametrar 134,217,728
# SwiGLU    parametrar 135,266,304   ← jämförbart tack vare 8/3-justeringen

Behärskning innebär

  • Beskriver feed-forward-blocket och dess dimensioner
  • Förklarar SwiGLU och varför expansionsfaktorn justeras
  • Räknar parametrar i blocket

Logga in för att göra övningarna och bygga upp din behärskning.

Källor

Alla källor och licenser