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

Батчевые операции

Где живёт ось батча, что такое контракция и почему форма — это документация

Шаг 77 из 117 · ~26 мин

Батчевая ось — это ось, по которой ничего не сворачивается

Из урока про einsum: индекс, который есть во входах и в выходе, не суммируется. Он просто повторяет операцию.

bhqd,bhkd->bhqk\texttt{'}\htmlData{k=batch}{bh}\htmlData{k=contract}{q d}\texttt{,}\htmlData{k=batch}{bh}\htmlData{k=contract}{k d}\texttt{->}\htmlData{k=batch}{bh}qk\texttt{'}

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

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

Правило batched matmul

torch.matmul для тензоров порядка выше двух работает так:

  1. последние две оси трактуются как матрица;
  2. все остальные — как батчевые и подчиняются broadcasting (урок 020);
  3. свёртка идёт по последней оси первого и предпоследней второго.
ABрезультат
(b,n,k)(b, n, k)(b,k,m)(b, k, m)(b,n,m)(b, n, m)
(b,h,n,k)(b, h, n, k)(b,h,k,m)(b, h, k, m)(b,h,n,m)(b, h, n, m)
(b,n,k)(b, n, k)(k,m)(k, m)(b,n,m)(b, n, m) — B растянута
(b,1,n,k)(b, 1, n, k)(h,k,m)(h, k, m)(b,h,n,m)(b, h, n, m) — обе растянуты

Третья строка — самый частый случай на практике: линейный слой применяется к батчу, и матрица весов (k,m)(k, m) растягивается по всем батчевым осям без копирования.

Четвёртая строка — там, где стоит быть внимательным: broadcasting по батчевым осям означает, что пропущенная ось не вызовет ошибки, если формы окажутся совместимыми. То же предупреждение, что в уроке 020, но теперь внутри matmul.

Что видно в спецификации

bnk, km bnm

индексразмерроль
b
8
батчевая или свободная
n
64
батчевая или свободная
k
512
контракция
m
512
батчевая или свободная
форма результата
[8, 64, 512]
элементов
262,144
умножений-сложений
134,217,728
операций на элемент
512
Линейный слой на батче последовательностей. Батчевых осей две (b и n — обе просто повторяют операцию), контракция одна (k). Матрица весов не имеет ни b, ни n — потому и растягивается по ним бесплатно.

Сравните первые две спецификации: 'bnk,km->bnm' и 'bnk,bkm->bnm'. Работа одинаковая, а вес второй операции в bb раз больше по памяти — потому что во втором случае у каждого элемента батча своя матрица. Это буквально разница между nn.Linear и групповой свёрткой, и в нотации она видна как наличие или отсутствие одной буквы.

Свёртка как батчевая операция

Полезное упражнение для интуиции: свёртка 1×11\times1 — это в точности линейный слой, применённый к каждой позиции.

’bchw,dc->bdhw’\texttt{'bchw,dc->bdhw'}

Каналы сворачиваются (cc), пространственные оси ведут себя как батчевые (hh, ww). Отсюда и практическое наблюдение: свёртки 1×11\times1 в архитектурах вроде ResNet — это не «упрощённые свёртки», а способ смешать каналы, ничего не делая с пространством.

Настоящая свёртка 3×33\times3 в эту нотацию не укладывается, потому что требует перекрывающихся окон, а einsum умеет только независимые индексы. Реализуется она через unfold (материализация окон) или специализированные ядра — и это единственное место в блоке, где нотация индексов не покрывает операцию.

Форма как документация

Практический вывод из всего урока. Формы — самая дешёвая доступная документация, и её стоит писать явно:

# q, k, v: (batch, heads, seq, head_dim)
scores = einsum('bhqd,bhkd->bhqk', q, k) / math.sqrt(head_dim)
assert scores.shape == (batch, heads, seq, seq), scores.shape

Три строки, из которых:

  • комментарий фиксирует смысл осей, чего форма не сообщает;
  • спецификация фиксирует, что именно сворачивается;
  • assert ловит несовпадение там, где оно произошло.

Из уроков 010 и 020 известно, что ошибки в осях чаще всего не падают. Эти три строки — то, что превращает молчаливую ошибку в громкую.

Источники

Проверки

0 из 2
  1. Батчевые оси и контракция

    Отметьте все верные утверждения о батчевых операциях.

  2. Форма результата batched matmul

    Реализуйте matmul_shape(a, b) по правилам torch.matmul для тензоров порядка не меньше двух. Верните [1.0, ...shape] при успехе и [0.0] при несовместимости.

    Правила:

    1. последние две оси каждого операнда — это матрица: у a она (n,k1)(n, k_1), у b(k2,m)(k_2, m). Если k1k2k_1 \ne k_2, форма несовместима;
    2. все предшествующие оси — батчевые, и они подчиняются broadcasting: выравнивание справа, совместимы при равенстве или единице, результат — максимум;
    3. итоговая форма — батчевые оси, затем (n,m)(n, m).

    Порядок меньше двух считайте несовместимым.

    Обратите внимание, что несовместимость возникает по двум разным причинам — внутренняя ось и батчевые оси, — и обе должны давать [0.0].

    функция matmul_shape

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

    Ctrl/⌘ + Enter