Математика глубокого обучения
einsum
Одно правило вместо десяти функций — и способ считать цену операции глазами
Одно правило
Нотация Эйнштейна выглядит как набор рецептов, а является одним правилом:
Отсюда весь «список рецептов» перестаёт быть списком:
| спецификация | что получается | почему |
|---|---|---|
'ij,jk->ik' | матричное умножение | j пропал → сумма по j |
'ij,j->i' | матрица на вектор | то же |
'i,i->' | скалярное произведение | i пропал, выход скаляр |
'i,j->ij' | внешнее произведение | ничего не пропало → суммы нет |
'ij->ji' | транспонирование | ничего не пропало, только порядок |
'ij->i' | суммы по строкам | j пропал |
'ij->' | сумма всего | пропали оба |
'ii->i' | диагональ | повтор индекса = выбор диагонали |
'ii->' | след | диагональ, затем сумма |
Последние две строки показывают вторую половину нотации: повторённый индекс внутри одного
операнда означает диагональ. 'ii->' — это , то есть след.
Цена операции читается глазами
Практическая ценность нотации не в краткости, а в том, что по ней сразу виден объём работы.
Число умножений-сложений равно произведению размеров всех индексов — и суммируемых, и остающихся. Причина простая: одно умножение-сложение происходит на каждый набор значений всех индексов.
Проверьте на виджете: у 'ij,jk->ik' при выходит операции, и это
. Результат содержит элементов, значит на каждый приходится по
операции — ровно длина суммируемой оси.
Столбец «операций на элемент» полезен сам по себе. Он равен произведению суммируемых размеров, и по нему видно, чем ограничена операция:
- отношение — суммирования нет, операция только двигает данные. Транспонирование, перестановка осей, поэлементные операции. Ограничены пропускной способностью памяти;
- большое отношение — много арифметики на каждое прочитанное число. Матричные умножения. Ограничены вычислительной мощностью.
Это и есть различие между memory-bound и compute-bound, увиденное из нотации.
Почему порядок свёртки важен
Возьмём произведение трёх матриц с размерами , , , .
| порядок | промежуточный результат | операций |
|---|---|---|
Здесь одинаково. А теперь , , , :
| порядок | операций |
|---|---|
Разница в раз. Причина — узкая ось : свернув по ней первым делом, вы избегаете создания матрицы .
Отсюда важное предупреждение: einsum с тремя и более операндами не гарантирует
оптимального порядка. NumPy принимает optimize=True, opt_einsum подбирает порядок
поиском; наивная реализация может выбрать худший вариант. Для критичных мест порядок стоит
задавать явно, разбив выражение на два вызова.
Что einsum не делает
Честные ограничения:
- не оптимизирует за вас порядок свёртки при трёх и более операндах, если не попросить;
- не всегда быстрее специализированных вызовов:
matmulможет использовать более подходящее ядро, чем то, во что развернётсяeinsum; - не поддерживает всё, что нужно:
einopsпоявился потому, чтоeinsumне умеет разбивать и склеивать оси ('b (h d) -> b h d').
Но у него есть то, чего нет ни у одной специализированной функции: спецификация читается
как утверждение о формах. 'bhqd,bhkd->bhqk' говорит, что головы и батч сохраняются, а
свёртка идёт по размеру головы, — и если вы ошиблись осью, ошибка видна в строке, а не в
отладчике.
Источники
- Rogozhnikov — Einops tutorial — Именованные оси, rearrange и reduce
- Olah — Understanding Einsum — Разбор нотации на примерах
Проверки
0 из 2Одно правило einsum
Отметьте все верные утверждения о нотации einsum.
Разобрать einsum и посчитать цену
Реализуйте
einsum_cost(spec, sizes)— разберите спецификацию и верните[n_summed, out_size, flops, flops_per_element]:spec— строка вида"ij,jk->ik"; выход всегда указан явно (после->, возможно пустой);sizes— объект (в Python — словарь), сопоставляющий букве её размер;n_summed— сколько различных букв встречается во входах, но не в выходе;out_size— произведение размеров букв выхода (для скаляра — );flops— произведение размеров всех различных букв, встречающихся во входах;flops_per_element=flops / out_size.
Каждую букву считайте один раз, даже если она повторяется: в
"ii->"букваiодна, иflopsравен , а не .Проверить себя легко на матричном умножении:
"ij,jk->ik"при должно дать операции () и операции на элемент — длину свёртки.Загрузка редактора…
Ctrl/⌘ + Enter