Материал подготовлен автоматически по первоисточникам: ссылки на них — в конце статьи.
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 при одинаковом порядке данных. Эти результаты дают основание пересмотреть проверку долгих прогонов на таком стеке, но не описывают поведение других форматов точности и аппаратных платформ.
Источники
Похоже на вашу задачу?
Расскажите, что собираете. За полчаса разложим на этапы и назовём сроки — это бесплатно и ни к чему не обязывает.



