Функция bmm
Функция bmm (batch matrix multiplication) выполняет пакетное перемножение матриц. Первым параметром функция принимает трёхмерный тензор input, вторым параметром - трёхмерный тензор mat2. Функция возвращает трёхмерный тензор, содержащий результат умножения каждой пары матриц из входных тензоров.
Синтаксис
torch.bmm(input, mat2)
Параметры
Первый параметр input - трёхмерный тензор формы (batch_size, n, m).
Второй параметр mat2 - трёхмерный тензор формы (batch_size, m, p).
Пример
Давайте выполним пакетное умножение двух тензоров размера 2x2:
import torch
t1 = torch.tensor([
[[1, 2],
[3, 4]],
[[5, 6],
[7, 8]]
])
t2 = torch.tensor([
[[9, 10],
[11, 12]],
[[13, 14],
[15, 16]]
])
res = torch.bmm(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[[ 31, 34],
[ 71, 78]],
[[155, 166],
[211, 226]]
])
Пример
Давайте выполним пакетное умножение тензоров разных размерностей 2x3 и 3x2:
import torch
t1 = torch.tensor([
[[1, 2, 3],
[4, 5, 6]],
[[7, 8, 9],
[10, 11, 12]]
])
t2 = torch.tensor([
[[13, 14],
[15, 16],
[17, 18]],
[[19, 20],
[21, 22],
[23, 24]]
])
res = torch.bmm(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[[ 94, 100],
[229, 244]],
[[508, 532],
[733, 766]]
])
Пример
Давайте сравним bmm с обычным умножением через цикл:
import torch
t1 = torch.tensor([
[[1, 2],
[3, 4]],
[[5, 6],
[7, 8]]
])
t2 = torch.tensor([
[[9, 10],
[11, 12]],
[[13, 14],
[15, 16]]
])
res1 = torch.bmm(t1, t2)
res2 = torch.stack([
torch.mm(t1[0], t2[0]),
torch.mm(t1[1], t2[1])
])
print(res1)
print(res2)
Результат выполнения кода:
tensor([
[[ 31, 34],
[ 71, 78]],
[[155, 166],
[211, 226]]
])
tensor([
[[ 31, 34],
[ 71, 78]],
[[155, 166],
[211, 226]]
])