Flash Attention 4 на MXFP8: 2.85 PF/s на Blackwell

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 на длинных контекстах при нейтральном вкладе в численную ошибку. Код лежит в открытой библиотеке, и это, пожалуй, главное: техники можно переиспользовать, а не переизобретать заново.

Если вы обучаете модели с длинным контекстом, начните с двух вещей: посмотрите на свою раскладку масштабов и измерьте, сколько времени уходит на конвертацию форматов. Незабранное ускорение обычно лежит именно там.

← Все записи