Обучение сетей: оптимизаторы, инициализация, нормализация, регуляризация
В прошлой статье мы разобрались, как сеть считает предсказание и как обратное распространение выдаёт градиент функции потерь по каждому весу. Это ответ на вопрос «куда шагать».
Эта статья — про всё остальное, и «остальное» здесь больше, чем кажется:
- как шагать (оптимизатор: 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. Нормализация: перестроить ландшафт, а не бороться с ним
Пока мы улучшали спуск. Нормализация улучшает сам ландшафт — делает его ближе к сферическому, снижая число обусловленности.
Общая формула у всех одна, отличаются только оси, по которым считаются $\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
Обратите внимание на две вещи, которые чаще всего делают неправильно:
loss / accum_steps— без деления эффективный LR неявно умножается наaccum_steps.- В чекпойнт кладём состояние оптимизатора и шедулера. Без них 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 на трёх батчах перед долгим запуском.
Источники
- I. Goodfellow, Y. Bengio, A. Courville. Deep Learning, главы 7–8 — бесплатно онлайн.
- A. Karpathy. A Recipe for Training Neural Networks — karpathy.github.io. Лучший чек-лист по методологии обучения, читается за 20 минут.
- S. Ruder. An overview of gradient descent optimization algorithms — arxiv.org/abs/1609.04747.
- Google Research. Deep Learning Tuning Playbook — github.com/google-research/tuning_playbook.
- Официальная документация: torch.optim, torch.amp, torch.nn.init.
Что дальше
Мы научились обучать произвольную сеть. Дальше начинаются архитектуры, каждая со своим индуктивным смещением под конкретный тип данных. Первая и самая наглядная — свёрточные сети: они встраивают в архитектуру знание о том, что у изображения есть локальность и трансляционная инвариантность, и заодно дают отличный пример того, как все приёмы этой статьи (He-инициализация, BatchNorm, аугментация, SGD+momentum) работают вместе.