Математика глубокого обучения
Батчевые операции
Где живёт ось батча, что такое контракция и почему форма — это документация
Батчевая ось — это ось, по которой ничего не сворачивается
Из урока про einsum: индекс, который есть во входах и в выходе, не суммируется. Он просто повторяет операцию.
Отсюда полное определение:
Практическое следствие: батчевых осей может быть несколько, и «батч» в смысле примеров данных — лишь одна из них. У внимания их две: батч примеров и голова. Обе ведут себя одинаково — операция просто повторяется по каждому значению.
Правило batched matmul
torch.matmul для тензоров порядка выше двух работает так:
- последние две оси трактуются как матрица;
- все остальные — как батчевые и подчиняются broadcasting (урок 020);
- свёртка идёт по последней оси первого и предпоследней второго.
| A | B | результат |
|---|---|---|
| — B растянута | ||
| — обе растянуты |
Третья строка — самый частый случай на практике: линейный слой применяется к батчу, и матрица весов растягивается по всем батчевым осям без копирования.
Четвёртая строка — там, где стоит быть внимательным: broadcasting по батчевым осям означает, что пропущенная ось не вызовет ошибки, если формы окажутся совместимыми. То же предупреждение, что в уроке 020, но теперь внутри matmul.
Что видно в спецификации
Сравните первые две спецификации: 'bnk,km->bnm' и 'bnk,bkm->bnm'. Работа одинаковая, а
вес второй операции в раз больше по памяти — потому что во втором случае у каждого
элемента батча своя матрица. Это буквально разница между nn.Linear и групповой свёрткой, и
в нотации она видна как наличие или отсутствие одной буквы.
Свёртка как батчевая операция
Полезное упражнение для интуиции: свёртка — это в точности линейный слой, применённый к каждой позиции.
Каналы сворачиваются (), пространственные оси ведут себя как батчевые (, ). Отсюда и практическое наблюдение: свёртки в архитектурах вроде ResNet — это не «упрощённые свёртки», а способ смешать каналы, ничего не делая с пространством.
Настоящая свёртка в эту нотацию не укладывается, потому что требует
перекрывающихся окон, а 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 известно, что ошибки в осях чаще всего не падают. Эти три строки — то, что превращает молчаливую ошибку в громкую.
Источники
- PyTorch — Broadcasting semantics for matmul — Правила батчевого умножения
- Rogozhnikov — Einops tutorial — Явные оси в батчевых операциях
Проверки
0 из 2Батчевые оси и контракция
Отметьте все верные утверждения о батчевых операциях.
Форма результата batched matmul
Реализуйте
matmul_shape(a, b)по правиламtorch.matmulдля тензоров порядка не меньше двух. Верните[1.0, ...shape]при успехе и[0.0]при несовместимости.Правила:
- последние две оси каждого операнда — это матрица: у
aона , уb— . Если , форма несовместима; - все предшествующие оси — батчевые, и они подчиняются broadcasting: выравнивание справа, совместимы при равенстве или единице, результат — максимум;
- итоговая форма — батчевые оси, затем .
Порядок меньше двух считайте несовместимым.
Обратите внимание, что несовместимость возникает по двум разным причинам — внутренняя ось и батчевые оси, — и обе должны давать
[0.0].Загрузка редактора…
Ctrl/⌘ + Enter- последние две оси каждого операнда — это матрица: у