Загружаем научный разбор
Подготавливаем текст, источники и редакционные примечания без изменения разметки страницы.
Подготавливаем текст, источники и редакционные примечания без изменения разметки страницы.
Автор: Казачкин Даниил Михайлович · Обновлено
Как FlashAttention сокращает обмен с памятью GPU: собственный блочный softmax на Python, проверка точности и границы вывода об ускорении.
При увеличении контекста модель может исчерпать память ещё до того, как время вычислений станет неприемлемым. Возникает исследовательский вопрос: обязательно ли хранить все промежуточные попарные оценки attention, чтобы получить тот же результат? FlashAttention показывает, что порядок вычислений и движение данных между уровнями памяти сами по себе могут существенно менять стоимость алгоритма.
Значение query, key и value уже разобрано в [уроке о Transformer](/lessons/without-university/generative-ai-foundations/generative-ai-03), а память инференса — в [уроке об обслуживании моделей](/lessons/without-university/generative-ai-foundations/generative-ai-15). Здесь разберём другую задачу: как нормировать веса порциями, не потеряв общий знаменатель softmax, и какой эксперимент способен подтвердить выигрыш конкретной реализации.
Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra и Christopher Ré представили FlashAttention на NeurIPS 2022. Алгоритм разбивает вычисления на блоки и сокращает обмен между большой памятью GPU, HBM, и быстрой памятью на чипе, SRAM. Для плотного варианта вычисляется обычный точный attention; в работе отдельно рассматривается блочно-разреженное расширение. При обратном проходе часть промежуточных значений вычисляют повторно вместо их длительного хранения. В статье, например, сообщается трёхкратное ускорение обучения GPT-2 при длине 1024 в указанной авторами конфигурации. Это результат определённого сравнения, не обещание такого ускорения для любого GPU или сервиса. Первичная публикация.
Дальше построим собственное маленькое доказательство вычислительной идеи. Оно не повторяет CUDA-код статьи и не претендует на производительность GPU. Нам нужен алгоритм, правильность которого можно проверить на ноутбуке без установки библиотек.
Для одной строки оценок обозначим её элементы через s, а значения через v. Требуется взвешенная сумма sum(exp(s) * v) / sum(exp(s)). Прямое вычисление экспонент может переполниться. Вычитание максимума из всех оценок сохраняет отношение и делает экспоненты не больше единицы.
Пусть для уже прочитанной части сохранены максимум m, сумма экспонент l и ненормированный взвешенный вектор a. Новый блок может содержать больший максимум M. Тогда старые l и a нужно умножить на exp(m - M), прежде чем добавлять значения нового блока, рассчитанные относительно M. Это просто приведение двух сумм к общему масштабу. Нормировать результат достаточно в конце: a / l.
Популярная ошибка — сначала независимо получить softmax каждого блока, затем усреднить результаты. Так каждый блок получает одинаковый общий вес независимо от величины своих оценок. Если один блок содержит оценки около 100, а другой около 0, эта процедура теряет решающее различие. Сохранённый знаменатель как раз позволяет не потерять его.
from math import exp, isclose
def reference(scores, values):
maximum = max(scores)
weights = [exp(s - maximum) for s in scores]
return [sum(w * v[j] for w, v in zip(weights, values)) / sum(weights)
for j in range(len(values[0]))]
def blocked(scores, values, width):
maximum = float("-inf")
denominator = 0.0
numerator = [0.0] * len(values[0])
for start in range(0, len(scores), width):
block = scores[start:start + width]
next_maximum = max(maximum, max(block))
rescale = exp(maximum - next_maximum)
denominator *= rescale
numerator = [value * rescale for value in numerator]
for score, vector in zip(block, values[start:start + width]):
weight = exp(score - next_maximum)
denominator += weight
for j, value in enumerate(vector):
numerator[j] += weight * value
maximum = next_maximum
return [value / denominator for value in numerator]
scores = [1000.0, 1001.0, 999.0, 1004.0, 997.0]
values = [[2.0, -1.0], [0.0, 3.0], [4.0, 1.0],
[-2.0, 2.0], [1.0, -4.0]]
for visible in range(1, len(scores) + 1):
# Префикс соответствует уже разрешённым causal-mask позициям строки.
expected = reference(scores[:visible], values[:visible])
for width in (1, 2, 3, 8):
actual = blocked(scores[:visible], values[:visible], width)
assert all(isclose(a, b, rel_tol=1e-12, abs_tol=1e-12)
for a, b in zip(actual, expected))
whole = reference([0.0, 0.0, 100.0, 100.0], [[0.0]] * 2 + [[1.0]] * 2)
wrong = (reference([0.0, 0.0], [[0.0]] * 2)[0]
+ reference([100.0, 100.0], [[1.0]] * 2)[0]) / 2
assert whole[0] > 0.999
assert wrong == 0.5
print("Все размеры блока совпали; неверное усреднение:", wrong)Проверки охватывают неполный последний блок, единственную видимую позицию, разные размеры блока и большие положительные оценки. В Python exp(1004) переполнился бы, но код работает с разностями. Начальный максимум -inf обнуляет вклад ещё не существующей старой суммы. Предполагается, что в строке есть хотя бы один разрешённый элемент; поведение полностью замаскированной строки нужно определять отдельно.
Наш код заранее получает список оценок, поэтому сам по себе не устраняет хранение матрицы во всём Transformer. Он проверяет только операцию объединения. В GPU-реализации оценки нужно получать из блоков Q и K, сразу использовать с V и освобождать локальную память; размещение блоков и параллельное исполнение определяют реальную скорость.
Возьмём собственный расчёт: 8192 токена, одна голова, по два байта на элемент оценки. Одна полная матрица содержит 8192² элементов и занимает 128 MiB. Для 32 голов это уже 4 GiB, если одновременно материализовать такую матрицу для каждой головы. Здесь не учтены градиенты, веса, активации и служебные буферы. Этот расчёт описывает конкретный промежуточный массив, а не полное потребление модели.
Отказ от массива не отменяет сравнение каждой разрешённой пары query и key. У плотного полного attention арифметическая работа остаётся квадратичной по длине. Также никуда не исчезают сами Q, K, V и выход. Поэтому фразу «линейная память» нужно относить к памяти соответствующего вычисления, а не обещать постоянный расход для модели целиком.
Различайте обработку длинного входного запроса и пошаговую генерацию с KV-cache. У них разные формы матриц, объём повторного использования и узкие места. Небольшой выигрыш на одной фазе не опровергает выигрыш на другой. Для продукта нужны раздельные измерения задержки первого токена и последующего декодирования, а не одна средняя цифра на весь запрос.
Зафиксируйте модель GPU, версии библиотек, dtype, длины последовательностей, число голов, размер головы, mask и dropout. Сначала проверьте численную близость выходов и градиентов при допустимом для dtype отклонении. «Точный» означает отсутствие алгоритмического приближения attention, но не побитовое совпадение при другой группировке операций с плавающей точкой.
После прогрева измеряйте завершённые операции с корректной синхронизацией GPU. Отдельно фиксируйте пиковую память и время; перенос данных и компиляция не должны случайно входить только в один вариант сравнения. Обязательно укажите baseline: современная библиотека уже может выбирать эффективное ядро автоматически. Проверяйте фактически запущенную реализацию, иначе сравниваются два имени одной операции.
В репозитории авторов разделены поколения FlashAttention и требования конкретных реализаций. Современный пакет не следует считать неизменным артефактом статьи 2022 года. Номер реализации, оборудование и параметры вызова входят в описание результата так же, как длина контекста.
Самостоятельная задача: удалите rescale из функции и найдите самый короткий пример, на котором ответ станет неверным. Затем восстановите формулу и переставьте блоки местами. В пределах погрешности результат должен сохраниться. Если получен выигрыш в скорости, объясните отдельно, связан ли он с меньшим обменом с памятью, другим числом операций или изменённой задачей.
Самостоятельный русскоязычный разбор ЯдроКода. Результаты исследований отделены от авторских учебных данных, вычислительных опытов и выводов. Материал не является переводом или перепечаткой.
Учебные данные, расчёты и программные примеры созданы для этой публикации. Изображения, таблицы и программный код первоисточников не воспроизводятся.