torch.einsum
-
torch.einsum(equation, *operands) → Tensor[исходный код] -
Суммирует произведения элементов входного
operandsпо измерениям, заданным с помощью обозначений, основанных на соглашении Эйнштейна о суммировании.Einsum позволяет вычислять многие распространённые операции линейной алгебры над многомерными массивами, представляя их в краткой форме, основанной на соглашении Эйнштейна о суммировании и задаваемой выражением
equation. Подробное описание этого формата приведено ниже, но общий принцип заключается в том, чтобы пометить каждое измерение входногоoperandsнекоторым индексом и определить, какие индексы входят в результат. Результат вычисляется суммированием произведений элементовoperandsпо измерениям, индексы которых не входят в результат. Например, умножение матриц можно вычислить с помощью einsum какtorch.einsum(“ij,jk->ik”, A, B). Здесь j — индекс суммирования, а i и k — индексы результата (подробнее об этом см. в разделе ниже).Уравнение:
Строка
equationзадаёт индексы (буквы в[a-zA-Z]) для каждого измерения входногоoperandsв том же порядке, что и измерения; индексы каждого операнда разделяются запятой (‘,’). Например,‘ij,jk’задаёт индексы для двух двумерных операндов. Размеры измерений с одинаковыми индексами должны быть совместимы для широковещательного распространения, то есть должны совпадать или быть равны1. Исключение составляют случаи, когда индекс повторяется для одного и того же входного операнда: тогда размеры измерений этого операнда с данным индексом должны совпадать, а вместо операнда будет использована его диагональ по этим измерениям. Индексы, встречающиеся вequationровно один раз, войдут в результат в алфавитном порядке. Результат вычисляется поэлементным умножением входныхoperandsс выравниванием их измерений по индексам, а затем суммированием по измерениям, индексы которых не входят в результат.При необходимости индексы результата можно задать явно, добавив в конце уравнения стрелку (‘->’), за которой следуют индексы результата. Например, следующее уравнение вычисляет транспонированное произведение матриц: ‘ij,jk->ki’. Каждый индекс результата должен встречаться хотя бы один раз в одном из входных операндов и не более одного раза в результате.
Многоточие (’…’) можно использовать вместо индексов для широковещательного распространения измерений, охватываемых многоточием. Каждый входной операнд может содержать не более одного многоточия, которое охватывает измерения, не обозначенные индексами. Например, для входного операнда с пятью измерениями многоточие в уравнении
‘ab…c’охватывает третье и четвёртое измерения. Многоточие может охватывать разное количество измерений в разныхoperands, однако «форма» многоточия (размеры охватываемых им измерений) должна быть совместима для широковещательного распространения. Если результат не задан явно с помощью стрелки (‘->’), многоточие будет первым в результате (самые левые измерения), перед индексами, которые встречаются ровно один раз во входных операндах. Например, следующее уравнение выполняет пакетное умножение матриц‘…ij,…jk’.Несколько заключительных замечаний: между элементами уравнения (индексами, многоточием, стрелкой и запятой) могут быть пробелы, но выражение вроде
‘…’недопустимо. Пустая строка‘’допустима для скалярных операндов.Примечание
torch.einsumобрабатывает многоточие (’…’) иначе, чем NumPy: он позволяет суммировать измерения, охватываемые многоточием, то есть многоточие не обязательно должно входить в результат.Примечание
Установите opt-einsum (https://optimized-einsum.readthedocs.io/en/stable/), чтобы воспользоваться более производительной реализацией einsum. Его можно установить вместе с torch следующим образом:
pip install torch[opt-einsum]или отдельно с помощьюpip install opt-einsum.Если opt-einsum доступен, эта функция автоматически ускорит вычисления и/или снизит потребление памяти, оптимизируя порядок свёртки с помощью нашей серверной части opt_einsum
torch.backends.opt_einsum(Знаю, подчёркивание вместо дефиса сбивает с толку). Эта оптимизация выполняется при наличии не менее трёх входных данных, поскольку в противном случае порядок не имеет значения. Обратите внимание: поискtheоптимального пути — NP-трудная задача, поэтому opt-einsum использует разные эвристики для получения близких к оптимальным результатов. Если opt-einsum недоступен, по умолчанию свёртка выполняется слева направо.Чтобы отключить это поведение по умолчанию, добавьте следующее, чтобы выключить opt_einsum и пропустить вычисление пути:
torch.backends.opt_einsum.enabled = FalseЧтобы указать стратегию, которую opt_einsum будет использовать для вычисления пути свёртки, добавьте следующую строку:
torch.backends.opt_einsum.strategy = 'auto'. По умолчанию используется стратегия ‘auto’; также поддерживаются ‘greedy’ и ‘optimal’. Обратите внимание, что время выполнения стратегии ‘optimal’ растёт факториально с увеличением числа входных данных! Подробнее см. в документации opt_einsum (https://optimized-einsum.readthedocs.io/en/stable/path_finding.html).Примечание
Начиная с PyTorch 1.10
torch.einsum()также поддерживает формат подсписков (см. примеры ниже). В этом формате индексы каждого операнда задаются подсписками — списками целых чисел в диапазоне [0, 52). Эти подсписки следуют за соответствующими операндами; в конце входных данных может быть указан дополнительный подсписок для задания индексов результата, напримерtorch.einsum(op1, sublist1, op2, sublist2, …, [subslist_out]). ОбъектEllipsisв Python можно указать в подсписке, чтобы включить широковещательное распространение, как описано выше в разделе «Уравнение».- Параметры:
- Тип возвращаемого значения:
Примеры:
>>> # trace >>> torch.einsum('ii', torch.randn(4, 4)) tensor(-1.2104) >>> # diagonal >>> torch.einsum('ii->i', torch.randn(4, 4)) tensor([-0.1034, 0.7952, -0.2433, 0.4545]) >>> # outer product >>> x = torch.randn(5) >>> y = torch.randn(4) >>> torch.einsum('i,j->ij', x, y) tensor([[ 0.1156, -0.2897, -0.3918, 0.4963], [-0.3744, 0.9381, 1.2685, -1.6070], [ 0.7208, -1.8058, -2.4419, 3.0936], [ 0.1713, -0.4291, -0.5802, 0.7350], [ 0.5704, -1.4290, -1.9323, 2.4480]]) >>> # batch matrix multiplication >>> As = torch.randn(3, 2, 5) >>> Bs = torch.randn(3, 5, 4) >>> torch.einsum('bij,bjk->bik', As, Bs) tensor([[[-1.0564, -1.5904, 3.2023, 3.1271], [-1.6706, -0.8097, -0.8025, -2.1183]], [[ 4.2239, 0.3107, -0.5756, -0.2354], [-1.4558, -0.3460, 1.5087, -0.8530]], [[ 2.8153, 1.8787, -4.3839, -1.2112], [ 0.3728, -2.1131, 0.0921, 0.8305]]]) >>> # with sublist format and ellipsis >>> torch.einsum(As, [..., 0, 1], Bs, [..., 1, 2], [..., 0, 2]) tensor([[[-1.0564, -1.5904, 3.2023, 3.1271], [-1.6706, -0.8097, -0.8025, -2.1183]], [[ 4.2239, 0.3107, -0.5756, -0.2354], [-1.4558, -0.3460, 1.5087, -0.8530]], [[ 2.8153, 1.8787, -4.3839, -1.2112], [ 0.3728, -2.1131, 0.0921, 0.8305]]]) >>> # batch permute >>> A = torch.randn(2, 3, 4, 5) >>> torch.einsum('...ij->...ji', A).shape torch.Size([2, 3, 5, 4]) >>> # equivalent to torch.nn.functional.bilinear >>> A = torch.randn(3, 5, 4) >>> l = torch.randn(2, 5) >>> r = torch.randn(2, 4) >>> torch.einsum('bn,anm,bm->ba', l, A, r) tensor([[-0.3430, -5.2405, 0.4494], [ 0.3311, 5.5201, -3.0356]])
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.einsum.html