Запись архива

FlashAttention-3: почему внимание стало быстрее на Hopper

Разбираем статью FlashAttention-3 от 11 июля 2024 года: как авторы переписали exact attention под NVIDIA Hopper, зачем нужны TMA, WGMMA и FP8, и что реально означают заявленные 1.5–2.0x ускорения.

Схема работы FlashAttention-3 на Hopper: TMA копирует плитки Q, K и V в shared memory, warpgroups выполняют WGMMA, параллельно считается online softmax, затем результат умножается на V и записывается обратно

FlashAttention-3 — работа Jay Shah и соавторов, впервые выложенная на arXiv 11 июля 2024 года и затем опубликованная в материалах NeurIPS 2024. Она не меняет формулу attention как таковую: авторы переписали exact attention под особенности NVIDIA Hopper, прежде всего H100.

Практический смысл в том, что узкое место здесь ищут уже не только в доступе к памяти, как в ранних версиях FlashAttention, но и в асинхронности конвейера, Tensor Cores и устойчивости FP8. Ниже — разбор именно результатов статьи 2024 года, без подмены их более поздними реализациями.

Коротко

  • FlashAttention-3 — это не новый тип внимания, а новая реализация exact attention под архитектуру Hopper.
  • Ключевые идеи три: асинхронное перекрытие копирования и вычислений, перекрытие GEMM и softmax, а также FP8 с блоковой квантизацией и incoherent processing.
  • В статье и официальных постах авторов заявлено ускорение на H100 на 1.5–2.0x относительно FlashAttention-2 в FP16, до 740 TFLOPs/s и около 75% утилизации теоретического максимума H100.
  • Заявление про FP8 важно не только из-за скорости: авторы также утверждают, что их вариант даёт в 2.6x меньшую численную ошибку, чем baseline FP8 attention.
  • Но scope у работы узкий: основные результаты показаны именно для Hopper/H100, а не для всех GPU и не как гарантированное ускорение полного обучения модели.

Контекст: почему attention стало узким местом

Исходный FlashAttention 2022 года был важен тем, что не приближал внимание, а менял порядок вычислений и работу с памятью. Вместо материализации большой матрицы attention в HBM авторы использовали тайлинг между HBM и on-chip SRAM, из-за чего память по длине последовательности в практическом исполнении сокращалась с квадратичной до линейной, а на нескольких задачах были показаны ускорения уровня 2–4x по wall-clock времени. Эти тезисы повторяются и в статье 2022 года, и в официальном разборе Stanford CRFM.

FlashAttention-2 в 2023 году сместил фокус: уже не столько I/O, сколько параллелизм и разбиение работы между thread blocks и warps. В статье и официальном посте Tri Dao сказано, что FlashAttention-2 был примерно в 2 раза быстрее первой версии и доходил до 225 TFLOPs/s на A100 при обучении GPT-подобных моделей, что авторы описывали как 72% model FLOP utilization.

Но именно здесь появляется мотивация FlashAttention-3. Авторы работы 2024 года прямо пишут, что на H100 FlashAttention-2 использовал лишь около 35% теоретического максимума. Иными словами, после перехода от Ampere к Hopper выяснилось, что старый алгоритмический дизайн плохо использует новые аппаратные возможности, даже если сама идея FlashAttention остаётся правильной.

Метод: как работает FlashAttention-3

1. Асинхронность Hopper как часть алгоритма

Главная смена оптики в FlashAttention-3 такая: attention здесь проектируется уже не под абстрактный GPU, а под конкретную модель исполнения Hopper. В официальном Hopper Tuning Guide NVIDIA описывает TMA (Tensor Memory Accelerator) как более развитый асинхронный механизм копирования между global memory и shared memory. В документации CUTLASS про WGMMA сказано, что Hopper вводит асинхронные warpgroup-инструкции wgmma.mma_async, где одна warpgroup из 128 потоков коллективно исполняет матричное умножение.

FlashAttention-3 строит kernel так, чтобы эти специализированные блоки работали параллельно. Упрощённо: часть warps готовит и подаёт данные через TMA, а часть держит Tensor Cores занятыми через WGMMA. В статье это называется использованием асинхронности Tensor Cores и TMA через warp-specialization.

2. Перекрытие GEMM и softmax

Во второй идее авторы обращают внимание на то, что attention — это не только матричные умножения QK^T и PV, но и online softmax. Если считать softmax строго после завершения GEMM, часть железа простаивает. Поэтому FlashAttention-3 пытается перекрывать эти стадии: пока одни warpgroups считают GEMM, другие выполняют softmax и связанные с ним операции перенормировки.

В блоге авторов на PyTorch эта логика разобрана через два уровня конвейера: меж-warpgroup ping-pong scheduling и внутригрупповое перекрытие GEMM и softmax. В статье те же идеи сформулированы компактнее: interleave block-wise matmul and softmax operations.

3. FP8 без наивной потери точности

Третья идея связана с low precision. По официальным материалам NVIDIA для Hopper, FP8 на Tensor Cores даёт примерно вдвое больший throughput по сравнению с FP16/BF16 на уровне матричных операций. Но на практике простая замена precision быстро упирается в ошибку квантизации, особенно если в активациях есть выбросы.

Чтобы смягчить эту проблему, авторы используют блоковую квантизацию и incoherent processing. Интуиция простая: перед квантизацией признаки в Q и K преобразуются так, чтобы большие по модулю выбросы были «размазаны» по координатам и хуже портили масштабирование. В блоге и статье эта часть реализована через Hadamard transform со случайными знаками.

Результаты: что показала статья FlashAttention-3

Здесь важна оговорка о scope: ключевые числа в статье относятся к kernel-бенчмаркам на Hopper/H100, а не к универсальному ускорению любого LLM-стека. Именно так их и нужно читать.

Что измерялось Что утверждают авторы Как это интерпретировать
Отправная точка: FlashAttention-2 на H100 Около 35% использования теоретического максимума H100. Это аргумент, что у Hopper оставался большой запас, который FA2 не выбирал.
FlashAttention-3 в FP16 Ускорение на 1.5–2.0x относительно FlashAttention-2, до 740 TFLOPs/s и около 75% утилизации H100. Основной выигрыш приходит не из новой математики attention, а из более плотного конвейера под Hopper.
FlashAttention-3 в FP8 Скорость близка к 1.2 PFLOPs, а численная ошибка ниже в 2.6x относительно baseline FP8 attention. FP8 в статье подаётся не просто как «быстрее», а как более аккуратный низкоточный вариант по сравнению с наивным baseline.

Все три строки выше подтверждаются как самой статьёй FlashAttention-3, так и официальными постами авторов для PyTorch и NVIDIA. Это один из редких случаев, когда paper и lab-post почти дословно совпадают по верхнеуровневым числам.

Есть и более детальные, но уже одноисточниковые результаты. В таблице 2 самой статьи авторы приводят абляцию конвейера: полная версия FlashAttention-3 даёт 661 TFLOPs/s, вариант без GEMM-softmax pipelining, но с warp-specialization — 582 TFLOPs/s, а вариант с GEMM-softmax pipelining, но без warp-specialization — 570 TFLOPs/s. Эти конкретные числа мы можем приписать только статье, потому что отдельного независимого первичного источника для таблицы 2 нет.

То же относится к таблице 3 с численной ошибкой. В единственном первоисточнике этих конкретных значений baseline FP8 attention имеет RMSE 2.4e-2, а FlashAttention-3 FP8 — 9.1e-3; для FP16 baseline указан как 3.2e-4, а FlashAttention-2 и FlashAttention-3 — как 1.9e-4. Это полезные цифры, но их надо читать именно как авторское измерение внутри статьи, а не как внешне подтверждённый консенсус.

Интерпретация: что это значит для практики

Наш комментарий

На наш взгляд, FlashAttention-3 важен не только самими числами, а сменой инженерного фокуса. Первый FlashAttention был прежде всего I/O-aware: он доказывал, что exact attention можно сделать быстрым, если правильно обращаться с памятью. FlashAttention-3 показывает следующую ступень: после I/O-оптимизации узким местом становятся уже асинхронность исполнения, загрузка специализированных блоков GPU и численная устойчивость низкой точности.

Для практиков это означает две вещи. Во-первых, ускорение attention всё меньше похоже на «подключил новую библиотеку и получил бесплатные 2x»; всё больше — на со-дизайн kernel-а и конкретной архитектуры GPU. Во-вторых, paper даёт сильный сигнал, что на Hopper ещё оставался существенный запас даже после FlashAttention-2, но не доказывает, что тот же выигрыш автоматически переносится на любой inference-server, любой training stack или любой другой ускоритель.

Ограничения и критика

  • Жёсткая аппаратная привязка. Основной вклад статьи привязан к Hopper, особенно к H100. Если вы работаете на A100, MI300, RTX-картах или TPU, из этой работы нельзя напрямую вывести те же проценты ускорения.
  • Это не асимптотический прорыв. FlashAttention-3 остаётся реализацией exact attention. Он не отменяет квадратичную по длине последовательности вычислительную стоимость самого attention; он уменьшает накладные расходы и лучше использует железо.
  • В центре статьи — kernel-бенчмарки. В paper есть forward/backward benchmarks attention-kernel-ов, но нет столь же развёрнутой end-to-end картины обучения больших моделей, как в некоторых более ранних публикациях семейства FlashAttention. Поэтому итоговое ускорение полного пайплайна может быть заметно меньше, если bottleneck у вас в MLP, коммуникациях, KV-cache или data pipeline.
  • FP8-часть проверена в ограниченном сценарии. Главный численный аргумент про 2.6x меньшую ошибку опирается на baseline FP8 attention и на synthetic setup с редкими выбросами в Q, K, V. Это разумный, но всё же узкий тест, а не широкая независимая валидация на десятках production-моделей.
  • Сложность реализации высока. WGMMA, TMA, регистровое давление, warp-specialization и ручное конвейерирование делают такой kernel существенно труднее в сопровождении и переносе, чем «обычную» fused attention-реализацию.

Заключение

FlashAttention-3 — это не новая формула Transformer и не отказ от exact attention. Это аккуратная перепаковка уже известного attention под архитектуру Hopper, где выигрыш рождается из асинхронного исполнения, более плотного конвейера и более осмысленного обращения с FP8.

Если вы работаете на H100, статья 2024 года показывает, что attention можно подвинуть гораздо ближе к аппаратному потолку, чем это делал FlashAttention-2. Если нет, главный урок всё равно полезен: сегодня производительность LLM всё чаще определяется не только архитектурой модели, но и тем, насколько глубоко kernel понимает конкретное железо.

Источники

FAQ

FlashAttention-3 — это новый тип attention?

Нет. Как и предыдущие версии семейства, это реализация exact attention. Новизна не в новой формуле, а в том, как attention раскладывается на tiles, как конвейерится на Hopper и как использует FP8.

Меняет ли FlashAttention-3 алгоритмическую сложность внимания?

Нет. Он не убирает квадратичную вычислительную стоимость attention по длине последовательности. Основной эффект — в уменьшении I/O-накладных расходов и в лучшей загрузке аппаратных блоков GPU.

Нужен ли Hopper/H100, чтобы получить результаты из статьи?

Для основных чисел статьи — да, потому что они показаны именно на Hopper/H100. Некоторые идеи семейства FlashAttention полезны и шире, но конкретные цифры FlashAttention-3 нельзя честно переносить на другие архитектуры без отдельного бенчмарка.

Это ускорение для обучения или для inference?

И для того, и для другого на уровне attention-kernel-а: статья показывает как forward, так и backward бенчмарки. Но итоговый end-to-end выигрыш зависит от всей системы, поэтому на практике он не обязан совпадать с ускорением самого kernel-а.

Читайте также