Неделя 10. Обучение
Учиться в приложении: тьютор, задачи с кодом →
Ядро: packing и маска между документами, bf16 и master weights, диагностика лосса, чекпоинты · Глубина: проверки здоровья и дашборды в деталях · ≈ 12 ч ядро / 22 ч всё
Первая настоящая тренировка, и с ней первые тихие ошибки. Модель с багом в данных или в лоссе обычно всё равно обучается: лосс падает, текст получается, просто хуже, чем мог бы. Задача недели: научиться отличать баг от свойства данных и читать логи так, чтобы видеть проблему заранее.
Шаг 1. Данные и лосс
На пальцах. Документы разной длины склеивают в одну ленту и режут на окна длины S. Это packing
(упаковка без паддинга). Окно может начинаться концом рецепта и продолжаться началом новости.
Без специальной маски первое слово новости «видит» рецепт и учится его продолжать. Модель тратит
силы на связь, которой в данных нет. Маска внимания между документами запрещает смотреть через границу.
- Data loading: packing, границы документов, маска внимания между документами. Между документами
ставят
EOS; блочно-диагональная маска (внимание только внутри документа) и сброс позиций RoPE остаются опциями. Llama 3 маскирует внимание через границы и отмечает, что важнее всего это на длинном контексте
На пальцах. В батче два примера: в первом 2 токена с лоссом 4, во втором 8 токенов с лоссом 1.
Среднее по токенам: (2·4 + 8·1)/10 = 1.6. Среднее по примерам: (4 + 1)/2 = 2.5. Это разные цели:
во второй каждый токен короткого примера весит в 4 раза больше. Выбор нормировки не деталь, а решение.
- Лосс считается как среднее по учитываемым токенам всего батча (
masked_cross_entropyделит наmask.sum()). При накоплении градиента (grad_accumмикробатчей) лосс каждого делят на их число. Если число токенов в микробатчах разное, среднее средних ≠ среднему по токенам: нормируй на общее число токенов
Шаг 2. Точность вычислений
На пальцах. Пусть вес хранится с тремя значащими цифрами: 1.00. Прибавь 0.001: выйдет 1.001,
и округление вернёт 1.00. Прибавь так тысячу раз, и вес не сдвинется, хотя должен был стать 2.00.
bf16 хранит всего 2–3 значащие цифры, а шаг оптимизатора это как раз такие крошечные добавки.
Поэтому обновления накапливают в fp32-копии весов (master weights), а в bf16 только считают.
- Mixed precision: bf16 vs fp16, master weights в fp32, loss scaling (и почему bf16 его не требует).
fp16: 5 бит порядка, 10 мантиссы, максимум 65 504, и маленькие градиенты уходят в ноль. Loss scaling
умножает лосс на
s(например, 2¹⁶), градиенты сдвигаются в представимую зону, перед клиппингом и шагом их делят обратно; приinfшаг пропускают и уменьшаютs. bf16: 8 бит порядка, как у fp32, и 7 бит мантиссы: диапазон тот же, underflow нет, но точность грубее (D12) autocastоставляет веса в fp32 и выбирает точность по операции: матмулы в bf16, softmax и нормы в fp32. Загрузка модели целиком в bf16 это другое: проще, но для обучения опаснее (см.05-ГЛУБИНА)
Шаг 3. Диагностика и чекпоинты
- Диагностика обучения: скачки лосса, NaN, взрыв нормы градиента, мёртвые нейроны. Скачку лосса обычно
предшествует скачок grad norm. Частые причины: конец warmup при слишком большом LR, плохой батч,
рост логитов внимания (лечат QK-norm, неделя 7). NaN:
log(0), переполнениеexpв fp16, строка маски целиком из−inf(softmax делит 0 на 0) - Что смотреть в wandb: loss, grad norm, LR, норма весов по слоям, доля клиппинга, токенов в секунду
- Чекпоинты, возобновление, детерминизм. В чекпоинт идут веса, состояние оптимизатора (
m,vи шагtдля bias correction), шаг расписания, состояния всех RNG, позиция в данных, масштаб loss scaler. Бит-в-бит воспроизводимость на GPU требуетtorch.use_deterministic_algorithms(True); на MPS её нет - Проверки здоровья до долгого запуска: лосс на старте ≈
ln V; модель запоминает один батч. У лосса есть пол, энтропия данных: вnanolm/README.mdперплексия упирается в 3.4, и это не баг
Типичные ошибки
- Weight decay на нормы и bias: ловит
test_param_groups_exclude_norms_and_biases. Лосс на промпте или паддинге: ловитtest_masked_cross_entropy_ignores_masked_positions - Возобновить без состояния оптимизатора:
tсбросится,mиvобнулятся, и первые шаги пойдут полным размером (test_bias_correction_makes_first_step_full_sizeпоказывает, насколько велик такой шаг). Итог: скачок лосса - Забыть делить на
grad_accumили клиппировать до конца всех backward. Теста нет: сравни кривыеbatch × 4иbatchприgrad_accum = 4, они обязаны совпасть - Утечка между документами при packing. Теста в nanolm нет, напиши его по образцу
test_model_is_causal: поменяй токен в документе A, логиты документа B не должны измениться
Код → nanolm/train.py: обучить свою модель на TinyStories / Shakespeare. Получить осмысленную генерацию.
Это твой первый настоящий LM. Внутри: get_batch (packing без учёта границ, прочитай докстринг),
estimate_loss, train (cosine_lr, AdamW(get_param_groups(...)), grad_accum, clip_grad_norm_ перед step).
Быстрый прогон: python -m nanolm.train --data data/toy.txt --steps 300 --vocab-size 400 --lr 3e-3.
Проверки: pytest tests/test_model.py tests/test_optim.py -v, ключевая из них test_model_can_overfit_one_batch.
Математика (трек D): D11: MLE на трёх распределениях и почему кросс-энтропия LM это и есть MLE; D12: почему обновления теряются в bf16, диапазоны форматов, стохастическое округление.
Интервью-вопрос недели: «Лосс вырос в 10 раз на шаге 40 000 и медленно возвращается. Что делаешь?» Структура на 3 минуты: grad norm и LR перед скачком → воспроизвести с чекпоинта на том же батче (данные или численность?) → логиты внимания и выход за диапазон → меры: откат и пропуск батча, ниже LR, QK-norm, клиппинг → профилактика: мониторинг grad norm по слоям.
Источники: Micikevicius et al., Mixed Precision Training (2018); Kalamkar et al., bfloat16 (2019); Llama 3 technical report (2024), где описана маска между документами.
Результаты недели
- Могу обучить модель до связной генерации и объяснить каждую кривую в логах обучения.
- Могу диагностировать скачок лосса по grad norm, LR и нормам весов по слоям.
- Могу реализовать packing с маской внимания между документами.
- Могу объяснить, почему fp16 требует loss scaling, а bf16 нет.
- Могу написать тест, который ловит утечку внимания между документами.
Самопроверка
- Лосс стал NaN на шаге 3000. Твои действия по порядку?
- Зачем master weights в fp32, если forward и backward идут в bf16?
- Что нужно сохранить в чекпоинте, чтобы возобновление было неотличимо от непрерывного обучения?
✅ Контрольная точка 2
- Трансформер с пустого файла за ≤40 минут, проходит тесты
- Обученная модель генерирует связный текст
- Устно за 10 минут: полный расчёт памяти и FLOPs для модели 7B