Flash Attention 4 на MXFP8: 2.85 PF/s на Blackwell
2,85 PF/s в прямом проходе и 1,98 PF/s в обратном. Именно такие цифры получила команда Meta Ads, расширив FlashAttention-4 поддержкой формата MXFP8 для обеих фаз вычисления. Интересна тут не сама скорость, а то, что вскрывается по дороге: на Blackwell переход внимания в FP8 не сводится к замене типа данных. Это переделка распределения памяти, порядка барьеров и квантование тензоров, которые раньше никто не квантовал.
Код опубликован в открытой библиотеке Ads Model Kernel Library на GitHub, а сам модуль уже используется внутри Meta для обучения GEM.
Что такое MXFP8 и Flash Attention 4
MXFP8 это формат микроскейлинга, где блок из 32 значений делит один общий масштаб, записанный в E8M0. Тензорные ядра Blackwell получили инструкцию tcgen05.mma.block_scale, которая работает с MXFP8, MXFP6, MXFP4 и NVFP4 напрямую и выдаёт от 2 до 4 раз больше пропускной способности, чем BF16 MMA.
Flash Attention 4 это четвёртое поколение kernel'ов внимания, оптимизированное под ту же архитектуру. До этой работы FA4 в проде считал в BF16, и вопрос стоял не в том, можно ли вообще считать внимание в FP8, а в том, как это сделать, не потеряв ни такта на конвертацию форматов.
Почему замена типа данных не работает
Наивный путь выглядит так: поменяли тип, получили ускорение. На Blackwell он разбивается о три стены, и авторы разбирают каждую.
Первая стена это масштабы. Блочное MMA требует, чтобы scale factors лежали в TMEM, тензорной памяти. У Blackwell её ровно 512 колонок, и в существующих FA-ядрах они уже заняты под операнды и аккумуляторы целиком. Свободного места под масштабы просто нет.
Вторая стена это K-измерение. Для Q и K масштабы считаются вдоль эмбеддинг-размерности, но для V их приходится считать вдоль последовательности. А P и dS вообще появляются внутри kernel'а, на лету, и квантовать их нужно онлайн, прямо в процессе счёта.
Третья стена это накладные расходы на конвертацию. Если квантование встаёт в критический путь, MMA перестаёт выполняться на полной скорости, и весь выигрыш от FP8 исчезает на пересыпке байтов между регистрами.
Почему блочный микроскейлинг, а не один масштаб на тензор
Обычное per-tensor квантование в FP8 держит один масштаб на весь тензор. Для активаций трансформера это плохо работает: значения внутри одного тензора различаются на порядки, и общий масштаб подгоняется под выбросы, обесценивая точность на всём остальном.
Микроскейлинг режет тензор на блоки по 32 значения и даёт каждому свой масштаб. Выброс портит точность только внутри своего блока, а не всего тензора. Плата за это, разумеется, сами масштабы: их нужно считать, хранить и подкладывать к тензорным ядрам там, где для них нет ни байта свободной памяти.
Как масштабы поместились в 512 колонок TMEM
Прямой проход FA4 использует пинг-понг между двумя тайлами Q размером 128 на 128. Ядро загружает Q0 и Q1 и в цикле идёт по тайлам K и V вдоль размерности N, чередуя два набора аккумуляторов.
Авторы не стали расширять TMEM, а начали перекрывать регионы. Для пролога, где считаются S(i) = Q(i) @ K(0), масштабы S(i) кладутся в регион O(i), потому что O(i) ещё не начал считаться. Масштабы O(i) живут в регионе S(i), но в области, не занятой P(i). Это проходит потому, что в FP8 P занимает 32 колонки, а S в FP32 требует 128.
Для масштабов S(i) нашлось место в регионе S(1-i). Здесь пришлось добавить барьер между MMA- и softmax-warp'ами: MMA выполняется асинхронно, и возможна запись-в-запись между аккумулятором S(1-i) и масштабами S(i). Барьер почти ничего не стоит, потому что параллельно всё равно идёт GEMM O(i), который длиннее, чем чтение TMEM в регистры.
Квантование dS без транспонирования
Обратный проход сложнее. Тензор dS нужно квантовать квадратными блоками, и обычная схема требует транспонирования: масштабы обязаны быть инвариантны к нему. Авторы обошли это через warp-wide редукцию redux.sync.max.abs.f32. Максимум по модулю считается прямо в той раскладке, в которой dS лежит в регистрах, и транспонировать не нужно ничего.
Для dQ выбрали компромисс между точностью и скоростью. При накоплении dQ в FP32 обратный проход даёт 1,82 PF/s, при накоплении в FP16 с масштабом 2^9 уже 1,98 PF/s. Это 1,21x и 1,32x к BF16-базлайну cuDNN на 1,50 PF/s.
Квантование внутри производителей
Отдельная часть работы это убрать квантование из критического пути, спрятав его внутри тех ядер, которые и так создают тензор.
Fused GEMM+Quantize пишет FP8 на выходе GEMM сразу в двух раскладках масштабов, чтобы потребителю не пришлось ничего пересчитывать. На формах B200 это поднимает пропускную способность с примерно 0,20 PF/s при раздельных GEMM и квантовании до 0,91-0,98 PF/s. Ускорение в 4,4-4,7 раза без изменения математики.
Fused RMSNorm+Quantize делает то же самое для нормализации. При N=512 эффективная пропускная способность растёт с 0,75-0,87 TB/s до 2,6-4,0 TB/s, то есть до 4,7x. Нормализация перестаёт быть отдельным проходом по памяти и становится эпилогом того, что и так читает данные.
Jagged-модуль без gather
Внутренние нагрузки Meta отличаются от академических бенчмарков: больше батч, меньше голов, короче запросы и сильно неравномерное распределение длин последовательностей. Это классический jagged-случай, где тензоры приходится выравнивать, а выравнивание означает копирование памяти.
Авторы не стали выравнивать FP8-данные. Данные остаются на непакованных позициях, а разбрасываются, паддятся и переставляются только масштабы, которые в 32 раза меньше. После этого адреса становятся выровненными по 128, как того требует TMA. Экономия налицо: пересылать метаданные вместо самих тензоров дешевле в десятки раз.
Переменная длина последовательностей потребовала отдельной работы с TMA. Адресацию пришлось строить так, чтобы дескрипторы указывали на реальные, а не на выровненные начало и конец каждой последовательности, иначе на коротких строках ядро работало бы по пустым данным.
Результаты на LLM-формах
Forward MXFP8 выходит на 2,85 PF/s. Это сопоставимо с cuDNN 9.24 MX8 на 2,82 PF/s и даёт 1,43x к внутренней BF16-реализации на 2,00 PF/s.
Обратный проход даёт 1,82 и 1,98 PF/s против 1,50 у BF16-базлайна. Замеры сделаны на GPU GB300 внутри Meta.
На внутренних формах с jagged K при sparsity 0,5 и максимуме K в 16384 токенов forward достигает 2,54 PF/s. Это 1,59x к BF16 на 1,60 PF/s и отставание в 4% от cuDNN MX8. При равномерном K результат 2,59 PF/s и 1,51x соответственно. Обратный проход выдаёт 1,42 PF/s с FP32-накоплением dQ и 1,58 PF/s с FP16, то есть 1,37x и 1,52x к BF16 на 1,04 PF/s.
Отдельно про честность сравнения. При неоднородных V-масштабах E8M0 в MXFP8-backward у cuDNN 9.24 результаты dQ и dK не совпали с эталоном, и расхождение ушло только при переходе к однородным масштабам. Свой результат в 1,76 PF/s авторы поэтому помечают как чисто производительностный, без подтверждённой точности. Ещё одно расхождение: cuDNN возвращает градиенты dK и dV в BF16, тогда как FA4 MX8 отдаёт их квантованными.
Сквозной модуль
Самая практичная таблица это латентность готового модуля. Замер на одном GB300, батч 768, 512 seed-запросов, D = 384, H = 3, медиана прямого и обратного прохода, 1000 итераций после 200 прогревочных.
| Токенов KV | BF16 FA4, мс | MXFP8, мс | Ускорение |
|---|---|---|---|
| 4096 | 10,327 | 10,281 | 1,00x |
| 8192 | 17,348 | 14,663 | 1,18x |
| 16384 | 31,705 | 24,479 | 1,30x |
На коротком контексте выигрыша нет вообще. Он появляется на 8 тысячах токенов и доходит до 1,30x на 16 тысячах. Причина в том, что на малых длинах модуль упирается не в вычисление, а в накладные расходы запуска и планирования, и снижение точности тут не помогает ничем.
Что с точностью
В работе считают SQNR, отношение сигнала к шуму квантования, где шум это разница между результатом MXFP8 и BF16-эталоном. Метрика удобнее обычного SNR, потому что амплитуды активаций и градиентов различаются на порядки, и единая шкала для них мало что значит.
На продакшен-подобном синтетическом прогоне (768 в батче, 5 856 019 токенов KV, максимум 14 980, D = 384, H = 3, Q = 512) показатели держатся в диапазоне от 18,84 до 28,02 дБ по разным тензорам. Строгая сводка выглядит как 23 / 1 / 0: двадцать три тензора в норме, один на границе, ни одного выхода за допуск.
Практический вывод авторов: MXFP8-внимание дало нейтральный вклад в численную ошибку на обучении GEM. Но устойчивость не досталась бесплатно, её выстраивали приёмами, вдохновлёнными SageAttention3 и работой Qiu и коллег.
Где это уже применяется
Модуль встроен в cross-attention для рекламных задач Meta, где формы далеки от тех, на которых отлаживали оригинальный FA4. Это важная деталь: работа проверялась не на удобных прямоугольных тензорах из статьи, а на реальной продакшен-нагрузке с рваными длинами.
Авторы отмечают, что это одна из первых SoTA-реализаций MXFP8 для forward и backward, реально используемых в обучении, а не только в бенчмарках. Разница принципиальная: в бенчмарке достаточно показать пиковую пропускную способность, в обучении нужно ещё не развалить градиенты на тысячах шагов.
Часто задаваемые вопросы
Можно ли просто взять MXFP8-ядро и заменить BF16 в своём обучении? Нет. Ядро решает только часть задачи. Без fused GEMM+Quantize и без перестройки раскладки масштабов квантование встанет в критический путь, и выигрыш от низкой точности съестся накладными расходами на конвертацию.
Почему ускорение растёт с длиной контекста? На коротких последовательностях модуль ограничен накладными расходами и запуском ядер, а не пропускной способностью тензорных ядер. Чем длиннее контекст, тем большую долю времени занимает собственно MMA, где FP8 и даёт выигрыш.
Значит ли нейтральный NE, что FP8-внимание безопасно для любого обучения? Нет. Результат получен на конкретном модуле и конкретных формах, с техниками стабилизации из SageAttention3 и отдельным режимом накопления dQ. Переносить его на другой домен без собственных замеров SQNR не стоит.
Реально ли повторить это вне Blackwell?
Нет, и не только из-за инструкции tcgen05.mma.block_scale. Вся конструкция опирается на конкретный размер TMEM в 512 колонок и на перекрытие регионов под масштабы. На другой архитектуре эти решения придётся искать заново.
Итог
Работа Meta показывает, что block-scaled attention на Blackwell это системная задача. Масштабы пришлось втискивать в полностью занятую TMEM через перекрытие регионов и дополнительный барьер, dS квантовать онлайн без транспонирования, квантование прятать внутрь GEMM и RMSNorm, а jagged-тензоры оставлять на месте, переставляя только метаданные масштабов.
На выходе Flash Attention 4 на MXFP8 даёт до 1,6x к BF16 в forward, до 1,52x в backward и сквозное ускорение модуля до 1,30x на длинных контекстах при нейтральном вкладе в численную ошибку. Код лежит в открытой библиотеке, и это, пожалуй, главное: техники можно переиспользовать, а не переизобретать заново.
Если вы обучаете модели с длинным контекстом, начните с двух вещей: посмотрите на свою раскладку масштабов и измерьте, сколько времени уходит на конвертацию форматов. Незабранное ускорение обычно лежит именно там.