Внимание в LLM · часть 3

Causal attention

Обычный self-attention видит всю последовательность. Для понимания текста это удобно, но при предсказании следующего токена превращается в подсказку: модель может посмотреть на сам ответ. Causal attention (причинное внимание) закрывает будущее до того, как softmax распределит веса.

Модель видит правильный ответ

Во время обучения модели показывают готовую фразу целиком. Позиция «потому» должна предсказать следующий токен — «что». Если оставить attention без ограничений, запрос «потому» сможет направить вес прямо на будущее слово «что». Задача будет решена нечестно: ответ уже лежит во входе.

Оставьте запрос «потому» и сначала рассмотрите строку без маски. Какие позиции маска должна исключить? Включите её, проверьте оставшуюся сумму весов, затем сравните первую и последнюю строки последовательности.

Causal mask

будущее видно
Режим attention
Строка запроса
Запрос «потому» отдаёт будущему 46% веса. Правильный следующий токен «что» уже получает 32%.
кошка
лакает
молоко
потому
что
она
голодна
33%
7%
5%
4%
4%
27%
19%
7%
48%
19%
5%
5%
7%
10%
6%
22%
49%
5%
5%
6%
7%
5%
5%
5%
39%
32%
8%
6%
6%
6%
6%
36%
30%
9%
7%
29%
7%
5%
7%
7%
26%
18%
25%
12%
7%
7%
7%
23%
19%
  • доступно
  • будущее
  • следующий токен
УтечкаМодель видит правильный ответ во входе.
Доступ к будущему до и после маскиМаска зануляет веса будущих позиций и заново нормирует веса доступных. При последнем запросе будущих позиций уже нет.

Переключите режим и выберите несколько строк. Causal mask закрывает всё, что находится правее текущего токена; граница всегда проходит по диагонали.

Запрет до softmax

Causal mask (причинная маска) — матрица M, в которой разрешённые позиции содержат ноль, а будущее — . Её прибавляют к attention scores (оценкам внимания) до softmax:

S=QKdk+M
Attention(Q,K,V)=softmax(S)V
Mij={0,ji,,j>i

После softmax: e=0. Запрещённая связь не может получить вес.

Обучение остаётся параллельным

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

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