Нейронные сети Обучение сетей: оптимизаторы, инициализация, нормализация, регуляризация
0%

Обучение сетей: оптимизаторы, инициализация, нормализация, регуляризация

Обучение сетей: оптимизаторы, инициализация, нормализация, регуляризация

В прошлой статье мы разобрались, как сеть считает предсказание и как обратное распространение выдаёт градиент функции потерь по каждому весу. Это ответ на вопрос «куда шагать».

Эта статья — про всё остальное, и «остальное» здесь больше, чем кажется:

  • как шагать (оптимизатор: SGD, momentum, Adam, AdamW);
  • насколько крупно шагать и как менять размер шага во времени (learning rate schedule, warmup);
  • откуда стартовать (инициализация весов);
  • как сделать сам ландшафт удобным для спуска (нормализация);
  • как не дать модели выучить шум (регуляризация);
  • как не получить NaN на шаге 3000 (клиппинг, mixed precision, стабильность).

Формально всё это — «гиперпараметры». На практике именно они отделяют модель, которая сходится к 92% accuracy за 40 минут, от той же самой архитектуры, которая три дня болтается на уровне случайного угадывания. Backprop даёт градиент; обучение — это инженерная дисциплина о том, что с этим градиентом делать.


1. Что мы вообще оптимизируем

Задача формулируется как минимизация эмпирического риска по параметрам $\theta$:

$$ L(\theta) ;=; \frac{1}{N}\sum_{i=1}^{N} \ell\big(f_\theta(x_i),, y_i\big) ;+; \Omega(\theta) $$

где $\ell$ — потеря на одном примере, $\Omega$ — регуляризатор. Выглядит как обычная задача оптимизации, но у неё три свойства, которые ломают классическую интуицию.

1.1. Функция невыпуклая — и это почти не мешает

$L(\theta)$ для сети с нелинейностями невыпукла: локальных минимумов экспоненциально много. Классический курс оптимизации сказал бы, что задача безнадёжна. Практика говорит иначе, и причина известна: в высокой размерности плохие локальные минимумы редки.

Интуиция: чтобы точка была локальным минимумом, гессиан должен быть положительно определён — все $d$ собственных значений положительны. Если знаки собственных значений «примерно случайны», вероятность такого события падает экспоненциально с $d$, а $d$ у нас миллионы. Гораздо вероятнее, что часть значений положительна, часть отрицательна — это седловая точка. Именно седловые точки и окружающие их плато, а не локальные минимумы, — основная проблема (Dauphin et al., 2014).

Практический вывод: не бойтесь застрять в «плохой яме». Бойтесь плато, где градиент близок к нулю по многим координатам и обучение стоит.

1.2. Ландшафт плохо обусловлен

Даже локально, вблизи минимума, $L$ ведёт себя как квадратичная форма с гессианом $H$. Скорость сходимости градиентного спуска определяется числом обусловленности $\kappa = \lambda_{\max}/\lambda_{\min}$. Для градиентного спуска с оптимальным шагом число итераций до точности $\varepsilon$ растёт как $O(\kappa \log \frac{1}{\varepsilon})$.

У нейросетей $\kappa$ легко достигает $10^4$–$10^6$: спектр гессиана содержит несколько огромных собственных значений и длинный хвост около нуля. Геометрически это овраг — вытянутый каньон, где по одной оси функция крутая, по другой почти плоская.

Траектории оптимизаторов в овраге функции потерь

Размер шага приходится выбирать по самому крутому направлению (иначе разлетимся), а прогресс идёт по самому пологому — отсюда зигзаг. Почти вся история оптимизаторов глубокого обучения — это борьба именно с этим эффектом.

1.3. Градиент, который мы видим, — шумный

Считать градиент по всем $N$ примерам (full-batch) слишком дорого. Мы берём мини-батч размера $B$ и получаем несмещённую, но шумную оценку:

$$ g_B = \frac{1}{B}\sum_{i \in \mathcal{B}} \nabla_\theta \ell_i, \qquad \mathbb{E}[g_B] = \nabla L, \qquad \operatorname{Var}[g_B] \propto \frac{1}{B} $$

Шум — не только цена за скорость, он ещё и полезен: помогает выбираться с седловых точек и, по распространённой гипотезе, смещает решение в сторону широких минимумов, которые лучше обобщаются (Keskar et al., 2016). Именно поэтому «увеличить батч до размера датасета» — плохая идея, даже если памяти хватает.


2. Оптимизаторы: от SGD до AdamW

2.1. Эволюция

2.2. SGD: базовая линия, которую рано списывать

θ ← θ − η · g

Один гиперпараметр, ноль дополнительной памяти. Главный недостаток — тот самый зигзаг: шаг одинаков по всем координатам, а кривизна по ним разная.

2.3. Momentum: инерция вместо реакции на каждый градиент

Идея физическая: шарик, катящийся по оврагу, не разворачивается на каждой стенке — он накапливает скорость вдоль долины, а поперечные толчки взаимно гасятся.

v ← μ·v + g          # PyTorch-конвенция
θ ← θ − η·v

При $\mu = 0.9$ эффективный шаг вдоль устойчивого направления умножается примерно на $\frac{1}{1-\mu} = 10$, а осцилляции затухают. Momentum превращает сходимость $O(\kappa)$ в $O(\sqrt{\kappa})$ для квадратичных задач — это не косметика, а смена порядка.

Nesterov momentum дополнительно считает градиент в «предвосхищённой» точке $\theta - \eta\mu v$: если мы уже летим к стенке, поправка приходит на шаг раньше. На практике даёт небольшой, но стабильный выигрыш.

2.4. AdaGrad и RMSProp: свой шаг каждой координате

AdaGrad накапливает сумму квадратов градиентов и делит шаг на её корень:

G ← G + g²
θ ← θ − η · g / (√G + ε)

Редко обновляемые параметры (например, эмбеддинги редких слов) получают большой шаг, частые — маленький. Проблема: $G$ монотонно растёт, шаг неизбежно затухает до нуля. Для выпуклых задач это корректно, для глубоких сетей — смерть обучения на середине.

RMSProp заменяет сумму на экспоненциальное скользящее среднее — «забывающий» AdaGrad:

v ← β·v + (1−β)·g²
θ ← θ − η · g / (√v + ε)

2.5. Adam: momentum по первому и второму моменту

Adam (Kingma & Ba, 2014) = momentum + RMSProp

  • поправка на смещение старта из нуля:

$$ m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t, \qquad v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2 $$ $$ \hat m_t = \frac{m_t}{1-\beta_1^t}, \qquad \hat v_t = \frac{v_t}{1-\beta_2^t}, \qquad \theta_t = \theta_{t-1} - \eta \frac{\hat m_t}{\sqrt{\hat v_t} + \varepsilon} $$

Bias correction важнее, чем кажется: на первых шагах $m_1 = (1-\beta_1)g_1 = 0.1 g_1$, и без деления на $1-\beta_1^t$ первые сотни шагов оптимизатор двигался бы неправдоподобно медленно — а потом резко разгонялся, что и порождает нестабильность в начале обучения.

Реализация «с нуля», чтобы формулы перестали быть абстракцией:

import torch

class Adam:
    """Учебная реализация Adam. Совпадает с torch.optim.Adam до float-погрешности."""

    def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8):
        self.params = list(params)
        self.lr, (self.b1, self.b2), self.eps = lr, betas, eps
        self.t = 0
        # два буфера состояния на каждый параметр -> память x2 от размера модели
        self.m = [torch.zeros_like(p) for p in self.params]
        self.v = [torch.zeros_like(p) for p in self.params]

    @torch.no_grad()
    def step(self):
        self.t += 1
        # поправки на смещение зависят только от шага, считаем один раз
        bc1 = 1 - self.b1 ** self.t
        bc2 = 1 - self.b2 ** self.t
        for p, m, v in zip(self.params, self.m, self.v):
            if p.grad is None:
                continue
            g = p.grad
            m.mul_(self.b1).add_(g, alpha=1 - self.b1)          # первый момент
            v.mul_(self.b2).addcmul_(g, g, value=1 - self.b2)   # второй момент
            step_size = self.lr / bc1
            denom = (v / bc2).sqrt_().add_(self.eps)
            p.addcdiv_(m, denom, value=-step_size)              # θ -= lr * m̂ / (√v̂ + ε)

    def zero_grad(self):
        for p in self.params:
            p.grad = None   # None дешевле, чем зануление тензора

Сложность. Все перечисленные оптимизаторы — $O(d)$ по времени на шаг, где $d$ — число параметров: только поэлементные операции. Разница в памяти:

Оптимизатор Буферы состояния Доп. память (fp32)
SGD нет 0
SGD + momentum v 4 байта/параметр
Adam / AdamW m, v 8 байт/параметр
Adafactor факторизованный v ~ $O(\sqrt{d})$
8-bit Adam квантованные m, v 2 байта/параметр

Для модели на 7 млрд параметров в mixed precision бюджет выглядит так: 2 байта (веса bf16) + 4 (fp32 master-копия) + 4 (m) + 4 (v) ≈ 14 байт на параметр, то есть ~98 ГБ только под состояние — до всяких активаций. Отсюда и популярность Adafactor и 8-bit оптимизаторов. Подробнее про память на инференсе — в статье про деплой.

2.6. AdamW: почему weight_decay в Adam был сломан

L2-регуляризация обычно реализуется как добавка к градиенту: g ← g + λθ. В SGD это в точности эквивалентно затуханию весов $\theta \leftarrow (1-\eta\lambda)\theta$. В Adam — нет: добавка проходит через нормировку на $\sqrt{\hat v}$, и веса с большими градиентами регуляризуются слабее, чем веса с маленькими. То есть сила регуляризации непредсказуемо зависит от истории градиентов.

Loshchilov & Hutter (2017) предложили развязать их: применять затухание напрямую к весам, минуя адаптивный знаменатель.

# AdamW: ключевое отличие — эта строка вне адаптивного апдейта
p.mul_(1 - lr * weight_decay)                  # развязанный weight decay
p.addcdiv_(m_hat, v_hat.sqrt() + eps, value=-lr)

Практическое следствие, которое стоит запомнить намертво: для Adam всегда берите torch.optim.AdamW, а не Adam(weight_decay=...). Это самый частый источник «регуляризация вроде включена, а модель всё равно переобучается».

И второе правило: не применяйте weight decay к bias и к параметрам нормализации. Затухание масштаба $\gamma$ в LayerNorm к нулю — это просто порча слоя.

def param_groups(model, weight_decay=0.1):
    """Эвристика из nanoGPT: decay только для многомерных тензоров —
    матриц весов и ядер свёрток; bias, γ и β остаются без затухания."""
    decay, no_decay = [], []
    for name, p in model.named_parameters():
        if p.requires_grad:
            (no_decay if p.ndim < 2 else decay).append(p)
    return [{"params": decay, "weight_decay": weight_decay},
            {"params": no_decay, "weight_decay": 0.0}]

2.7. Что выбрать

  • Трансформеры, LLM, всё с эмбеддингами — AdamW. Практически безальтернативно: разреженные градиенты эмбеддингов требуют покоординатной адаптации.
  • Свёрточные сети на изображениях — SGD + Nesterov momentum 0.9 всё ещё часто даёт лучший финальный accuracy, чем Adam, хотя сходится медленнее. Все канонические цифры ResNet на ImageNet получены именно так.
  • Табличные данные, маленькие сети — разницы почти нет, берите AdamW и не думайте.

Никогда не подбирайте оптимизатор раньше, чем настроите learning rate: разница между оптимизаторами почти всегда меньше, чем разница между хорошим и плохим LR.


3. Learning rate — гиперпараметр номер один

Если у вас есть время настроить ровно один гиперпараметр, настраивайте LR. Порядки:

LR Что происходит
слишком большой loss растёт, уходит в NaN, веса взрываются
чуть больше нужного loss падает и застревает на плато, шумно колеблется
оптимальный быстрый спад, гладкая кривая
слишком маленький сходится, но в 10 раз медленнее и часто в худший минимум

3.1. LR range test

Дешёвый способ найти диапазон: за один короткий проход экспоненциально увеличивать LR и смотреть, где loss начинает расти (Smith, 2015).

def lr_range_test(model, loader, loss_fn, lr_min=1e-7, lr_max=1.0, num_steps=100):
    """Прогон с экспоненциально растущим LR. Рабочий LR берём примерно на порядок
    меньше того, при котором loss достиг минимума перед взрывом."""
    opt = torch.optim.AdamW(model.parameters(), lr=lr_min)
    sched = torch.optim.lr_scheduler.ExponentialLR(opt, (lr_max / lr_min) ** (1 / num_steps))
    history = []
    for _, (x, y) in zip(range(num_steps), loader):
        opt.zero_grad(set_to_none=True)
        loss = loss_fn(model(x), y)
        loss.backward()
        opt.step(); sched.step()
        history.append((sched.get_last_lr()[0], loss.item()))
        if loss.item() > 4 * history[0][1]:   # loss взорвался — дальше смысла нет
            break
    return history

3.2. Warmup: зачем начинать медленно

На первых шагах статистики второго момента $v$ ещё не собраны, а веса случайны — оценка $\hat v$ ненадёжна, и Adam может сделать огромный неоправданный шаг, из которого сеть уже не восстановится. Отсюда linear warmup: первые 1–5% шагов LR растёт линейно от нуля до целевого. Для трансформеров warmup не опция, а необходимость — без него post-LN модели просто расходятся (Liu et al., 2019).

3.3. Расписание

Индустриальный стандарт сегодня — warmup + cosine decay до примерно 10% от пикового LR:

import math
from torch.optim.lr_scheduler import LambdaLR

def warmup_cosine(optimizer, warmup_steps, total_steps, min_ratio=0.1):
    def lr_lambda(step):
        if step < warmup_steps:
            return step / max(1, warmup_steps)          # линейный разогрев
        # косинусное затухание от 1.0 до min_ratio
        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
        progress = min(1.0, progress)
        cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
        return min_ratio + (1 - min_ratio) * cosine
    return LambdaLR(optimizer, lr_lambda)

Важная деталь: cosine рассчитан на заранее известное число шагов. Если вы остановите обучение на половине расписания, вы получите модель заметно хуже, чем при cosine, рассчитанном на эту половину. Для запусков с неизвестной длительностью удобнее WSD (warmup–stable–decay): долгая полка на постоянном LR и короткий спад в конце, что позволяет «отрезать» модель в любой момент.

3.4. Связь LR и размера батча

Увеличили батч вдвое — градиент стал вдвое менее шумным, можно увеличить шаг. Два эмпирических правила:

  • Linear scaling: $\eta \propto B$. Работает для SGD до батчей ~8k (Goyal et al., 2017, ImageNet за час).
  • Square-root scaling: $\eta \propto \sqrt{B}$. Обычно лучше для Adam, потому что адаптивный знаменатель уже частично поглощает изменение масштаба.

За пределами «критического размера батча» (McCandlish et al., 2018) дальнейшее увеличение $B$ перестаёт ускорять обучение по числу шагов — вы просто сжигаете вычисления.


4. Инициализация: старт определяет, сойдётесь ли вы вообще

4.1. Почему нельзя нулями

Если все веса слоя равны, все нейроны получают одинаковый вход, одинаковый выход и одинаковый градиент. Они останутся одинаковыми навсегда — слой шириной 512 эквивалентен слою шириной 1. Это называется проблемой симметрии, и её ломает только случайность.

Нулями инициализируют только смещения (bias) — там симметрии нет.

4.2. Масштаб важнее распределения

Рассмотрим линейный слой $y = Wx$ с $n_{in}$ входами, независимыми весами с нулевым средним и дисперсией $\sigma^2$. Тогда

$$ \operatorname{Var}(y_j) = n_{in} \cdot \sigma^2 \cdot \operatorname{Var}(x) $$

Если $n_{in}\sigma^2 > 1$, дисперсия активаций растёт от слоя к слою экспоненциально — взрыв. Если $< 1$ — экспоненциально убывает, и к десятому слою сигнал вырождается в ноль (затухание). Единственный устойчивый режим: $n_{in}\sigma^2 \approx 1$.

  • Xavier/Glorot (2010) — компромисс между прямым и обратным проходом: $\sigma^2 = \dfrac{2}{n_{in}+n_{out}}$. Выведено для симметричных насыщающих функций (tanh).
  • He/Kaiming (2015) — для ReLU, которая зануляет половину активаций и тем самым режет дисперсию вдвое; компенсируем множителем 2: $\sigma^2 = \dfrac{2}{n_{in}}$. Это дефолт для всего, что использует ReLU/GELU.
  • Ортогональная — сохраняет нормы точно, а не в среднем; полезна для глубоких RNN, где произведение матриц применяется многократно (см. рекуррентные сети).
import torch.nn as nn

def init_weights(module):
    """Инициализация в стиле современных трансформеров."""
    if isinstance(module, nn.Linear):
        # нормальное с std=0.02 — конвенция GPT; для CNN чаще kaiming_normal_
        nn.init.normal_(module.weight, mean=0.0, std=0.02)
        if module.bias is not None:
            nn.init.zeros_(module.bias)
    elif isinstance(module, nn.Conv2d):
        nn.init.kaiming_normal_(module.weight, mode="fan_out", nonlinearity="relu")
    elif isinstance(module, (nn.LayerNorm, nn.BatchNorm2d)):
        nn.init.ones_(module.weight)    # γ = 1
        nn.init.zeros_(module.bias)     # β = 0

model.apply(init_weights)

# Отдельный приём для residual-сетей: масштабируем выходную проекцию каждого блока
# на 1/sqrt(2 * n_layers), чтобы дисперсия не накапливалась вдоль остаточного потока.
for name, p in model.named_parameters():
    if name.endswith("out_proj.weight"):
        nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * n_layers))

Последний трюк (residual-aware scaling) взят из GPT-2 и заметно стабилизирует ранние шаги глубоких трансформеров. Ещё жёстче работают Fixup и T-Fixup, позволяющие обучать глубокие residual-сети вообще без нормализации.

4.3. Инициализация bias последнего слоя

Малозаметный, но очень полезный приём. При сильном дисбалансе классов инициализируйте bias выходного слоя логитом априорной вероятности: $b = \log\frac{p}{1-p}$. Тогда модель на первом же шаге предсказывает базовую частоту и не тратит сотни шагов на то, чтобы «догадаться», что положительный класс встречается в 1% случаев. Тот же приём — основа focal loss в детекции объектов.


5. Нормализация: перестроить ландшафт, а не бороться с ним

Пока мы улучшали спуск. Нормализация улучшает сам ландшафт — делает его ближе к сферическому, снижая число обусловленности.

Оси усреднения для BatchNorm, LayerNorm, GroupNorm и InstanceNorm

Общая формула у всех одна, отличаются только оси, по которым считаются $\mu$ и $\sigma$:

$$ \hat x = \frac{x - \mu}{\sqrt{\sigma^2 + \varepsilon}}, \qquad y = \gamma \hat x + \beta $$

Обучаемые $\gamma, \beta$ обязательны: без них слой лишается возможности выдавать ненулевое среднее и произвольный масштаб, что для многих задач вредно.

5.1. BatchNorm

Ioffe & Szegedy, 2015. Нормализует каждый канал по всему батчу. Работает отлично, но имеет неприятную особенность: вычисление для одного объекта зависит от остальных объектов батча. Отсюда практически все проблемы:

  • на инференсе батча нет — приходится хранить скользящие средние (running_mean, running_var) и переключать режим через model.eval(). Забытый model.eval() — классическая причина «в тестах метрика отличная, в проде мусор»;
  • при батче 1–4 (детекция, сегментация, 3D) статистики шумные, качество падает;
  • в распределённом обучении статистики считаются на каждом GPU отдельно — нужен SyncBatchNorm;
  • в RNN длины последовательностей разные, применимость сомнительная.

Исходное объяснение эффекта — «уменьшение internal covariate shift» — сегодня считается неверным: Santurkar et al. (2018) показали, что дело в сглаживании ландшафта потерь, а не в стабилизации распределений.

Кстати: если перед свёрткой/линейным слоем стоит нормализация, bias в этом слое не нужен — его роль полностью берёт на себя $\beta$. Отсюда bias=False во всех современных CNN и трансформерах.

5.2. LayerNorm и RMSNorm

LayerNorm (Ba et al., 2016) нормализует по признакам внутри одного объекта — от батча не зависит вовсе. Это стандарт для трансформеров (см. статью про внимание).

RMSNorm (Zhang & Sennrich, 2019) выбрасывает вычитание среднего:

$$ y = \gamma \cdot \frac{x}{\sqrt{\frac{1}{d}\sum_i x_i^2 + \varepsilon}} $$

Качество то же, операций меньше — поэтому LLaMA, Mistral, Gemma и почти все современные LLM используют именно RMSNorm.

5.3. Pre-LN против Post-LN

В Post-LN нормализация стоит на остаточном пути, и градиент, идущий назад через десятки слоёв, каждый раз перемасштабируется. Такие модели требуют аккуратного warmup и легко расходятся. В Pre-LN остаточный путь остаётся «чистой магистралью» от выхода ко входу — градиент течёт беспрепятственно, обучение устойчиво почти без warmup (Xiong et al., 2020). Ценой небольшого проигрыша в финальном качестве, который компенсируют финальным LayerNorm перед головой.

5.4. Что выбирать

Ситуация Норма
CNN, батч ≥ 32 BatchNorm
Детекция/сегментация, батч 2–8 GroupNorm (G=32)
Трансформеры, NLP, любые последовательности LayerNorm / RMSNorm
Style transfer, GAN-генераторы InstanceNorm / AdaIN
Нужен строго батч-независимый инференс что угодно, кроме BatchNorm

6. Регуляризация: борьба за обобщение

Обучение минимизирует ошибку на трейне. Нас интересует ошибка на новых данных. Всё, что сокращает этот разрыв ценой небольшого ухудшения на трейне, — регуляризация.

6.1. Weight decay

Штраф $\frac{\lambda}{2}|\theta|^2$ тянет веса к нулю, ограничивая сложность функции. В байесовской интерпретации это гауссов априор на веса. Типичные значения: 0.01–0.1 для AdamW в трансформерах, 1e-4–5e-4 для SGD в CNN. Обратите внимание на разницу в три порядка — она не случайна: в AdamW множитель развязан и не делится на $\sqrt{v}$, поэтому масштаб совершенно другой. Не переносите значения между оптимизаторами.

6.2. Dropout

Srivastava et al., 2014. На обучении случайно зануляем долю $p$ активаций и делим остальные на 1 $-p$ (inverted dropout, чтобы матожидание сохранялось и на инференсе ничего не менять). На инференсе выключен.

Две интерпретации: (1) неявный ансамбль из $2^n$ подсетей с общими весами; (2) запрет на коадаптацию — нейрон не может полагаться на конкретного соседа, приходится учить избыточные, самостоятельно полезные признаки.

Важно: dropout и BatchNorm плохо уживаются — dropout меняет дисперсию активаций между train и eval, статистики BN оказываются смещёнными (Li et al., 2018). В современных CNN dropout почти исчез, его роль заняли BatchNorm и аугментация. В трансформерах dropout 0.0–0.1 жив, но в больших LLM-предобучениях его часто ставят в 0: данных столько, что переобучение не является проблемой, а dropout только замедляет сходимость.

6.3. Label smoothing

Вместо жёсткой цели [0,0,1,0] используем [ε/K, ε/K, 1−ε+ε/K, ε/K] при $\varepsilon \approx 0.1$. Модель перестаёт стремиться к бесконечным логитам, калибровка улучшается, уверенность становится честнее. Стандарт в классификации изображений и в машинном переводе. Оговорка: label smoothing ухудшает качество эмбеддингов для задач поиска и дистилляции (Müller et al., 2019) — если модель потом будет учителем, не включайте.

6.4. EMA весов

Дешёвый и почти всегда полезный приём: параллельно с обучением ведите экспоненциальное скользящее среднее весов и на валидации используйте именно его.

class EMA:
    """Скользящее среднее весов. Часто даёт +0.3-1% метрики бесплатно,
    и всегда — заметно более гладкую кривую валидации."""

    def __init__(self, model, decay=0.999):
        self.decay = decay
        self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()}

    @torch.no_grad()
    def update(self, model):
        for k, v in model.state_dict().items():
            if v.dtype.is_floating_point:
                self.shadow[k].mul_(self.decay).add_(v, alpha=1 - self.decay)
            else:
                self.shadow[k].copy_(v)   # счётчики BN и целочисленные буферы копируем как есть

EMA — де-факто обязательная часть обучения диффузионных моделей (см. генеративные модели); без неё сэмплы заметно хуже.

6.5. Early stopping

Останавливаемся, когда валидационная метрика не улучшается $k$ эпох подряд, и возвращаем лучший чекпойнт, а не последний. Формально это тоже ограничение на норму весов: чем меньше шагов сделано от малой инициализации, тем меньше веса успели уйти от нуля.


7. Стабильность: клиппинг, precision и почему всё превращается в NaN

7.1. Gradient clipping

Иногда попадается «плохой» батч и норма градиента подскакивает на порядки. Один такой шаг способен разрушить обучение целиком. Лечится обрезанием глобальной нормы:

# ВАЖНО: после backward(), до optimizer.step(); при AMP — после scaler.unscale_()
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

Обрезаем именно глобальную норму, а не покомпонентно (clip_grad_value_): это сохраняет направление градиента и меняет только длину шага. max_norm=1.0 — практически универсальный дефолт для трансформеров. Полезно логировать total_norm: устойчивый рост означает надвигающуюся расходимость, резкие пики — проблемные данные.

7.2. Mixed precision

Считаем forward/backward в bf16 или fp16, а веса и апдейт оптимизатора храним в fp32. Выигрыш — 2× по памяти активаций и 2–5× по скорости на тензорных ядрах.

from torch.amp import autocast, GradScaler

scaler = GradScaler("cuda")   # нужен только для fp16; для bf16 можно обойтись без него

for x, y in loader:
    opt.zero_grad(set_to_none=True)
    with autocast("cuda", dtype=torch.bfloat16):
        loss = loss_fn(model(x), y)
    scaler.scale(loss).backward()      # масштабируем loss, чтобы малые градиенты не стали нулём
    scaler.unscale_(opt)               # возвращаем истинный масштаб перед клиппингом
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    scaler.step(opt)                   # шаг пропускается, если обнаружены inf/nan
    scaler.update()                    # адаптируем масштаб

Разница между форматами принципиальна. У fp16 всего 5 бит экспоненты — динамический диапазон узкий, малые градиенты схлопываются в ноль, отсюда и нужен loss scaling. У bf16 экспонента как у fp32 (8 бит) при меньшей мантиссе: диапазон тот же, точность хуже, но для градиентного спуска точность мантиссы почти не важна. Если железо поддерживает bf16 (Ampere и новее) — берите bf16 и забудьте про GradScaler.

7.3. Диагностика: жизненный цикл запуска

7.4. Типичные ошибки

  • Забытый optimizer.zero_grad(). PyTorch накапливает градиенты — вы учитесь на сумме всех прошлых батчей. NaN при этом не возникает, поэтому ошибку долго не замечают.
  • Забытый model.eval() на валидации. Dropout активен, BatchNorm обновляет статистики по валидационным данным: метрика занижена и нестабильна.
  • Забытый torch.no_grad() на инференсе. Строится граф вычислений, память течёт.
  • Shuffle выключен. Батчи, отсортированные по классу, — прямой путь к осцилляциям.
  • Статистики нормализации входа считаются по всему датасету, включая тест. Утечка данных; считать только по трейну.
  • Метрику накапливают тензорами, а не числами — OOM к концу эпохи. Всегда .item().
  • LR-шедулер вызывается по эпохам, а настроен на шаги (или наоборот). Проверьте, что число вызовов sched.step() совпадает с total_steps в расписании.
  • Sanity check пропущен. Перед долгим запуском всегда переобучайте модель на 2–3 батчах до loss ≈ 0. Не получается — баг в модели или в данных, и никакие гиперпараметры не помогут.

8. Полный цикл обучения

Собираем всё вместе. Ниже — цикл, который можно брать за основу продакшн-обучения: он включает accumulation, AMP, клиппинг, шедулер, EMA и чекпойнты.

import math
import torch
from torch.amp import autocast, GradScaler


def train(model, train_loader, val_loader, loss_fn, *,
          epochs=10, lr=3e-4, weight_decay=0.1, accum_steps=4,
          max_norm=1.0, warmup_ratio=0.03, device="cuda"):

    model.to(device)
    opt = torch.optim.AdamW(param_groups(model, weight_decay), lr=lr, betas=(0.9, 0.95))

    # шедулер считаем в ШАГАХ оптимизатора, а не в батчах
    steps_per_epoch = math.ceil(len(train_loader) / accum_steps)
    total_steps = steps_per_epoch * epochs
    sched = warmup_cosine(opt, int(total_steps * warmup_ratio), total_steps)

    scaler = GradScaler("cuda", enabled=(device == "cuda"))
    ema = EMA(model, decay=0.999)
    best_val, global_step = float("inf"), 0

    for epoch in range(epochs):
        model.train()
        opt.zero_grad(set_to_none=True)

        for i, (x, y) in enumerate(train_loader):
            x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)

            with autocast("cuda", dtype=torch.bfloat16, enabled=(device == "cuda")):
                loss = loss_fn(model(x), y) / accum_steps   # делим, чтобы масштаб не зависел от accum

            scaler.scale(loss).backward()

            if (i + 1) % accum_steps == 0:
                scaler.unscale_(opt)                        # ОБЯЗАТЕЛЬНО до клиппинга
                grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
                scaler.step(opt)
                scaler.update()
                opt.zero_grad(set_to_none=True)
                sched.step()
                ema.update(model)
                global_step += 1

                if global_step % 50 == 0:
                    print(f"step {global_step}  loss {loss.item() * accum_steps:.4f}  "
                          f"lr {sched.get_last_lr()[0]:.2e}  |g| {grad_norm:.2f}")

        val_loss = evaluate(model, val_loader, loss_fn, device)
        if val_loss < best_val:                             # сохраняем ЛУЧШИЙ, не последний
            best_val = val_loss
            torch.save({"model": model.state_dict(),
                        "ema": ema.shadow,
                        "opt": opt.state_dict(),            # нужен для корректного resume
                        "sched": sched.state_dict(),
                        "step": global_step}, "best.pt")
        print(f"epoch {epoch}: val_loss={val_loss:.4f} (best {best_val:.4f})")


@torch.no_grad()
def evaluate(model, loader, loss_fn, device):
    model.eval()                                # ключевая строка: dropout off, BN в режиме inference
    total, n = 0.0, 0
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        with autocast("cuda", dtype=torch.bfloat16, enabled=(device == "cuda")):
            total += loss_fn(model(x), y).item() * x.size(0)
        n += x.size(0)
    model.train()
    return total / n

Обратите внимание на две вещи, которые чаще всего делают неправильно:

  1. loss / accum_steps — без деления эффективный LR неявно умножается на accum_steps.
  2. В чекпойнт кладём состояние оптимизатора и шедулера. Без них resume после сбоя стартует с нулевыми моментами и сбитым расписанием — на кривой обучения это видно как отчётливая «ступенька».

9. Как это выглядит в проде

Память. Градиентное накопление даёт большой эффективный батч на маленькой карте (effective_batch = batch_size × accum_steps × num_gpus) — ценой времени, но не памяти. Gradient checkpointing (Chen et al., 2016) не хранит активации, а пересчитывает их на обратном проходе: память падает с $O(L)$ до $O(\sqrt{L})$ по числу слоёв ценой ~30% времени.

Распределённое обучение. DDP реплицирует модель на каждый GPU и синхронизирует градиенты через all-reduce. Когда модель не влезает даже в один GPU — FSDP или ZeRO (Rajbhandari et al., 2019), шардирующие параметры, градиенты и состояние оптимизатора между устройствами.

Воспроизводимость. Фиксируйте seed для torch, numpy, random и torch.backends.cudnn.deterministic = True. Битовой воспроизводимости на GPU вы всё равно не получите — атомарные операции недетерминированы по порядку суммирования, — но разброс между запусками сократится до приемлемого.

Поиск гиперпараметров. Случайный поиск почти всегда лучше сеточного (Bergstra & Bengio, 2012): сетка тратит бюджет на многократный перебор неважных измерений. Для дорогих запусков — байесовская оптимизация или Hyperband (Optuna, Ray Tune). Порядок настройки: LR → размер батча → weight decay → всё остальное.

Мониторинг. Логируйте не только loss: кривую LR, глобальную норму градиента, норму весов, долю «мёртвых» ReLU-нейронов, пропускную способность и утилизацию GPU. Половина проблем диагностируется по этим графикам за минуту, а по одному loss — за день.


10. Мини-итог

  • Ландшафт невыпуклый и плохо обусловленный; беда — не локальные минимумы, а овраги, плато и седловые точки.
  • Momentum гасит зигзаг, адаптивность (RMSProp/Adam) даёт каждой координате свой масштаб. Трансформеры — AdamW; CNN на изображениях — SGD+Nesterov всё ещё конкурентен.
  • AdamW, а не Adam(weight_decay=...). И никакого weight decay на bias и нормализации.
  • Learning rate — гиперпараметр номер один. Warmup + cosine decay; LR масштабируется с размером батча.
  • Инициализация держит дисперсию активаций около единицы: He для ReLU, Xavier для tanh, плюс масштабирование выходных проекций residual-блоков на $1/\sqrt{2L}$.
  • Нормализация переделывает ландшафт: BatchNorm для CNN с большими батчами, LayerNorm/RMSNorm для последовательностей, GroupNorm при малых батчах. Pre-LN устойчивее Post-LN.
  • Регуляризация — weight decay, dropout, аугментация, label smoothing, EMA, early stopping. Начинайте с данных: аугментация почти всегда сильнее любого штрафа на веса.
  • Стабильностьclip_grad_norm_ = 1.0, bf16 вместо fp16 где возможно, логирование нормы градиента. И обязательный sanity check на трёх батчах перед долгим запуском.

Источники


Что дальше

Мы научились обучать произвольную сеть. Дальше начинаются архитектуры, каждая со своим индуктивным смещением под конкретный тип данных. Первая и самая наглядная — свёрточные сети: они встраивают в архитектуру знание о том, что у изображения есть локальность и трансляционная инвариантность, и заодно дают отличный пример того, как все приёмы этой статьи (He-инициализация, BatchNorm, аугментация, SGD+momentum) работают вместе.

Свёрточные сети и компьютерное зрение

Нашли неточность? Выделите фрагмент текста — рядом появится жучок.

Нужен разбор именно вашей ситуации?

Статья описывает общий случай. Если у вас частный — можно разобрать его отдельно, платно. А если не хватает целого материала, предложите тему: её оплачивают вскладчину, и она выходит открытой для всех.

Доска запросов