Архитектура трансформеров · Эффективное внимание

Разреженное и линейное внимание, варианты голов, FlashAttention и эволюция активаций

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

Проблема и разреженное/линейное внимание

Классическое self-attention считает матрицу n×n для каждой пары токенов — O(n²) по времени и памяти. При росте контекста до десятков и сотен тысяч токенов это становится главным узким местом, и вокруг него выросло целое семейство приближённых схем.

Во сколько раз растёт объём вычислений при переходе с 1K к 64K токенов (схематично)
×4096
Полное внимание
O(n²)
×103
Reformer
O(n·log n)
×64
BigBird / Linformer /
Performer / cosFormer · O(n)
Схематичная иллюстрация порядка величин (не бенчмарк): при 64-кратном росте длины последовательности квадратичная сложность даёт 4096-кратный рост вычислений, а линейные схемы — ровно 64-кратный.
BigBirdразреженноеZaheer et al., 2020

Комбинирует три типа связей вместо полного графа внимания: локальное окно (каждый токен видит соседей), случайные связи (несколько случайных токенов на каждый запрос) и глобальные токены (небольшой набор токенов, которые видят всех и которых видят все — аналог CLS). Число связей на токен константно, поэтому сложность падает до O(n).

BigBird: окно + глобальные + случайные
Полное внимание — для сравнения
глобальные токеныокнослучайные связи

Авторы показали теоретически, что при достаточном числе глобальных токенов такая разреженная схема остаётся универсальным аппроксиматором последовательностей — то есть не теряет выразительности полного внимания, несмотря на линейную сложность. Хорошо работает на задачах с длинными документами: QA, суммаризация, геномные последовательности.

Reformerразреженное · O(n log n)Kitaev, Kaiser, Levskaya, 2020

Строится на LSH-attention (locality-sensitive hashing): вместо сравнения каждого запроса со всеми ключами, запросы и ключи хэшируются в бакеты так, что похожие векторы с высокой вероятностью попадают в один бакет. Внимание считается только внутри бакета — токены сортируются по хэшу и разбиваются на чанки, что даёт O(n log n).

Блоки-бакеты после LSH-сортировки

Второе ключевое решение — обратимые (reversible) остаточные слои, заимствованные из RevNet: активации каждого слоя восстанавливаются из следующего слоя при обратном проходе вместо хранения в памяти. Это убирает необходимость хранить активации всех слоёв для backprop — память перестаёт расти с глубиной сети. В сумме с чанкованным FFN это позволяло обрабатывать последовательности до 64k токенов на одном GPU.

Хэширование недетерминированно: для надёжности используют несколько раундов хэширования (multi-round LSH) и берут объединение бакетов — иначе релевантная пара запрос/ключ может случайно не попасть в один бакет.
Linformerлинейное · O(n)Wang et al., 2020

Отталкивается от эмпирического наблюдения: матрица внимания в реальных моделях почти всегда низкого ранга. Вместо того чтобы проецировать K и V по размерности признаков, Linformer проецирует их по оси длины последовательности: обучаемые матрицы E, F размера n×k сжимают K и V с n строк до k ≪ n строк перед вычислением внимания.

K' = E·K   # (k×d), было (n×d)
V' = F·V   # (k×d), было (n×d)
attn = softmax(Q·K'ᵀ / √d) · V'   # O(n·k) вместо O(n²)
k фиксирован и не зависит от n (например, k=256), поэтому сложность становится линейной по n. Ограничение: проекционные матрицы E, F подогнаны под максимальную длину обучения — модель хуже экстраполируется на длины, сильно превышающие обучающие.
Performerлинейное · ядерная аппроксимацияChoromanski et al., 2020

Использует механизм FAVOR+ (Fast Attention Via positive Orthogonal Random features): softmax-ядро exp(q·k) приближается через случайные признаковые отображения φ(·), такие что softmax(Q·Kᵀ) ≈ φ(Q)·φ(K)ᵀ. Это превращает внимание в обычное линейное произведение, которое можно переставить по ассоциативности:

# вместо: (φ(Q) φ(K)ᵀ) V  — снова матрица n×n
# считаем в другом порядке:
KV = φ(K)ᵀ @ V        # (r×d), не зависит от n
out = φ(Q) @ KV       # O(n·d·r), матрица n×n не строится никогда

В отличие от Linformer, аппроксимация ядра не завязана на фиксированную длину последовательности — она обобщается на произвольные n, а оценка несмещённая (unbiased) за счёт положительных ортогональных случайных признаков, что даёт низкую дисперсию приближения softmax.

cosFormerлинейное · без softmaxQin et al., 2022

Исходит из двух эмпирических свойств softmax-внимания, которые стоит сохранить при линеаризации: неотрицательность весов и локальность (внимание концентрируется на близких токенах). cosFormer заменяет softmax на ReLU (гарантирует неотрицательность весов линейного внимания) и добавляет косинусное перевзвешивание по относительной позиции токенов, имитируя локальность без явной нормализации softmax.

Q' = ReLU(Q), K' = ReLU(K)
weight(i,j) = cos((i−j)·π / 2M)     # локальность через относительную позицию
attn(i) = Σⱼ weight(i,j) · Q'ᵢ·K'ⱼ · Vⱼ   # считается за O(n), как и Performer

Результат — строго линейная сложность O(n) и рекуррентная форма вычисления, удобная для авторегрессивной генерации (можно писать как RNN-подобное накопление состояния), при качестве, сопоставимом или превосходящем Linformer/Performer на ряде задач.

МетодСложностьКлючевая идеяКомпромисс
BigBirdO(n)окно + глобальные + случайные токенынужны global-токены, шаблон внимания фиксирован
ReformerO(n log n)LSH-бакеты + обратимые слоинедетерминированность хэширования, нужны раунды
LinformerO(n)низкоранговая проекция длины K/Vпроекция подогнана под фикс. длину
PerformerO(n)ядерная аппроксимация softmax (FAVOR+)приближение, не точное внимание
cosFormerO(n)ReLU + косинусное перевзвешиваниеявная локальность зависит от гиперпараметра M
Архитектурные варианты голов внимания

Отдельная от разреженности ось оптимизации — сколько независимых наборов ключей/значений держит модель. Это не меняет асимптотическую сложность внимания, но напрямую влияет на объём KV-кэша и пропускную способность памяти при генерации.

Multi-Head Attention (MHA)baselineVaswani et al., 2017

Каждая из h голов имеет собственные проекции Q, K, V в подпространство размерности d_model/h, считает внимание независимо, результаты конкатенируются и проецируются обратно. Разные головы обучаются захватывать разные типы зависимостей (синтаксис, кореференции, позиционные паттерны).

Q1
Q2
Q3
Q4
каждая query-голова — своя собственная KV-голова
Multi-Query Attention (MQA)максимальное сжатие KVShazeer, 2019

Все query-головы используют одну общую KV-голову. Число различных проекций K/V падает с h до 1 — KV-кэш при генерации сжимается во столько же раз, во сколько раз h больше 1, что резко снижает нагрузку на память при инкрементальном декодировании.

Q1
Q2
Q3
Q4
все головы делят одну KV-голову

Цена — некоторая потеря выразительности и качества по сравнению с MHA, из-за чего MQA использовался ограниченно (PaLM, ранние версии Falcon) до появления промежуточного варианта.

Grouped-Query Attention (GQA)компромисс MHA/MQAAinslie et al., 2023

Query-головы делятся на g групп, внутри каждой группы головы делят одну общую KV-голову (g между 1 — это MQA, и h — это MHA). При g=8 для модели с h=64 головами KV-кэш сжимается в 8 раз при качестве, близком к полному MHA.

Q1
Q2
Q3
Q4
пары query-голов делят одну KV-голову (g=2 из 4)
GQA можно получить не только обучением с нуля, но и «uptraining»: усреднить KV-головы существующей MHA-модели внутри каждой группы и дообучить на небольшой доле исходных данных — так собирали GQA для Llama 2 70B.

Сегодня это стандарт для моделей с длинным контекстом: Llama 2/3, Mistral, большинство современных открытых LLM используют GQA как золотую середину между качеством MHA и эффективностью MQA.

ВариантKV-головКэш относительно MHAКачество
MHAh (= числу query-голов)1×эталон
GQAg, 1<g<hg/h ×близко к MHA
MQA11/h ×заметнее теряет в качестве
FlashAttention: точное внимание, но IO-осознанное

Важно понимать: FlashAttention не аппроксимирует внимание, как BigBird или Performer — она считает ровно тот же результат, что обычное softmax-внимание, но иначе организует обращения к памяти GPU.

FlashAttentionточное · IO-awareDao et al., 2022

Ключевое наблюдение: обычная реализация внимания упирается не в число FLOPs, а в скорость обмена данными между быстрой SRAM (на кристалле, мало памяти) и медленной HBM (много памяти, но на порядок медленнее). Материализация полной матрицы n×n внимания в HBM и есть главный источник задержки.

Тайлинг: Q, K, V разбиваются на блоки, которые целиком помещаются в SRAM. Внимание считается блок за блоком, результат накапливается прямо на кристалле — полная матрица n×n никогда не материализуется в HBM.

Online softmax: softmax по строке обычно требует видеть всю строку сразу (для нормировки). FlashAttention считает softmax инкрементально, блок за блоком, постоянно корректируя частичную сумму и максимум — математически эквивалентно обычному softmax, но без необходимости держать всю строку в памяти одновременно.

Пересчёт вместо хранения: при обратном проходе матрица внимания не хранится, а пересчитывается заново из Q, K, V — это дороже по FLOPs, но дешевле по IO, а именно IO — узкое место.

# схематично: один блок тайлинга
for block_k, block_v in tiles(K, V):        # блоки помещаются в SRAM
    s_block = Q_block @ block_k.T          # маленькая матрица, не n×n
    m_new = max(m_running, s_block.max())  # online softmax: обновляем максимум
    # переcчитываем накопленную сумму под новый максимум и добавляем блок
    acc = acc * exp(m_running - m_new) + exp(s_block - m_new) @ block_v
    m_running = m_new
FlashAttention и разреженные/линейные методы решают разные проблемы и комбинируются, а не конкурируют: FlashAttention ускоряет вычисление точного полного внимания, а BigBird/Performer/GQA/MQA снижают либо число вычисляемых пар, либо объём кэша.

FlashAttention-2 (2023) улучшает распределение работы между потоками GPU: лучше параллелит вычисления по оси длины последовательности (а не только по батчу и головам), сокращает долю немatmul-операций — почти двукратное ускорение над первой версией. FlashAttention-3 (2024) нацелен на архитектуру Hopper (H100): использует асинхронное выполнение и специализацию варпов между matmul и softmax, а также низкую точность FP8 — дополнительный прирост на новом железе.

Эволюция функций активации

Параллельная линия развития — как менялась нелинейность внутри FFN-блока трансформера: от простого порога до обучаемого гейтинга.

ReLU SiLU / Swish GELU
ReLUmax(0, x)

Простейшая нелинейность: обнуляет отрицательные значения, пропускает положительные без изменений. Дешева в вычислении, не страдает от насыщения градиента на положительной части (в отличие от sigmoid/tanh), поэтому долго была стандартом. Проблема — «умирающие» нейроны: если вход стабильно отрицателен, градиент через нейрон равен нулю, и он перестаёт обучаться.

GELUx·Φ(x)Hendrycks & Gimpel, 2016

Взвешивает вход его перцентилем по стандартной нормальной функции распределения Φ(x) — по сути, стохастическое обоснование: вход «пропускается» с вероятностью, зависящей от его величины. Гладкая и немонотонная (небольшой провал в отрицательной области вместо жёсткого нуля), что даёт более мягкий сигнал градиента, чем ReLU. Стала стандартом в BERT, GPT-2, GPT-3.

GELU(x) ≈ 0.5x · (1 + tanh[√(2/π)·(x + 0.044715x³)])   # практическая аппроксимация
SiLU / Swishx·σ(x)Ramachandran et al., 2017

Тот же принцип самогейтинга, что и у GELU, но вентиль — обычная сигмоида, а не нормальная CDF: SiLU(x) = x · sigmoid(x). Кривая почти неотличима от GELU визуально (см. график выше), но дешевле вычислять. Используется в EfficientNet, а главное — стал строительным блоком для гейтинга в современных FFN-слоях LLM.

SwiGLUгейтинг, не просто нелинейностьShazeer, «GLU Variants Improve Transformer», 2020

Это не отдельная функция активации, а вариант Gated Linear Unit, применённый ко всему FFN-блоку: вход проецируется двумя независимыми линейными слоями, один из результатов пропускается через SiLU и служит «воротами», которые поэлементно умножаются на второй, линейный, результат.

FFN_SwiGLU(x) = (SiLU(x·W) ⊙ (x·V)) · W2   # W, V, W2 — три отдельные матрицы
x
→
SiLU(x·W)
⊙
x·V
→
·W2 → выход
Три матрицы вместо двух увеличивают число параметров FFN примерно в 1.5 раза при том же скрытом размере — поэтому на практике скрытую размерность FFN уменьшают (обычно до 2/3 от «стандартной»), чтобы сохранить бюджет по параметрам и FLOPs сопоставимым с обычным FFN.

Эмпирически SwiGLU-FFN стабильно превосходит варианты на чистом ReLU или GELU при равном вычислительном бюджете, поэтому стал де-факто стандартом FFN в PaLM, LLaMA/Llama 2/3, Mistral и большинстве современных открытых LLM.

ФункцияФормулаГде используется
ReLUmax(0,x)ранние трансформеры, CNN, классический FFN
GELUx·Φ(x)BERT, GPT-2, GPT-3
SiLU/Swishx·σ(x)EfficientNet, компонент гейтинга в SwiGLU
SwiGLU(SiLU(xW)⊙xV)·W2PaLM, LLaMA/Llama 2/3, Mistral

Итоги: четыре независимых рычага эффективности

Все четыре темы решают разные части одной большой задачи — сделать трансформер быстрее и дешевле без потери качества, — но действуют на разных уровнях и потому свободно комбинируются в одной модели.