Журнал · Rit.work

FlashAttention-3 искажает градиенты BF16 без NaN: как GProj устраняет ошибку

GProj восстанавливает нарушенное округлением свойство градиента внимания и предотвращает поздний сбой обучения с FlashAttention-3.

Rit.work
Студия разработки
30 сентября 2026 г.3 мин чтения

Материал подготовлен автоматически по первоисточникам: ссылки на них — в конце статьи.

BF16-версия FlashAttention-3 может незаметно исказить обратный проход лишь на поздней стадии обучения: функция потерь ухудшается, хотя вычисления не дают NaN. В нерецензированном препринте Rutgers University, Carnegie Mellon University, Oracle, New York University и MBZUAI, где все числа получили сами авторы, GProj обучила модель до той же функции потерь, что внимание в FP32. Короткий тест ядра на случайных входах такую ошибку не обнаружит.

Почему ошибка проявляется к концу обучения

Проблему нашли при обучении Transformer на 450 млн параметров и 50 млрд токенов. Первые 25 млрд токенов прогон выглядел штатно, затем норма градиента выросла в тысячу раз, а итоговая функция потерь оказалась на 0,2 нат выше результата с вниманием в FP32.

При этом обучение не остановилось: переполнений и NaN не было. Если оставить прямой проход неизменным, но пересчитать обратный проход наиболее пострадавших слоёв в FP32, почти весь лишний градиент исчезает. Значит, источник сбоя находится не в данных, оптимизаторе или самой архитектуре, а в вычислительном ядре внимания.

Первую часть ошибки создаёт известная оптимизация FlashAttention-3. Ядро объединяет умножение на масштаб и вычитание максимума строки в одну операцию. Из-за округления максимальная оценка до softmax не всегда превращается точно в ноль, а обратный проход получает неточно сохранённый выход внимания.

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

Как нулевая сумма превращается в ложный градиент

У точного градиента softmax есть сохраняемое свойство: сумма элементов в каждой строке равна нулю. Благодаря этому градиент запроса зависит от различий между ключами, но не меняется, если ко всем ключам прибавить один и тот же вектор.

FlashAttention-3 вычисляет градиент оценок, затем округляет его до BF16 перед умножением на ключи. BF16 — формат с пониженной точностью, поэтому после округления нулевая сумма превращается в небольшой остаток. При умножении этот остаток добавляет к градиенту компоненту, направленную вдоль среднего ключа.

В начале обучения добавка мала. Позже ключи растут, а внимание концентрируется почти на одной позиции. Правильный градиент в таком режиме приближается к нулю, но ошибка округления умножается на крупный средний ключ и начинает превосходить полезный сигнал. Медианная относительная ошибка градиента запроса достигала 219%, хотя прямой проход оставался работоспособным.

GProj восстанавливает нулевую сумму уже после приведения к BF16, то есть на тех значениях, которые ядро действительно подаёт в матричное умножение. Метод измеряет оставшуюся массу строки и вычитает соответствующую долю вероятностей внимания. Поправка для запросов помещается в существующий обратный проход, а для ключей требует дополнительного прохода.

После такой проекции медианная ошибка градиента запроса снизилась до 0,34%, то есть до уровня внимания в FP32. Исправление прямого прохода, напротив, не может убрать эту составляющую: ошибка появляется позже, при округлении уже сформированного градиента.

Что менять в плане обучения модели

Для команд, которые обучают модели с BF16 FlashAttention-3 на Hopper GPU, работа меняет прежде всего способ проверки вычислительного стека. Сравнения на случайных входах и короткого пробного обучения недостаточно: ошибка усиливается только после того, как ключи становятся крупными, а внимание — почти дискретным.

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

GProj добавила 4,7% ко времени полного шага обучения. Для сравнения, объединённое ядро внимания в FP32 добавило 33,4%. Вычитание среднего ключа и исправление прямого прохода стоили дешевле концептуально, но в сопоставленных прогонах не предотвратили переход слоёв в режим крупных оценок внимания.

Проверка охватывает один Transformer, BF16 FlashAttention-3 на Hopper GPU, фиксированные размерность головы и длину последовательности. Сравнивали исходное ядро, исправленный прямой проход, сглаживание ключей, GProj и внимание в FP32 при одинаковом порядке данных. Эти результаты дают основание пересмотреть проверку долгих прогонов на таком стеке, но не описывают поведение других форматов точности и аппаратных платформ.

Источники

Пауза в чтении

Похоже на вашу задачу?

Расскажите, что собираете. За полчаса разложим на этапы и назовём сроки — это бесплатно и ни к чему не обязывает.

Rit.work

Студия разработки

Собираем мобильные приложения и помогаем командам получать от AI реальную пользу. Основатель и команда, работаем удалённо — с клиентами в России и за рубежом.

← Ко всем материалам
Понравилось? Обсудим вашу задачу