Внимание и трансформеры
Трансформер — самая влиятельная архитектура последнего десятилетия: на ней работают языковые модели, распознавание речи, генерация изображений, предсказание структуры белка. При этом её ядро — одна формула на три матрицы, выводимая за пятнадцать минут из простого вопроса: «как дать элементу последовательности самому решить, на что смотреть?».
Начнём с боли рекуррентных сетей («Рекуррентные сети, LSTM и GRU»), выведем внимание как мягкий поиск по словарю, соберём блок по кирпичику, разберём сложность и позиционные кодировки, напишем рабочий код и дойдём до вещей, определяющих цену инференса в проде: KV-кэша, GQA и FlashAttention.
1. Откуда взялось внимание
Схема машинного перевода 2014 года — encoder-decoder на LSTM. Энкодер читает предложение и сжимает его в один вектор $c = h_T^{\text{enc}}$; декодер разворачивает из него перевод. Проблема очевидна: предложение любой длины должно поместиться в вектор фиксированного размера. Для пяти слов нормально, для абзаца — катастрофа; качество таких моделей заметно падало после ~30 токенов.
Богданов, Чо и Бенджио в «Neural Machine Translation by Jointly Learning to Align and Translate» предложили не сжимать вообще. Пусть энкодер отдаёт все состояния $h_1,\dots,h_T$, а декодер на каждом шаге собирает контекст сам:
$$e_{tj} = a(s_{t-1}, h_j), \qquad \alpha_{tj} = \frac{\exp e_{tj}}{\sum_k \exp e_{tk}}, \qquad c_t = \sum_j \alpha_{tj} h_j$$
Посчитать релевантность каждого входного слова текущему шагу, превратить в распределение через softmax, взять взвешенную сумму. Побочный эффект оказался ценным: матрица $\alpha$ визуализировалась как таблица выравнивания между словами двух языков — осмысленная, хотя выравнивания никто не размечал.
Главный сдвиг мышления здесь такой. В RNN путь информации от токена $i$ к токену $j$ проходит через $|i-j|$ рекуррентных шагов, затухая на каждом. Во внимании путь — одно матричное умножение, длина $O(1)$ независимо от расстояния. Это, а не количество параметров, и есть причина победы трансформеров: короткий путь градиента между любыми двумя позициями.
2. Интуиция: мягкий поиск по словарю
Забудьте про нейросети. Питоновский словарь d["кот"] — жёсткий поиск: запрос сравнивается с ключами на равенство, возвращается ровно одно значение. Внимание делает то же самое, но непрерывно и дифференцируемо: запрос $q$ сравнивается с каждым ключом $k_j$ не на равенство, а на похожесть (скалярное произведение); похожести превращаются в веса через softmax — вместо одного попадания получаем распределение по всем ключам; результат — взвешенная сумма всех значений.
$$\text{lookup}(q) = \sum_j \underbrace{\frac{\exp(q \cdot k_j)}{\sum_l \exp(q \cdot k_l)}}_ {\text{насколько ключ } j \text{ подошёл}} ; v_j$$
Отсюда три роли, которые каждый токен играет одновременно: Query — «что я ищу», Key — «по какому признаку меня находить», Value — «что я отдам, если меня выбрали».
Почему K и V — разные векторы, а не один? Потому что «как меня адресовать» и «что я сообщаю» — разные вещи. В фразе «банк реки был крутым» токен «реки» должен находиться по запросу «уточни смысл слова банк», но отдавать содержание «это география, не финансы». Ключ отвечает за адресацию, значение — за содержание, и обучать их независимо выгоднее.
Все три получаются из одних и тех же представлений тремя обучаемыми проекциями: $Q = XW_Q$, $K = XW_K$, $V = XW_V$, где $X \in \mathbb{R}^{n \times d_{\text{model}}}$. Обучение подбирает проекции так, чтобы полезные пары «запрос — ключ» давали большое скалярное произведение.
3. Scaled dot-product attention
Собираем в одну формулу — ту самую, из «Attention Is All You Need»:
$$\text{Attention}(Q, K, V) = \text{softmax}!\left(\frac{QK^{\top}}{\sqrt{d_k}} + M\right)V$$
$M$ — аддитивная маска ($0$ где смотреть можно, $-\infty$ где нельзя). Размерности: $Q \in \mathbb{R}^{n \times d_k}$, $K \in \mathbb{R}^{m \times d_k}$, $V \in \mathbb{R}^{m \times d_v}$, результат — $n \times d_v$.
Почему делим на $\sqrt{d_k}$. Это не косметика, а вопрос обучаемости. Пусть компоненты $q$ и $k$ независимы, с нулевым средним и единичной дисперсией. Тогда $q \cdot k = \sum_i q_i k_i$ имеет $\mathbb{E} = 0$ и $\operatorname{Var} = d_k$, то есть разброс логитов растёт как $\sqrt{d_k}$. При $d_k = 128$ логиты гуляют в пределах $\pm 11$, softmax от них практически one-hot, а его якобиан в насыщении близок к нулю: $\partial p_i / \partial x_j = p_i(\delta_{ij} - p_j) \to 0$ при $p_i \to 1$. Градиент не течёт, обучение стоит. Деление на $\sqrt{d_k}$ возвращает дисперсию логитов к единице при любой размерности головы. Практическое следствие: «оптимизировав» самописное внимание выкидыванием масштаба, вы получите модель, которая обучается при малых $d_k$ и внезапно перестаёт при увеличении головы.
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
def scaled_dot_product_attention(q, k, v, mask=None):
"""Внимание в чистом виде — ровно формула выше.
q: (B, H, Tq, Dh), k: (B, H, Tk, Dh), v: (B, H, Tk, Dv)
mask: булев тензор, broadcastable к (B, H, Tq, Tk); True = «смотреть можно».
"""
d_k = q.size(-1)
# Единственное место во всей архитектуре, где токены вообще взаимодействуют.
scores = (q @ k.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
# -inf, а не «большое отрицательное»: после softmax даёт РОВНО ноль,
# и веса разрешённых позиций суммируются точно в единицу.
scores = scores.masked_fill(~mask, float("-inf"))
# Нормируем по последней оси, то есть ПО КЛЮЧАМ: каждая строка — распределение.
attn = torch.softmax(scores, dim=-1)
return attn @ v, attn
Сложность: две матричных операции по $O(nmd)$ плюс softmax $O(nm)$. Для self-attention ($n = m$) это $O(n^2 d)$ времени и $O(n^2)$ памяти — вторая цифра и есть главная боль архитектуры, к ней вернёмся в разделе 9.
4. Многоголовое внимание
Одна голова даёт одно распределение на строку. Но токену нужно несколько вещей сразу: «кто здесь подлежащее», «где закрывающая скобка», «какое имя я упоминал двести токенов назад». Одно softmax-распределение — это бюджет, который приходится делить. Решение: $h$ независимых внимания на подпространствах размерности $d_h = d_{\text{model}}/h$, результаты склеиваются:
$$\text{MHA}(X) = \big[\text{head}_ 1 | \dots | \text{head}_ h\big] W_O, \qquad \text{head}_ i = \text{Attention}(XW_Q^i, XW_K^i, XW_V^i)$$
Ключевая деталь: суммарная стоимость не растёт. Восемь голов по 64 измерения стоят столько же, сколько одна голова на 512, потому что $h \cdot d_h = d_{\text{model}}$. Мы не тратим больше вычислений — тратим их разнообразнее. Редкий случай бесплатного обеда в архитектурах.
class MultiHeadAttention(nn.Module):
"""Все четыре проекции без bias: в связке с нормализацией он избыточен."""
def __init__(self, d_model: int, n_heads: int, dropout: float = 0.0):
super().__init__()
assert d_model % n_heads == 0, "d_model должно делиться на число голов"
self.h, self.dh = n_heads, d_model // n_heads
self.wq = nn.Linear(d_model, d_model, bias=False)
self.wk = nn.Linear(d_model, d_model, bias=False)
self.wv = nn.Linear(d_model, d_model, bias=False)
self.wo = nn.Linear(d_model, d_model, bias=False)
self.p = dropout
def _split(self, x):
# (B, T, d_model) -> (B, H, T, Dh): голова становится ещё одной осью батча
B, T, _ = x.shape
return x.view(B, T, self.h, self.dh).transpose(1, 2)
def forward(self, x, causal: bool = True):
B, T, _ = x.shape
q, k, v = self._split(self.wq(x)), self._split(self.wk(x)), self._split(self.wv(x))
# В проде — только так: функция сама выберет FlashAttention или
# memory-efficient ядро под конкретное железо и размерности.
out = F.scaled_dot_product_attention(
q, k, v, is_causal=causal, dropout_p=self.p if self.training else 0.0)
out = out.transpose(1, 2).contiguous().view(B, T, self.h * self.dh)
return self.wo(out) # без Wo головы никогда не общаются между собой
Про $W_O$ часто забывают, а он содержателен: без него головы просто конкатенируются. $W_O$ — место, где информация из разных голов смешивается перед возвратом в остаточный поток. В анализе Transformer Circuits удобно думать о паре $W_V W_O$ как о «что голова пишет в поток», а о $W_Q W_K^{\top}$ — как о «куда голова смотрит».
5. Три режима: self, cross, causal
Формула одна, конфигурации три, и путать их — источник трудноуловимых багов.
| Режим | Q | K, V | Маска | Где встречается |
|---|---|---|---|---|
| Self-attention двунаправленное | X | та же X | только padding | BERT, ViT, энкодеры |
| Causal self-attention | X | та же X | нижнетреугольная + padding | GPT и все LLM |
| Cross-attention | декодер | энкодер | padding энкодера | перевод, Whisper, text-to-image |
Causal-маска превращает трансформер в языковую модель: позиция $i$ не должна видеть позиции $> i$, иначе предсказание следующего токена тривиально — ответ уже на входе. Технически это torch.ones(T, T, dtype=torch.bool).tril(). Её красота — в параллельном обучении: за один проход модель получает градиент от всех $n$ предсказаний сразу, потому что позиция $i$ уже видит ровно тот контекст, который увидит на инференсе. RNN так не умеет — там $n$ последовательных шагов. Именно это, а не сама формула внимания, сделало возможным обучение на триллионах токенов.
Padding-маска нужна при разной длине в батче. Классическая ошибка — замаскировать padding в потерях, но не во внимании: тогда реальные токены подмешивают в себя мусор. Вторая — строка, состоящая целиком из запрещённых позиций: softmax от вектора $-\infty$ даёт NaN, который расползается по всей модели.
Cross-attention обусловливает одну последовательность на другую: перевод, image captioning, text-to-image диффузия (см. «Генеративные модели»), где картинка задаёт Q, а текстовый эмбеддинг — K и V.
6. Позиционное кодирование: главная дыра механизма
Поищите в формуле внимания хоть что-нибудь, зависящее от порядка токенов. Его там нет. Внимание перестановочно эквивариантно: переставьте строки $X$ — выход переставится так же, но значения не изменятся. Для модели «собака укусила человека» и «человека укусила собака» — одно и то же множество. Порядок вносят отдельно, и подходов четыре.
1. Синусоидальное абсолютное (оригинальный трансформер): к эмбеддингу прибавляется фиксированный вектор $PE_{(pos,2i)} = \sin(pos / 10000^{2i/d})$, $PE_{(pos,2i+1)} = \cos(pos / 10000^{2i/d})$. Разные измерения — синусоиды разных частот, вместе дающие уникальную «подпись» позиции, похожую на двоичную запись числа.
2. Обучаемое абсолютное (BERT, GPT-2): таблица nn.Embedding(max_len, d). Работает, но жёстко ограничивает контекст — за пределами max_len эмбеддингов просто нет.
3. ALiBi: никаких эмбеддингов, к логитам добавляется линейный штраф за расстояние $s_{ij} \mathrel{-}= m_h |i-j|$ со своей константой на голову. Просто и хорошо экстраполируется за пределы обучающих длин (arXiv:2108.12409).
4. RoPE — де-факто стандарт (LLaMA, Qwen, Mistral, Gemma). Вместо того чтобы прибавлять позицию, повернём пары координат Q и K на угол, пропорциональный позиции. Тогда скалярное произведение зависит только от разности позиций: $\langle R_m q, R_n k \rangle = \langle q, R_{n-m} k \rangle$. Абсолютные позиции применяются, а видит модель относительные (RoFormer, arXiv:2104.09864).
def build_rope_cache(seq_len: int, d_head: int, base: float = 10000.0, device=None):
"""cos/sin для всех позиций — считаем один раз при инициализации.
Частоты убывают геометрически: младшие пары координат вращаются быстро
(различают соседние токены), старшие — медленно (различают далёкий контекст).
"""
inv_freq = 1.0 / (base ** (torch.arange(0, d_head, 2, device=device).float() / d_head))
freqs = torch.outer(torch.arange(seq_len, device=device).float(), inv_freq) # (T, Dh/2)
return torch.cos(freqs), torch.sin(freqs)
def apply_rope(x, cos, sin, offset: int = 0):
"""Поворачиваем пары (x_2i, x_2i+1) как точки на плоскости. x: (B, H, T, Dh).
offset обязателен на инференсе: при декодировании подаётся один токен,
но его настоящая позиция — не ноль, а текущая длина KV-кэша.
"""
T = x.size(-2)
cos, sin = cos[offset:offset + T][None, None], sin[offset:offset + T][None, None]
x1, x2 = x[..., 0::2], x[..., 1::2]
return torch.stack((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1).flatten(-2)
RoPE применяют только к Q и K, но не к V: позиция должна влиять на то, кто на кого смотрит, а не на передаваемое содержание. Забыть offset при декодировании с кэшем — классическая ошибка: генерация начинается нормально и постепенно сходит с ума.
Отдельная тема — растяжение контекста. Модель, обученная на 4k, при подаче 32k ломается: углы уходят в область, которой не было в обучении. Position Interpolation (сжать позиции в обученный диапазон) и NTK-aware scaling (изменить base) позволяют расширить окно дообучением на небольшом объёме данных.
7. Блок трансформера целиком
Внимание — это только смешивание токенов, причём линейное относительно V: между позициями информация течёт, но внутри позиции ничего сложного не происходит. Поэтому блок всегда состоит из двух чередующихся операций: attention смешивает информацию между токенами, FFN обрабатывает каждый токен независимо. Разделение труда фундаментально: attention решает «что релевантно», FFN — «что с этим делать». По объёму параметров FFN обычно вдвое больше внимания, и именно в нём, по данным интерпретируемости, лежит основная часть фактических знаний модели.
B × n × d_model"]:::io IN --> N1["RMSNorm"] subgraph ATTN["Смешивание между токенами"] direction TB N1 --> PROJ["проекции Wq, Wk, Wv
+ разбиение на головы"] PROJ --> ROPE["RoPE: поворот Q и K по позиции"] ROPE --> SDPA["causal scaled dot-product
softmax по ключам"] SDPA --> WO["склейка голов и проекция Wo"] end WO --> ADD1(("+")) IN -- "остаточная связь" --> ADD1 ADD1 --> N2["RMSNorm"] subgraph FFN["Обработка каждого токена отдельно"] direction TB N2 --> GATE["SwiGLU: d → 8/3·d"] GATE --> DOWN["проекция обратно 8/3·d → d"] end DOWN --> ADD2(("+")) ADD1 -- "остаточная связь" --> ADD2 ADD2 --> OUT["выход блока
та же форма"]:::io classDef io fill:#3b82f6,fill-opacity:0.15,stroke:#3b82f6
Pre-LN против post-LN. В оригинале нормализация стояла после остаточной связи: $x \leftarrow \text{LN}(x + \text{Sublayer}(x))$. Такая схема без аккуратного warmup расходится за первые сотни шагов. Современные модели нормализуют до подслоя: $x \leftarrow x + \text{Sublayer}(\text{LN}(x))$. Разница принципиальна: в pre-LN есть чистый остаточный путь от входа до выхода без единой нормализации, поэтому градиент доходит до нижних слоёв неискажённым и сотня блоков обучается без warmup (arXiv:2002.04745). Цена — чуть худшее финальное качество при прочих равных, отсюда гибриды вроде sandwich-LN в очень больших моделях.
RMSNorm (arXiv:1910.07467) выбрасывает центрирование, оставляя масштабирование: $x / \sqrt{\tfrac{1}{d}\sum x_i^2 + \varepsilon} \odot g$. Качество то же, вычислений меньше — а норм в модели много. SwiGLU (arXiv:2002.05202) заменяет Linear → ReLU → Linear на вентильную схему $(\text{SiLU}(xW_{\text{gate}}) \odot xW_{\text{up}})W_{\text{down}}$; матриц три вместо двух, поэтому скрытую размерность берут $\tfrac{2}{3}\cdot 4d$, сохраняя число параметров.
class Block(nn.Module):
"""Pre-LN блок decoder-only трансформера. nn.RMSNorm — начиная с PyTorch 2.4."""
def __init__(self, d_model: int, n_heads: int, mult: int = 4, dropout: float = 0.0):
super().__init__()
self.n1, self.n2 = nn.RMSNorm(d_model), nn.RMSNorm(d_model)
self.attn = MultiHeadAttention(d_model, n_heads, dropout)
# 2/3 * 4d компенсирует третью матрицу вентиля: параметров столько же,
# сколько у обычного FFN с четырёхкратным расширением.
hidden = 64 * ((int(mult * d_model * 2 / 3) + 63) // 64) # + выравнивание под тензорные ядра
self.w_gate = nn.Linear(d_model, hidden, bias=False)
self.w_up = nn.Linear(d_model, hidden, bias=False)
self.w_down = nn.Linear(hidden, d_model, bias=False)
self.drop = nn.Dropout(dropout)
def forward(self, x):
x = x + self.drop(self.attn(self.n1(x))) # токены смотрят друг на друга
h = self.n2(x)
return x + self.drop(self.w_down(F.silu(self.w_gate(h)) * self.w_up(h)))
Остаточный поток как разделяемая шина. Полезная ментальная модель: вектор, текущий через все блоки, — не «представление токена», а общая шина, в которую каждый подслой дописывает свой вклад (x = x + ...), а читает через нормализацию и проекции. Ничего не стирается, только добавляется. Из этой картины растёт весь анализ в «Интерпретируемости»: логиты раскладываются на вклады отдельных голов и слоёв.
8. Три семейства архитектур
| Семейство | Внимание | Обучение | Сильная сторона |
|---|---|---|---|
| Encoder-only (BERT, ViT) | двунаправленное | masked LM | понимание: классификация, поиск, эмбеддинги |
| Decoder-only (GPT, LLaMA) | причинное | следующий токен | генерация и универсальность |
| Encoder-decoder (T5, Whisper) | оба + cross | seq2seq | вход и выход разной природы |
Decoder-only победило по прагматичной причине: сигнал обучения плотнее — потеря на каждой позиции, а не на 15% замаскированных, — и данные не требуют разметки. Подробно об этом в «Больших языковых моделях», про эмбеддинги из энкодеров — в «Эмбеддингах».
Отдельно стоит ViT (arXiv:2010.11929): картинка режется на патчи 16×16, каждый линейно проецируется в токен — дальше обычный encoder-only трансформер. Это показало, что архитектура вообще не про язык, а про множества элементов с обучаемыми связями. Сравнение индуктивных смещений со свёрточными сетями: CNN зашивает локальность и трансляционную инвариантность в архитектуру, трансформер их не имеет и выучивает из данных — поэтому на малых датасетах ViT проигрывает, а на больших выигрывает. Формально self-attention эквивалентно передаче сообщений на полном графе — прямая связь с графовыми сетями.
9. Сложность и trade-offs
Таблица из оригинальной статьи ($n$ — длина, $d$ — размерность, $k$ — свёрточное ядро):
| Слой | Сложность на слой | Последовательных операций | Макс. длина пути |
|---|---|---|---|
| Self-attention | $O(n^2 d)$ | $O(1)$ | $O(1)$ |
| Рекуррентный | $O(n d^2)$ | $O(n)$ | $O(n)$ |
| Свёрточный | $O(k n d^2)$ | $O(1)$ | $O(\log_k n)$ |
| Self-attention в окне $r$ | $O(r n d)$ | $O(1)$ | $O(n/r)$ |
Два правых столбца объясняют победу: $O(1)$ последовательных операций (всё параллелится) и $O(1)$ длина пути (градиент не затухает между любыми токенами). Цена — $n^2$ слева.
Полезная арифметика на слой: attention даёт около $4n^2 d$ FLOPs, проекции и FFN — около $24 n d^2$. Точка перелома: $4n^2 d > 24 n d^2 \iff n > 6d$. При $d_{\text{model}} = 4096$ внимание начинает доминировать только после ~25 тысяч токенов. Это часто понимают неверно: на типичных длинах 2–8k квадратичность почти не влияет на время обучения, львиную долю FLOPs съедают полносвязные слои. Зато она беспощадно бьёт по памяти — матрица $n \times n$ на голову на слой при наивной реализации даёт гигабайты активаций. Ровно это решает FlashAttention, причём точно, а не приближённо.
10. Обучение: что именно ломается
Общие вопросы оптимизации разобраны в «Обучении сетей»; здесь — только специфика трансформеров.
Warmup. Post-LN без разогрева расходится почти гарантированно. Оригинальное расписание: $\text{lr} = d_{\text{model}}^{-0.5}\min(\text{step}^{-0.5}, \text{step}\cdot\text{warmup}^{-1.5})$ при warmup = 4000. С pre-LN warmup нужен уже не для выживания, а для качества — обычно 1–2% шагов, дальше косинусный спад.
Инициализация по глубине. Остаточный поток накапливает вклады всех слоёв, и его дисперсия растёт с глубиной. Стандартный приём (из GPT-2): делить инициализацию выходных проекций подслоёв ($W_O$ и w_down) на $\sqrt{2L}$, где $L$ — число блоков.
Численная устойчивость. Внимание считают в bf16, но softmax — обязательно в fp32: в fp16 экспонента от логита ~12 уже близка к переполнению. Вычитание максимума строки перед экспонентой встроено во все библиотечные реализации. Отдельный эффект — attention logit growth: у очень больших моделей логиты склонны неограниченно расти, загоняя softmax в насыщение; лечится QK-нормализацией (RMSNorm на Q и K перед скалярным произведением).
Dropout. В эпоху малых данных ставили 0.1 везде. При обучении на триллионах неповторяющихся токенов переобучения нет, и при предобучении dropout часто выставляют в ноль — он только замедляет. При дообучении на маленьком датасете возвращают.
11. Инференс: KV-кэш и всё, что из него следует
Генерация авторегрессивна: чтобы получить токен $t+1$, нужен токен $t$. Наивная реализация на каждом шаге прогоняет всю последовательность заново — $O(n^2)$ работы на токен и $O(n^3)$ на ответ. Абсурд, потому что K и V прошлых позиций не меняются: причинная маска гарантирует, что новый токен не влияет на старые. Отсюда KV-кэш: посчитанные K и V складываются в память и переиспользуются.
def attend_with_cache(self, x, cache=None, rope=None):
"""Шаг внимания с кэшем. cache = (k_prev, v_prev) формы (B, H, T_past, Dh)."""
B, T, _ = x.shape
q, k, v = self._split(self.wq(x)), self._split(self.wk(x)), self._split(self.wv(x))
past_len = 0 if cache is None else cache[0].size(2)
if rope is not None:
cos, sin = rope
# Настоящая позиция нового токена = длина уже накопленного кэша.
q, k = apply_rope(q, cos, sin, past_len), apply_rope(k, cos, sin, past_len)
if cache is not None:
k = torch.cat([cache[0], k], dim=2) # растим по оси времени
v = torch.cat([cache[1], v], dim=2)
# При T == 1 маска не нужна вовсе: единственный запрос имеет полное право
# смотреть на весь накопленный кэш.
out = F.scaled_dot_product_attention(q, k, v, is_causal=(T > 1))
out = out.transpose(1, 2).contiguous().view(B, T, self.h * self.dh)
return self.wo(out), (k, v)
Отсюда две принципиально разные фазы. Prefill — обработка промпта целиком: матрицы большие, GPU занят арифметикой, фаза compute-bound, определяет TTFT (время до первого токена). Decode — по одному токену: матричное умножение вырождается в умножение матрицы на вектор, арифметическая интенсивность падает, время уходит на чтение весов и кэша из HBM, фаза memory-bound, определяет скорость генерации. Понимание этой асимметрии — половина работы инженера по инференсу: оптимизации, помогающие prefill, для decode почти бесполезны, и наоборот.
def kv_cache_bytes(n_layers, n_kv_heads, d_head, seq_len, batch, dtype_bytes=2):
"""Двойка — потому что храним и K, и V."""
return 2 * n_layers * n_kv_heads * d_head * seq_len * batch * dtype_bytes
# Llama-3-8B: 32 слоя, 8 KV-голов благодаря GQA, d_head = 128, bfloat16
print(kv_cache_bytes(32, 8, 128, 1, 1) / 1024) # 128.0 KiB на один токен
print(kv_cache_bytes(32, 8, 128, 8192, 1) / 2**30) # 1.0 GiB на запрос с 8k контекста
print(kv_cache_bytes(32, 32, 128, 8192, 1) / 2**30) # 4.0 GiB — та же модель БЕЗ GQA
Гигабайт на запрос — и на карте с 80 ГБ, из которых 16 заняты весами, помещается около шестидесяти параллельных запросов. Это и есть потолок пропускной способности, и определяется он кэшем, а не числом параметров.
MQA и GQA. Раз кэш пропорционален числу KV-голов — уменьшим их. Multi-Query Attention оставляет одну общую пару K,V на все головы запросов: кэш падает в $h$ раз ценой лёгкой деградации качества. Grouped-Query Attention — компромисс: Q-головы делятся на группы, каждая со своей парой K,V. При 32 Q-головах и 8 KV-головах кэш меньше вчетверо почти без потери качества. Сейчас это стандарт.
новые подсаживаются на свободные слоты end B->>G: освободить страницы KV-кэша
PagedAttention. Наивно кэш выделяют непрерывным блоком на максимальную длину — и теряют 60–80% памяти на фрагментации, потому что реальные ответы короче лимита. vLLM (arXiv:2309.06180) применил идею виртуальной памяти ОС: кэш хранится страницами фиксированного размера с таблицей трансляции. Утилизация выросла почти до 100%, а общие префиксы (системный промпт!) стали физически разделяемыми между запросами. Подробнее — в «Инференсе и деплое».
12. Варианты внимания: как борются с квадратичностью
FlashAttention (arXiv:2205.14135) — важнейшая из этих работ, потому что она ничего не приближает. Наблюдение: узкое место не арифметика, а чтение и запись матрицы $n \times n$ в HBM. Решение: разбить Q, K, V на блоки, влезающие в SRAM мультипроцессора, и считать softmax инкрементально (online softmax: поддерживаем текущий максимум и сумму экспонент, корректируя накопленный результат при появлении нового максимума). Матрица внимания целиком не материализуется никогда. Итог: память $O(n)$ вместо $O(n^2)$, ускорение в 2–4 раза, результат математически идентичен наивному. Отличная иллюстрация общего принципа: на современных GPU выигрывают алгоритмы, экономящие обращения к памяти, а не арифметику.
Скользящее окно. Каждый токен видит $w$ предыдущих — сложность $O(nw)$. Кажется грубым, но за $L$ слоёв рецептивное поле растёт до $L \cdot w$, ровно как в свёрточных сетях: Mistral 7B с окном 4096 и 32 слоями формально охватывает 131k. Бонус — кэш перестаёт расти, хватает кольцевого буфера на $w$ позиций.
Линейное внимание. Заменив $\exp(q\cdot k)$ на $\phi(q)\cdot\phi(k)$, можно переставить скобки: $(\phi(Q)\phi(K)^{\top})V = \phi(Q)(\phi(K)^{\top}V)$ — и получить $O(nd^2)$ вместо $O(n^2 d)$ (arXiv:2006.16236). Формально это возвращает нас к RNN с матричным состоянием. На практике качество на языковом моделировании стабильно ниже точного softmax-внимания, поэтому чистые линейные варианты в больших LLM не прижились — а вот гибриды (несколько слоёв точного внимания среди многих линейных или SSM) выглядят живо.
Attention sink. Модели сваливают «лишнюю» массу внимания на первые несколько токенов, используя их как сток для случая «мне сейчас не на что смотреть». Если при потоковой генерации выкидывать самые старые элементы кэша, качество рушится — именно потому, что выкидываются стоки. StreamingLLM (arXiv:2309.17453) предлагает держать первые 4 токена всегда, и тогда скользящее окно работает на бесконечных потоках.
13. Типичные ошибки
- Забыт масштаб $1/\sqrt{d_k}$ — обучение то идёт, то нет в зависимости от размера головы; симптом — почти one-hot веса внимания с самого начала.
- Маска большим отрицательным числом вместо $-\infty$:
-1e9в fp16 переполняется,-1e4оставляет запрещённым позициям заметный вес. - Полностью замаскированная строка → softmax от $-\infty$ →
NaN. Возникает на пустых последовательностях и при пересечении causal- и padding-масок. - Softmax по неправильной оси. Нормировать надо по ключам, то есть по последней; ошибка даёт модель, которая обучается, но заметно хуже, и это тяжело заметить.
- Потерянный
offsetв RoPE при декодировании с кэшем — генерация деградирует по мере роста контекста. - Утечка будущего. Подозрительно низкий loss и мусор на инференсе. Проверка: смените последний токен и убедитесь, что логиты первых позиций не изменились.
- Забыта финальная нормализация перед выходной головой в pre-LN архитектуре — логиты живут в неограниченном масштабе.
- Экономия на padding. Сортировка по длине и бакетизация нередко ускоряют обучение на 20–40% просто потому, что модель перестаёт считать внимание на пустоте.
14. Как это применяют в проде
Никогда не пишите внимание руками для продакшена — используйте F.scaled_dot_product_attention, она сама выберет подходящее ядро. Собственная реализация нужна ровно для двух вещей: понять механизм и исследовать нестандартные маски.
- Стек обучения: decoder-only, pre-LN, RMSNorm, SwiGLU, RoPE, GQA, без bias, AdamW с $\beta_2 = 0.95$, косинусный спад с коротким warmup, обрезка градиента по норме 1.0, bf16, активационный чекпоинтинг на длинных контекстах. Отклоняться стоит осознанно — набор выстрадан тысячами GPU-часов.
- Стек инференса: vLLM / TensorRT-LLM / SGLang, continuous batching (запросы входят и выходят из батча независимо, а не ждут самого медленного), paged KV-cache, prefix caching для общего системного промпта, квантизация весов в INT8/FP8 и кэша в FP8, спекулятивное декодирование там, где важна задержка.
- Что мерить: TTFT и TPOT (время на последующий токен) — разные метрики, упирающиеся в разные ресурсы; пропускная способность в токенах в секунду на GPU — метрика денег.
- За пределами NLP: Whisper (речь), ViT и DINOv2 (зрение), AlphaFold (белки), Decision Transformer (RL), временные ряды, диффузионные трансформеры DiT, рекомендательные системы.
Мини-итог
- Внимание — мягкий поиск по словарю: похожесть запроса с ключами → softmax → взвешенная сумма значений. Всё остальное — инженерная обвязка вокруг этой идеи.
- $1/\sqrt{d_k}$ удерживает softmax вне насыщения; многоголовость даёт несколько параллельных распределений внимания бесплатно по вычислениям.
- Внимание не знает о порядке — его вносит позиционное кодирование; стандарт RoPE поворачивает Q и K так, что скалярное произведение зависит от относительной позиции.
- Блок = смешивание между токенами (attention) + независимая обработка каждого (FFN), оба через остаточные связи и pre-нормализацию.
- $O(n^2)$ по времени, но на типичных длинах FLOPs съедает FFN; квадратичность бьёт прежде всего по памяти, и FlashAttention решает это точно, а не приближённо.
- Цена инференса определяется KV-кэшем, а не числом параметров; GQA, paged-кэш и continuous batching — три главных рычага.
Источники
- Vaswani et al. Attention Is All You Need — arXiv:1706.03762
- Bahdanau et al. NMT by Jointly Learning to Align and Translate — arXiv:1409.0473
- The Annotated Transformer (Harvard NLP), построчный разбор кодом — nlp.seas.harvard.edu
- Jay Alammar. The Illustrated Transformer — jalammar.github.io
- Andrej Karpathy. nanoGPT, минимальная полноценная реализация — github.com/karpathy/nanoGPT
- Xiong et al. On Layer Normalization in the Transformer Architecture — arXiv:2002.04745
- Su et al. RoFormer: Rotary Position Embedding — arXiv:2104.09864
- Press et al. Train Short, Test Long: ALiBi — arXiv:2108.12409
- Dao et al. FlashAttention — arXiv:2205.14135; FlashAttention-2 — arXiv:2307.08691
- Ainslie et al. GQA — arXiv:2305.13245
- Kwon et al. PagedAttention / vLLM — arXiv:2309.06180
- Elhage et al. A Mathematical Framework for Transformer Circuits — transformer-circuits.pub
- Dosovitskiy et al. An Image is Worth 16x16 Words (ViT) — arXiv:2010.11929
- PyTorch:
scaled_dot_product_attention— pytorch.org/docs
Что дальше
Архитектуру разобрали. Следующий шаг — что происходит, когда её масштабируют до сотен миллиардов параметров и триллионов токенов: как устроено предобучение, откуда берутся законы масштабирования, что такое SFT и RLHF и почему инференс больших моделей стал отдельной инженерной дисциплиной.