Перейти к содержимому

Механизм внимания

Внимание (attention) — центральная операция трансформера: именно через неё токены обмениваются информацией. Все остальные слои блока декодера (нормализация, FFN) обрабатывают каждую позицию отдельно. Все шесть моделей репозитория используют одно и то же causal self-attention и отличаются тремя независимыми «ручками»:

  1. сколько голов K/V приходится на головы Q: MHA, GQA или MQA;
  2. какие позиции видит токен: всё прошлое или только скользящее окно;
  3. как в attention попадает позиция: через слагаемое к эмбеддингам (GPT) или поворотом Q и K (RoPE).
  • Откуда взялось внимание и почему его описывают как «мягкий поиск по словарю» с запросами, ключами и значениями.
  • Формулу scaled dot-product attention, вывод множителя 1/dh1/\sqrt{d_h} и что ломается без него.
  • Как устроено многоголовое внимание (multi-head attention), сколько у него параметров и сколько оно стоит по времени и памяти.
  • Чем отличаются MHA, GQA и MQA и почему это в первую очередь вопрос размера KV-кэша.
  • Как работают скользящее окно и KV-кэш и как всё это реализовано в llm/core.

До трансформеров машинный перевод решали моделями «кодировщик — декодировщик» (encoder-decoder) на рекуррентных сетях (Sutskever et al., 2014; Cho et al., 2014). Кодировщик читал исходное предложение слово за словом и сжимал его в один вектор фиксированной длины — последнее скрытое состояние RNN. Декодировщик генерировал перевод, опираясь только на этот вектор.

Узкое место очевидно: предложение из 5 слов и предложение из 50 слов упаковываются в вектор одного и того же размера. Bahdanau, Cho, Bengio (2015) показали, что качество такого перевода падает с ростом длины предложения, и предложили не сжимать вход в один вектор. Кодировщик сохраняет скрытые состояния h1,…,hn\mathbf{h}_1, \dots, \mathbf{h}_n всех слов, а декодировщик на каждом шаге ii сам выбирает, на какие из них смотреть:

eij=a(si−1,hj),αij=exp⁡(eij)∑k=1nexp⁡(eik),ci=∑j=1nαijhje_{ij} = a(\mathbf{s}_{i-1}, \mathbf{h}_j), \qquad \alpha_{ij} = \frac{\exp(e_{ij})}{\sum_{k=1}^{n} \exp(e_{ik})}, \qquad \mathbf{c}_i = \sum_{j=1}^{n} \alpha_{ij} \mathbf{h}_j

где:

  • si−1\mathbf{s}_{i-1} — скрытое состояние декодировщика перед шагом ii («что я сейчас ищу»);
  • hj\mathbf{h}_j — состояние кодировщика для jj-го слова источника;
  • a(⋅,⋅)a(\cdot,\cdot) — функция сходства; у Bahdanau это маленькая сеть с одним скрытым слоем и tanh⁡\tanh (так называемое аддитивное внимание);
  • αij\alpha_{ij} — вес внимания: насколько слово jj важно на шаге ii; веса неотрицательны и в сумме дают 1;
  • ci\mathbf{c}_i — вектор контекста, взвешенное среднее состояний кодировщика.

Это и есть внимание: вместо одного фиксированного вектора — своя смесь входов для каждого шага. Vaswani et al. (2017) сделали следующий шаг: убрали рекуррентность совсем и построили модель только из внимания («Attention Is All You Need»). Функцию сходства они заменили на скалярное произведение — оно считается одним умножением матриц для всех пар сразу. Внимание, в котором запросы и ключи берутся из одной и той же последовательности, называется самовниманием (self-attention); именно оно используется во всех моделях репозитория.

Обычный словарь (Python dict) хранит пары «ключ → значение». Поиск по запросу qq находит ключ, точно равный qq, и возвращает его значение:

словарь: "кот" → 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) q\mathbf{q} — «что я ищу»;
  • Ключ (key) k\mathbf{k} — «по какому признаку меня находить»;
  • Значение (value) v\mathbf{v} — «что я отдаю, если меня нашли».

Мягкость важна по двум причинам. Во-первых, результат дифференцируем по запросам и ключам: жёсткий выбор argmax не пропускает градиент, а softmax пропускает, и модель можно обучать градиентным спуском. Во-вторых, токен может собрать информацию сразу из нескольких мест.

В self-attention каждый токен последовательности играет все три роли: он задаёт свой запрос, выставляет свой ключ и отдаёт своё значение.

Пусть на вход слоя пришли скрытые состояния TT токенов, записанные строками матрицы XX. Запросы, ключи и значения — три разные линейные проекции одного и того же входа:

Q=XWQ,K=XWK,V=XWVQ = X W_Q, \qquad K = X W_K, \qquad V = X W_V

где:

  • X∈RT×dX \in \mathbb{R}^{T \times d} — входные векторы токенов (строка tt — токен на позиции tt);
  • WQ,WK,WV∈Rd×dhW_Q, W_K, W_V \in \mathbb{R}^{d \times d_h} — обучаемые матрицы проекций (пока рассматриваем одну голову);
  • Q,K,V∈RT×dhQ, K, V \in \mathbb{R}^{T \times d_h} — запросы, ключи и значения; строки qt,kt,vt∈Rdh\mathbf{q}_t, \mathbf{k}_t, \mathbf{v}_t \in \mathbb{R}^{d_h};
  • dhd_h — размер головы (head_size).

Зачем три разные проекции, а не сам XX? Если бы запрос и ключ совпадали (Q=K=XQ = K = X), сходство xi⋅xj\mathbf{x}_i \cdot \mathbf{x}_j было бы симметричным, и больше всего токен смотрел бы на самого себя (xi⋅xi=∥xi∥2\mathbf{x}_i \cdot \mathbf{x}_i = \lVert \mathbf{x}_i \rVert^2). Раздельные WQW_Q и WKW_K позволяют отношению «ii ищет jj» быть несимметричным: глагол ищет подлежащее, а подлежащее глагол — не обязательно. Отдельная WVW_V отделяет «по чему ищут» от «что передают».

В коде с батчем все тензоры получают ведущее измерение BB: XX имеет форму [B, T, d], и nn.Linear применяет одну и ту же матрицу ко всем токенам и всем примерам. Линейные слои в репозитории могут иметь смещение (bias): Q=XWQ+bQQ = XW_Q + \mathbf{b}_Q, где bQ∈Rdh\mathbf{b}_Q \in \mathbb{R}^{d_h} прибавляется к каждой строке. В GPT-1 и GPT-2 bias есть, в оригинальных LLaMA, Mistral, Mixtral и Gemma — нет. В репозитории у LLaMA, Mistral, Mixtral и Gemma его включает ключ конфига bias, по умолчанию true (ради совместимости со старыми чекпоинтами); чтобы получить архитектуру статей, задают "bias": false.

Основная формула (Vaswani et al., 2017, разд. 3.2.1):

Attention⁡(Q,K,V)=softmax⁡ ⁣(QK⊤dh+M)V\operatorname{Attention}(Q, K, V) = \operatorname{softmax}\!\left(\frac{Q K^\top}{\sqrt{d_h}} + M\right) V

где:

  • Q∈RT×dhQ \in \mathbb{R}^{T \times d_h}, K∈RTkv×dhK \in \mathbb{R}^{T_{kv} \times d_h}, V∈RTkv×dhV \in \mathbb{R}^{T_{kv} \times d_h} — запросы, ключи, значения; TkvT_{kv} — число ключей (без кэша Tkv=TT_{kv} = T, с кэшем — длина кэша плюс TT);
  • S=QK⊤/dh∈RT×TkvS = QK^\top / \sqrt{d_h} \in \mathbb{R}^{T \times T_{kv}} — матрица оценок (scores): Sij=qi⋅kj/dhS_{ij} = \mathbf{q}_i \cdot \mathbf{k}_j / \sqrt{d_h} — сходство запроса ii с ключом jj;
  • M∈{0,−∞}T×TkvM \in \{0, -\infty\}^{T \times T_{kv}} — маска: 00 там, где смотреть можно, −∞-\infty там, где нельзя (подробно — в главе Маски);
  • softmax⁡\operatorname{softmax} применяется к каждой строке отдельно: Pij=exp⁡(Sij+Mij)/∑kexp⁡(Sik+Mik)P_{ij} = \exp(S_{ij} + M_{ij}) / \sum_{k} \exp(S_{ik} + M_{ik});
  • P∈RT×TkvP \in \mathbb{R}^{T \times T_{kv}} — матрица весов внимания; каждая строка — распределение вероятностей по ключам;
  • результат O=PV∈RT×dhO = PV \in \mathbb{R}^{T \times d_h}; строка oi=∑jPijvj\mathbf{o}_i = \sum_j P_{ij} \mathbf{v}_j.

Словами, для каждой позиции ii:

  1. сравнить её запрос со всеми разрешёнными ключами (скалярное произведение);
  2. поделить на dh\sqrt{d_h};
  3. превратить оценки в веса softmax’ом;
  4. взять взвешенную сумму значений.

Интуиция. Выход oi\mathbf{o}_i — выпуклая комбинация векторов значений: он лежит «между» ними. Внимание не создаёт новые признаки, а перераспределяет существующие между позициями; новые признаки создают проекции и FFN. Если убрать softmax и оставить QK⊤VQK^\top V, веса перестанут быть неотрицательными и нормированными, а масштаб выхода будет расти с длиной последовательности.

Название «scaled dot-product» — «масштабированное скалярное произведение» — описывает именно шаги 1–2.

Предположим, что компоненты запроса и ключа — независимые случайные величины со средним 0 и дисперсией 1 (примерно так и есть в начале обучения при стандартной инициализации и нормализованном входе). Посчитаем дисперсию их скалярного произведения.

Вывод

Скалярное произведение — сумма dhd_h слагаемых:

q⋅k=∑m=1dhqmkm\mathbf{q} \cdot \mathbf{k} = \sum_{m=1}^{d_h} q_m k_m

Шаг 1. Среднее одного слагаемого. Так как qmq_m и kmk_m независимы,

E[qmkm]=E[qm] E[km]=0⋅0=0\mathbb{E}[q_m k_m] = \mathbb{E}[q_m]\,\mathbb{E}[k_m] = 0 \cdot 0 = 0

Шаг 2. Дисперсия одного слагаемого. По определению Var⁡[Z]=E[Z2]−(E[Z])2\operatorname{Var}[Z] = \mathbb{E}[Z^2] - (\mathbb{E}[Z])^2, а среднее мы уже нашли:

Var⁡[qmkm]=E[qm2km2]−0=E[qm2] E[km2]=1⋅1=1\operatorname{Var}[q_m k_m] = \mathbb{E}[q_m^2 k_m^2] - 0 = \mathbb{E}[q_m^2]\,\mathbb{E}[k_m^2] = 1 \cdot 1 = 1

Здесь снова использована независимость (qm2q_m^2 и km2k_m^2 тоже независимы), а E[qm2]=Var⁡[qm]+(E[qm])2=1\mathbb{E}[q_m^2] = \operatorname{Var}[q_m] + (\mathbb{E}[q_m])^2 = 1.

Шаг 3. Дисперсия суммы. Слагаемые с разными mm независимы, поэтому дисперсии складываются:

Var⁡[q⋅k]=∑m=1dhVar⁡[qmkm]=dh\operatorname{Var}[\mathbf{q} \cdot \mathbf{k}] = \sum_{m=1}^{d_h} \operatorname{Var}[q_m k_m] = d_h

Шаг 4. Масштабирование. Для любой константы cc верно Var⁡[cZ]=c2Var⁡[Z]\operatorname{Var}[cZ] = c^2 \operatorname{Var}[Z]. При c=1/dhc = 1/\sqrt{d_h}:

Var⁡ ⁣[q⋅kdh]=1dh⋅dh=1\operatorname{Var}\!\left[\frac{\mathbf{q} \cdot \mathbf{k}}{\sqrt{d_h}}\right] = \frac{1}{d_h} \cdot d_h = 1

Итог: без масштаба стандартное отклонение оценок равно dh\sqrt{d_h} — при dh=128d_h = 128 (LLaMA, Mistral) это около 11,3, при dh=256d_h = 256 (Gemma) — 16. После деления на dh\sqrt{d_h} оно равно 1 при любом размере головы. Это же объяснение дано в сноске 4 статьи Vaswani et al. Проверка моделированием (100 000 пар случайных векторов) даёт дисперсию 2,0 / 64,0 / 128,7 до деления и 1,00 / 1,00 / 1,01 после для dh=2,64,128d_h = 2, 64, 128.

Предположение о независимости и единичной дисперсии — идеализация: после обучения q\mathbf{q} и k\mathbf{k} коррелированы (для того внимание и обучают). Масштаб задаёт правильный порядок величин в начале обучения, а дальше модель сама подстраивает нормы через WQW_Q и WKW_K.

Что происходит без масштаба: насыщение softmax

Заголовок раздела «Что происходит без масштаба: насыщение softmax»

Большие по модулю оценки делают softmax почти one-hot (распределением с одной единицей): вес максимального элемента близок к 1, остальные — к 0. Пример — одни и те же оценки до и после умножения на 128≈11,3\sqrt{128} \approx 11{,}3:

Оценки zzsoftmax⁡(z)\operatorname{softmax}(z)
(0, 1, 2)(0,\ 1,\ 2)(0,090, 0,245, 0,665)(0{,}090,\ 0{,}245,\ 0{,}665)
(0, 11,3, 22,6)(0,\ 11{,}3,\ 22{,}6)(1,5⋅10−10, 1,2⋅10−5, 0,99999)(1{,}5 \cdot 10^{-10},\ 1{,}2 \cdot 10^{-5},\ 0{,}99999)

Для 64 ключей со случайными оценками средний максимальный вес строки — 0,11 при стандартном отклонении оценок 1 и 0,85 при стандартном отклонении 11,3: без масштаба внимание почти всегда «смотрит в одну точку».

Почему это плохо для обучения, видно из производной softmax. Пусть p=softmax⁡(z)\mathbf{p} = \operatorname{softmax}(\mathbf{z}), pi=ezi/∑kezkp_i = e^{z_i} / \sum_k e^{z_k}. Тогда

∂pi∂zj=pi (δij−pj)\frac{\partial p_i}{\partial z_j} = p_i\,(\delta_{ij} - p_j)

где δij\delta_{ij} — символ Кронекера (1 при i=ji = j, иначе 0).

Вывод

Обозначим Σ=∑kezk\Sigma = \sum_k e^{z_k}, тогда pi=ezi/Σp_i = e^{z_i}/\Sigma и ∂Σ/∂zj=ezj\partial \Sigma / \partial z_j = e^{z_j}. По правилу дифференцирования частного:

∂pi∂zj=∂ezi∂zj Σ−ezi ezjΣ2=δij eziΣ−eziΣ⋅ezjΣ=δij pi−pi pj=pi (δij−pj)\frac{\partial p_i}{\partial z_j} = \frac{\dfrac{\partial e^{z_i}}{\partial z_j}\,\Sigma - e^{z_i}\,e^{z_j}}{\Sigma^2} = \frac{\delta_{ij}\,e^{z_i}}{\Sigma} - \frac{e^{z_i}}{\Sigma}\cdot\frac{e^{z_j}}{\Sigma} = \delta_{ij}\,p_i - p_i\,p_j = p_i\,(\delta_{ij} - p_j)

Если p\mathbf{p} почти one-hot с единицей на позиции kk, то:

  • для i=ki = k: pk(1−pk)≈1⋅0=0p_k(1 - p_k) \approx 1 \cdot 0 = 0;
  • для i≠ki \ne k: pi(δij−pj)≈0p_i(\delta_{ij} - p_j) \approx 0, потому что pi≈0p_i \approx 0.

Все элементы якобиана близки к нулю, и градиент почти не доходит до QQ и KK — обучение внимания останавливается. В примере выше наибольший по модулю элемент якобиана — 0,22 для (0,1,2)(0, 1, 2) и 1,2⋅10−51{,}2 \cdot 10^{-5} для (0,11,3,22,6)(0, 11{,}3, 22{,}6): в 18 000 раз меньше.

Этим scaled dot-product отличается от «простого» скалярного внимания. Vaswani et al. (разд. 3.2.1) отмечают, что без масштаба скалярное внимание при больших dhd_h проигрывает аддитивному, и объясняют это именно малыми градиентами softmax.

Возьмём T=3T = 3 токена и голову размера dh=2d_h = 2 (тогда dh≈1,4142\sqrt{d_h} \approx 1{,}4142):

Q=K=(100111),V=(100233)Q = K = \begin{pmatrix} 1 & 0 \\ 0 & 1 \\ 1 & 1 \end{pmatrix}, \qquad V = \begin{pmatrix} 1 & 0 \\ 0 & 2 \\ 3 & 3 \end{pmatrix}

Шаг 1. Оценки S=QK⊤/2S = QK^\top / \sqrt{2}. Например, S21=(1⋅0+1⋅1)/1,4142=0,7071S_{21} = (1 \cdot 0 + 1 \cdot 1)/1{,}4142 = 0{,}7071 (строки и столбцы нумеруем с 0):

k₀ k₁ k₂
q₀ 0.7071 0.0000 0.7071
q₁ 0.0000 0.7071 0.7071
q₂ 0.7071 0.7071 1.4142

Шаг 2. Causal-маска запрещает j>ij > i — элементы над диагональю становятся −∞-\infty.

Шаг 3. Softmax по строкам.

  • Строка 0: разрешён один ключ, вес 11.
  • Строка 1: e0=1e^{0} = 1, e0,7071=2,0281e^{0{,}7071} = 2{,}0281, сумма 3,02813{,}0281; веса (0,3302, 0,6698)(0{,}3302,\ 0{,}6698).
  • Строка 2: e0,7071=2,0281e^{0{,}7071} = 2{,}0281 (дважды), e1,4142=4,1133e^{1{,}4142} = 4{,}1133, сумма 8,16958{,}1695; веса (0,2483, 0,2483, 0,5035)(0{,}2483,\ 0{,}2483,\ 0{,}5035).
P = 1.0000 0 0
0.3302 0.6698 0
0.2483 0.2483 0.5035

Шаг 4. Выход O=PVO = PV:

  • o0=1⋅(1,0)=(1, 0)\mathbf{o}_0 = 1 \cdot (1, 0) = (1,\ 0);
  • o1=0,3302⋅(1,0)+0,6698⋅(0,2)=(0,3302, 1,3396)\mathbf{o}_1 = 0{,}3302 \cdot (1, 0) + 0{,}6698 \cdot (0, 2) = (0{,}3302,\ 1{,}3396);
  • o2=0,2483⋅(1,0)+0,2483⋅(0,2)+0,5035⋅(3,3)=(1,7588, 2,0071)\mathbf{o}_2 = 0{,}2483 \cdot (1, 0) + 0{,}2483 \cdot (0, 2) + 0{,}5035 \cdot (3, 3) = (1{,}7588,\ 2{,}0071).

Токен 2 больше всего смотрит на себя: его запрос (1,1)(1, 1) сильнее всего совпадает с ключом (1,1)(1, 1). Проверка в 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]])

(Последний знак отличается от ручного счёта из-за округления промежуточных весов.)

Маска MM в формуле решает, какие пары «запрос — ключ» разрешены. Во всех моделях репозитория есть causal-маска: позиция ii не видит будущих позиций j>ij > i. Без неё модель, обучаясь предсказывать токен i+1i + 1, видела бы его во входе. −∞-\infty перед softmax даёт запрещённым ключам ровно нулевой вес, а разрешённые веса по-прежнему в сумме дают 1.

Почему именно −∞-\infty и именно перед softmax, как маска выглядит со скользящим окном и с KV-кэшем, что делать с паддингом и что случится во float16, — в отдельной главе Маски.

Одна голова даёт на каждую позицию одно распределение весов. Многоголовое внимание (multi-head attention, MHA) считает HH голов параллельно, каждую — со своими проекциями размера dhd_h:

headh=softmax⁡ ⁣(QhKh⊤dh+M)Vh,Qh=XWQ(h), Kh=XWK(h), Vh=XWV(h)\text{head}_h = \operatorname{softmax}\!\left(\frac{Q_h K_h^\top}{\sqrt{d_h}} + M\right) V_h, \qquad Q_h = X W_Q^{(h)},\ K_h = X W_K^{(h)},\ V_h = X W_V^{(h)} MultiHead⁡(X)=Concat⁡(head1,…,headH) WO\operatorname{MultiHead}(X) = \operatorname{Concat}(\text{head}_1, \dots, \text{head}_H)\, W_O

где:

  • h=1,…,Hh = 1, \dots, H — номер головы, HH — число голов (num_heads);
  • WQ(h),WK(h),WV(h)∈Rd×dhW_Q^{(h)}, W_K^{(h)}, W_V^{(h)} \in \mathbb{R}^{d \times d_h} — проекции головы hh;
  • headh∈RT×dh\text{head}_h \in \mathbb{R}^{T \times d_h} — выход головы;
  • Concat⁡\operatorname{Concat} склеивает выходы голов по последнему измерению: RT×Hdh\mathbb{R}^{T \times H d_h};
  • WO∈RHdh×dW_O \in \mathbb{R}^{H d_h \times d} — выходная проекция, возвращающая результат в размерность модели;
  • MultiHead⁡(X)∈RT×d\operatorname{MultiHead}(X) \in \mathbb{R}^{T \times d} — выход слоя; к нему затем прибавляется residual-связь (см. Нормализация).

В коде головы не хранятся отдельно. Матрицы всех голов поставлены рядом по столбцам: WQ=[WQ(1)∣⋯∣WQ(H)]∈Rd×HdhW_Q = [W_Q^{(1)} \mid \dots \mid W_Q^{(H)}] \in \mathbb{R}^{d \times H d_h} — это один nn.Linear(emb_size, num_heads * head_size). Одно умножение даёт запросы всех голов сразу, а reshape разрезает последнее измерение на HH последовательных кусков по dhd_h — ровно Q1,…,QHQ_1, \dots, Q_H:

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] оценки всех голов одним matmul
weights @ 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] это и есть Concat
self._layer(...) [B, T, d] W_O

Схема в виде блоков — в gpt.md.

Softmax одной строки — одно распределение. Если токену нужно одновременно взять информацию из двух мест (например, из предыдущего слова и из подлежащего в начале предложения), одна голова вынуждена усреднить их, и оба сигнала размываются. Vaswani et al. (разд. 3.2.2) формулируют это так: несколько голов позволяют совместно обращать внимание на информацию из разных подпространств представлений в разных позициях, а при одной голове этому мешает усреднение.

С HH головами у каждой позиции HH независимых распределений внимания, каждое в своём подпространстве размера dhd_h. При Hdh=dH d_h = d это почти не стоит дополнительных вычислений: число параметров и FLOPs проекций то же, что у одной головы размера dd. Анализ обученных моделей показывает, что головы действительно специализируются (одни смотрят на соседний токен, другие — на синтаксически связанные слова), хотя заметную часть голов можно удалить почти без потери качества (Voita et al., 2019; Michel et al., 2019).

Зачем WOW_O, если можно просто склеить головы? Разобьём WOW_O на блоки строк по dhd_h, WO(h)∈Rdh×dW_O^{(h)} \in \mathbb{R}^{d_h \times d}. По правилу умножения блочных матриц

WO=(WO(1)⋮WO(H)),Concat⁡(head1,…,headH) WO=∑h=1Hheadh WO(h)W_O = \begin{pmatrix} W_O^{(1)} \\ \vdots \\ W_O^{(H)} \end{pmatrix}, \qquad \operatorname{Concat}(\text{head}_1, \dots, \text{head}_H)\, W_O = \sum_{h=1}^{H} \text{head}_h\, W_O^{(h)}

То есть WOW_O складывает вклады голов, предварительно переводя каждый из подпространства dhd_h в общее пространство модели dd. Без неё координаты 1…dh1 \dots d_h выхода всегда принадлежали бы голове 1 и так далее — головы не смешивались бы до следующего слоя. Кроме того, WOW_O нужна, когда Hdh≠dH d_h \ne d.

Обычно dh=d/Hd_h = d / H, и Hdh=dH d_h = d. Но это не обязательно: WQW_Q отображает d→Hdhd \to H d_h, WOW_O — обратно Hdh→dH d_h \to d, и ничто не требует равенства. Пример — Gemma 7B: H=16H = 16 голов по dh=256d_h = 256 при d=3072d = 3072, так что Hdh=4096≠3072H d_h = 4096 \ne 3072. Проекции Q, K, V расширяют пространство до 4096, а WOW_O проецирует обратно в 3072 (см. gemma.md).

В конфиге это ключ head_size; без него head_size = embed_dim // число голов, и тогда embed_dim обязан делиться на число голов (функция resolve_head_size в core/config_checks.py; при RoPE она также требует чётного head_size).

При MHA каждая из четырёх матриц WQ,WK,WV,WOW_Q, W_K, W_V, W_O имеет d⋅Hdhd \cdot H d_h элементов. Обозначим через NN число параметров слоя attention (буква PP в этой главе занята матрицей весов внимания). При Hdh=dH d_h = d:

NMHA=4d2  (+ 4d при наличии bias)N_{\text{MHA}} = 4 d^2 \;(+\,4d \text{ при наличии bias})

Для GPT-1 (d=768d = 768, bias есть) — 4⋅7682+4⋅768=2 362 3684 \cdot 768^2 + 4 \cdot 768 = 2\,362\,368 параметров на слой; для LLaMA 7B (d=4096d = 4096, без bias) — 4⋅40962=67 108 8644 \cdot 4096^2 = 67\,108\,864.

Если у K и V только GG голов (GQA, см. ниже), то WK,WV∈Rd×GdhW_K, W_V \in \mathbb{R}^{d \times G d_h}, и

NGQA=d⋅Hdh⏟WQ+2⋅d⋅Gdh⏟WK, WV+Hdh⋅d⏟WO=2d dh(H+G)N_{\text{GQA}} = \underbrace{d \cdot H d_h}_{W_Q} + \underbrace{2 \cdot d \cdot G d_h}_{W_K,\,W_V} + \underbrace{H d_h \cdot d}_{W_O} = 2 d\, d_h (H + G)

без учёта bias. При Hdh=dH d_h = d это 2d2(1+G/H)2d^2(1 + G/H): от 4d24d^2 (MHA, G=HG = H) до 2d2+2d2/H2d^2 + 2d^2/H (MQA, G=1G = 1).

МодельddHH / GGdhd_hПараметров attention на слой
LLaMA 7B409632 / 3212867 108 864
Mistral 7B409632 / 812841 943 040 (при MHA было бы 67 108 864)
Gemma 2B20488 / 12569 437 184
Gemma 7B307216 / 1625650 331 648 (4⋅3072⋅40964 \cdot 3072 \cdot 4096)

Проверка в коде:

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 (операции с плавающей точкой), Hdh=dH d_h = d, одна последовательность длины TT без кэша:

ОперацияФормаFLOPs
проекции Q, K, V, O4 умножения [T×d]⋅[d×d][T \times d] \cdot [d \times d]8Td28 T d^2
оценки QK⊤QK^\topHH умножений [T×dh]⋅[dh×T][T \times d_h] \cdot [d_h \times T]2T2d2 T^2 d
взвешенная сумма PVPVHH умножений [T×T]⋅[T×dh][T \times T] \cdot [T \times d_h]2T2d2 T^2 d
softmax, маска, масштабHT2H T^2 элементовO(HT2)O(H T^2)

Главное:

  • Ядро внимания стоит O(T2d)O(T^2 d) — квадратично по длине, в отличие от проекций и FFN, которые линейны по TT. Отношение ядра к проекциям равно 4T2d/8Td2=T/2d4T^2 d / 8Td^2 = T / 2d. Для GPT-1 (T=512T = 512, d=768d = 768) это 1/3 — проекции дороже; для T=32 768T = 32\,768 и d=4096d = 4096 — уже 4.
  • Матрица весов занимает O(T2)O(T^2) памяти на каждую голову: тензор [B, H, T, T]. Для T=4096T = 4096, H=32H = 32 во float16 это 32⋅40962⋅232 \cdot 4096^2 \cdot 2 байт =1= 1 ГиБ на слой и на одну последовательность — и при обучении её (вместе с оценками) нужно хранить до обратного прохода.
  • Causal-маска обнуляет почти половину матрицы, но явная реализация всё равно вычисляет её целиком.

Квадратичная память — главное практическое ограничение длины контекста; от неё избавляют эффективные реализации (см. ниже). Квадратичные вычисления ограничивает скользящее окно.

В attention используются два разных dropout:

  1. На весах внимания — сразу после softmax: P′=Dropout⁡(P)P' = \operatorname{Dropout}(P). Случайно выключает отдельные связи «токен → токен», чтобы модель не полагалась на одну позицию. Так делает GPT-1 (разд. 4.1 статьи: «attention dropouts» с вероятностью 0,1; attn_pdrop в HuggingFace) и GPT-2. При обучении оставшиеся веса умножаются на 1/(1−p)1/(1-p), поэтому строка P′P' в сумме даёт 1 лишь в среднем. В репозитории — параметр attention_dropout у MultiHeadAttention (ключ конфига attention_dropout у GPT и GPT-2, по умолчанию 0,0 — dropout выключен и не расходует генератор случайных чисел).
  2. После выходной проекции — Dropout⁡(MultiHead⁡(X))\operatorname{Dropout}(\operatorname{MultiHead}(X)) перед residual-сложением (resid_pdrop в GPT). В репозитории — параметр dropout у всех трёх классов attention (ключ конфига dropout).

LLaMA, Mistral, Mixtral и Gemma при предобучении dropout не используют; в репозитории у них есть только второй вид, управляемый ключом dropout. У GroupedQueryAttention и MultiQueryAttention dropout на весах внимания нет.

Головы Q всегда свои. Меняется только то, сколько отдельных K и V на них приходится. Пусть HH — число голов Q, GG — число голов 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 AttentionG=HG = H: у каждой головы Q свои K и VVaswani et al., 2017GPT-1, GPT-2, LLaMA-1, Gemma 7B
GQA — Grouped Query Attention1<G<H1 < G < H: одна пара K/V на группу из H/GH / G голов QAinslie et al., 2023Mistral 7B (32 Q, 8 K/V), Mixtral 8x7B, LLaMA-2 70B
MQA — Multi-Query AttentionG=1G = 1: одна пара K/V на все головы QShazeer, 2019Gemma 2B (8 Q, 1 K/V), PaLM

Формально при GQA голова hh (нумерация с 0) использует K/V-голову номер

g(h)=⌊hH/G⌋,headh=softmax⁡ ⁣(QhKg(h)⊤dh+M)Vg(h)g(h) = \left\lfloor \frac{h}{H / G} \right\rfloor, \qquad \text{head}_h = \operatorname{softmax}\!\left(\frac{Q_h K_{g(h)}^\top}{\sqrt{d_h}} + M\right) V_{g(h)}

где H/GH/G — размер группы (сколько голов Q делят одну пару K/V), а ⌊⋅⌋\lfloor\cdot\rfloor — округление вниз. Для H=8H = 8, G=2G = 2: головы Q 0–3 используют K/V-голову 0, головы 4–7 — K/V-голову 1.

MHA и MQA — крайние случаи GQA: G=HG = H (g(h)=hg(h) = h) и G=1G = 1 (g(h)=0g(h) = 0). Поэтому в репозитории один класс GroupedQueryAttention описывает все три, а num_q_heads должно делиться на num_kv_heads (иначе конструктор бросает ValueError).

История: MQA предложил Shazeer (2019), чтобы ускорить инкрементальную генерацию: при декодировании по одному токену время уходит не на арифметику, а на чтение K и V всех прошлых токенов из памяти, и уменьшение их объёма в HH раз прямо ускоряет шаг. Ainslie et al. (2023) заметили, что MQA теряет в качестве, и предложили промежуточный вариант — GQA; они же показали, что готовую MHA-модель можно превратить в GQA, усреднив K/V-головы внутри группы и дообучив на небольшой доле исходного объёма данных.

Чтобы использовать обычное батчевое умножение q @ k.transpose(-2, -1) с формой [B, H, T, d_h], K и V нужно привести к HH головам. GroupedQueryAttention._repeat_kv_heads повторяет каждую K/V-голову H/GH/G раз подряд:

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» реализует g(h)=⌊h/(H/G)⌋g(h) = \lfloor h / (H/G) \rfloor и совпадает с repeat_interleave(H // G, dim=1). Он важен при загрузке чужих весов: строки WQW_Q должны быть сгруппированы так же.

При G=1G = 1 копирования нет. Тензор K формы [B, 1, T_kv, d_h] участвует в умножении [B, H, T, d_h] @ [B, 1, d_h, T_kv] напрямую: по правилам трансляции (broadcasting) PyTorch измерение размера 1 растягивается на HH без выделения памяти. Так же устроен MultiQueryAttention, поэтому Gemma с num_kv_heads: 1 даёт побитово тот же результат, что и MultiQueryAttention с теми же весами (это проверено: torch.equal на выходах).

Вычислений в самом softmax⁡(QK⊤)V\operatorname{softmax}(QK^\top)V GQA не экономит: каждая голова Q по-прежнему считает свои веса по всем ключам. Экономятся проекции WKW_K, WVW_V и, главное, память кэша.

При генерации каждый новый токен смотрит на K и V всех предыдущих, и их хранят в KV-кэше, чтобы не пересчитывать (подробно — ниже). На один токен в одном слое кэш — 2⋅G⋅dh2 \cdot G \cdot d_h чисел (K и V), то есть он пропорционален числу голов K/V, а не Q. Для контекста 4096 токенов во float16 (2 байта на число), одна последовательность:

МодельГолов Q / K/Vdhd_hСлоёвKV-кэш на 4096 токеновБыл бы при MHA
LLaMA 7B32 / 32 (MHA)128322 ГиБ2 ГиБ
Mistral 7B, Mixtral 8x7B32 / 8 (GQA)12832512 МиБ2 ГиБ
Gemma 2B8 / 1 (MQA)2561872 МиБ576 МиБ
Gemma 7B16 / 16 (MHA)256281,75 ГиБ1,75 ГиБ

Например, для Mistral 7B: 2⋅32 слоя⋅8⋅128⋅4096⋅2 байта=536 870 9122 \cdot 32 \text{ слоя} \cdot 8 \cdot 128 \cdot 4096 \cdot 2 \text{ байта} = 536\,870\,912 байт =512= 512 МиБ.

Меньше кэш — больше последовательностей в батче и длиннее контекст на той же памяти, а генерация, которая упирается в чтение кэша из памяти, идёт быстрее. Цена — качество: у MQA все головы Q читают одни и те же K и V. GQA — компромисс: по качеству близка к MHA, по скорости — к MQA (Ainslie et al.). Поэтому MQA встречается в маленьких моделях (Gemma 2B), а GQA стала стандартом для больших.

Эта ось не зависит от числа голов. С полным (causal) вниманием токен видит всё прошлое. Со скользящим окном (sliding window attention; Longformer, Mistral 7B v0.1) — только ближайшее прошлое. В репозитории пара «запрос ii, ключ jj» разрешена, если

0≤i−j≤W0 \le i - j \le W

где WW — ширина окна (window_size), i,ji, j — абсолютные позиции. Условие i−j≥0i - j \ge 0 — это causal-маска, i−j≤Wi - j \le W — отсечение дальнего прошлого. Токен видит W+1W + 1 позиций вместе с собой. В HuggingFace то же окно определено как i−j<Wi - j < W, то есть на одну позицию уже, — почему так и как это учитывать при загрузке весов, см. mistral.md. Картинки масок — в главе Маски.

Стоимость. Каждая строка матрицы оценок содержит не больше W+1W + 1 разрешённых элементов, поэтому полезная работа ядра внимания — O(TWd)O(T W d) вместо O(T2d)O(T^2 d). (Явная реализация в репозитории этим не пользуется при обработке всей последовательности сразу: она считает полную матрицу и маскирует лишнее. Экономия появляется при генерации — за счёт короткого кэша.)

Рецептивное поле. За один слой информация перемещается максимум на WW позиций назад. Но выход слоя ℓ\ell на позиции jj уже содержит информацию о позициях до j−Wj - W, поэтому слой ℓ+1\ell + 1, глядя на jj, косвенно видит и их. По индукции после LL слоёв позиция ii зависит от позиций до i−L⋅Wi - L \cdot W. Для Mistral 7B (L=32L = 32, W=4096W = 4096) это 32⋅4096=131 07232 \cdot 4096 = 131\,072 позиции — «теоретический охват внимания около 131K токенов» (Jiang et al., 2023, разд. 2). Дальние зависимости передаются через слои, хотя и с потерями.

Кэш ограничен окном. Ключи старше WW позиций больше никогда не понадобятся, поэтому их можно выбросить. GroupedQueryAttention хранит в кэше только последние WW позиций K и V; вместе с новым токеном это ровно W+1W + 1 видимых позиций. Размер кэша перестаёт расти с длиной текста: у 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-модели скрытое состояние любого слоя на позиции jj зависит только от токенов x0,…,xjx_0, \dots, x_j.

Доказательство по индукции по слоям. На входе (эмбеддинги) состояние позиции jj зависит только от xjx_j и позиции jj. Пусть утверждение верно для входа слоя ℓ\ell. Attention на позиции jj смешивает значения только с позиций k≤jk \le j (causal-маска), а каждое из них по предположению зависит от x0,…,xk⊆x0,…,xjx_0, \dots, x_k \subseteq x_0, \dots, x_j. Нормализация, FFN и residual-связи работают с каждой позицией отдельно. Значит, выход слоя ℓ\ell на позиции jj тоже зависит только от x0,…,xjx_0, \dots, x_j.

Следствие: дописывание новых токенов не меняет K и V старых позиций ни в одном слое. Их можно посчитать один раз и хранить. Это и есть KV-кэш: для каждого слоя — тензоры K и V всех уже обработанных позиций. Запросы Q старых позиций не нужны: на шаге генерации нужен выход только новой позиции, и её запрос сравнивается со всеми ключами.

В двунаправленных моделях (BERT) это неверно: там старые позиции смотрят и на новые, поэтому их K и V меняются при каждом добавлении токена.

Тонкость: утверждение предполагает, что позиции токенов не меняются. Когда текст перерастает max_position_embeddings и модель берёт последние Tmax⁡T_{\max} токенов, абсолютные позиции всех токенов сдвигаются, и закэшированные K (с позиционной информацией внутри) устаревают. Поэтому generate в этот момент сбрасывает кэш и пересчитывает окно без него (next_generation_input в core/generation.py).

Память KV-кэша=2⋅L⋅G⋅dh⋅T⋅B⋅b\text{Память KV-кэша} = 2 \cdot L \cdot G \cdot d_h \cdot T \cdot B \cdot b

где:

  • 22 — отдельно K и V;
  • LL — число слоёв (у каждого слоя свой кэш);
  • GG — число голов K/V, dhd_h — размер головы;
  • TT — число закэшированных позиций (со скользящим окном — не больше WW);
  • BB — число последовательностей в батче;
  • bb — байт на число (2 для float16/bfloat16, 4 для float32).

Кэш растёт линейно с длиной контекста и размером батча и при длинных контекстах легко превышает объём весов модели. Отсюда интерес к уменьшению GG (GQA, MQA) и TT (скользящее окно).

Генерация с кэшем состоит из двух фаз:

  1. Prefill (заполнение). Весь промпт длины TpT_p обрабатывается за один проход, как при обучении: матрица оценок Tp×TpT_p \times T_p с causal-маской. Побочный результат — K и V всех позиций промпта во всех слоях, то есть заполненный кэш. Эта фаза упирается в вычисления: много токенов, большие матричные умножения.
  2. Decode (декодирование). На каждом шаге подаётся один новый токен (T=1T = 1). Для него считаются q,k,v\mathbf{q}, \mathbf{k}, \mathbf{v} (стоимость O(d2)O(d^2) на слой), новые k,v\mathbf{k}, \mathbf{v} дописываются в кэш, и запрос сравнивается со всеми t+1t + 1 ключами — одна строка матрицы оценок, O(t⋅d)O(t \cdot d). Без кэша этот шаг стоил бы O(t d2+t2d)O(t\, d^2 + t^2 d). Эта фаза упирается в пропускную способность памяти: на каждом шаге нужно прочитать весь кэш и все веса ради небольшого объёма арифметики.

Длинный промпт можно заполнять и кусками по несколько токенов (chunked prefill). Тогда внутри куска снова нужна causal-маска — новые токены не должны видеть друг друга «вперёд»; как берётся нужный кусок маски, см. Маски.

import torch
from 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

Кэш хранит G=2G = 2 головы (а не 8) и только W=4W = 4 позиции, хотя обработано 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 — intnext_pos

Зачем третий элемент. Без окна длина кэша равна числу обработанных токенов, то есть абсолютной позиции следующего. С окном кэш обрезается до WW позиций, и длина кэша перестаёт совпадать с позицией: после 7 токенов при W=4W = 4 в кэше 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, то есть с GG, а не HH головами — ради этого GQA и нужна.

Скалярное произведение qi⋅kj\mathbf{q}_i \cdot \mathbf{k}_j само по себе порядка не знает: если переставить токены, веса переставятся вместе с ними, и выход каждого токена не изменится (attention эквивариантно к перестановкам). Causal-маска вносит частичную информацию о порядке, но недостаточную. Поэтому позиция подаётся явно (подробно — в главе Позиционное кодирование):

  • GPT-1, GPT-2 прибавляют обучаемый эмбеддинг позиции к эмбеддингу токена на входе — attention получает позицию косвенно, через XX.
  • LLaMA, Mistral, Mixtral, Gemma поворачивают q\mathbf{q} и k\mathbf{k} внутри attention на угол, зависящий от позиции (RoPE): тогда qi⋅kj\mathbf{q}_i \cdot \mathbf{k}_j зависит только от разности i−ji - j. V не поворачивается. Схема головы с RoPE — в llama.md.
КлассФайлГоловы K/VRoPEОкноKV-кэш слояМодели
MultiHeadAttentioncore/multi_head_attention.py= num_headsнеобязательнонет(K, V)GPT, GPT-2 (без RoPE), LLaMA (с RoPE)
GroupedQueryAttentioncore/group_query_attention.pynum_kv_headsнеобязательноwindow_size или нет(K, V, next_pos)Mistral, Mixtral, Gemma
MultiQueryAttentioncore/multi_query_attention.py1необязательнонет(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)Q=XWQQ = XW_Q, K=XWKK = XW_K, V=XWVV = XW_V, форма [B, T, H·d_h]
3. Разбиение на головы.reshape(B, T, H, d_h), .transpose(1, 2)Q1,…,QHQ_1, \dots, Q_H, форма [B, H, T, d_h]
4. Позицияq = self._rope(q, start_pos=start_pos, positions=positions), то же для kRoPE для Q и K (если задан); при паддинге позиции — padding.positions
5. Кэшk = torch.cat([k_cache, k], dim=2), то же для vK и V всех позиций 0 … start_pos + T − 1
6. Оценкиscores = q @ k.transpose(-2, -1) / (self._head_size ** 0.5)S=QhKh⊤/dhS = Q_h K_h^\top / \sqrt{d_h}, форма [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"))S+MS + M
8. Softmax и dropout весовweights = self._attn_dropout(F.softmax(scores, dim=-1))P=softmax⁡(S+M)P = \operatorname{softmax}(S + M) по строкам
9. Взвешенная суммаx_out = weights @ vheadh=PVh\text{head}_h = P V_h
10. Склейка голов.transpose(1, 2).contiguous().reshape(B, T, H·d_h)Concat⁡(head1,…,headH)\operatorname{Concat}(\text{head}_1, \dots, \text{head}_H)
11. Выходная проекция и dropoutself._dropout(self._layer(...))Dropout⁡(Concat⁡(… )WO)\operatorname{Dropout}(\operatorname{Concat}(\dots) W_O)
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 == 1 K и V транслируются как есть, иначе _repeat_kv_heads доводит их до HH голов;
  • шаг 7: маска построена _create_sliding_window_mask (условие 0≤i−j≤W0 \le i - j \le W; без окна вместо WW подставляется max_seq_len, и это обычная causal-маска), а срез столбцов начинается с start_pos - cache_len — с позиции самого старого ключа в кэше;
  • шаг 8: dropout на весах внимания нет;
  • шаг 12: при window_size K и 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_pdrop GPT-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) вычисляет тот же самый результат, не храня матрицу T×TkvT \times T_{kv} целиком:

  • QQ, KK, VV разбиваются на блоки, которые помещаются в быструю память на кристалле (SRAM);
  • softmax считается «онлайн»: для каждой строки поддерживаются текущий максимум и текущая сумма экспонент, и при обработке очередного блока ключей накопленный результат пересчитывается с новым максимумом;
  • при обратном проходе матрица весов не читается из памяти, а пересчитывается по блокам.

Память на матрицу внимания становится O(T)O(T) вместо O(T2)O(T^2), а число обращений к HBM сокращается в разы; FLOPs остаются O(T2d)O(T^2 d), результат совпадает с точностью до округлений (это точный алгоритм, а не приближение).

В PyTorch (начиная с 2.0) есть torch.nn.functional.scaled_dot_product_attention: он сам выбирает ядро — FlashAttention, memory-efficient attention или обычную реализацию — в зависимости от устройства, типа данных и аргументов:

import torch
import 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) @ v
out = 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.

  • Забыть масштаб или взять не тот. Делить нужно на dh\sqrt{d_h} (размер головы), а не на d\sqrt{d} (размер модели) — иначе при H>1H > 1 оценки окажутся сжатыми.
  • Softmax не по той оси. Нормировать нужно по ключам (dim=-1), чтобы каждый запрос получил распределение. При softmax по запросам (dim=-2) вес пары (i,j)(i, j) зависел бы от оценок более поздних запросов i′>ii' > i — это уже не внимание, и вдобавок утечка информации из будущего.
  • 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 даёт веса, результат — взвешенная сумма значений.
  • Attention⁡(Q,K,V)=softmax⁡(QK⊤/dh+M) V\operatorname{Attention}(Q, K, V) = \operatorname{softmax}(QK^\top / \sqrt{d_h} + M)\,V. Деление на dh\sqrt{d_h} держит дисперсию оценок около 1; без него softmax насыщается и его градиент pi(δij−pj)p_i(\delta_{ij} - p_j) почти исчезает.
  • Multi-head: HH голов в подпространствах размера dhd_h, склейка и WOW_O; в коде — одна проекция и reshape. HdhH d_h может не равняться dd (Gemma 7B).
  • Параметры: 4d24d^2 при MHA, 2d dh(H+G)2 d\, d_h (H + G) при GQA. Вычисления ядра O(T2d)O(T^2 d), память на веса O(T2)O(T^2).
  • MHA, GQA и MQA различаются числом голов K/V; это определяет размер KV-кэша 2LGdhTBb2 L G d_h T B b.
  • Скользящее окно 0≤i−j≤W0 \le i - j \le W ограничивает кэш WW позициями, а рецептивное поле через LL слоёв — L⋅WL \cdot W.
  • KV-кэш корректен, потому что при causal-маске прошлые K и V не зависят от будущих токенов. Генерация — prefill промпта и decode по одному токену.
  • В репозитории — три явные реализации (MultiHeadAttention, GroupedQueryAttention, MultiQueryAttention); на практике используют FlashAttention и F.scaled_dot_product_attention.
  1. В численном примере замените VV на (200200)\begin{pmatrix} 2 & 0 \\ 0 & 2 \\ 0 & 0 \end{pmatrix}, оставив QQ и KK. Найдите выход токена 2.

    Ответ

    Веса строки 2 не зависят от VV: (0,2483, 0,2483, 0,5035)(0{,}2483,\ 0{,}2483,\ 0{,}5035). Выход: 0,2483⋅(2,0)+0,2483⋅(0,2)+0,5035⋅(0,0)=(0,4966, 0,4966)0{,}2483 \cdot (2, 0) + 0{,}2483 \cdot (0, 2) + 0{,}5035 \cdot (0, 0) = (0{,}4966,\ 0{,}4966).

  2. Компоненты q\mathbf{q} и k\mathbf{k} независимы, со средним 0 и дисперсией σ2\sigma^2. Чему равна дисперсия q⋅k/dh\mathbf{q} \cdot \mathbf{k} / \sqrt{d_h}? Что это говорит о роли нормализации входа слоя?

    Ответ

    Var⁡[qmkm]=σ2⋅σ2=σ4\operatorname{Var}[q_m k_m] = \sigma^2 \cdot \sigma^2 = \sigma^4, сумма dhd_h слагаемых — dhσ4d_h \sigma^4, после деления на dhd_h — σ4\sigma^4. Масштаб 1/dh1/\sqrt{d_h} убирает зависимость от размера головы, но не от масштаба самих векторов: если нормы Q и K вырастут, softmax снова насытится. Поэтому важно, что вход attention нормализован (pre-LN, см. Нормализация).

  3. Посчитайте число параметров attention одного слоя Mistral 7B (d=4096d = 4096, H=32H = 32, G=8G = 8, dh=128d_h = 128, без bias) и сравните с MHA той же ширины.

    Ответ

    2d dh(H+G)=2⋅4096⋅128⋅40=41 943 0402 d\, d_h (H + G) = 2 \cdot 4096 \cdot 128 \cdot 40 = 41\,943\,040. При MHA — 4⋅40962=67 108 8644 \cdot 4096^2 = 67\,108\,864; GQA экономит 37,5 % параметров attention.

  4. Сколько памяти займёт KV-кэш Gemma 2B (18 слоёв, G=1G = 1, dh=256d_h = 256) на 8192 токена во float16 для батча из 4 последовательностей? А у Mistral 7B со скользящим окном W=4096W = 4096 на 32 768 токенов, одна последовательность?

    Ответ

    Gemma 2B: 2⋅18⋅1⋅256⋅8192⋅4⋅2=603 979 7762 \cdot 18 \cdot 1 \cdot 256 \cdot 8192 \cdot 4 \cdot 2 = 603\,979\,776 байт =576= 576 МиБ. Mistral 7B: в кэше не больше W=4096W = 4096 позиций, поэтому 2⋅32⋅8⋅128⋅4096⋅22 \cdot 32 \cdot 8 \cdot 128 \cdot 4096 \cdot 2 байт =512= 512 МиБ независимо от длины; без окна было бы 8×5128 \times 512 МиБ =4= 4 ГиБ.

  5. Покажите, что якобиан softmax вырожден: его строки в сумме по jj дают 0. Какой смысл у этого факта?

    Ответ

    ∑jpi(δij−pj)=pi−pi∑jpj=pi−pi=0\sum_j p_i(\delta_{ij} - p_j) = p_i - p_i \sum_j p_j = p_i - p_i = 0. Смысл: прибавление одной и той же константы ко всем оценкам строки не меняет softmax, поэтому производная вдоль направления (1,1,…,1)(1, 1, \dots, 1) равна нулю. По той же причине маскирование можно делать и большим отрицательным числом, и −∞-\infty, а softmax в коде стабилизируют вычитанием максимума строки.

  6. Для Mistral 7B (L=32L = 32) и учебного конфига с W=16W = 16, L=4L = 4: на сколько позиций назад теоретически может дотянуться информация к последнему слою?

    Ответ

    L⋅WL \cdot W: 32⋅4096=131 07232 \cdot 4096 = 131\,072 для Mistral 7B и 4⋅16=644 \cdot 16 = 64 для учебного конфига.

  7. Почему в GroupedQueryAttention при num_kv_heads == 1 не вызывается _repeat_kv_heads, и почему результат от этого не меняется?

    Ответ

    При умножении [B, H, T, d_h] @ [B, 1, d_h, T_kv] измерение голов размера 1 транслируется на HH — каждая голова Q умножается на одни и те же K. Результат тот же, что после явного копирования, но без выделения HH копий K и V в памяти.

  8. (Для размышления.) Энкодер 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