LayerNorm, RMSNorm och pre-/post-norm
Kunna jämföra normaliseringsvarianter och placeringar i moderna transformers.
Förkunskaper
Intuition
LayerNorm normaliserar över features inom varje exempel: dra bort medelvärdet, dela med standardavvikelsen, skala och flytta med inlärda γ och β.
RMSNorm hoppar över medelvärdescentreringen och normaliserar bara med kvadratiska medelvärdet:
En statistik i stället för två, ingen β. Cirka 10–15 % snabbare, och empiriskt lika bra — därför använder Llama, Mistral och Qwen RMSNorm.
Formellt
Placeringen spelar större roll än varianten.
Post-norm (originaltransformern):
Pre-norm (alla moderna):
Skillnaden verkar kosmetisk men är avgörande. I pre-norm finns en ren residualväg från indata till utdata som inte passerar genom någon normalisering. Gradienten kan flöda tillbaka längs den utan att skalas om lager för lager.
Konsekvensen (Xiong m.fl. 2020): post-norm-transformers kräver warmup för att inte divergera, och blir instabila vid stort djup. Pre-norm tränar stabilt utan warmup och skalar till hundratals lager.
Priset är en liten kvalitetsförlust vid samma djup — post-norm ger något bättre resultat när den väl konvergerar. Lösningen i praktiken är pre-norm plus en slutlig normalisering före utmatningen.
Varianter värda att känna till: DeepNorm (skalad residual som möjliggör mycket djupa post-norm-nät) och sandwich-norm (normalisering både före och efter delblocket).
Kod
import torch, torch.nn as nn
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__()
self.gamma = nn.Parameter(torch.ones(d)); self.eps = eps
def forward(self, x):
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return x * rms * self.gamma
class PreNormBlock(nn.Module):
"""Modern ordning: normalisera INNAN delblocket, addera till residualen."""
def __init__(self, d, n_heads):
super().__init__()
self.n1, self.n2 = RMSNorm(d), RMSNorm(d)
self.attn = nn.MultiheadAttention(d, n_heads, batch_first=True)
self.mlp = nn.Sequential(nn.Linear(d, 4*d), nn.GELU(), nn.Linear(4*d, d))
def forward(self, x, mask=None):
h = self.n1(x)
x = x + self.attn(h, h, h, attn_mask=mask, need_weights=False)[0]
return x + self.mlp(self.n2(x))
x = torch.randn(2, 8, 64)
print(PreNormBlock(64, 8)(x).shape) # torch.Size([2, 8, 64])
print(RMSNorm(64)(x).pow(2).mean(-1).mean().item()) # ≈ 1.0 (RMS normaliserad)
Lägg märke till att residualadditionen sker på x, inte på det normaliserade h — det är just det som bevarar den rena residualvägen.
Behärskning innebär
- Jämför LayerNorm och RMSNorm
- Förklarar skillnaden mellan pre-norm och post-norm
- Vet varför pre-norm blev standard
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — On Layer Normalization in the Transformer Architecture — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Root Mean Square Layer Normalization — arXiv (öppen åtkomst; licens per artikel)