Spec-Zone.ru › PyTorch 2.14

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 можно указать в подсписке, чтобы включить широковещательное распространение, как описано выше в разделе «Уравнение».

Параметры:
  • equation (str) – Индексы для суммирования по соглашению Эйнштейна.
  • operands (List[Tensor]) – Тензоры, для которых выполняется суммирование по соглашению Эйнштейна.
Тип возвращаемого значения:

Tensor

Примеры:

>>> # 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API