Математика глубокого обучения

Многоголовость и маски

Головы как блочная структура, причинная маска как −∞ — и почему она именно так

Шаг 81 из 117 · ~28 мин

Головы — это блочная структура

Многоголовое внимание не добавляет вычислений. Оно делит существующую ширину на части и считает внимание в каждой независимо.

MHA(X)=concat(head1,,headh)WO,d=D/h\text{MHA}(X) = \htmlData{k=concat}{\text{concat}}\big(\htmlData{k=split}{\text{head}_1, \dots, \text{head}_h}\big) W^O, \qquad d = D/h

Из прошлого урока: стоимость равна 2n2D+4nD22n^2D + 4nD^2 и не содержит hh. Двадцать четыре головы по 3232 и двенадцать по 6464 стоят одинаково — меняется только то, как ширина поделена.

В нотации индексов — это переформовка одной оси в две:

’b n (h d) -> b h n d’\texttt{'b n (h d) -> b h n d'}

Именно поэтому einops удобнее, чем view и permute: здесь видно, что hh и dd — это разложение одной оси, а не две независимые. Порядок множителей в скобках важен, и перепутать (h d) с (d h) — распространённая ошибка, которая не падает.

Зачем несколько голов

Если стоимость та же, что даёт разбиение? Одна голова может внимать только «в одном направлении»: её распределение по ключам одно. Несколько голов дают несколько распределений одновременно.

Ограничение, из которого это следует: у одной головы результат — выпуклая комбинация значений, то есть точка внутри их выпуклой оболочки, с одним набором весов. Чтобы собрать информацию из двух разных мест по разным правилам, нужны два набора весов.

Цена — размер головы. При D=768D = 768:

hhddчто происходит
11768768одно распределение, максимум выразительности внутри него
12126464стандартный компромисс
969688оценки почти вырождены: qkq\cdot k по восьми компонентам очень шумно

Последняя строка — реальное ограничение. При очень малом dd скалярное произведение вычисляется по нескольким числам, и его дисперсия относительно велика: внимание становится шумным. Отсюда и практический диапазон d[32,128]d \in [32, 128] почти во всех моделях.

Причинная маска

Для авторегрессионной модели токен не должен видеть будущее. Реализуется это прибавлением -\infty к запрещённым оценкам до softmax:

sqk{sqk,kq,k>qs_{qk} \leftarrow \begin{cases} s_{qk}, & k \le q \\ -\infty, & k > q \end{cases}

Почему именно -\infty, а не удаление или ноль:

  • ноль не подходит: нулевая оценка даёт вес e0=1e^0 = 1, то есть не запрет, а средний вес. Маскировать нужно до softmax, а не после;
  • -\infty даёт ровно нуль после softmax: e=0e^{-\infty} = 0, и знаменатель тоже не включает это слагаемое. Нормировка остаётся корректной автоматически;
  • удаление не подходит технически: тензор должен остаться прямоугольным, чтобы операции оставались батчевыми.

На практике вместо -\infty берут большое отрицательное число (например 109-10^9 или минимальное значение типа), потому что -\infty в арифметике может дать nan при умножении на нуль. Но с точки зрения математики это именно -\infty.

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

ключи

запросы
<s>котселнаковрик.
<s>
кот
сел
на
коврик
.
размер головы d 64
температура T 1
разброс оценок
4.38
энтропия строки
1.209 бит
максимальный вес
1
масштаб градиента
2.47e-1

Обратите внимание на первую строку: у неё энтропия равна нулю независимо от оценок. Первый токен имеет ровно один допустимый ключ — себя, — и после нормировки его вес равен единице принудительно. Никакой информации в этой строке нет, и это не дефект: так устроена авторегрессия.

Сколько работы экономит маска

Наивная реализация считает всю матрицу, а потом половину выбрасывает. Экономии нет вовсе — только трата.

nnиспользуетсядоля
441010 из 161662.5%62.5\%
12812882568256 из 163841638450.4%50.4\%
10241024524800524\,800 из 10485761\,048\,57650.05%50.05\%

Доля равна n+12n\frac{n+1}{2n} и стремится к половине сверху. То есть при честной реализации причинное внимание должно быть примерно вдвое дешевле полного — и FlashAttention с блочными ядрами эту экономию действительно получает, просто не вычисляя запрещённые блоки.

Другие маски

маскачто запрещаетгде
причиннаябудущеедекодеры, языковые модели
paddingзаполнитель в коротких последовательностяхбатчи переменной длины
скользящее окновсё дальше ww токеновLongformer, Mistral
блочно-разреженнаявсё вне выбранных блоковдлинный контекст

Padding-маска стоит отдельного упоминания как источник ошибок: если её забыть, модель будет внимать заполнителям, и результат станет зависеть от того, как батч был собран. Ошибка не падает — просто качество на коротких последовательностях оказывается хуже, чем на длинных, и причина неочевидна.

Вариации по числу наборов ключей

Из урока про KV-кэш (140) станет понятно, зачем это нужно, но структура относится сюда:

схемазапросовключей и значений
MHAhh наборовhh наборов
MQAhh набороводин набор на все головы
GQAhh наборовgg групп, 1<g<h1 < g < h

Вычислений это почти не меняет, а вот размер кэша делит на hh (для MQA) — и именно кэш, а не арифметика, ограничивает генерацию. Поэтому современные модели почти все используют GQA: компромисс между выразительностью MHA и памятью MQA.

Источники

Проверки

0 из 2
  1. Головы и маски

    Отметьте все верные утверждения о многоголовом внимании и масках.

  2. Маска, головы и размер кэша

    Реализуйте mask_facts(n, heads, model_dim, kv_groups) — верните [allowed_fraction, head_dim, kv_cache_ratio, masked_cells]:

    • allowed_fraction — доля матрицы оценок, разрешённая причинной маской. Разрешено n(n+1)2\frac{n(n+1)}{2} ячеек из n2n^2;
    • head_dim = D/hD / h;
    • kv_cache_ratio = g/hg / h — во сколько раз кэш ключей и значений меньше, чем при полном MHA. При g=hg = h это единица (обычное MHA), при g=1g = 1 — MQA;
    • masked_cells — сколько ячеек замаскировано, то есть n2n(n+1)2n^2 - \frac{n(n+1)}{2}.

    Проверить себя можно двумя тождествами: allowed_fraction обязана равняться n+12n\frac{n+1}{2n}, а masked_cells — числу n(n1)2\frac{n(n-1)}{2}. Оба следуют из того, что разрешённая часть — треугольник с диагональю, а запрещённая — треугольник без неё.

    функция mask_facts

    Загрузка редактора…

    Ctrl/⌘ + Enter