Эпоха 6 · Генеративка и системы · 2022

48 FlashAttention

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness · Dao, Fu, Ermon, Rudra & Ré · Stanford · NeurIPS
🟧 оригинал выборочно~2 чоригинал ↗
Суть за 20 секунд. ТОЧНОЕ внимание, минимизирующее чтения/записи между медленной (HBM) и быстрой (SRAM) памятью GPU через tiling. Узкое место attention — пропускная способность памяти, не FLOPs. Результат тот же, но в разы быстрее и память O(N) вместо O(N²).

Контекст

Attention в Transformer (#32) стоит O(N²) по памяти/времени для длины N — и упирается не туда, куда думали.

Идея и механизм

Инсайт: стандартная реализация материализует полную матрицу оценок N×N в медленной памяти GPU (HBM), а узкое место современных GPU — не FLOPs, а пропускная способность памяти (трафик HBM↔SRAM). Решение: IO-aware точный attention через tiling — считаем блоками, держим промежуточное в быстром SRAM, никогда не материализуя полную матрицу.

алгоритмы · системы Tiling и online softmax: почему память O(N)

Память GPU иерархична: HBM большая, но медленная; SRAM крошечная, но очень быстрая. Наивный attention пишет в HBM матрицу N×N — это O(N²) трафика, и именно он тормозит, а не умножения.

Tiling. Разбиваем Q, K, V на блоки, грузим их в SRAM и считаем attention по кусочкам. Проблема: softmax требует нормировки по всей строке, а мы видим лишь блок. Решает online softmax — храним бегущие максимум m и сумму ℓ и пересчитываем накопленный результат при каждом новом блоке (численно устойчиво):

m ← max(m, mblock),   ℓ, O пересчитываются с новым m

Полная матрица N×N никогда не лежит в HBM — память O(N). Результат побитово тот же (exact), но трафик памяти резко падает → ускорение в разы. Урок: на современном железе оптимизируй движение данных, а не арифметику.

Backward — тоже без N×N (вторая половина идеи). При обучении обычный attention хранит матрицу N×N с forward, чтобы посчитать градиенты. FlashAttention её НЕ хранит: в backward он пересчитывает нужные блоки оценок заново из Q, K, V и сохранённых нормировщиков (m, ℓ). Это рематериализация — классический размен «лишние FLOPs ради памяти»: чуть больше счёта в обмен на память O(N) вместо O(N²). Именно это делает эффективным не только инференс, но и ОБУЧЕНИЕ на длинных контекстах.

PyTorch FlashAttention под капотом sdpa
import torch.nn.functional as F
# PyTorch выбирает FlashAttention-ядро автоматически:
out = F.scaled_dot_product_attention(Q, K, V)
# результат идентичен softmax(QKᵀ/√d)·V, но память O(N), не O(N²),
# потому что полная матрица оценок не материализуется в HBM
HBM (медленно, большая) наивно: матрица N×N SRAMбыстрая блокипо кусочкам,без N×N в HBM
Вместо записи полной матрицы N×N в медленную HBM — счёт блоками в быстрой SRAM. Тот же результат, но трафик памяти и потолок O(N).
Аналогия. Готовить большой заказ, бегая за каждым ингредиентом в дальний склад (HBM), — медленно из-за беготни, а не готовки. FlashAttention держит нужные продукты на столе под рукой (SRAM) и обрабатывает заказ небольшими порциями, не раскладывая весь склад на кухне. Узкое место было не в скорости рук, а в дороге до склада.

Почему это важно

Сделало длинные контексты практичными, ускорило обучение/инференс всех трансформеров; де-факто стандарт — встроено в PyTorch, дефолт в vLLM, есть в HF. Урок шире: на современном железе оптимизируй ДВИЖЕНИЕ ДАННЫХ по иерархии памяти, а не только арифметику.

Связи

← ускоряет32. Transformer

FlashAttention не меняет математику attention (#32) — он считает то же самое быстрее и экономнее по памяти. Квадратичная стоимость, которая ограничивала длину контекста трансформеров, во многом снята именно здесь.

→ используется в49. vLLM / PagedAttention

vLLM строит эффективный сёрвинг поверх FlashAttention-ядер. Два уровня оптимизации инференса: FlashAttention — внутри одного attention, PagedAttention — в управлении памятью между запросами.

↔ контраст11. LSTM

Любопытный разворот истории: рекуррентность (LSTM) убрали ради параллелизма, но получили квадратичную стоимость attention. FlashAttention возвращает эффективность, не возвращая рекуррентность — оптимизируя память, а не меняя архитектуру.

Вопросы пытливого ума

Если FlashAttention точный, почему вообще существуют приближённые attention?

Они атакуют другое: приближённые методы (Linformer, Performer, sparse) снижают асимптотику с O(N²) до линейной по числу операций, что важно на очень длинных последовательностях. FlashAttention оставляет O(N²) FLOPs, но убирает квадратичную память и трафик. Для большинства практических длин FlashAttention выигрывает без потери точности; приближения нужны там, где даже N² FLOPs неподъёмны.

Почему «memory-bound, не compute-bound» — это вообще про что?

У GPU соотношение «арифметика / пропускная способность памяти» очень высокое: посчитать дешевле, чем подвезти данные. Если операция мало считает на каждый загруженный байт (низкая arithmetic intensity), её скорость определяет память, а не АЛУ. Attention с его чтениями/записями N×N — ровно такой случай. Понимать, упирается ли ядро в compute или в память — базовый навык GPU-оптимизации.

Online softmax звучит хитро — зачем он, нельзя ли просто посчитать softmax потом?

Чтобы посчитать softmax обычно, нужна вся строка сразу (для нормировки и стабильного вычитания максимума) — а её-то мы и не хотим материализовать. Online softmax поддерживает бегущие максимум и сумму и корректно «дореживает» уже накопленный результат при каждом новом блоке. Это та деталь, без которой tiling сломал бы численную устойчивость softmax.

Что читать в оригинале

Читать ключевое — IO-aware постановка и tiling/online-softmax; детали CUDA-ядра можно пропустить, если не пишешь kernels.