Математика глубокого обучения

Нормализация в сети

LayerNorm против RMSNorm, по какой оси и почему pre-norm вытеснил post-norm

Шаг 82 из 117 · ~26 мин

Повтор: одна арифметика, разные оси

В блоке 5 (урок 100) нормализация разбиралась как предобусловливание — средство против плохой обусловленности. Здесь тот же объект нужен с другой стороны: как элемент архитектуры, у которого есть выбор оси и выбор места.

x^=xμAσA2+ε,y=γx^+β\hat{x} = \frac{x - \mu_{\htmlData{k=axis}{A}}}{\sqrt{\sigma^2_{\htmlData{k=axis}{A}} + \varepsilon}}, \qquad y = \gamma\hat{x} + \beta

AA — это всё, что отличает BatchNorm от LayerNorm от RMSNorm. Проверьте по цифрам: в каждом режиме выделенная группа имеет среднее нуль и дисперсию единицу, а соседняя ячейка из другой группы — нет.

нормированные значения

batch ↓
признаки →
1.002.003.0010.00
2.004.001.008.00
3.001.005.0012.00
0.003.002.009.00
усреднение по
групп
16
Исходные числа. Последний столбец на порядок больше остальных — так выглядит вход, у которого признаки разного масштаба.

Почему трансформеры выбрали LayerNorm

Решающий признак — зависимость от батча, а не качество нормировки.

BatchNormLayerNorm
статистика побатчупризнакам примера
примеры связаныданет
нужны running statisticsданет
train ≠ evalданет
переменная длинапроблемане проблема
батч из одного примераломаетсяработает

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

Плюс model.eval(): у BatchNorm обучение и вывод — разные функции, и забытый вызов даёт тихую деградацию. У LayerNorm такого различия нет вовсе, потому что ей нечего запоминать.

RMSNorm: что даёт отказ от центрирования

RMSNorm(x)=x1nixi2+εγ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac1n\sum_i x_i^2 + \varepsilon}} \cdot \gamma

Отличие от LayerNorm в одном вычитании. Экономия — один проход по тензору вперёд и один назад, плюс отсутствие β\beta.

Из блока 5 известен точный ответ на вопрос, когда это безопасно: на центрированных данных два слоя дают побитово одинаковый результат, потому что при μ=0\mu = 0 дисперсия равна среднему квадрату. Активации в обученной сети обычно близки к центрированным, поэтому замена проходит почти бесплатно — и современные модели (LLaMA, Mistral, Gemma) используют RMSNorm.

На нецентрированных данных разница реальна: для строки [1,2,3,10][1,2,3,10] LayerNorm даёт среднее нуль, RMSNorm оставляет 0.7490.749. Но последующий γ\gamma подстраивается под то представление, которое ему дали, и качество от этого не страдает.

Pre-norm против post-norm

Второй архитектурный выбор, и он важнее первого.

x+Sublayer(LN(x))pre-normпротивLN(x+Sublayer(x))post-norm\underbrace{x + \text{Sublayer}(\text{LN}(x))}_{\text{pre-norm}} \qquad\text{против}\qquad \underbrace{\text{LN}\big(x + \text{Sublayer}(x)\big)}_{\text{post-norm}}

Разница в том, проходит ли residual-путь через нормировку.

pre-normpost-norm
чистый путь от входа к выходуестьнет
нужен warmupобычно нетпочти всегда
устойчивость на большой глубиневысокаяпадает
качество при удачной настройкечуть нижечуть выше

У pre-norm градиент может дойти от выхода до входа, ни разу не пройдя через нормировку, — поэтому глубокие стеки обучаются без тщательного подбора расписания. У post-norm каждый residual-путь проходит через LN, и на большой глубине это накапливается.

Практический итог: почти все современные модели — pre-norm, и цена этого выбора известна (немного худшее качество при идеальной настройке), но настройка при глубине сто слоёв идеальной не бывает.

Что нормализация делает с масштабом весов

Свойство, из которого растёт неожиданное следствие. Если слой предшествует нормировке, то умножение его весов на cc не меняет выход — нормировка поделит на то же cc.

Значит:

  • выразительность слоя не зависит от нормы его весов, только от направления;
  • weight decay в такой сети не уменьшает «сложность» — он управляет эффективным learning rate, потому что меньшая норма означает больший относительный шаг;
  • learning rate и weight decay перестают быть независимыми гиперпараметрами.

Это тот же вывод, что в блоке 5, урок 100, и он остаётся одним из самых недооценённых следствий нормализации.

Источники

Проверки

0 из 2
  1. Оси, места и следствия

    Отметьте все верные утверждения о нормализации в сети.

  2. Что остаётся после нормировки

    Реализуйте norm_stats(row, c). Сначала умножьте все элементы row на c, затем примените к получившейся строке два слоя (возьмите ε=0\varepsilon = 0, γ=1\gamma = 1, β=0\beta = 0):

    LN(x)i=xiμσ,RMS(x)i=xix2\text{LN}(x)_i = \frac{x_i - \mu}{\sigma}, \qquad \text{RMS}(x)_i = \frac{x_i}{\sqrt{\overline{x^2}}}

    где μ\mu, σ2\sigma^2 и x2\overline{x^2} считаются по этой же строке (дисперсия — с делителем nn, не n1n-1). Верните [ln_mean, ln_var, rms_mean, rms_second_moment, rms_var] — среднее и дисперсию выхода LayerNorm, затем среднее, средний квадрат и дисперсию выхода RMSNorm.

    Два первых числа известны заранее, они и есть определение LayerNorm. Интересны остальные три, и проверить их можно тождеством: rms_var обязана равняться 1rms_mean21 - \texttt{rms\_mean}^2, потому что средний квадрат выхода RMSNorm равен единице по построению.

    Множитель c в задаче для того, чтобы вы убедились: он не влияет ни на одно из пяти чисел.

    функция norm_stats

    Загрузка редактора…

    Ctrl/⌘ + Enter