Внимание в LLM · часть 5

FlashAttention

KV cache убрал повторные проекции во время генерации. Но при обработке целого промпта и во время обучения точный attention по-прежнему работает с большой матрицей оценок. FlashAttention вычисляет то же самое внимание, не сохраняя матрицу N×N целиком. Секрет — в online softmax (потоковом вычислении softmax).

Главное ограничение — обмен с памятью

Внимание для последовательности длины N строит матрицу оценок QK размера N×N, применяет к ней softmax и умножает результат на V. Число операций растёт как N2, но вычислительная сложность — не единственная проблема. Квадратную матрицу нужно записать в память и прочитать обратно.

У GPU есть небольшая быстрая память рядом с вычислителями (SRAM) и большая, но более медленная память (HBM). Прямая реализация многократно переносит матрицу N×N между вычислителями и HBM: записывает после QK, читает для softmax, снова записывает. При N=8000 это 64 миллиона чисел для каждой головы в каждом слое. Скорость начинает определяться обменом с памятью: внимание становится memory-bound (ограниченным пропускной способностью памяти).

Увеличьте длину с 2K до 8K — в четыре раза. Во сколько раз должны вырасти число оценок и память матрицы? Проверьте ответ, затем повторите рассуждение для следующего доступного размера.

Рост квадратной матрицы

N = 2 048
Длина последовательности
Матрица содержит 4 194 304 оценок и занимает 8 МиБ в выбранном формате.
N = 2 0484 194 3048 МиБ
N = 8 19267 108 864128 МиБ
N = 32 7681 073 741 8242.0 ГиБ
N = 131 07217 179 869 18432.0 ГиБ
O(N2)FlashAttention сокращает хранение и перенос, но не число вычисляемых оценок
Квадратичный рост промежуточной матрицыПри увеличении длины в четыре раза число оценок и память полной матрицы возрастают в шестнадцать раз. Показан объём одной матрицы оценок одной головы, а не вся память модели.

Вычисление softmax по блокам

Обычный softmax вычитает максимум (ради устойчивости) и делит на сумму экспонент — для этого будто бы нужно увидеть всю строку. Но максимум и сумму можно накапливать потоково, блок за блоком. Достаточно поддерживать два скаляра и один вектор:

  • mбегущий максимум уже просмотренных оценок;
  • бегущая суммаesm;
  • oбегущий выходesmv (ненормированный).

Если в новом блоке находится более высокий максимум, прежние накопления нужно привести к новому масштабу. Для этого достаточно одного множителя α=emoldmnew:

mmax(m,mблок),α=emoldm,
α+jблокesjm
oαo+jблокesjmvj

Итоговый результат равен o/; полная строка весов при этом не хранится.

Блочная обработка

Разбиение на тайлы (tiling) позволяет одному запросу получать ключи блоками. При увеличении текущего максимума множитель α перенормирует накопленные ранее величины. Размер тайла и запрос можно выбрать до начала опыта; их смена начинает обработку заново.

При исходном размере тайла обработайте все четыре блока по одному. Перед каждым шагом решите, потребуется ли пересчитать старые вклады из-за нового максимума. Затем выберите другой размер тайла и проверьте, изменился ли итоговый результат.

Online softmax по тайлам

тайл 0/4
Размер тайла
Токен запроса
2
кошка0.06ещё не прочитан
лакает1.40ещё не прочитан
молоко2.18ещё не прочитан
потому0.00ещё не прочитан
что0.00ещё не прочитан
она0.06ещё не прочитан
голодна0.18ещё не прочитан
До первого тайлаВ памяти не более 2 оценокПолная строка весов не материализуется.
Точный softmax без полной строки в памятиПересчёт накопленных величин при смене максимума сохраняет результат softmax с точностью вычислений. Размер тайла меняет порядок работы и объём временной памяти.

Обратите внимание: в этой демонстрации в памяти в каждый момент находятся лишь текущий тайл, два скаляра и вектор-аккумулятор — но не матрица N×N. Рабочие данные помещаются в быстрой SRAM, а обмен с HBM резко сокращается.

Путь данных

Теперь сопоставим два пути. Прямая реализация сохраняет квадратную матрицу в HBM, затем читает её для softmax и следующего умножения. FlashAttention переносит в SRAM только текущие блоки Qi,Kj,Vj, вычисляет тайл S=QiKj и обновляет аккумулятор строки.

Пройдите обработку от первого до последнего тайла. Перед каждым шагом различайте уже выполненную работу и данные, которые ещё нужно хранить. Должны ли эти два числа расти одинаково?

Два пути через память

тайл 0/16
Вычислено 0 из 64 оценок; в SRAM одновременно 4, в HBM записано 0.
тайлы 1–2ожидают8 оценок
тайлы 3–4ожидают8 оценок
тайлы 5–6ожидают8 оценок
тайлы 7–8ожидают8 оценок
тайлы 9–10ожидают8 оценок
тайлы 11–12ожидают8 оценок
тайлы 13–14ожидают8 оценок
тайлы 15–16ожидают8 оценок
0/8 строк результатаПолная матрица S не записывается в HBM
Два маршрута данных через память GPUЧисло вычисленных оценок растёт до полного квадрата. Временное хранение ограничено тайлом; полная матрица оценок в HBM не записывается, хотя обмен другими данными сохраняется.

В этом и состоит экономия: прямая реализация переносит через HBM всю квадратную матрицу S, а FlashAttention держит в SRAM лишь текущий тайл и записывает обратно только результат. Меньше обмена с медленной памятью — выше скорость и нет квадратного расхода памяти на матрицу внимания.

Как менялось следующее ограничение

Базовая математика внимания оставалась прежней. Каждое поколение устраняло одно ограничение и тем самым делало заметнее следующее: сначала обмен с HBM, затем загрузку GPU, наконец — паузы между softmax и матричными умножениями.

Сравните версии по порядку, начиная с FA-1. Для каждого перехода назовите устранённое ограничение и оставшуюся проблему. Проверьте также оборудование: можно ли напрямую сравнивать показанные ускорения между разными устройствами?

Эволюция FlashAttention

FA-1 · 2022
Версия FlashAttention
Убрать матрицу N×N из памяти. Tiling, потоковый softmax и пересчёт весов на обратном проходе.
результат 1линейная дополнительная памятьAmpere · A100
результат 2резко меньше обращений к HBMAmpere · A100
результат 3точный алгоритмAmpere · A100
загрузка вычислителей32%
в 2–4 раза быстрее обычного attentionОстаётся: неполная загрузка GPU и много вспомогательных операций
Что менялось между версиями FlashAttentionВерсии последовательно меняют организацию вычислений и обмена данными. Показатели ускорения относятся к своим условиям измерения и не образуют единую шкалу для разных устройств.

FA-1 (2022) — основа алгоритма

Первая версия объединяет всё, что мы разобрали в главах 1–4: тайлы, потоковый softmax и пересчёт весов на обратном проходе. Она убрала матрицу N×N из медленной памяти и резко сократила обмен с HBM. Но GPU всё ещё использовался не полностью. Работу делили только между парами «последовательность × голова», а внутри блока вычислители часто ждали друг друга. Отсюда — второе поколение.

FA-2 (2023) — больше параллелизма, меньше операций

Все три приёма направлены на то, чтобы GPU меньше простаивал. Матричные умножения выполняют Tensor Cores, а экспоненты, деления и поиск максимума выполняются на обычных CUDA-ядрах. Чем меньше вспомогательной работы и ожидания, тем полнее используется быстрая матричная часть GPU.

  • Отложить деление. FA-1 нормировал выход — делил на бегущую сумму — на каждом блоке. В FA-2 деление выполняется один раз в самом конце. Стоимость одного деления невелика, но во внутреннем цикле эта операция повторяется множество раз. Её перенос сокращает долю вспомогательных вычислений.
  • Больше параллельных блоков. FA-1 распределял работу по парам «последовательность × голова». Но при длинном контексте последовательностей мало, а сами они велики — независимых задач не хватает, чтобы загрузить все вычислительные блоки. FA-2 дополнительно делит работу по длине запроса: даже одна длинная последовательность даёт достаточно независимых задач для загрузки GPU.
  • Каждому варпу — свои строки. Внутри блока работу делят варпы (группы потоков). FA-1 делил между ними ключи K: варпы вычисляли части одного результата, а затем объединяли их через разделяемую память и ждали друг друга. FA-2 делит запросы Q — теперь у каждого варпа свои строки выхода целиком, поэтому им реже приходится объединять результаты и ждать друг друга.

FA-3 (2024) — совмещение этапов вычисления

Даже в FA-2 шаги идут по очереди: матричное умножение (Tensor Cores) → softmax (обычные ядра) → следующее матричное умножение. Пока обычные ядра считают экспоненты одного тайла, Tensor Cores ждут — и наоборот. На GPU архитектуры Hopper FA-3 выполняет эти этапы внахлёст: пока для тайла N считается softmax, для тайла N+1 уже идёт матричное умножение. Асинхронные инструкции позволяют одновременно считать и загружать данные, а специализация варпов разделяет эти обязанности между группами потоков. Получается конвейер, в котором ожидание сведено к минимуму.

Второе нововведение — FP8, восьмибитный числовой формат. Такие матричные умножения быстрее, но менее точны, поэтому значения квантуют поблочно, контролируя численную ошибку. В этом режиме производительность на H100 достигает примерно 1,2 PFLOPS — уже с компромиссом между скоростью и численной точностью.

Эта идея получила дальнейшее развитие. Например, Flash-Decoding ускоряет генерацию с большим контекстом, распараллеливая работу по ключам. Основной принцип остаётся прежним: не хранить то, что можно пересчитать, и сокращать простои GPU.