Ну чё, малютки, вопрос на миллион: почему attention, у которого FLOPs растут всего лишь квадратично (то есть точно так же, как у любого другого механизма с матрицей N×N), на практике оказывается самым тормозным местом в трансформере? Ответ не в арифметике. GPU считает быстро, а вот таскать данные туда-обратно через память не успевает. Naive-реализация attention гоняет через VRAM полную матрицу очков внимания четырежды за один forward pass, и именно это её топит. FlashAttention не меняет формулу attention ни на йоту: она считает ровно то же самое, что и обычный attention. Меняется только то, в каком порядке и через какую память идут вычисления.
Naive attention четырежды гоняет полную N×N-матрицу через медленную HBM-память GPU: записал очки, прочитал для softmax, записал вероятности, прочитал для умножения на V. FlashAttention считает всё блоками, прямо в быстрой on-chip SRAM, и никогда не материализует полную матрицу в HBM целиком. Результат математически точный: это не аппроксимация, а другой порядок вычислений с online-softmax поверх.
Тот же тезис уже звучал в статье про KV-кэш (decode упирается в bandwidth, не в FLOPs) и в статье про спекулятивный декодинг (GPU простаивает, пока ждёт память). FlashAttention закрывает третий угол той же истории: сам forward pass attention.
Иерархия памяти GPU
Чтобы понять, откуда берётся проблема, надо на секунду вспомнить, что GPU — это не один кусок памяти, а несколько уровней с очень разным балансом «объём vs скорость». Прямо рядом с вычислительными блоками (Streaming Multiprocessor) сидит крошечная, но очень быстрая SRAM. Основная память (та самая VRAM, о которой думаешь, когда видишь OOM) — это HBM: объёма много, но по скорости она на порядок медленнее SRAM.
На NVIDIA A100 счёт такой: SRAM — это 192 КБ на каждый из 108 SM (совокупно ~20 МБ) при пропускной способности порядка 19 ТБ/с. HBM — это уже 40–80 ГБ, но пропускная способность падает до 1.5–2 ТБ/с. То есть SRAM быстрее HBM примерно на порядок, а по объёму — в тысячи раз меньше. Это фундаментальный компромисс архитектуры: быстрая память физически не может быть большой. Ключевой вопрос для производительности attention — сколько байт и сколько раз ты гоняешь через медленный уровень.
Куда утекает трафик: naive attention
Смотри, что делает наивная реализация attention на входе Q, K, V размера N×D. Три шага, и на каждом — полный проход через HBM:
- S = QKᵀ — считаем матрицу очков размера N×N и пишем её в HBM.
- P = softmax(S) — читаем S обратно из HBM, считаем softmax по строкам, пишем результат P (тоже N×N) обратно в HBM.
- O = P·V — читаем P из HBM ещё раз, умножаем на V, получаем итоговый выход.
Итого — два полных N×N-массива, каждый из которых прошёл через HBM не один раз, а дважды: запись S, чтение S, запись P, чтение P. И это всё до того, как ты вообще получил результат внимания для одного слоя одной головы.
Формально и naive, и FlashAttention делают O(N²) арифметики. Асимптотика вычислений не меняется. Но у naive-версии ещё и O(N²) трафика через HBM, причём с немаленькой константой (несколько полных проходов туда-обратно). А пропускная способность HBM на порядок ниже, чем скорость, с которой тензорные ядра готовы жрать числа. GPU в этой схеме большую часть времени не считает. Он ждёт, пока данные доедут по шине.
Тайлинг: считаем блоками, а не всё сразу
Идея FlashAttention обманчиво простая: если полная N×N-матрица не помещается в быструю SRAM, не нужно её туда и пытаться засунуть целиком. Вместо этого Q, K и V режутся на блоки, и обработка идёт блок за блоком: загрузил в SRAM кусочек K и V, посчитал для него частичный результат, накопил в выходе и выбросил, не дожидаясь, пока соберётся полная матрица очков.
Проблема в том, что softmax по определению — операция по всей строке: чтобы нормализовать, нужен знаменатель, который зависит от всех значений сразу. Как это сделать блоками, если ты ещё не видел все блоки? Именно для этого нужен online softmax. Компонент выше как раз проигрывает его вживую на одном query-ряду и 4 блоках по 4 ключа.
Механика по шагам компонента: на каждом новом блоке считается локальный максимум очков и сравнивается с бегущим максимумом m. Если новый блок принёс более крупное значение, максимум обновляется. Разница между старым и новым максимумом даёт коэффициент коррекции (correction), на который домножается уже накопленный результат: так учитывается, что раньше экспоненты считались относительно устаревшего максимума. Затем к бегущей сумме l прибавляется вклад нового блока, тоже с той же коррекцией. К последнему, четвёртому блоку бегущий максимум перестаёт меняться (correction = 1.0), и накопленный числитель, делённый на бегущую сумму l, даёт точно тот же результат, что обычный softmax по всем 16 очкам сразу за один проход. Полная строка внимания при этом ни разу не существовала в памяти целиком.
Формулы online softmax
Раз уж речь зашла о механике, вот сами формулы пересчёта на каждом блоке (в этом блоге нет рендера LaTeX, поэтому просто как код, без всякой магии):
m_new = max(m_old, m_block)
correction = exp(m_old - m_new)
l_new = correction * l_old + sum(exp(scores_block - m_new))
O_new = correction * O_old + exp(scores_block - m_new) @ V_block
m — бегущий максимум очков, нужен исключительно для численной стабильности экспоненты (без вычитания максимума exp() от больших чисел улетает в overflow). l — бегущая сумма знаменателя softmax. O — бегущий, ещё не нормализованный числитель выхода. Как только очередной блок приносит новый максимум, всё накопленное до этого домножается на correction: это ровно тот механизм, который в компоненте выше был подписан как «коэффициент коррекции». После последнего блока делишь O на l и получаешь тот же результат, что дал бы честный softmax по полной строке.
Backward pass: платим вычислениями, а не памятью
С обратным проходом та же логика, но вывернутая наизнанку. Для backward нужны матрицы S и P (те самые, которые FlashAttention принципиально не сохраняет в HBM после forward pass). Вместо этого он их просто пересчитывает заново, блок за блоком, прямо во время backward, используя те же Q, K, V и уже посчитанный выходной градиент. Это дороже по арифметике: часть forward-вычислений выполняется дважды. Но раз GPU всё равно простаивает в ожидании памяти, а не упирается в вычислительный потолок, лишние FLOPs почти бесплатны, а вот трафик через HBM остаётся низким на обоих проходах. Классический инженерный размен: чуть больше компьюта в обмен на память, которую и так неоткуда взять.
От v1 к v3: что менялось дальше
FlashAttention v1 (Dao et al., 2022) — тот самый тайлинг и online softmax, описанные выше. Уже он один убрал материализацию полной N×N-матрицы в HBM.
FlashAttention-2 сфокусировался на том, чтобы полнее занять GPU параллелизмом: v1 распараллеливал работу в основном по батчу и головам внимания, v2 добавил параллелизм ещё и по самой длине последовательности: это важно на длинном контексте с маленьким батчем, где иначе часть SM простаивала бы без работы. Заодно сократили долю нематричных (non-matmul) операций: они дешевле по FLOPs, но на GPU относительно дороги, потому что тензорные ядра заточены именно под матричные умножения.
FlashAttention-3 — заточка под архитектуру Hopper (H100): warp-specialization (разные группы потоков внутри SM параллельно занимаются разными частями вычислений: умножением и вспомогательными операциями, а не строгой последовательностью, как раньше) и поддержка более низкой точности FP8. Оба изменения бьют в ту же цель: полнее утилизировать конкретное железо нового поколения, не трогая математику самого attention.
Сколько это реально экономит
Квадратичный рост на графике выше — это точный расчёт по формуле (2 матрицы × N² × 2 байта в fp16), а не бенчмарк: он показывает, сколько байт naive-подход обязан прогнать через HBM только на промежуточные S и P, и как быстро это число взлетает с длиной контекста, пока у FlashAttention этот же трафик держится на нуле. А вот цифры про ускорение обучения — это уже цитата, а не что-то посчитанное локально: по данным авторов FlashAttention, 3× ускорение end-to-end обучения на GPT-2 (seq len 1K) и до 20× экономии памяти на длинных последовательностях при том же бюджете VRAM (источник: Dao-AILab/flash-attention). Разница между «посчитано точно по формуле» и «взято из чужого бенчмарка» здесь принципиальна, не путай одно с другим.
SRAM на A100 быстрее HBM на порядок, но её всего ~20 МБ против 40–80 ГБ. FlashAttention намеренно гоняет данные через быстрый, но крошечный уровень блоками, а не всё разом через медленный.
Naive attention пишет и читает две полные N×N-матрицы через HBM. FlashAttention никогда не материализует их там целиком: S и P живут только внутри SRAM, блок за блоком.
v1 убрал материализацию S/P в HBM. v2 добавил параллелизм по длине последовательности. v3 — warp-specialization и FP8 под Hopper. Математика attention не менялась ни разу.
Код: от игрушечного NumPy до реального PyTorch
Ниже — минимальная реализация обоих подходов на NumPy, адаптированная из скрипта, которым я проверял всю механику выше локально: 16 «токенов», размерность головы 8, блоки по 4. Наивный вариант считает честный softmax по всей строке за один проход; тайловый — блоками, с online softmax из раздела про формулы. Максимальное расхождение между ними на выходе — 4.44e-16, это машинный ноль, а не «почти совпадает»:
import numpy as np
np.random.seed(42)
N, D, BLOCK = 16, 8, 4
n_blocks = N // BLOCK
Q = np.random.randn(N, D).astype(np.float64)
K = np.random.randn(N, D).astype(np.float64)
V = np.random.randn(N, D).astype(np.float64)
scale = 1.0 / np.sqrt(D)
# ---- naive full attention ----
S = (Q @ K.T) * scale
S_max = S.max(axis=1, keepdims=True)
P = np.exp(S - S_max)
P_norm = P / P.sum(axis=1, keepdims=True)
O_naive = P_norm @ V
# ---- tiled / online-softmax attention ----
O_tiled = np.zeros((N, D))
m = np.full((N, 1), -np.inf) # бегущий максимум
l = np.zeros((N, 1)) # бегущая сумма exp
for j in range(n_blocks):
Kj, Vj = K[j*BLOCK:(j+1)*BLOCK], V[j*BLOCK:(j+1)*BLOCK]
Sij = (Q @ Kj.T) * scale
m_new = np.maximum(m, Sij.max(axis=1, keepdims=True))
P_ij = np.exp(Sij - m_new)
correction = np.where(np.isneginf(m), 0.0, np.exp(m - m_new))
l = correction * l + P_ij.sum(axis=1, keepdims=True)
O_tiled = correction * O_tiled + P_ij @ Vj
m = m_new
O_tiled_final = O_tiled / l
print("max abs diff:", np.max(np.abs(O_naive - O_tiled_final))) # 4.440892098500626e-16
assert np.allclose(O_naive, O_tiled_final, atol=1e-10), "tiled attention does not match naive!"
А вот как это выглядит в реальном коде — через встроенный в PyTorch scaled_dot_product_attention, который сам решает, использовать ли FlashAttention-совместимый backend:
import torch
import torch.nn.functional as F
# Реальный API, синтаксис по официальной документации PyTorch.
# Не гонялось локально при написании статьи — на этой машине сейчас нет
# рабочего CUDA-окружения (driver/library version mismatch).
out = F.scaled_dot_product_attention(query, key, value, is_causal=True)
# PyTorch сам выбирает FlashAttention-совместимый backend, если условия подходят
# (см. torch.backends.cuda.sdp_kernel для явного выбора backend'а).
Где встретишь на практике
Почти все современные инференс-движки используют FlashAttention (или его производные) под капотом. Я разбирал их в гайде по инференс-движкам. Не путай с PagedAttention: та механика — про то, как эффективно хранить KV-кэш в памяти (виртуальная память для внимания), а FlashAttention — про то, как эффективно этот кэш считать. Задачи разные, но на практике их обычно используют вместе.
TL;DR
FlashAttention — хороший пример того, что ускорение необязательно требует жертвовать точностью: иногда всё, что нужно, — это перестать таскать данные туда, куда их не нужно было таскать. 🫡

