Слогер Создать блог
AI

Почему Flash Attention быстрее: вся суть в трафике памяти

Обычный attention делает в несколько раз больше обращений к памяти, чем нужно. Flash Attention убирает лишнее и сохраняет точность — разбираем механику.

Flash Attention выполняет ровно те же операции с плавающей запятой, что и обычный attention, и выдаёт тот же результат с точностью до ошибок округления. Он быстрее по одной причине: меньше передаёт данных между чипом и памятью. Вся техника в этом и заключается.

Проблема — трафик, а не арифметика

Если записать attention прямо из определения — для последовательности из S токенов и головы размерности d — получится так: вычислить S×S скоров, записать их в память, прочитать обратно, применить softmax к каждой строке, записать снова, опять прочитать, умножить на значения. Четыре полных путешествия к памяти и обратно. А матрица растёт квадратично от длины последовательности.

Арифметика — O(S²d), и она неизбежна. Трафик — O(S²), и это артефакт способа записи вычисления. У attention при скромной размерности головы низкая арифметическая интенсивность: само по себе умножение дешёвое, а память — дорогая. На практике операция и стоит как трафик.

Почему вообще было написано так? Не из неаккуратности. Каждый шаг — отдельная библиотечная функция: матричное умножение, softmax, ещё одно умножение. Они оптимизированы по отдельности, но каждая обязана читать вход из памяти и писать выход обратно — таков контракт вызова. Неэффективность живёт в стыках между операциями, а не внутри них. Исправить это может только слитая реализация.

Считаем байты

Возьмём типичный случай: один слой, одна голова, длина контекста 8192, значение fp16. Вот что получается для стандартного и слитого вариантов.

ПоказательСтандартный attentionFlash Attention
Трафик на голову≈ 537 МБ (матрица S×S пишется и читается примерно четыре раза)O(S·d): данные Q, K, V и результат, плюс повторное чтение K и V для каждого тайла запросов
Память под матрицу S×SВыделяется полностью: 8192² элементов × 2 байта = 134 МБВ память устройства не пишется вообще
Рост занимаемой памятиКвадратичный от длины контекстаЛинейный

Экономия памяти не менее важна, чем скорость. Стандартный вариант обязан выделить матрицу S×S, поэтому память растёт квадратично. Именно это делало длинный контекст непрактичным ещё до появления таких ядер — независимо от того, сколько времени считалась операция.

Тилирование... и заминка

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

Но есть одно препятствие — softmax. Он нелокальный. Чтобы нормализовать строку, нужна сумма экспонент по всей строке. А для численной стабильности сначала вычитают максимум строки — так что нужен максимум всей строки. Получается, что для нормализации любого элемента нужны все скоры. Тилирование этого не даёт.

Online softmax: как обойти нелокальность

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

  1. Вычислить скоры блока и его локальный максимум.
  2. Пусть m_new = max(m_old, m_block).
  3. Пересчитать бегущую сумму и накопленный выход, умножив на exp(m_old − m_new) — это компенсирует то, что старые результаты нормированы на устаревший максимум.
  4. Добавить вклад текущего блока, нормализованный уже относительно m_new.

Каждая поправка — это скалярное умножение на строку, ничтожное по сравнению с матричной работой. В конце бегущая сумма окажется настоящим знаменателем, а накопленный выход — точным результатом attention.

Тот же принцип используется при вычислении скользящего среднего без хранения потока — только здесь его применили к softmax. Именно этот трюк делает весь подход возможным.

Важно понимать: результат не приближённый. Flash Attention — не аппроксимация в духе разреженных или низкоранговых вариантов. Просто переставлен порядок суммирования и пересчитаны частичные результаты, поэтому последние биты могут отличаться — как при любой перестановке операций с плавающей точкой. Но математическая функция та же.

Обратный проход несёт ту же идею с одним дополнительным нюансом. Градиенты требуют весов внимания, которые специально не сохранялись, поэтому они пересчитываются из сохранённых статистик во время обратного прохода. Небольшая лишняя арифметика вместо больших затрат на память — всегда правильный обмен, когда ядро упирается в память, и паттерн этот повторяется далеко за пределами attention.

Что это изменило на практике

  • Длинный контекст стал вопросом ёмкости памяти. Квадратичный член исчез, ограничением стала KV-кэш, которая растёт линейно. Это качественно другая задача — поэтому свежие работы по длинному контексту упираются в размер кэша.
  • Этап prefill стал заметно дешевле. Именно на префилле раньше вычислялась полная матрица внимания, поэтому экономия приходится на него. Это видно по времени до первого токена при длинных промптах.
  • Декодирование выигрывает иначе. При генерации одного токена есть длинная история и один запрос — матрицы S×S просто нет. Специализированные варианты таких ядер решают параллелизм сканирования ключей и значений.
  • Техника стала необходимостью для переносимости. Ядро пишется вручную под архитектуру, и вопрос «есть ли IO-aware реализация attention для этого бэкенда» — один из самых острых при выборе нестандартной платформы.

Практические выводы

Вот на что стоит ориентироваться, когда вы собираетесь использовать такие ядра:

  • Проверьте, поддерживает ли ваш фреймворк fused attention под вашу видеокарту. Если поддержки нет — придётся либо ставить готовые библиотеки, либо писать кастомный kernel.
  • Не считайте, что Flash Attention приближает математику. Он даёт тот же результат, что и обычный attention, просто быстрее. Если вы получаете иные выводы — смотрите в сторону типов данных и порядка операций, а не аппроксимации.
  • Помните про KV-кэш: когда attention перестаёт быть узким местом по памяти, упирается он. Оптимизация контекста — это уже задача управления кэшем.

Главный урок из всей истории — понятие IO-awareness: считайте не операции, а обращения к памяти. Этот принцип работает не только для attention, и именно он превращает «невозможный» длинный контекст в реальность.

По материалам: ai. Текст переработан редакцией Слогера.

← На главную

Рекламное место — Конец поста
Реклама · Слогер

Комментарии (0)

Войдите, чтобы комментировать.

Пока нет комментариев. Будьте первым.