Функция kron
Функция kron вычисляет произведение Кронекера для двух входных тензоров.
Результатом является тензор, составленный из всех возможных произведений элементов первого тензора на элементы второго.
Параметры функции: input и other - тензоры произвольной размерности.
Функция поддерживает автоматическое расширение размерностей (broadcasting) для согласования форм.
Синтаксис
torch.kron(input, other)
Пример с одномерными тензорами
Вычислим произведение Кронекера двух одномерных тензоров:
import torch
t1 = torch.tensor([1, 2])
t2 = torch.tensor([3, 4])
res = torch.kron(t1, t2)
print(res)
Результат выполнения кода:
tensor([3, 4, 6, 8])
Пример с двумерными тензорами
Применим функцию к двумерным тензорам размера 2x2. Результатом будет блочная матрица размером 4x4:
import torch
t1 = torch.tensor([
[1, 2],
[3, 4]
])
t2 = torch.tensor([
[0, 5],
[6, 7]
])
res = torch.kron(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[ 0, 5, 0, 10],
[ 6, 7, 12, 14],
[ 0, 15, 0, 20],
[18, 21, 24, 28]
])
Пример с тензорами разной размерности
Функция поддерживает работу с тензорами разного размера. В этом случае меньший тензор расширяется до размерности большего:
import torch
t1 = torch.tensor([1, 2])
t2 = torch.tensor([
[3, 4],
[5, 6]
])
res = torch.kron(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[3, 4, 6, 8],
[5, 6, 10, 12]
])
Пример с трехмерными тензорами
Рассмотрим работу функции с трехмерными тензорами. Итоговый тензор будет иметь размерность, равную сумме размерностей входных тензоров:
import torch
t1 = torch.tensor([
[[1, 2]],
[[3, 4]]
])
t2 = torch.tensor([
[[5, 6]],
[[7, 8]]
])
res = torch.kron(t1, t2)
print(res.shape)
print(res)
Результат выполнения кода:
torch.Size([4, 1, 4])
tensor([
[[ 5, 6, 10, 12]],
[[ 7, 8, 14, 16]],
[[15, 18, 20, 24]],
[[21, 24, 28, 32]]
])
Пример с вещественными числами
Функция корректно работает с тензорами вещественного типа. Результат сохраняет тип данных входных тензоров:
import torch
t1 = torch.tensor([1.5, 2.5])
t2 = torch.tensor([3.0, 4.0])
res = torch.kron(t1, t2)
print(res)
Результат выполнения кода:
tensor([4.5000, 6.0000, 7.5000, 10.0000])