Layer Normalization: как стабилизировать обучение LLM
opensourceaillmnormalizationdeep-learningit
Введение: проблема нестабильных активаций
При обучении глубоких сетей значения активаций могут "уплывать" — становиться слишком большими или слишком маленькими. Это ломает градиенты и обучение.
Слой 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