Функция mm
Функция mm выполняет матричное умножение двух двумерных тензоров.
Первым параметром передаётся левый множитель (тензор размера n×m),
вторым - правый множитель (тензор размера m×p).
Результатом является тензор размера n×p.
Функция не поддерживает broadcasting, тензоры должны быть строго двумерными.
Синтаксис
torch.mm(input, mat2)
Пример
Выполним умножение двух матриц размером 2×3 и 3×2:
import torch
t1 = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
t2 = torch.tensor([
[7, 8],
[9, 10],
[11, 12],
])
res = torch.mm(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[ 58, 64],
[139, 154],
])
Пример
Попробуем умножить квадратные матрицы 3×3:
import torch
t1 = torch.tensor([
[1, 0, 0],
[0, 1, 0],
[0, 0, 1],
])
t2 = torch.tensor([
[2, 3, 4],
[5, 6, 7],
[8, 9, 10],
])
res = torch.mm(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[2, 3, 4],
[5, 6, 7],
[8, 9, 10],
])
Пример
Умножение матрицы на вектор выполняется через функцию mv,
но с помощью mm тоже можно, если вектор представить как матрицу 3×1:
import torch
t1 = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
t2 = torch.tensor([
[7],
[8],
[9],
])
res = torch.mm(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[ 50],
[122],
])