Назад
283

FlashAttention-3 и FlashAttention-4

283

Сегодня мы завершим разбор оптимизаций attention в рамках исследований FlashAttention-3 и FlashAttention-4

Обе они объединены общей идеей: каждое новое поколение GPU смещает баланс эффективности алгоритма — его нужно каждый раз перепроектировать под конкретную архитектуру железа. Давайте разберём, что изменилось и как достигается ускорение.

Почему FlashAttention-2 перестало хватать?

В прошлом посте рассказали, как FlashAttention-2 довела утилизацию потоковых мультипроцессоров до 80–90% на архитектуре Ampere (A100). Но с появлением Hopper (H100) выяснилось — та же реализация показывает лишь около 35% утилизации, а эталонные оптимизированные GEMM-ядра (General Matrix Multiplication) выдают 80–90%.

Проблема просадки производительности была в появлении принципиально новых возможностей H100, которые алгоритм FA не использовал для оптимизации накладных расходов. Как именно третья версия FA оптимизировала работу — мы и разберём ниже 🙂

Новинки Hopper-архитектуры

Улучшения FA-3 зиждятся на двух ключевых нововведения архитектуры Hopper:

Tensor Memory Accelerator (TMA) — отдельный аппаратный блок для перемещения данных между глобальной памятью и Shared Memory. Раньше потоки сами занимались копированием данных, отвлекаясь от вычислений.

Асинхронные Tensor Cores — матричные умножения, запускаемые асинхронно, без ожидания их завершения. Это даёт возможность перекрывать разные операции во времени.

FA-2 ни одной такой возможности не задействовала, поэтому на новом железе значительная часть его мощности простаивала.

Кратко про архитектуру GPU можно прочитать здесь.

Ключевая идея FlashAttention-3

Основным направлением улучшений стало использование асинхронности железа Hopper, чтобы разные типы операций выполнялись одновременно, а не по очереди. Механизм attention при этом остаётся неизменным.

Warp-специализация

Вторая версия внутри блока делала всё последовательно: загружала тайл → вычисляла → загружала следующий → вычисляла. Третья версия разделила варпы (минимальные единицы выполнения GPU) блока на роли. Одни варпы («producer») занимаются только загрузкой данных через TMA, другие («consumer») — только вычислениями.

Пока consumer-варпы считывают текущий тайл, producer-варпы уже подгружают следующий. Ожидание данных исчезает, и их загрузка полностью перекрывается вычислениями.

Чередование matmul и softmax

Видеокарта H100 имеет около 989 TFLOPS для FP16 матричного умножения, но лишь около 3.9 TFLOPS для специальных функций. А softmax нужны вычисления экспоненты — специальных функций. Это приводило к простаиванию основных мощностей видеокарты во время просчёта экспоненты.

FA-3 перекрывает операции таким образом: пока тензорные ядра считают следующий блок матмула, отдельные потоки параллельно вычисляют softmax предыдущего блока. В результате медленный блок специальных функций перестаёт быть «узким местом».

Поддержка FP8 с контролем ошибок

Для FP8 FA-3 использует incoherent processing: перед квантованием данные проходят дешёвые преобразования, которые распределяют выбросы по координатам. Это уменьшает ошибку FP8-квантования и делает её менее сконцентрированной в отдельных элементах.

В результате FP8-версия FA-3 показывает в 2.6× меньшую численную ошибку, чем наивная FP8-реализация attention.

Результаты

Благодаря всем этим изменениям FA-3 достигает 740 TFLOPs/s с fp16 (75% утилизации H100) и 1.2 PFLOPs/s с FP8, что даёт ускорение в 1.5–2× относительно FA-2 на той же карте.

Рисунок 1. Масштаб повышения скорости работы forward прохода FA-3 (fp16) на новых H100 вычислителях

Ключевая идея FlashAttention-4

С выходом архитектуры Blackwell (B200, GB200) произошёл значительный рост тензорных ядер и меньший — пропускной способности Shared Memory и других функциональных блоков. Непропорциональность создаёт асимметрию, которая напрямую влияет на работу attention на GPU.

Ситуация напоминала появление H100: алгоритм, качественно работавший на более ранних GPU, на новых видеокартах работает недостаточно оптимально.

Переработанные расписания GPU-ядер

Blackwell поддерживает более глубокую асинхронность матричных умножений, чем Hopper. FA-4 перепроектировала расписания так, чтобы операции никогда не ждали друг друга, и увеличила размеры тайлов. Это значительно снизило относительные накладные расходы на управление и синхронизацию.

Программная эмуляция экспоненты

Поскольку аппаратный блок экспоненты не масштабировался пропорционально тензорным ядрам, FA-4 использует приближённую программную аппроксимацию экспоненты, вычисляемую через быстрые матричные / арифметические блоки. Это снижает зависимость softmax от медленного блока специальных функций.

Режим 2-CTA MMA

Ещё одна оптимизация FA-4 связана с новым режимом матричных умножений на архитектуре Blackwell — 2-CTA MMA.

CTA расшифровывается как Cooperative Thread Array. В терминах CUDA это примерно то же самое, что и thread block: группа потоков, которые выполняются вместе на одном SM, могут синхронизироваться между собой и обмениваться данными через Shared Memory.

MMA — Matrix Multiply-Accumulate, операция матричного умножения с накоплением:

\(D=A×B+C\)

Именно такие операции выполняют Tensor Cores. Attention в значительной степени состоит из MMA: сначала считается произведение \(QK^T\), после softmax — произведение attention weights на V.

В ранних версиях одна CTA обычно сама загружала нужные тайлы данных, выполняла MMA и сохраняла результат. На Blackwell появился режим 2-CTA MMA, где две CTA совместно участвуют в одной матричной операции. Это позволяет лучше переиспользовать загруженные данные и уменьшать лишний обмен через Shared Memory.

Для FlashAttention это особенно важно, потому что attention — не равен чистому GEMM. Между матричными умножениями есть softmax, нормализация, маскирование и дополнительные редукции. Поэтому накладные расходы на движение данных и синхронизацию становятся заметными.

В FA-4 режим 2-CTA MMA помогает снизить давление на Shared Memory и эффективнее загрузить Tensor Cores. В backward pass это также уменьшает накладные расходы, связанные с редукциями и атомарными обновлениями, что нужно при длинных последовательностях.

Новый фреймворк CuTe-DSL

Новое инженерное решение: FA-4 написана не на C++ с шаблонами (как 1 и 3), а на CuTe-DSL — Python-встраиваемом языке для описания GPU-ядер. Это даёт ускорение компиляции в 20–30× при сохранении полного контроля над железом и делает реализацию значительно более доступной для дальнейшей модификации.

Результаты

С помощью этих изменений FA-4 достигает до 1613 TFLOPs/s на B200 с BF16 (71% утилизации) — в 1.3× быстрее cuDNN и в 2.7× быстрее реализации на Triton на той же карте.

Рисунок 2. Масштаб повышения скорости работы forward прохода FA-4 (fp16) на B200 вычислителях
4 месяца
Large Language Models

Разобраться с агентами и LLM можно на нашем курсе! Для тех, кто знаком с DL и Pytorch: научитесь использовать LLM в приложениях: обучать, деплоить, ускорять и многое другое.

0/0

Телеграм-канал

DeepSchool

Короткие посты по теории ML/DL, полезные
библиотеки и фреймворки, вопросы с собеседований
и советы, которые помогут в работе

Открыть Телеграм

Увидели ошибку?

Напишите нам в Telegram!