ЯдроКодаподготовка к экзаменам
Научная библиотека

Загружаем научный разбор

Подготавливаем текст, источники и редакционные примечания без изменения разметки страницы.

Каталог статейМатериал и источники

FlashAttention: точный attention без хранения полной матрицы

Автор: · Обновлено

Как 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, и какой эксперимент способен подтвердить выигрыш конкретной реализации.

Что относится к исследованию 2022 года

Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra и Christopher Ré представили FlashAttention на NeurIPS 2022. Алгоритм разбивает вычисления на блоки и сокращает обмен между большой памятью GPU, HBM, и быстрой памятью на чипе, SRAM. Для плотного варианта вычисляется обычный точный attention; в работе отдельно рассматривается блочно-разреженное расширение. При обратном проходе часть промежуточных значений вычисляют повторно вместо их длительного хранения. В статье, например, сообщается трёхкратное ускорение обучения GPT-2 при длине 1024 в указанной авторами конфигурации. Это результат определённого сравнения, не обещание такого ускорения для любого GPU или сервиса. Первичная публикация.

Дальше построим собственное маленькое доказательство вычислительной идеи. Оно не повторяет CUDA-код статьи и не претендует на производительность GPU. Нам нужен алгоритм, правильность которого можно проверить на ноутбуке без установки библиотек.

Как объединять частичные softmax

Для одной строки оценок обозначим её элементы через 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 из функции и найдите самый короткий пример, на котором ответ станет неверным. Затем восстановите формулу и переставьте блоки местами. В пределах погрешности результат должен сохраниться. Если получен выигрыш в скорости, объясните отдельно, связан ли он с меньшим обменом с памятью, другим числом операций или изменённой задачей.

Источники

Формат и права

Формат
Авторский разбор

Атрибуция

Самостоятельный русскоязычный разбор ЯдроКода. Результаты исследований отделены от авторских учебных данных, вычислительных опытов и выводов. Материал не является переводом или перепечаткой.

Код, данные и иллюстрации

Учебные данные, расчёты и программные примеры созданы для этой публикации. Изображения, таблицы и программный код первоисточников не воспроизводятся.