Функция tensordot
Функция tensordot вычисляет тензорное произведение двух тензоров a и b по заданным осям. Она суммирует произведения элементов вдоль указанных размерностей, свёртывая их, и возвращает новый тензор, размерность которого определяется оставшимися осями. Это мощный инструмент для работы с многомерными данными, позволяющий выполнять обобщённое умножение матриц и тензоров.
Первый параметр a - это левый тензор. Второй параметр b - правый тензор. Третий параметр dims задаёт оси для свёртки и может быть передан в виде целого числа, списка или кортежа. Если передаётся целое число, то оно указывает количество последних осей первого тензора и первых осей второго тензора, которые будут свёрнуты. Если передаются списки или кортежи, они задают конкретные оси для каждого тензора соответственно.
Синтаксис
torch.tensordot(a, b, dims)
Пример
Выполним тензорное произведение двух двумерных тензоров (матриц) с размерностью (3, 4) и (4, 5) по одной оси (матричное умножение):
import torch
a = torch.tensor([
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12]
])
b = torch.tensor([
[1, 2, 3, 4, 5],
[6, 7, 8, 9, 10],
[11, 12, 13, 14, 15],
[16, 17, 18, 19, 20]
])
res = torch.tensordot(a, b, dims=1)
print(res)
Результат выполнения кода:
tensor([
[ 150, 160, 170, 180, 190],
[ 358, 384, 410, 436, 462],
[ 566, 608, 650, 692, 734]
])
Пример
Используем параметр dims в виде кортежей для указания конкретных осей для свёртки. Пусть первый тензор имеет размерность (2, 3, 4), а второй - (4, 5, 6). Свернём ось 2 первого тензора с осью 0 второго:
import torch
a = torch.randn(2, 3, 4)
b = torch.randn(4, 5, 6)
res = torch.tensordot(a, b, dims=([2], [0]))
print(res.shape)
Результат выполнения кода:
torch.Size([2, 3, 5, 6])
Пример
Выполним свёртку по двум осям одновременно. Первый тензор имеет размерность (3, 4, 5), второй - (4, 3, 6). Свернём оси (1, 0) первого тензора с осями (0, 1) второго:
import torch
a = torch.randn(3, 4, 5)
b = torch.randn(4, 3, 6)
res = torch.tensordot(a, b, dims=([1, 0], [0, 1]))
print(res.shape)
Результат выполнения кода:
torch.Size([5, 6])
Пример
Рассмотрим случай, когда dims задаётся целым числом. Для тензоров размерности (2, 3, 4) и (4, 3, 5) с dims=2 свёртываются две последние оси первого и две первые оси второго:
import torch
a = torch.randn(2, 3, 4)
b = torch.randn(4, 3, 5)
res = torch.tensordot(a, b, dims=2)
print(res.shape)
Результат выполнения кода:
torch.Size([2, 5])