Механизм внимания
Внимание (attention) — центральная операция трансформера: именно через неё токены обмениваются информацией. Все остальные слои блока декодера (нормализация, FFN) обрабатывают каждую позицию отдельно. Все шесть моделей репозитория используют одно и то же causal self-attention и отличаются тремя независимыми «ручками»:
- сколько голов K/V приходится на головы Q: MHA, GQA или MQA;
- какие позиции видит токен: всё прошлое или только скользящее окно;
- как в attention попадает позиция: через слагаемое к эмбеддингам (GPT) или поворотом Q и K (RoPE).
Что вы узнаете
Заголовок раздела «Что вы узнаете»- Откуда взялось внимание и почему его описывают как «мягкий поиск по словарю» с запросами, ключами и значениями.
- Формулу scaled dot-product attention, вывод множителя и что ломается без него.
- Как устроено многоголовое внимание (multi-head attention), сколько у него параметров и сколько оно стоит по времени и памяти.
- Чем отличаются MHA, GQA и MQA и почему это в первую очередь вопрос размера KV-кэша.
- Как работают скользящее окно и KV-кэш и как всё это реализовано в
llm/core.
Предварительные знания
Заголовок раздела «Предварительные знания»- Языковое моделирование: предсказание следующего токена, общая схема decoder-only трансформера.
- Эмбеддинги: как токен превращается в вектор размерности .
- Позиционное кодирование: обучаемые позиции и RoPE.
- Из математики: скалярное произведение, умножение матриц, softmax, дисперсия суммы независимых величин (см. Обозначения).
Зачем нужно внимание
Заголовок раздела «Зачем нужно внимание»До трансформеров машинный перевод решали моделями «кодировщик — декодировщик» (encoder-decoder) на рекуррентных сетях (Sutskever et al., 2014; Cho et al., 2014). Кодировщик читал исходное предложение слово за словом и сжимал его в один вектор фиксированной длины — последнее скрытое состояние RNN. Декодировщик генерировал перевод, опираясь только на этот вектор.
Узкое место очевидно: предложение из 5 слов и предложение из 50 слов упаковываются в вектор одного и того же размера. Bahdanau, Cho, Bengio (2015) показали, что качество такого перевода падает с ростом длины предложения, и предложили не сжимать вход в один вектор. Кодировщик сохраняет скрытые состояния всех слов, а декодировщик на каждом шаге сам выбирает, на какие из них смотреть:
где:
- — скрытое состояние декодировщика перед шагом («что я сейчас ищу»);
- — состояние кодировщика для -го слова источника;
- — функция сходства; у Bahdanau это маленькая сеть с одним скрытым слоем и (так называемое аддитивное внимание);
- — вес внимания: насколько слово важно на шаге ; веса неотрицательны и в сумме дают 1;
- — вектор контекста, взвешенное среднее состояний кодировщика.
Это и есть внимание: вместо одного фиксированного вектора — своя смесь входов для каждого шага. Vaswani et al. (2017) сделали следующий шаг: убрали рекуррентность совсем и построили модель только из внимания («Attention Is All You Need»). Функцию сходства они заменили на скалярное произведение — оно считается одним умножением матриц для всех пар сразу. Внимание, в котором запросы и ключи берутся из одной и той же последовательности, называется самовниманием (self-attention); именно оно используется во всех моделях репозитория.
Внимание как мягкий поиск по словарю
Заголовок раздела «Внимание как мягкий поиск по словарю»Обычный словарь (Python dict) хранит пары «ключ → значение». Поиск по запросу находит ключ, точно равный , и возвращает его значение:
словарь: "кот" → v₁, "пёс" → v₂, "дом" → v₃запрос: "пёс" → результат v₂ (веса 0, 1, 0)Внимание делает то же самое, но мягко: запрос сравнивается со всеми ключами, сходство превращается в веса через softmax, а результат — взвешенная сумма всех значений:
запрос q: сходство с ключами (0.2, 2.1, 0.4)softmax → веса (0.11, 0.75, 0.14)результат = 0.11·v₁ + 0.75·v₂ + 0.14·v₃У такой операции три роли для векторов:
- Запрос (query) — «что я ищу»;
- Ключ (key) — «по какому признаку меня находить»;
- Значение (value) — «что я отдаю, если меня нашли».
Мягкость важна по двум причинам. Во-первых, результат дифференцируем по запросам и ключам: жёсткий выбор argmax не пропускает градиент, а softmax пропускает, и модель можно обучать градиентным спуском. Во-вторых, токен может собрать информацию сразу из нескольких мест.
В self-attention каждый токен последовательности играет все три роли: он задаёт свой запрос, выставляет свой ключ и отдаёт своё значение.
Запросы, ключи и значения
Заголовок раздела «Запросы, ключи и значения»Пусть на вход слоя пришли скрытые состояния токенов, записанные строками матрицы . Запросы, ключи и значения — три разные линейные проекции одного и того же входа:
где:
- — входные векторы токенов (строка — токен на позиции );
- — обучаемые матрицы проекций (пока рассматриваем одну голову);
- — запросы, ключи и значения; строки ;
- — размер головы (
head_size).
Зачем три разные проекции, а не сам ? Если бы запрос и ключ совпадали (), сходство было бы симметричным, и больше всего токен смотрел бы на самого себя (). Раздельные и позволяют отношению « ищет » быть несимметричным: глагол ищет подлежащее, а подлежащее глагол — не обязательно. Отдельная отделяет «по чему ищут» от «что передают».
В коде с батчем все тензоры получают ведущее измерение : имеет форму [B, T, d], и nn.Linear применяет одну и ту же матрицу ко всем токенам и всем примерам. Линейные слои в репозитории могут иметь смещение (bias): , где прибавляется к каждой строке. В GPT-1 и GPT-2 bias есть, в оригинальных LLaMA, Mistral, Mixtral и Gemma — нет. В репозитории у LLaMA, Mistral, Mixtral и Gemma его включает ключ конфига bias, по умолчанию true (ради совместимости со старыми чекпоинтами); чтобы получить архитектуру статей, задают "bias": false.
Scaled dot-product attention
Заголовок раздела «Scaled dot-product attention»Основная формула (Vaswani et al., 2017, разд. 3.2.1):
где:
- , , — запросы, ключи, значения; — число ключей (без кэша , с кэшем — длина кэша плюс );
- — матрица оценок (scores): — сходство запроса с ключом ;
- — маска: там, где смотреть можно, там, где нельзя (подробно — в главе Маски);
- применяется к каждой строке отдельно: ;
- — матрица весов внимания; каждая строка — распределение вероятностей по ключам;
- результат ; строка .
Словами, для каждой позиции :
- сравнить её запрос со всеми разрешёнными ключами (скалярное произведение);
- поделить на ;
- превратить оценки в веса softmax’ом;
- взять взвешенную сумму значений.
Интуиция. Выход — выпуклая комбинация векторов значений: он лежит «между» ними. Внимание не создаёт новые признаки, а перераспределяет существующие между позициями; новые признаки создают проекции и FFN. Если убрать softmax и оставить , веса перестанут быть неотрицательными и нормированными, а масштаб выхода будет расти с длиной последовательности.
Название «scaled dot-product» — «масштабированное скалярное произведение» — описывает именно шаги 1–2.
Почему делить на корень из d_h
Заголовок раздела «Почему делить на корень из d_h»Предположим, что компоненты запроса и ключа — независимые случайные величины со средним 0 и дисперсией 1 (примерно так и есть в начале обучения при стандартной инициализации и нормализованном входе). Посчитаем дисперсию их скалярного произведения.
Вывод
Скалярное произведение — сумма слагаемых:
Шаг 1. Среднее одного слагаемого. Так как и независимы,
Шаг 2. Дисперсия одного слагаемого. По определению , а среднее мы уже нашли:
Здесь снова использована независимость ( и тоже независимы), а .
Шаг 3. Дисперсия суммы. Слагаемые с разными независимы, поэтому дисперсии складываются:
Шаг 4. Масштабирование. Для любой константы верно . При :
Итог: без масштаба стандартное отклонение оценок равно — при (LLaMA, Mistral) это около 11,3, при (Gemma) — 16. После деления на оно равно 1 при любом размере головы. Это же объяснение дано в сноске 4 статьи Vaswani et al. Проверка моделированием (100 000 пар случайных векторов) даёт дисперсию 2,0 / 64,0 / 128,7 до деления и 1,00 / 1,00 / 1,01 после для .
Предположение о независимости и единичной дисперсии — идеализация: после обучения и коррелированы (для того внимание и обучают). Масштаб задаёт правильный порядок величин в начале обучения, а дальше модель сама подстраивает нормы через и .
Что происходит без масштаба: насыщение softmax
Заголовок раздела «Что происходит без масштаба: насыщение softmax»Большие по модулю оценки делают softmax почти one-hot (распределением с одной единицей): вес максимального элемента близок к 1, остальные — к 0. Пример — одни и те же оценки до и после умножения на :
| Оценки | |
|---|---|
Для 64 ключей со случайными оценками средний максимальный вес строки — 0,11 при стандартном отклонении оценок 1 и 0,85 при стандартном отклонении 11,3: без масштаба внимание почти всегда «смотрит в одну точку».
Почему это плохо для обучения, видно из производной softmax. Пусть , . Тогда
где — символ Кронекера (1 при , иначе 0).
Вывод
Обозначим , тогда и . По правилу дифференцирования частного:
Если почти one-hot с единицей на позиции , то:
- для : ;
- для : , потому что .
Все элементы якобиана близки к нулю, и градиент почти не доходит до и — обучение внимания останавливается. В примере выше наибольший по модулю элемент якобиана — 0,22 для и для : в 18 000 раз меньше.
Этим scaled dot-product отличается от «простого» скалярного внимания. Vaswani et al. (разд. 3.2.1) отмечают, что без масштаба скалярное внимание при больших проигрывает аддитивному, и объясняют это именно малыми градиентами softmax.
Численный пример
Заголовок раздела «Численный пример»Возьмём токена и голову размера (тогда ):
Шаг 1. Оценки . Например, (строки и столбцы нумеруем с 0):
k₀ k₁ k₂q₀ 0.7071 0.0000 0.7071q₁ 0.0000 0.7071 0.7071q₂ 0.7071 0.7071 1.4142Шаг 2. Causal-маска запрещает — элементы над диагональю становятся .
Шаг 3. Softmax по строкам.
- Строка 0: разрешён один ключ, вес .
- Строка 1: , , сумма ; веса .
- Строка 2: (дважды), , сумма ; веса .
P = 1.0000 0 0 0.3302 0.6698 0 0.2483 0.2483 0.5035Шаг 4. Выход :
- ;
- ;
- .
Токен 2 больше всего смотрит на себя: его запрос сильнее всего совпадает с ключом . Проверка в PyTorch:
import math, torch
Q = torch.tensor([[1., 0.], [0., 1.], [1., 1.]])K = Q.clone()V = torch.tensor([[1., 0.], [0., 2.], [3., 3.]])
scores = Q @ K.T / math.sqrt(2)mask = torch.tril(torch.ones(3, 3, dtype=torch.bool))weights = torch.softmax(scores.masked_fill(~mask, float("-inf")), dim=-1)print(weights @ V)# tensor([[1.0000, 0.0000],# [0.3302, 1.3395],# [1.7587, 2.0070]])(Последний знак отличается от ручного счёта из-за округления промежуточных весов.)
Маскирование
Заголовок раздела «Маскирование»Маска в формуле решает, какие пары «запрос — ключ» разрешены. Во всех моделях репозитория есть causal-маска: позиция не видит будущих позиций . Без неё модель, обучаясь предсказывать токен , видела бы его во входе. перед softmax даёт запрещённым ключам ровно нулевой вес, а разрешённые веса по-прежнему в сумме дают 1.
Почему именно и именно перед softmax, как маска выглядит со скользящим окном и с KV-кэшем, что делать с паддингом и что случится во float16, — в отдельной главе Маски.
Multi-head attention
Заголовок раздела «Multi-head attention»Разбиение на головы
Заголовок раздела «Разбиение на головы»Одна голова даёт на каждую позицию одно распределение весов. Многоголовое внимание (multi-head attention, MHA) считает голов параллельно, каждую — со своими проекциями размера :
где:
- — номер головы, — число голов (
num_heads); - — проекции головы ;
- — выход головы;
- склеивает выходы голов по последнему измерению: ;
- — выходная проекция, возвращающая результат в размерность модели;
- — выход слоя; к нему затем прибавляется residual-связь (см. Нормализация).
В коде головы не хранятся отдельно. Матрицы всех голов поставлены рядом по столбцам: — это один nn.Linear(emb_size, num_heads * head_size). Одно умножение даёт запросы всех голов сразу, а reshape разрезает последнее измерение на последовательных кусков по — ровно :
x [B, T, d]self._q(x) [B, T, H·d_h] одно умножение на все головы.reshape(B, T, H, d_h) [B, T, H, d_h] разрезали на головы.transpose(1, 2) [B, H, T, d_h] головы — как ещё одно «батчевое» измерениеq @ k.transpose(-2,-1) [B, H, T, T_kv] оценки всех голов одним matmulweights @ v [B, H, T, d_h].transpose(1, 2) [B, T, H, d_h].reshape(B, T, H·d_h) [B, T, H·d_h] это и есть Concatself._layer(...) [B, T, d] W_OСхема в виде блоков — в gpt.md.
Зачем несколько голов
Заголовок раздела «Зачем несколько голов»Softmax одной строки — одно распределение. Если токену нужно одновременно взять информацию из двух мест (например, из предыдущего слова и из подлежащего в начале предложения), одна голова вынуждена усреднить их, и оба сигнала размываются. Vaswani et al. (разд. 3.2.2) формулируют это так: несколько голов позволяют совместно обращать внимание на информацию из разных подпространств представлений в разных позициях, а при одной голове этому мешает усреднение.
С головами у каждой позиции независимых распределений внимания, каждое в своём подпространстве размера . При это почти не стоит дополнительных вычислений: число параметров и FLOPs проекций то же, что у одной головы размера . Анализ обученных моделей показывает, что головы действительно специализируются (одни смотрят на соседний токен, другие — на синтаксически связанные слова), хотя заметную часть голов можно удалить почти без потери качества (Voita et al., 2019; Michel et al., 2019).
Выходная проекция
Заголовок раздела «Выходная проекция»Зачем , если можно просто склеить головы? Разобьём на блоки строк по , . По правилу умножения блочных матриц
То есть складывает вклады голов, предварительно переводя каждый из подпространства в общее пространство модели . Без неё координаты выхода всегда принадлежали бы голове 1 и так далее — головы не смешивались бы до следующего слоя. Кроме того, нужна, когда .
Когда H·d_h ≠ d
Заголовок раздела «Когда H·d_h ≠ d»Обычно , и . Но это не обязательно: отображает , — обратно , и ничто не требует равенства. Пример — Gemma 7B: голов по при , так что . Проекции Q, K, V расширяют пространство до 4096, а проецирует обратно в 3072 (см. gemma.md).
В конфиге это ключ head_size; без него head_size = embed_dim // число голов, и тогда embed_dim обязан делиться на число голов (функция resolve_head_size в core/config_checks.py; при RoPE она также требует чётного head_size).
Сколько это стоит
Заголовок раздела «Сколько это стоит»Параметры
Заголовок раздела «Параметры»При MHA каждая из четырёх матриц имеет элементов. Обозначим через число параметров слоя attention (буква в этой главе занята матрицей весов внимания). При :
Для GPT-1 (, bias есть) — параметров на слой; для LLaMA 7B (, без bias) — .
Если у K и V только голов (GQA, см. ниже), то , и
без учёта bias. При это : от (MHA, ) до (MQA, ).
| Модель | / | Параметров attention на слой | ||
|---|---|---|---|---|
| LLaMA 7B | 4096 | 32 / 32 | 128 | 67 108 864 |
| Mistral 7B | 4096 | 32 / 8 | 128 | 41 943 040 (при MHA было бы 67 108 864) |
| Gemma 2B | 2048 | 8 / 1 | 256 | 9 437 184 |
| Gemma 7B | 3072 | 16 / 16 | 256 | 50 331 648 () |
Проверка в коде:
from llm.core.multi_head_attention import MultiHeadAttention
attn = MultiHeadAttention(num_heads=8, emb_size=256, head_size=32, max_seq_len=64)print(sum(p.numel() for p in attn.parameters())) # 263168 = 4·256² + 4·256Вычисления и память
Заголовок раздела «Вычисления и память»Считаем умножение с последующим сложением за 2 FLOPs (операции с плавающей точкой), , одна последовательность длины без кэша:
| Операция | Форма | FLOPs |
|---|---|---|
| проекции Q, K, V, O | 4 умножения | |
| оценки | умножений | |
| взвешенная сумма | умножений | |
| softmax, маска, масштаб | элементов |
Главное:
- Ядро внимания стоит — квадратично по длине, в отличие от проекций и FFN, которые линейны по . Отношение ядра к проекциям равно . Для GPT-1 (, ) это 1/3 — проекции дороже; для и — уже 4.
- Матрица весов занимает памяти на каждую голову: тензор
[B, H, T, T]. Для , во float16 это байт ГиБ на слой и на одну последовательность — и при обучении её (вместе с оценками) нужно хранить до обратного прохода. - Causal-маска обнуляет почти половину матрицы, но явная реализация всё равно вычисляет её целиком.
Квадратичная память — главное практическое ограничение длины контекста; от неё избавляют эффективные реализации (см. ниже). Квадратичные вычисления ограничивает скользящее окно.
Dropout в attention
Заголовок раздела «Dropout в attention»В attention используются два разных dropout:
- На весах внимания — сразу после softmax: . Случайно выключает отдельные связи «токен → токен», чтобы модель не полагалась на одну позицию. Так делает GPT-1 (разд. 4.1 статьи: «attention dropouts» с вероятностью 0,1;
attn_pdropв HuggingFace) и GPT-2. При обучении оставшиеся веса умножаются на , поэтому строка в сумме даёт 1 лишь в среднем. В репозитории — параметрattention_dropoutуMultiHeadAttention(ключ конфигаattention_dropoutу GPT и GPT-2, по умолчанию 0,0 — dropout выключен и не расходует генератор случайных чисел). - После выходной проекции — перед residual-сложением (
resid_pdropв GPT). В репозитории — параметрdropoutу всех трёх классов attention (ключ конфигаdropout).
LLaMA, Mistral, Mixtral и Gemma при предобучении dropout не используют; в репозитории у них есть только второй вид, управляемый ключом dropout. У GroupedQueryAttention и MultiQueryAttention dropout на весах внимания нет.
Виды по числу голов K/V: MHA, GQA, MQA
Заголовок раздела «Виды по числу голов K/V: MHA, GQA, MQA»Головы Q всегда свои. Меняется только то, сколько отдельных K и V на них приходится. Пусть — число голов Q, — число голов K/V:
%%{init: {"flowchart": {"rankSpacing": 24, "nodeSpacing": 16}}}%%
flowchart TB
subgraph MHA["MHA · G = H"]
direction TB
a1["Q₁"]:::blue --- b1["K/V₁"]:::gold
a2["Q₂"]:::blue --- b2["K/V₂"]:::gold
a3["Q₃"]:::blue --- b3["K/V₃"]:::gold
a4["Q₄"]:::blue --- b4["K/V₄"]:::gold
end
subgraph GQA["GQA · 1 < G < H"]
direction TB
c1["Q₁"]:::blue --- d1["K/V₁"]:::gold
c2["Q₂"]:::blue --- d1
c3["Q₃"]:::blue --- d2["K/V₂"]:::gold
c4["Q₄"]:::blue --- d2
end
subgraph MQA["MQA · G = 1"]
direction TB
e1["Q₁"]:::blue --- f1["K/V₁"]:::gold
e2["Q₂"]:::blue --- f1
e3["Q₃"]:::blue --- f1
e4["Q₄"]:::blue --- f1
end
MHA ~~~ GQA ~~~ MQA
classDef blue fill:#dae8fc,stroke:#6c8ebf,color:#1a1a1a;
classDef gold fill:#fff2cc,stroke:#d6b656,color:#1a1a1a;
| Вид | Голов K/V | Статья | Модели |
|---|---|---|---|
| MHA — Multi-Head Attention | : у каждой головы Q свои K и V | Vaswani et al., 2017 | GPT-1, GPT-2, LLaMA-1, Gemma 7B |
| GQA — Grouped Query Attention | : одна пара K/V на группу из голов Q | Ainslie et al., 2023 | Mistral 7B (32 Q, 8 K/V), Mixtral 8x7B, LLaMA-2 70B |
| MQA — Multi-Query Attention | : одна пара K/V на все головы Q | Shazeer, 2019 | Gemma 2B (8 Q, 1 K/V), PaLM |
Формально при GQA голова (нумерация с 0) использует K/V-голову номер
где — размер группы (сколько голов Q делят одну пару K/V), а — округление вниз. Для , : головы Q 0–3 используют K/V-голову 0, головы 4–7 — K/V-голову 1.
MHA и MQA — крайние случаи GQA: () и (). Поэтому в репозитории один класс GroupedQueryAttention описывает все три, а num_q_heads должно делиться на num_kv_heads (иначе конструктор бросает ValueError).
История: MQA предложил Shazeer (2019), чтобы ускорить инкрементальную генерацию: при декодировании по одному токену время уходит не на арифметику, а на чтение K и V всех прошлых токенов из памяти, и уменьшение их объёма в раз прямо ускоряет шаг. Ainslie et al. (2023) заметили, что MQA теряет в качестве, и предложили промежуточный вариант — GQA; они же показали, что готовую MHA-модель можно превратить в GQA, усреднив K/V-головы внутри группы и дообучив на небольшой доле исходного объёма данных.
Как K/V-головы доходят до голов Q
Заголовок раздела «Как K/V-головы доходят до голов Q»Чтобы использовать обычное батчевое умножение q @ k.transpose(-2, -1) с формой [B, H, T, d_h], K и V нужно привести к головам. GroupedQueryAttention._repeat_kv_heads повторяет каждую K/V-голову раз подряд:
kv: [B, G, T, d_h].unsqueeze(2) [B, G, 1, T, d_h].repeat(1, 1, H/G, 1, 1) [B, G, H/G, T, d_h].reshape(B, H, T, d_h) [B, H, T, d_h]при H = 8, G = 2: [KV₀, KV₁] → [KV₀, KV₀, KV₀, KV₀, KV₁, KV₁, KV₁, KV₁]Порядок «сначала все копии головы 0, потом головы 1» реализует и совпадает с repeat_interleave(H // G, dim=1). Он важен при загрузке чужих весов: строки должны быть сгруппированы так же.
При копирования нет. Тензор K формы [B, 1, T_kv, d_h] участвует в умножении [B, H, T, d_h] @ [B, 1, d_h, T_kv] напрямую: по правилам трансляции (broadcasting) PyTorch измерение размера 1 растягивается на без выделения памяти. Так же устроен MultiQueryAttention, поэтому Gemma с num_kv_heads: 1 даёт побитово тот же результат, что и MultiQueryAttention с теми же весами (это проверено: torch.equal на выходах).
Вычислений в самом GQA не экономит: каждая голова Q по-прежнему считает свои веса по всем ключам. Экономятся проекции , и, главное, память кэша.
Зачем делить K/V: размер KV-кэша
Заголовок раздела «Зачем делить K/V: размер KV-кэша»При генерации каждый новый токен смотрит на K и V всех предыдущих, и их хранят в KV-кэше, чтобы не пересчитывать (подробно — ниже). На один токен в одном слое кэш — чисел (K и V), то есть он пропорционален числу голов K/V, а не Q. Для контекста 4096 токенов во float16 (2 байта на число), одна последовательность:
| Модель | Голов Q / K/V | Слоёв | KV-кэш на 4096 токенов | Был бы при MHA | |
|---|---|---|---|---|---|
| LLaMA 7B | 32 / 32 (MHA) | 128 | 32 | 2 ГиБ | 2 ГиБ |
| Mistral 7B, Mixtral 8x7B | 32 / 8 (GQA) | 128 | 32 | 512 МиБ | 2 ГиБ |
| Gemma 2B | 8 / 1 (MQA) | 256 | 18 | 72 МиБ | 576 МиБ |
| Gemma 7B | 16 / 16 (MHA) | 256 | 28 | 1,75 ГиБ | 1,75 ГиБ |
Например, для Mistral 7B: байт МиБ.
Меньше кэш — больше последовательностей в батче и длиннее контекст на той же памяти, а генерация, которая упирается в чтение кэша из памяти, идёт быстрее. Цена — качество: у MQA все головы Q читают одни и те же K и V. GQA — компромисс: по качеству близка к MHA, по скорости — к MQA (Ainslie et al.). Поэтому MQA встречается в маленьких моделях (Gemma 2B), а GQA стала стандартом для больших.
Скользящее окно
Заголовок раздела «Скользящее окно»Эта ось не зависит от числа голов. С полным (causal) вниманием токен видит всё прошлое. Со скользящим окном (sliding window attention; Longformer, Mistral 7B v0.1) — только ближайшее прошлое. В репозитории пара «запрос , ключ » разрешена, если
где — ширина окна (window_size), — абсолютные позиции. Условие — это causal-маска, — отсечение дальнего прошлого. Токен видит позиций вместе с собой. В HuggingFace то же окно определено как , то есть на одну позицию уже, — почему так и как это учитывать при загрузке весов, см. mistral.md. Картинки масок — в главе Маски.
Стоимость. Каждая строка матрицы оценок содержит не больше разрешённых элементов, поэтому полезная работа ядра внимания — вместо . (Явная реализация в репозитории этим не пользуется при обработке всей последовательности сразу: она считает полную матрицу и маскирует лишнее. Экономия появляется при генерации — за счёт короткого кэша.)
Рецептивное поле. За один слой информация перемещается максимум на позиций назад. Но выход слоя на позиции уже содержит информацию о позициях до , поэтому слой , глядя на , косвенно видит и их. По индукции после слоёв позиция зависит от позиций до . Для Mistral 7B (, ) это позиции — «теоретический охват внимания около 131K токенов» (Jiang et al., 2023, разд. 2). Дальние зависимости передаются через слои, хотя и с потерями.
Кэш ограничен окном. Ключи старше позиций больше никогда не понадобятся, поэтому их можно выбросить. GroupedQueryAttention хранит в кэше только последние позиций K и V; вместе с новым токеном это ровно видимых позиций. Размер кэша перестаёт расти с длиной текста: у Mistral 7B он не превышает 512 МиБ (как в таблице выше при 4096 токенах) при любой длине генерации. В статье Mistral это «кольцевой буфер» (rolling buffer cache); здесь — torch.cat и обрезка срезом, результат тот же.
В репозитории окно — необязательный ключ window_size у Mistral и Mixtral; без него внимание полное. Скользящее окно есть только в Mistral 7B v0.1; в Mixtral 8x7B и в Mistral v0.2+ его нет.
Что кэшируется и почему это корректно
Заголовок раздела «Что кэшируется и почему это корректно»При генерации модель выдаёт текст по одному токену: вычисляет логиты для последней позиции, выбирает токен, дописывает его и повторяет (см. Генерация). Наивно на каждом шаге вся последовательность прогоняется заново. Но почти вся эта работа повторяется: для старых позиций получаются те же самые K и V.
Утверждение. В causal-модели скрытое состояние любого слоя на позиции зависит только от токенов .
Доказательство по индукции по слоям. На входе (эмбеддинги) состояние позиции зависит только от и позиции . Пусть утверждение верно для входа слоя . Attention на позиции смешивает значения только с позиций (causal-маска), а каждое из них по предположению зависит от . Нормализация, FFN и residual-связи работают с каждой позицией отдельно. Значит, выход слоя на позиции тоже зависит только от .
Следствие: дописывание новых токенов не меняет K и V старых позиций ни в одном слое. Их можно посчитать один раз и хранить. Это и есть KV-кэш: для каждого слоя — тензоры K и V всех уже обработанных позиций. Запросы Q старых позиций не нужны: на шаге генерации нужен выход только новой позиции, и её запрос сравнивается со всеми ключами.
В двунаправленных моделях (BERT) это неверно: там старые позиции смотрят и на новые, поэтому их K и V меняются при каждом добавлении токена.
Тонкость: утверждение предполагает, что позиции токенов не меняются. Когда текст перерастает max_position_embeddings и модель берёт последние токенов, абсолютные позиции всех токенов сдвигаются, и закэшированные K (с позиционной информацией внутри) устаревают. Поэтому generate в этот момент сбрасывает кэш и пересчитывает окно без него (next_generation_input в core/generation.py).
Объём кэша
Заголовок раздела «Объём кэша»где:
- — отдельно K и V;
- — число слоёв (у каждого слоя свой кэш);
- — число голов K/V, — размер головы;
- — число закэшированных позиций (со скользящим окном — не больше );
- — число последовательностей в батче;
- — байт на число (2 для float16/bfloat16, 4 для float32).
Кэш растёт линейно с длиной контекста и размером батча и при длинных контекстах легко превышает объём весов модели. Отсюда интерес к уменьшению (GQA, MQA) и (скользящее окно).
Prefill и decode
Заголовок раздела «Prefill и decode»Генерация с кэшем состоит из двух фаз:
- Prefill (заполнение). Весь промпт длины обрабатывается за один проход, как при обучении: матрица оценок с causal-маской. Побочный результат — K и V всех позиций промпта во всех слоях, то есть заполненный кэш. Эта фаза упирается в вычисления: много токенов, большие матричные умножения.
- Decode (декодирование). На каждом шаге подаётся один новый токен (). Для него считаются (стоимость на слой), новые дописываются в кэш, и запрос сравнивается со всеми ключами — одна строка матрицы оценок, . Без кэша этот шаг стоил бы . Эта фаза упирается в пропускную способность памяти: на каждом шаге нужно прочитать весь кэш и все веса ради небольшого объёма арифметики.
Длинный промпт можно заполнять и кусками по несколько токенов (chunked prefill). Тогда внутри куска снова нужна causal-маска — новые токены не должны видеть друг друга «вперёд»; как берётся нужный кусок маски, см. Маски.
import torchfrom llm.core.group_query_attention import GroupedQueryAttention
torch.manual_seed(0)attn = GroupedQueryAttention(num_q_heads=8, num_kv_heads=2, emb_size=256, head_size=32, max_seq_len=64, window_size=4, dropout=0.0).eval()x = torch.randn(1, 10, 256)with torch.no_grad(): full, _ = attn(x) # весь текст за один проход out, cache = attn(x[:, :7]) # prefill: 7 токенов промпта print(cache[0].shape, cache[2]) # torch.Size([1, 2, 4, 32]) 7 outs = [out] for t in range(7, 10): # decode: по одному токену out, cache = attn(x[:, t:t + 1], cache=cache) outs.append(out)print(torch.allclose(torch.cat(outs, dim=1), full, atol=1e-5)) # TrueКэш хранит головы (а не 8) и только позиции, хотя обработано 7 токенов; результат совпадает с прогоном без кэша.
Формат кэша в репозитории
Заголовок раздела «Формат кэша в репозитории»Кэш модели — список по слоям; элемент списка — кэш одного слоя:
| Класс | Кэш слоя | Формы | Позиция следующего токена |
|---|---|---|---|
MultiHeadAttention | (K, V) | [B, H, T_cache, d_h] | K.size(2) — длина кэша |
MultiQueryAttention | (K, V) | [B, 1, T_cache, d_h] | K.size(2) |
GroupedQueryAttention | (K, V, next_pos) | [B, G, T_cache, d_h], next_pos — int | next_pos |
Зачем третий элемент. Без окна длина кэша равна числу обработанных токенов, то есть абсолютной позиции следующего. С окном кэш обрезается до позиций, и длина кэша перестаёт совпадать с позицией: после 7 токенов при в кэше 4 позиции, а следующий токен — седьмой (с нуля). Позиция же нужна RoPE (start_pos) и маске. Поэтому GroupedQueryAttention хранит её явно. Функция cache_start_pos в core/generation.py понимает оба формата: берёт cache[0][2], если элементов три, иначе cache[0][0].size(2).
Ещё две детали:
- В кэш кладутся K после RoPE: поворот зависит только от позиции ключа, поэтому его не нужно повторять.
- В
GroupedQueryAttentionкэш хранится до_repeat_kv_heads, то есть с , а не головами — ради этого GQA и нужна.
Как в attention попадает позиция
Заголовок раздела «Как в attention попадает позиция»Скалярное произведение само по себе порядка не знает: если переставить токены, веса переставятся вместе с ними, и выход каждого токена не изменится (attention эквивариантно к перестановкам). Causal-маска вносит частичную информацию о порядке, но недостаточную. Поэтому позиция подаётся явно (подробно — в главе Позиционное кодирование):
- GPT-1, GPT-2 прибавляют обучаемый эмбеддинг позиции к эмбеддингу токена на входе — attention получает позицию косвенно, через .
- LLaMA, Mistral, Mixtral, Gemma поворачивают и внутри attention на угол, зависящий от позиции (RoPE): тогда зависит только от разности . V не поворачивается. Схема головы с RoPE — в llama.md.
Реализация в репозитории
Заголовок раздела «Реализация в репозитории»| Класс | Файл | Головы K/V | RoPE | Окно | KV-кэш слоя | Модели |
|---|---|---|---|---|---|---|
MultiHeadAttention | core/multi_head_attention.py | = num_heads | необязательно | нет | (K, V) | GPT, GPT-2 (без RoPE), LLaMA (с RoPE) |
GroupedQueryAttention | core/group_query_attention.py | num_kv_heads | необязательно | window_size или нет | (K, V, next_pos) | Mistral, Mixtral, Gemma |
MultiQueryAttention | core/multi_query_attention.py | 1 | необязательно | нет | (K, V) | учебный модуль, моделями не используется |
Модули attention создаются внутри блоков декодера: GptDecoder, Gpt2Decoder, CachedDecoder (LLaMA) — MultiHeadAttention; MistralDecoder, MixtralDecoder, GemmaDecoder — GroupedQueryAttention.
Прямой проход по шагам
Заголовок раздела «Прямой проход по шагам»Все три класса устроены одинаково. Ниже — MultiHeadAttention.forward (вход x формы [B, T, d], необязательный cache) и соответствие формулам этой главы:
| Шаг | Код | Формула / смысл |
|---|---|---|
| 1. Позиция первого нового токена | start_pos = cache[0].size(2) if cache is not None else 0 | абсолютная позиция; проверка start_pos + seq_len > max_seq_len → ValueError |
| 2. Проекции | q = self._q(x), k = self._k(x), v = self._v(x) | , , , форма [B, T, H·d_h] |
| 3. Разбиение на головы | .reshape(B, T, H, d_h), .transpose(1, 2) | , форма [B, H, T, d_h] |
| 4. Позиция | q = self._rope(q, start_pos=start_pos, positions=positions), то же для k | RoPE для Q и K (если задан); при паддинге позиции — padding.positions |
| 5. Кэш | k = torch.cat([k_cache, k], dim=2), то же для v | K и V всех позиций 0 … start_pos + T − 1 |
| 6. Оценки | scores = q @ k.transpose(-2, -1) / (self._head_size ** 0.5) | , форма [B, H, T, T_kv] |
| 7. Маска | causal_mask = self._tril_mask[start_pos:start_pos + seq_len, :start_pos + seq_len]; при паддинге causal_mask = padding.apply(causal_mask, start_pos, key_start=0); scores.masked_fill(~causal_mask, float("-inf")) | |
| 8. Softmax и dropout весов | weights = self._attn_dropout(F.softmax(scores, dim=-1)) | по строкам |
| 9. Взвешенная сумма | x_out = weights @ v | |
| 10. Склейка голов | .transpose(1, 2).contiguous().reshape(B, T, H·d_h) | |
| 11. Выходная проекция и dropout | self._dropout(self._layer(...)) | |
| 12. Возврат | (final_output, (k, v)) или (final_output, None) | новый кэш слоя при use_cache=True |
Маска _tril_mask — нижнетреугольная булева матрица [max_seq_len, max_seq_len], построенная один раз в конструкторе (torch.tril) и зарегистрированная как буфер с persistent=False: она не попадает в чекпоинт.
В GroupedQueryAttention.forward те же шаги со следующими отличиями:
- шаг 1:
start_pos = cache[2]; - шаг 2:
self._kиself._vпроецируют вnum_kv_heads * head_size, а не вnum_q_heads * head_size; - после шага 5 — шаг 5а: при
num_kv_heads == 1K и V транслируются как есть, иначе_repeat_kv_headsдоводит их до голов; - шаг 7: маска построена
_create_sliding_window_mask(условие ; без окна вместо подставляетсяmax_seq_len, и это обычная causal-маска), а срез столбцов начинается сstart_pos - cache_len— с позиции самого старого ключа в кэше; - шаг 8: dropout на весах внимания нет;
- шаг 12: при
window_sizeK и V обрезаются до последнихwindow_sizeпозиций, возвращается(k, v, start_pos + seq_len).
Паддинг. Все три класса принимают необязательный padding — Padding(key_mask, positions) из core/padding.py, который модель строит по attention_mask. Позиции идут в RoPE вместо start_pos, start_pos + 1, …, а padding.apply добавляет маску ключей к causal-маске и окну: маска становится [B, 1, T, T_kv], своей у каждой строки батча. Подробно — в Маски.
MultiQueryAttention.forward отличается от MHA тем, что K и V проецируются в одну голову (nn.Linear(emb_size, head_size)) и транслируются на все головы Q на шаге 6.
Параметры конструкторов и ключи конфига
Заголовок раздела «Параметры конструкторов и ключи конфига»MultiHeadAttention:num_heads,emb_size,head_size,max_seq_len,rope,dropout,attention_dropout(на весах после softmax;attn_pdropGPT-1/GPT-2),bias(по умолчаниюTrue;Llamaпередаёт значение ключа конфигаbias, для архитектуры статьи —False).GroupedQueryAttention:num_q_heads,num_kv_heads,emb_size,head_size,max_seq_len,window_size,rope,dropout,bias(по умолчаниюTrue; модели передают ключ конфигаbias, в оригинальных Mistral, Mixtral и Gemma bias нет).MultiQueryAttention:num_q_heads,emb_size,head_size,max_seq_len,rope,dropout; bias у проекций всегда есть.
Ключи конфига моделей: GPT, GPT-2, LLaMA — num_heads; Mistral и Mixtral — num_q_heads и num_kv_heads; Gemma — num_q_heads и необязательный num_kv_heads (по умолчанию 1 — MQA, как Gemma 2B; num_kv_heads = num_q_heads — MHA, как Gemma 7B). Во всех моделях — необязательный head_size.
Внешняя attention_mask (паддинг) в модули attention не передаётся: её проверяет forward модели, см. Маски.
Проверка эквивалентностей, упомянутых в главе (веса копируются через load_state_dict): GroupedQueryAttention с num_kv_heads = num_q_heads совпадает с MultiHeadAttention (до 1e-6), с num_kv_heads = 1 — побитово с MultiQueryAttention; заполнение кэша кусками со скользящим окном совпадает с прогоном всей последовательности.
Эффективные реализации
Заголовок раздела «Эффективные реализации»Явная формула из этой главы материализует матрицу оценок [B, H, T, T_kv] и матрицу весов той же формы. При длинном контексте это гигабайты на слой (см. выше), а главное — каждое чтение и запись этих матриц идёт через сравнительно медленную память GPU (HBM), и время уходит на пересылку данных, а не на арифметику.
FlashAttention (Dao et al., 2022) вычисляет тот же самый результат, не храня матрицу целиком:
- , , разбиваются на блоки, которые помещаются в быструю память на кристалле (SRAM);
- softmax считается «онлайн»: для каждой строки поддерживаются текущий максимум и текущая сумма экспонент, и при обработке очередного блока ключей накопленный результат пересчитывается с новым максимумом;
- при обратном проходе матрица весов не читается из памяти, а пересчитывается по блокам.
Память на матрицу внимания становится вместо , а число обращений к HBM сокращается в разы; FLOPs остаются , результат совпадает с точностью до округлений (это точный алгоритм, а не приближение).
В PyTorch (начиная с 2.0) есть torch.nn.functional.scaled_dot_product_attention: он сам выбирает ядро — FlashAttention, memory-efficient attention или обычную реализацию — в зависимости от устройства, типа данных и аргументов:
import torchimport torch.nn.functional as F
q, k, v = torch.randn(3, 1, 4, 6, 8).unbind(0) # [B, H, T, d_h]mask = torch.tril(torch.ones(6, 6, dtype=torch.bool))ref = torch.softmax((q @ k.transpose(-2, -1) / 8 ** 0.5).masked_fill(~mask, float("-inf")), -1) @ vout = F.scaled_dot_product_attention(q, k, v, is_causal=True)print(torch.allclose(out, ref, atol=1e-6)) # TrueВ репозитории внимание написано явно — q @ k.transpose, masked_fill, softmax, weights @ v — ради наглядности: каждый шаг соответствует строке формулы, и промежуточные scores и weights можно напечатать и нарисовать. Для обучения на длинных контекстах эти строки заменяют одним вызовом F.scaled_dot_product_attention.
Типичные ошибки и тонкости
Заголовок раздела «Типичные ошибки и тонкости»- Забыть масштаб или взять не тот. Делить нужно на (размер головы), а не на (размер модели) — иначе при оценки окажутся сжатыми.
- Softmax не по той оси. Нормировать нужно по ключам (
dim=-1), чтобы каждый запрос получил распределение. При softmax по запросам (dim=-2) вес пары зависел бы от оценок более поздних запросов — это уже не внимание, и вдобавок утечка информации из будущего. reshapeбезtranspose. Перед склейкой голов нужно вернуть оси[B, H, T, d_h] → [B, T, H, d_h];reshapeбез этого перемешает токены и головы, не вызвав ошибки. Послеtransposeтензор не непрерывен в памяти, поэтому передreshapeстоит.contiguous()(или используетсяreshape, который сам сделает копию).- Маска без учёта кэша. При
cacheстроки маски — это позицииstart_pos …, а не0 …. Если взять_tril_mask[:T, :T], новый токен увидит не те ключи. - Позиция из длины кэша при окне. Со скользящим окном длина кэша ≠ позиция токена; из-за этого в
GroupedQueryAttentionпоявилсяnext_pos. - Кэш после сдвига окна
max_seq_len. Позиции всех токенов меняются, кэш нужно сбросить. - GQA не уменьшает FLOPs ядра. Экономия — в памяти кэша и в проекциях.
- Внимание — мягкий поиск по словарю: запрос сравнивается со всеми ключами, softmax даёт веса, результат — взвешенная сумма значений.
- . Деление на держит дисперсию оценок около 1; без него softmax насыщается и его градиент почти исчезает.
- Multi-head: голов в подпространствах размера , склейка и ; в коде — одна проекция и
reshape. может не равняться (Gemma 7B). - Параметры: при MHA, при GQA. Вычисления ядра , память на веса .
- MHA, GQA и MQA различаются числом голов K/V; это определяет размер KV-кэша .
- Скользящее окно ограничивает кэш позициями, а рецептивное поле через слоёв — .
- KV-кэш корректен, потому что при causal-маске прошлые K и V не зависят от будущих токенов. Генерация — prefill промпта и decode по одному токену.
- В репозитории — три явные реализации (
MultiHeadAttention,GroupedQueryAttention,MultiQueryAttention); на практике используют FlashAttention иF.scaled_dot_product_attention.
Вопросы и упражнения
Заголовок раздела «Вопросы и упражнения»-
В численном примере замените на , оставив и . Найдите выход токена 2.
Ответ
Веса строки 2 не зависят от : . Выход: .
-
Компоненты и независимы, со средним 0 и дисперсией . Чему равна дисперсия ? Что это говорит о роли нормализации входа слоя?
Ответ
, сумма слагаемых — , после деления на — . Масштаб убирает зависимость от размера головы, но не от масштаба самих векторов: если нормы Q и K вырастут, softmax снова насытится. Поэтому важно, что вход attention нормализован (pre-LN, см. Нормализация).
-
Посчитайте число параметров attention одного слоя Mistral 7B (, , , , без bias) и сравните с MHA той же ширины.
Ответ
. При MHA — ; GQA экономит 37,5 % параметров attention.
-
Сколько памяти займёт KV-кэш Gemma 2B (18 слоёв, , ) на 8192 токена во float16 для батча из 4 последовательностей? А у Mistral 7B со скользящим окном на 32 768 токенов, одна последовательность?
Ответ
Gemma 2B: байт МиБ. Mistral 7B: в кэше не больше позиций, поэтому байт МиБ независимо от длины; без окна было бы МиБ ГиБ.
-
Покажите, что якобиан softmax вырожден: его строки в сумме по дают 0. Какой смысл у этого факта?
Ответ
. Смысл: прибавление одной и той же константы ко всем оценкам строки не меняет softmax, поэтому производная вдоль направления равна нулю. По той же причине маскирование можно делать и большим отрицательным числом, и , а softmax в коде стабилизируют вычитанием максимума строки.
-
Для Mistral 7B () и учебного конфига с , : на сколько позиций назад теоретически может дотянуться информация к последнему слою?
Ответ
: для Mistral 7B и для учебного конфига.
-
Почему в
GroupedQueryAttentionприnum_kv_heads == 1не вызывается_repeat_kv_heads, и почему результат от этого не меняется?Ответ
При умножении
[B, H, T, d_h] @ [B, 1, d_h, T_kv]измерение голов размера 1 транслируется на — каждая голова Q умножается на одни и те же K. Результат тот же, что после явного копирования, но без выделения копий K и V в памяти. -
(Для размышления.) Энкодер BERT видит последовательность в обе стороны. Можно ли ускорить его пошаговую обработку KV-кэшем так же, как декодер? Почему?
Ответ
Нет. Утверждение о независимости прошлых состояний от будущих токенов опирается на causal-маску. Без неё добавление токена меняет веса внимания всех старых позиций, а значит, и их скрытые состояния во всех слоях, кроме первого; закэшированные K и V устаревают.
Литература
Заголовок раздела «Литература»- Sutskever, Vinyals, Le. Sequence to Sequence Learning with Neural Networks. 2014. arXiv:1409.3215 — encoder-decoder на RNN
- Cho et al. Learning Phrase Representations using RNN Encoder-Decoder for Statistical Machine Translation. 2014. arXiv:1406.1078 — encoder-decoder на RNN
- Bahdanau, Cho, Bengio. Neural Machine Translation by Jointly Learning to Align and Translate. ICLR 2015. arXiv:1409.0473 — внимание в машинном переводе
- Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762 — scaled dot-product и multi-head attention (разд. 3.2)
- Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019. arXiv:1911.02150 — MQA
- Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023. arXiv:2305.13245 — GQA
- Beltagy, Peters, Cohan. Longformer: The Long-Document Transformer. 2020. arXiv:2004.05150 — sliding window attention
- Jiang et al. Mistral 7B. 2023. arXiv:2310.06825 — GQA + скользящее окно, rolling buffer cache
- Su et al. RoFormer: Enhanced Transformer with Rotary Position Embedding. 2021. arXiv:2104.09864 — RoPE
- Voita, Talbot, Moiseev, Sennrich, Titov. Analyzing Multi-Head Self-Attention: Specialized Heads Do the Heavy Lifting, the Rest Can Be Pruned. 2019. arXiv:1905.09418 — специализация голов
- Michel, Levy, Neubig. Are Sixteen Heads Really Better than One? 2019. arXiv:1905.10650 — избыточность голов
- Dao, Fu, Ermon, Rudra, Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. 2022. arXiv:2205.14135 — FlashAttention