Layer Normalization: как стабилизировать обучение LLM

opensourceaillmnormalizationdeep-learningit
← Back to Blog

Введение: проблема нестабильных активаций

При обучении глубоких сетей значения активаций могут "уплывать" — становиться слишком большими или слишком маленькими. Это ломает градиенты и обучение.

Слой 1:  активации = [0.5, -0.3, 1.2, 0.8, ...]  (mean≈0, std≈1)
Слой 10: активации = [500, -300, 1200, 800, ...]  (mean≈0, std≈1000) ← взрыв!
Слой 20: активации = [inf, -inf, nan, ...]       ← крах

Layer Normalization решает эту проблему, нормализуя активации по каждому элементу батча.


Уровень 1: No Normalization (без нормализации)

Что происходит без нормализации?

# Простой MLP без нормализации
class SimpleMLP(nn.Module):
    def __init__(self, hidden_size):
        super().__init__()
        self.fc1 = nn.Linear(hidden_size, hidden_size * 4)
        self.fc2 = nn.Linear(hidden_size * 4, hidden_size)
        self.act = nn.GELU()
    
    def forward(self, x):
        x = self.fc1(x)   # активации могут "улететь"
        x = self.act(x)   # GELU насыщается при больших значениях
        x = self.fc2(x)   # градиенты могут взорваться или затухнуть
        return x

# Проблема:
# Вход: [0.1, -0.2, 0.5, ...]
# После fc1: [15.3, -8.7, 22.1, ...]  ← умножение на веса
# После act: [1.0, 0.0, 1.0, ...]     ← GELU насытился!
# Градиенты = 0 ← обучение остановлено

Exploding/Vanishing Gradients

Exploding gradients:
  вес = 0.1
  градиент = 100
  новый_вес = вес - 0.001 * 100 = вес - 0.1  ← большой скачок

Vanishing gradients:
  вес = 0.1
  градиент = 0.00001
  новый_вес = вес - 0.001 * 0.00001 = вес - 0.00000000001  ← нет изменений

Уровень 2: Batch Normalization

Идея: нормализовать по батчу

class BatchNorm1d(nn.Module):
    def __init__(self, num_features):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(num_features))   # масштаб
        self.beta = nn.Parameter(torch.zeros(num_features))   # сдвиг
        self.running_mean = torch.zeros(num_features)
        self.running_var = torch.ones(num_features)
    
    def forward(self, x):  # x: [batch, features]
        # Нормализация по батчу
        mean = x.mean(dim=0)    # [features]
        var = x.var(dim=0)      # [features]
        x_norm = (x - mean) / torch.sqrt(var + eps)
        
        # Масштабирование
        return self.gamma * x_norm + self.beta

Проблема Batch Norm для NLP

Batch Normalization:
  + Работает отлично для CNN (нормализуем по пространству)
  - Плохо работает для RNN/Transformer
  
Почему для NLP плохо?
1. Маленький батч: нормализация по 8-32 примерам = шум
2. Разная длина последовательностей: нельзя нормализовать по time dim
3. Sequence-dependent: статистика зависит от входных данных

Уровень 3: Layer Normalization

Идея: нормализовать по всем признакам элемента

class LayerNorm(nn.Module):
    def __init__(self, hidden_size, eps=1e-5):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(hidden_size))
        self.beta = nn.Parameter(torch.zeros(hidden_size))
        self.eps = eps
    
    def forward(self, x):  # x: [batch_size, seq_len, hidden_size]
        # Нормализуем по последнему измерению (hidden_size)
        mean = x.mean(dim=-1, keepdim=True)    # [batch, seq, 1]
        var = x.var(dim=-1, keepdim=True, unbiased=False)  # [batch, seq, 1]
        
        x_norm = (x - mean) / torch.sqrt(var + self.eps)
        
        return self.gamma * x_norm + self.beta

Визуализация Layer Norm

Вход [1, 4, 768]:
  token_0: [0.5, -0.3, 1.2, 0.8, -0.1, ..., 0.4]  (768 чисел)
  
  mean(token_0) = 0.02
  std(token_0) = 0.87
  
  normalized(token_0) = [0.54, -0.37, 1.21, 0.77, -0.12, ...]
  
  output = gamma * normalized + beta

Layer Norm vs Batch Norm

                | Layer Norm              | Batch Norm
----------------|-------------------------|-------------------------
Нормализация по | hidden dim (последний)  | batch dim (первый)
Размер батча    | не важен                | должен быть большим
RNN/Transformer | отлично                 | плохо
CNN             | работает                | отлично
Память          | меньше                  | нужно хранить running stats

Уровень 4: RMS Normalization

Проблема: вычисление mean

Layer Norm требует:
  1. Вычислить mean    ← одно reduction
  2. Вычислить var     ← второе reduction
  3. Нормализовать     ← деление

RMS Norm (Zhang & Sennrich, 2019):
  1. Вычислить RMS     ← одно reduction
  2. Нормализовать     ← деление

RMS = sqrt(mean(x^2))

RMS Normalization

class RMSNorm(nn.Module):
    def __init__(self, hidden_size, eps=1e-5):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(hidden_size))
        self.eps = eps
    
    def forward(self, x):  # x: [batch, seq, hidden]
        # Вычисляем RMS по последнему измерению
        # RMS = sqrt(mean(x^2))
        variance = x.pow(2).mean(dim=-1, keepdim=True)
        
        x_norm = x * torch.rsqrt(variance + self.eps)
        
        return self.weight * x_norm

Почему RMS Norm работает?

Layer Norm:
  x_norm = (x - mean) / sqrt(var + eps)
  
RMS Norm:
  x_norm = x / sqrt(mean(x^2) + eps)

Разница:
  Layer Norm центрирует (вычитает mean)
  RMS Norm не центрирует
  
На практике: разница минимальна, RMS Norm быстрее!

Использование в LLM

Модель           | Normalization
-----------------|---------------------------
GPT-2            | Layer Norm
GPT-3            | Layer Norm
Llama 1          | RMS Norm
Llama 2          | RMS Norm
Llama 3          | RMS Norm
Mistral          | RMS Norm
Mixtral          | RMS Norm

Все современные LLM используют RMS Norm!


Где применяется Normalization в Transformer

Full Transformer с Layer Norm

Encoder:
  Input → Embedding → LayerNorm → Attention → LayerNorm → MLP → LayerNorm → Output

Decoder:
  Input → Embedding → LayerNorm → Self-Attention → LayerNorm → Cross-Attention → LayerNorm → MLP → LayerNorm → Output

Pre-Norm vs Post-Norm

Post-Norm (BERT):
  output = x + LayerNorm(Attention(x))
  output = output + LayerNorm(MLP(output))

Pre-Norm (GPT-2, Llama):
  output = x + Attention(LayerNorm(x))
  output = output + MLP(LayerNorm(output))

Pre-Norm лучше для глубоких сетей:
  + Градиенты лучше текут через skip connections
  + Стабильнее обучение
  + Используется во всех современных LLM

Pre-Norm в коде

class PreNormTransformerBlock(nn.Module):
    def __init__(self, hidden_size, num_heads, ff_size):
        super().__init__()
        self.attn_norm = RMSNorm(hidden_size)
        self.self_attn = MultiHeadAttention(hidden_size, num_heads)
        self.ff_norm = RMSNorm(hidden_size)
        self.ff = FeedForward(hidden_size, ff_size)
    
    def forward(self, x):
        # Pre-Norm
        x = x + self.self_attn(self.attn_norm(x))
        x = x + self.ff(self.ff_norm(x))
        return x

Влияние на градиенты

Градиенты с Layer Norm

Прямой проход:
  y = γ * (x - μ) / √(σ² + ε) + β
  
Обратный проход:
  ∂L/∂x = ∂L/∂y * γ / √(σ² + ε) * (1 - 1/n * Σ(1 - (x-μ)²/(σ²+ε)))
  
  → Градиенты масштабируются, но не взрываются

Стабильность обучения

Без нормализации:
  loss = [10.2, 8.5, 12.1, inf, nan, ...]  ← нестабильно

С Layer Norm:
  loss = [10.2, 7.8, 6.5, 5.9, 5.2, ...]   ← стабильное снижение

С RMS Norm:
  loss = [10.2, 7.6, 6.3, 5.7, 5.0, ...]   ← ещё стабильнее (меньше операций)

Практические аспекты

Выбор eps

eps = 1e-5 (стандарт)
eps = 1e-6 (Llama)
eps = 1e-12 (максимальная точность)

Больше eps → больше стабильность, меньше точность
Меньше eps → больше точность, риск деления на ноль

Выбор между Layer Norm и RMS Norm

Layer Norm:
  + Лучше для маленьких моделей и RNN
  + Стандарт для BERT-подобных архитектур
  
RMS Norm:
  + Быстрее (нет вычисления mean)
  + Меньше операций
  + Стандарт для современных LLM
  + Лучше масштабирование

Квантование и нормализация

При квантовании:
  Layer Norm нужно выполнять в float32
  Квантовать можно только после нормализации
  
  x (float32) → LayerNorm (float32) → Quantize → int8
  
  Иначе нормализация сломается при потере точности

Итоги

  • Layer Normalization стабилизирует обучение, нормализуя активации
  • Batch Norm плохо работает для NLP (маленькие батчи, переменная длина)
  • Layer Norm нормализует по hidden dim каждого элемента
  • RMS Norm — упрощённая версия, используется во всех современных LLM
  • Pre-Norm лучше Post-Norm для глубоких сетей
  • Normalization критически важна для стабильного обучения LLM