Ответ на вопрос
Возможные причины (кратко, с объяснением)
- Неподходящая скорость обучения: слишком большая или слишком малая. Приведённый оптимизатор с \(\eta=0.01\) может быть либо чересчур большим (колебания/разбег), либо слишком маленьким (почти отсутствие обновлений).
- Неверно заданные параметры оптимизатора (optimizer не содержит параметров модели): оптимизатор создан не от \(\text{model.parameters()}\).
- Градиенты равны нулю или NaN/Inf: из‑за замороженных параметров (\(\text{requires\_grad}=False\)), неправильных операций (inplace), или численной нестабильности (overflow).
- Неправильная функция потерь / формат таргетов: например, \(\text{CrossEntropyLoss}\) ожидает логиты размера \((N, C)\) и таргеты целочисленные \((N,)\) с dtype=\(\text{torch.long}\). Накладка softmax перед CrossEntropy даёт неправильный результат.
- Режим модели: забыли вызвать \(\text{model.train()}\) (или, наоборот, включили train при валидации — dropout/batchnorm ведут себя иначе).
- Ошибки в датасете: перепутаны входы и метки, метки не соответствуют классам, сильно несбалансированные классы, повреждённые данные, неправильная нормализация.
- Неправильный порядок вызовов: забыли \(\text{optimizer.zero\_grad()}\) или \(\text{loss.backward()}\) / \(\text{optimizer.step()}\).
- Неправильное усреднение/редукция потери (reduction) — если берёте медиану/макс вместо mean, или мешаете разные масштабы.
- Плохая инициализация или архитектура: сигмоиды/таны в глубоких сетях — затухающие градиенты; слишком глубокая сеть без нормализации.
- Моментум/шедулер/weight decay: неверные гиперпараметры мешают сходимости.
- Проблемы при mixed precision: неправильно настроенный loss scaling вызывает нули/инфы.
Порядок действий для диагностики (пошагово)
1. Быстрые базовые проверки
- Убедиться, что вызываются в правильном порядке: \(\text{optimizer.zero\_grad()}\) → \(\text{loss.backward()}\) → \(\text{optimizer.step()}\).
- Проверить режим: вызвать \(\text{model.train()}\) при обучении.
- Проверить, что оптимизатор действительно содержит параметры: посмотреть \(\text{len(list(model.parameters()))}>0\) и \(\text{optimizer.param\_groups}\).
2. Проверить форму и тип выходов/таргетов
- Убедиться, что output имеет форму \((N, C)\), target — \((N,)\) и dtype=\(\text{torch.long}\) (для CrossEntropy).
- Убедиться, что вы НЕ подаёте softmax в модель перед \(\text{CrossEntropyLoss}\).
3. Наблюдения за loss по мини‑батчам
- Печатать loss для первых 10 батчей и смотреть тренд внутри эпохи. Если loss не изменяется вообще — проблема с градиентами/обновлениями или данными.
4. Проверить градиенты и обновления параметров
- Перед и после шага оптимизатора измерять нормы:
- градиент: \(\|\nabla_\theta L\|_2\),
- параметры: \(\|\theta\|_2\),
- изменение параметров: \(\|\Delta\theta\|_2 = \|\theta_{\text{new}}-\theta_{\text{old}}\|_2\).
- Если \(\|\nabla_\theta L\|_2\approx 0\) — искать frozen-параметры/ошибки вычисления графа; если \(\|\nabla_\theta L\|_2\) огромен или NaN — численность/overflow.
5. Проверить наличие NaN/Inf
- Проверить loss, градиенты и параметры на isnan/isinf. При обнаружении — уменьшить \(\eta\), отключить/проверить loss scaling (AMP), проверить входные нормировки.
6. Попробовать простейшую проверку — overfit на маленьком наборе
- Попытаться полностью переобучить модель на очень маленьком наборе (например, 10 примеров). Если модель не может добиться почти нулевого loss — проблема в реализации/оптимизаторе/градиента. Если может — проблема в данных/регуляризации или сложности задачи.
7. Эксперименты с гиперпараметрами/оптимизатором
- Попробовать очень малую \(\eta\) (например \(\eta=10^{-5}\)) и очень большую (например \(\eta=1\)) для диагностики; также попробовать другой оптимизатор (Adam) с базовыми настройками.
- Выключить weight decay/momentum/сложные scheduler-ы.
8. Проверить порядок и тип операций в forward/backward
- Искать in-place операции (операторы с постфиксом \_ ) которые ломают граф.
- Проверить, не применяется ли torch.no_grad() вокруг вычисления loss.
9. Проверить даталоадер и аугментации
- Убедиться, что таргеты не случайно перемешиваются отдельно от входов, нормализация правильная (mean/std), аугментации не портят метки.
10. Локализовать проблему по слоям
- Печатать градиенты/веса по слоям (либо их нормы/статистику). Это поможет найти слой, где градиент пропадает или взрывается.
- Временно упростить модель (убрать слои) и посмотреть, где поведение меняется.
11. Логирование и повторяемость
- Закрепить seed, логировать гиперпараметры и ключевые величины (loss, lr, градиенты, нормы весов). Это упростит воспроизведение и диагностику.
Короткие практические советы
- Для CrossEntropy: вход — логиты, target — long; не применять softmax.
- Для проверки параметров: убедиться, что optimizer стойт после создания модели и получает актуальные параметры.
- Overfit на 10 примерах — самый быстрый тест работоспособности обучения.
- Печать \(\|\nabla_\theta L\|_2\) и \(\|\Delta\theta\|_2\) даст быстрый ответ: градиенты есть? параметры меняются?
Если нужно, могу дать краткий чеклист команд/фрагментов PyTorch для каждого шага диагностики.
Еще