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

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:

RMSNorm(x)=x1d∑ixi2+ϵ⊙γ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_i x_i^2 + \epsilon}}\odot\gamma

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): x←LN(x+Sublayer(x))x \leftarrow \text{LN}(x + \text{Sublayer}(x))

Pre-norm (alla moderna): x←x+Sublayer(LN(x))x \leftarrow x + \text{Sublayer}(\text{LN}(x))

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

Alla källor och licenser