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

einsum

Одно правило вместо десяти функций — и способ считать цену операции глазами

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

Одно правило

Нотация Эйнштейна выглядит как набор рецептов, а является одним правилом:

cik=jaijbjk’ij,jk->ik’\htmlData{k=rule}{c_{ik} = \sum_j a_{ij} b_{jk}} \qquad\longleftrightarrow\qquad \texttt{'ij,jk->ik'}

целиком: индекс, который есть во входах и отсутствует в выходе, суммируется. Всё остальное — следствия.

Отсюда весь «список рецептов» перестаёт быть списком:

спецификациячто получаетсяпочему
'ij,jk->ik'матричное умножениеj пропал → сумма по j
'ij,j->i'матрица на векторто же
'i,i->'скалярное произведениеi пропал, выход скаляр
'i,j->ij'внешнее произведениеничего не пропало → суммы нет
'ij->ji'транспонированиеничего не пропало, только порядок
'ij->i'суммы по строкамj пропал
'ij->'сумма всегопропали оба
'ii->i'диагональповтор индекса = выбор диагонали
'ii->'следдиагональ, затем сумма

Последние две строки показывают вторую половину нотации: повторённый индекс внутри одного операнда означает диагональ. 'ii->' — это iaii\sum_i a_{ii}, то есть след.

Цена операции читается глазами

Практическая ценность нотации не в краткости, а в том, что по ней сразу виден объём работы.

Число умножений-сложений равно произведению размеров всех индексов — и суммируемых, и остающихся. Причина простая: одно умножение-сложение происходит на каждый набор значений всех индексов.

Проверьте на виджете: у 'ij,jk->ik' при i=2,j=3,k=4i=2, j=3, k=4 выходит 2424 операции, и это 2342 \cdot 3 \cdot 4. Результат содержит 88 элементов, значит на каждый приходится по 33 операции — ровно длина суммируемой оси.

ij, jk ik

индексразмерроль
i
2
остаётся
j
3
суммируется
k
4
остаётся
форма результата
[2, 4]
элементов
8
умножений-сложений
24
операций на элемент
3
Матричное умножение. Операций i·j·k, элементов на выходе i·k, значит на каждый элемент приходится j операций — длина свёртки. Потяните j: выход не меняется, а работа растёт линейно.

Столбец «операций на элемент» полезен сам по себе. Он равен произведению суммируемых размеров, и по нему видно, чем ограничена операция:

  • отношение 11 — суммирования нет, операция только двигает данные. Транспонирование, перестановка осей, поэлементные операции. Ограничены пропускной способностью памяти;
  • большое отношение — много арифметики на каждое прочитанное число. Матричные умножения. Ограничены вычислительной мощностью.

Это и есть различие между memory-bound и compute-bound, увиденное из нотации.

Почему порядок свёртки важен

Возьмём произведение трёх матриц AijBjkCklA_{i j} B_{j k} C_{k l} с размерами i=1000i = 1000, j=10j = 10, k=10k = 10, l=1000l = 1000.

порядокпромежуточный результатопераций
(AB)C(AB)Ci×k=104i \times k = 10^4ijk+ikl=105+107i j k + i k l = 10^5 + 10^7
A(BC)A(BC)j×l=104j \times l = 10^4jkl+ijl=105+107j k l + i j l = 10^5 + 10^7

Здесь одинаково. А теперь i=1000i = 1000, j=1000j = 1000, k=2k = 2, l=1000l = 1000:

порядокопераций
(AB)C(AB)Cijk+ikl=2106+2106=4106i j k + i k l = 2\cdot10^6 + 2\cdot10^6 = 4\cdot10^6
A(BC)A(BC)jkl+ijl=2106+109j k l + i j l = 2\cdot10^6 + 10^9

Разница в 250250 раз. Причина — узкая ось kk: свернув по ней первым делом, вы избегаете создания матрицы j×lj \times l.

Отсюда важное предупреждение: einsum с тремя и более операндами не гарантирует оптимального порядка. NumPy принимает optimize=True, opt_einsum подбирает порядок поиском; наивная реализация может выбрать худший вариант. Для критичных мест порядок стоит задавать явно, разбив выражение на два вызова.

Что einsum не делает

Честные ограничения:

  • не оптимизирует за вас порядок свёртки при трёх и более операндах, если не попросить;
  • не всегда быстрее специализированных вызовов: matmul может использовать более подходящее ядро, чем то, во что развернётся einsum;
  • не поддерживает всё, что нужно: einops появился потому, что einsum не умеет разбивать и склеивать оси ('b (h d) -> b h d').

Но у него есть то, чего нет ни у одной специализированной функции: спецификация читается как утверждение о формах. 'bhqd,bhkd->bhqk' говорит, что головы и батч сохраняются, а свёртка идёт по размеру головы, — и если вы ошиблись осью, ошибка видна в строке, а не в отладчике.

Источники

Проверки

0 из 2
  1. Одно правило einsum

    Отметьте все верные утверждения о нотации einsum.

  2. Разобрать einsum и посчитать цену

    Реализуйте einsum_cost(spec, sizes) — разберите спецификацию и верните [n_summed, out_size, flops, flops_per_element]:

    • spec — строка вида "ij,jk->ik"; выход всегда указан явно (после ->, возможно пустой);
    • sizes — объект (в Python — словарь), сопоставляющий букве её размер;
    • n_summed — сколько различных букв встречается во входах, но не в выходе;
    • out_size — произведение размеров букв выхода (для скаляра — 11);
    • flops — произведение размеров всех различных букв, встречающихся во входах;
    • flops_per_element = flops / out_size.

    Каждую букву считайте один раз, даже если она повторяется: в "ii->" буква i одна, и flops равен ii, а не i2i^2.

    Проверить себя легко на матричном умножении: "ij,jk->ik" при 2,3,42, 3, 4 должно дать 2424 операции (=234= 2 \cdot 3 \cdot 4) и 33 операции на элемент — длину свёртки.

    функция einsum_cost

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

    Ctrl/⌘ + Enter