Функция einsum
Функция einsum реализует нотацию суммирования Эйнштейна.
Она позволяет записывать сложные операции над тензорами (умножение, транспонирование, суммирование по осям и перестановку размерностей)
в компактной строковой форме.
Первым аргументом передаётся строка с обозначением размерностей (например, "ij,jk->ik"),
последующие аргументы - входные тензоры.
Вторым параметром можно передать дополнительные аргументы, такие как dtype.
Синтаксис
torch.einsum(equation, *operands)
Пример
Давайте выполним матричное умножение двух матриц с помощью einsum:
import torch
t1 = torch.tensor([
[1, 2],
[3, 4],
])
t2 = torch.tensor([
[5, 6],
[7, 8],
])
res = torch.einsum("ij,jk->ik", t1, t2)
print(res)
Результат выполнения кода:
tensor([
[19, 22],
[43, 50],
])
Пример
Давайте вычислим скалярное произведение двух векторов:
import torch
v1 = torch.tensor([1, 2, 3])
v2 = torch.tensor([4, 5, 6])
res = torch.einsum("i,i->", v1, v2)
print(res)
Результат выполнения кода:
tensor(32)
Пример
Выполним перестановку осей трёхмерного тензора:
import torch
t = torch.tensor([
[
[1, 2],
[3, 4],
],
[
[5, 6],
[7, 8],
],
])
res = torch.einsum("ijk->kji", t)
print(res)
Результат выполнения кода:
tensor([
[
[1, 5],
[3, 7],
],
[
[2, 6],
[4, 8],
],
])
Пример
Вычислим сумму всех элементов тензора по определённой оси:
import torch
t = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
res = torch.einsum("ij->i", t)
print(res)
Результат выполнения кода:
tensor([6, 15])
Пример
Умножение матрицы на вектор с использованием einsum:
import torch
mat = torch.tensor([
[1, 2],
[3, 4],
])
vec = torch.tensor([5, 6])
res = torch.einsum("ij,j->i", mat, vec)
print(res)
Результат выполнения кода:
tensor([17, 39])