MLP-blocket: GELU, SwiGLU
Kunna förklara feed-forward-blocket och gated-varianter i Llama-stil.
Förkunskaper
- DAktiveringsfunktionerkrävs
- DTransformers — arkitekturenkrävs
Intuition
Varje transformerblock har två delar: attention (flyttar information mellan positioner) och MLP (bearbetar varje position för sig).
Klassiskt MLP-block: 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:
Där (ä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 , så att parameterantalet blir jämförbart:
| Variant | Matriser | Dold dim (d = 4096) | Parametrar per block |
|---|---|---|---|
| GELU-MLP | 2 | 16 384 | M |
| SwiGLU | 3 | 11 008 | 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 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
- arXiv — GLU Variants Improve Transformer — arXiv (öppen åtkomst; licens per artikel)
- arXiv — LLaMA: Open and Efficient Foundation Language Models — arXiv (öppen åtkomst; licens per artikel)