Jagged Flash Attention: ядро Meta быстрее FA4 на 50%
3 200 строк кода против примерно 10 000 у FlashAttention-4. Ядро внимания, собранное из них, обгоняет эталон на 13% в прямом проходе и на 50% в обратном, и это не синтетический бенчмарк, а самая тяжёлая операция в рекламной модели Meta.
Команда Meta Ads пересобрала Jagged Flash Attention (JFA) на TLX, надстройке над Triton с доступом к низкоуровневым примитивам Blackwell. История полезна каждому, кто ускоряет собственные модели: она показывает, что даёт управляемый высокоуровневый код там, где раньше был только рукописный CUDA.
Что такое Jagged Flash Attention
Рекламные модели Meta, генеративная рекламная модель GEM и архитектура Kunlun, работают не с текстом, а с историями пользовательских взаимодействий. Длины этих историй различаются на порядки, и именно такую нагрузку в индустрии называют jagged, то есть рваной.
Стандартный ответ на неравномерность прост: выровнять все последовательности по самой длинной с помощью паддинга. На пустых токенах при этом может сгорать до половины бюджета обучения. Поэтому GEM пакует последовательности вплотную, а границы каждой хранит в тензоре смещений (offsets).
Jagged Flash Attention это ядро, которое применяет алгоритм FlashAttention напрямую к упакованным Q, K и V с учётом смещений. Паддинг не материализуется ни на одном шаге. В продакшн-режиме, который авторы называют broadcast-Q, один плотный Q размножается на весь батч, а его градиент dQ суммируется по всем последовательностям сразу. Из этой детали вырастает половина оптимизаций ниже.
Почему обычный Triton упирался в потолок
Исходное ядро было написано на Triton и работало алгоритмически корректно. Но ни одно решение о планировании человек в нём не принимал: глубину конвейера, порядок загрузок и переиспользование буферов определял компилятор. Загрузки, softmax и матричные умножения выполняла одна однородная группа варпов, поэтому тензорные ядра простаивали всякий раз, когда варпы уходили считать softmax.
Для Blackwell это приговор: пик там достигается, только если тензорные ядра получают непрерывный поток MMA-операций. TLX (Triton Low-level Extensions) добавляет к высокоуровневой модели Triton первоклассные низкоуровневые примитивы: явное выделение shared- и тензорной памяти, специализацию варпов (warp specialization), барьеры, асинхронные TMA и MMA, а также Cluster Launch Control. Раньше такой контроль означал рукописный CuteDSL или CUDA, то есть тысячи строк на языке, который читают только специалисты по ядрам. Саму систему плагинов TLX мы разбирали в июле; здесь речь о том, что она даёт на реальных нагрузках.
Каркас: специализация варпов и персистентность
Первое, что даёт TLX, это возможность разложить вычислительный блок CTA на роли. За TMA-загрузки, матричные умножения, softmax и запись результатов отвечают отдельные варпы, а в обратном проходе появляется ещё и выделенный варп для редукции dQ. Тензорные ядра получают сплошной поток MMA, пока softmax считается параллельно.
Память тоже переходит под ручное управление. K и V держатся в тройном буфере, чтобы варп загрузок убегал на несколько шагов вперёд. Буферы тензорной памяти переиспользуются по временам жизни: результаты QK, матрица вероятностей P и статистики softmax делят один регион, а аккумулятор PV получает собственный. Данные ходят между варпами через явные producer/consumer барьеры, а само ядро становится персистентным: один CTA на SM без остановки перебирает тайлы вместо того, чтобы завершаться после каждого.
Персистентность превращает ядро в площадку для собственных решений о планировании. Все оптимизации ниже выросли именно из неё.
Оптимизации: что забрало оставшийся запас
Узкие места авторы искали методично: снимали счётчики Nsight Compute (утилизация конвейеров тензорных ядер и TMEM, загруженность SM, проливы регистров в локальную память), сверялись с дампами ptxas, а каждую правку до и после прогоняли через TritonBench в режиме профилировщика. Ниже четыре направления, каждое против своего боттлнека.
Балансировка тайлов и Cluster Launch Control
Рваные длины перекашивают нагрузку между потоковыми мультипроцессорами: пока один SM перемалывает длинную последовательность, соседние простаивают. Тепловая карта занятости показывает это наглядно.
В прямом проходе стоимость тайла пропорциональна длине ключей его батча, поэтому хост сортирует тайлы по убыванию нагрузки и раздаёт их зигзагом: чётные проходы слева направо, нечётные справа налево. Такая змеевидная раздача почти бесплатна и даёт каждому SM сбалансированную смесь длинных и коротких тайлов. Один этот приём добавил прямому проходу около 20%.
В обратном проходе тайлы почти равны по стоимости, но их количество зависит от длины истории, и сортировка не помогает. Здесь хост заранее вычисляет список валидных тайлов и раздаёт их по кругу, а порядок подобран так, чтобы соседние тайлы переиспользовали одни и те же блоки K/V в кэше L2.
Предсказуемый перекос закрывает программная балансировка, а остаточный, который виден только в рантайме, забирает Cluster Launch Control. Это фича Blackwell: CLC выдаёт SM следующий номер тайла по запросу, поэтому кто освободился, тот и берёт работу. Приёмы дополняют друг друга, ведь CLC не умеет отличать пустые тайлы от настоящих, и предварительная фильтрация на хосте продолжает экономить циклы.
Эпилог dQ: главный боттлнек обратного прохода
В режиме broadcast-Q градиент dQ суммируется по всему батчу, и все SM пишут в одни и те же адреса. Редукция с добавлением превращается в точку конкуренции, а высокоуровневый API выполняет её последовательно, колонка за колонкой, без конвейера. По замерам, этот эпилог съедал 9-11% утилизации тензорных ядер, больше любого другого места в ядре.
Решение это ручная двухбуферная постановка в shared-память: пока одна колонка редуцируется в HBM, следующая копируется из тензорной памяти, и хотя бы одна запись всегда в полёте. Техника обобщается на любой эпилог с редукцией, не только на attention.
Вторая правка бьёт в ту же точку с другой стороны. Варп редукции держит буфер dQ в тензорной памяти, пока не выгрузит все срезы, и MMA-варп ждёт его освобождения, прежде чем начать умножение следующего тайла. Если предварительно выгрузить последние один-два среза в регистры и освободить буфер раньше, матричные умножения поедут дальше. Контрольный замер, в котором запись dQ просто отключали, показывал запас в 8-11% утилизации, и большая его часть так и вернулась. Сколько срезов освобождать заранее, ядро подбирает само: один или два, дальше растёт давление на регистры и становится только хуже.
Разрезание цикла против проливов регистров
Профилировка прямого прохода показала неожиданное: ядро голодало по выдаче MMA. Конвейер тензорной памяти был загружен на 57% против 82% у FA4 на той же форме, а потери шли от softmax-варпа, перегруженного регистрами. Это подтверждали два сигнала: дампы ptxas показывали проливы в локальную память, а в одном из ранних прогонов простое увеличение лимита регистров срезало этот трафик на 40% и дало 6% к скорости.
Корень проблемы прятался в ветке маскирования. Внутри горячего цикла жила проверка, срабатывавшая только на последнем неполном тайле. Регистры распределяются статически, поэтому переменные редкой ветки (смещения колонок, тензор маски, select) держали свои регистры на каждой итерации и провоцировали проливы. Лечится это разрезанием цикла (loop peeling): цикл делится на ветвесвободную основную часть и крошечный хвост с маской. Маска становится константой времени компиляции, регистры освобождаются, и планировщик собирает более плотный поток MMA. В обратном проходе тот же приём вернул около 9% латентности, которую съедали проверки корректности.
Парные CTA: два SM на одно умножение
Обратный проход это прежде всего матричные умножения: пять GEMM на каждый блок K/V, и одного CTA не хватает, чтобы загрузить тензорные ядра Blackwell целиком. Авторы переняли у FA4 схему парных CTA: два SM в кластере считают один широкий матмул, поделив строки аккумулятора, а промежуточный результат обменивается через распределённую shared-память, минуя HBM. В продакшн-конфигурации с broadcast-Q и head_dim=128 это добавляет около 12% пропускной способности в обратном проходе. Новое здесь то, что схема портирована на рваную раскладку и работает поверх персистентного планировщика с CLC.
Цифры на B200
Замеры шли в bf16 на B200 в двух режимах. Первый, продакшн-режим с broadcast-Q, повторяет реальный профиль рекламных моделей. Второй, плотный LLM-стиль с формой B=768, H=4 и head_dim=128, нужен для сравнения с FA4 на его территории. Соперник это актуальная на май 2026 года версия FlashAttention-4 на CuteDSL.
На рваных формах TLX-ядро выигрывает оба прохода: прямой быстрее в среднем на 13%, уступая только на самых длинных последовательностях при высокой плотности, а обратный быстрее везде и сразу на 50%. Диапазон длин при этом заметает весь продакшн-разброс: от сильной рваности до почти выровненных последовательностей.
На плотных формах картина честнее: forward держится на уровне около 87% от FA4, зато backward снова впереди, примерно на 17%. Речь не о том, что высокоуровневый код победил везде, а о том, что на профиле, ради которого ядро строилось, он победил с заметным запасом.
MXFP8 и разреженность: форк вместо переписывания
Отдельный довод в пользу нового слоя в том, как ядро переживает смену требований.
Для обучения, терпимого к FP8, авторы собрали вариант на MXFP8: умножения BF16 заменены на блочно-масштабируемые MMA с данными E4M3 и масштабами E8M0, при этом специализация варпов, барьеры и раскладка памяти остались прежними. Вероятности P квантуются на лету, а их масштабы живут прямо в тензорной памяти и подкладываются в MMA без промежуточных шагов. После настройки forward обходит собственное FP8-ядро FA4, а backward на плотных формах идёт на паритете.
Для длинных историй из того же кода выращен блочно-разреженный вариант. Двухстадийная схема: дешёвое ядро усредняет блоки Q и K и выбирает top-k самых релевантных, после чего основное ядро считает внимание только по ним. При коэффициенте отбора 0,5 прямой проход ускоряется в 1,3-1,5 раза, а поддержка broadcast-Q, GQA и окон сохраняется. В обоих вариантах переписывалась математика, а не машинерия. Именно этот размен и есть главный результат работы.
Что это значит за пределами рекламы
Рваные нагрузки давно вышли за пределы рекламы. RL-посттренировка порождает роллауты переменной длины, агентные цепочки копят истории разной глубины, рекомендательные системы и поиск устроены так же. Везде выравнивание паддингом съедает десятки процентов вычислений, а инструментов с нужным уровнем контроля было ровно два: Triton, которому не хватает управления, и рукописный CuteDSL, которому не хватает скорости разработки. TLX встаёт между ними третьим путём.
Показательна и методика. Авторы не изобретали новых алгоритмов, а снимали счётчики, находили боттлнеки и закрывали их по одному. Инженерные приёмы вроде двухбуферного эпилога с редукцией или раннего освобождения тензорной памяти переносятся на другие ядра почти без изменений. Код целиком открыт в репозитории ads_model_kernel_library, а сам подход вырос из двух предыдущих публикаций про TLX, о Cluster Launch Control и о блочной разреженности.
Кстати, ту же команду Meta Ads мы разбирали несколько дней назад в другом контексте: как FlashAttention-4 получила MXFP8 и 2,85 PF/s на Blackwell. Две работы хорошо дополняют друг друга: одна про квантование внутри существующего ядра, вторая про то, как построить такое ядро заново на новом слое абстракции.
Частые вопросы
Чем jagged-внимание отличается от обычного?
Обычное внимание ожидает тензоры, выровненные по длине. Jagged-вариант работает с упакованными последовательностями и тензором смещений, поэтому не тратит вычисления на паддинг. Для нагрузок с сильным разбросом длин это разница до 50% вычислительного бюджета.
TLX это замена CUDA?
Нет, это уровень поверх Triton с доступом к низкоуровневым примитивам: явной памяти, барьерам, асинхронным TMA и MMA. Рукописные ядра на CuteDSL по-прежнему держат рекорды на плотных формах, где TLX-JFA отстаёт от FA4 в прямом проходе. Но когда кода втрое меньше и обновлять его нужно каждую неделю, размен уходит в другую сторону.
Где посмотреть код и бенчмарки?
Код лежит в открытом репозитории facebookresearch/ads_model_kernel_library, директория tlx_jfa. Графики по формам и плотностям, детали каждой оптимизации и полные таблицы замеров разобраны в блоге PyTorch.
Итог
Ядро внимания для рекламных моделей Meta стало короче втрое и быстрее эталонного FlashAttention-4 на целевых формах: плюс 13% в прямом проходе, плюс 50% в обратном. Разгадка не в единственной находке, а в уровне контроля: TLX позволил выразить специализацию варпов, персистентность и точное управление памятью в читаемом Triton-коде, а дальше осталось закрывать боттлнеки по данным профайлера.
Для тех, кто ускоряет внимание под собственный профиль нагрузки, это готовый набор инженерных приёмов и открытый код: 3 200 строк в репозитории читаются за один вечер.