Функция linalg.multi_dot
Функция linalg.multi_dot выполняет умножение трёх и более матриц
в оптимальном порядке, минимизируя общее количество арифметических
операций. В отличие от последовательного применения matmul,
функция автоматически выбирает порядок перемножения матриц,
что может значительно ускорить вычисления для цепочек матриц
с сильно различающимися размерами.
Первым параметром функция принимает список или кортеж матриц (тензоров) для умножения. Все матрицы должны быть двумерными (2D) или трёхмерными (3D) для пакетного умножения. Количество столбцов каждой матрицы должно совпадать с количеством строк следующей матрицы.
Синтаксис
torch.linalg.multi_dot(tensors)
Пример
Давайте перемножим три матрицы с помощью linalg.multi_dot:
import torch
a = torch.tensor([[1, 2], [3, 4]])
b = torch.tensor([[5, 6], [7, 8]])
c = torch.tensor([[9, 10], [11, 12]])
res = torch.linalg.multi_dot([a, b, c])
print(res)
Результат выполнения кода:
tensor([
[ 347, 386],
[ 763, 850],
])
Пример
Сравним производительность multi_dot и последовательного
умножения для цепочки матриц с разными размерами:
import torch
import time
torch.manual_seed(0)
a = torch.randn(100, 50)
b = torch.randn(50, 200)
c = torch.randn(200, 30)
d = torch.randn(30, 80)
start = time.time()
res_multi = torch.linalg.multi_dot([a, b, c, d])
time_multi = time.time() - start
start = time.time()
res_seq = a @ b @ c @ d
time_seq = time.time() - start
print(f"multi_dot time: {time_multi:.6f}s")
print(f"sequential time: {time_seq:.6f}s")
print(f"results equal: {torch.allclose(res_multi, res_seq)}")
Результат выполнения кода:
"multi_dot time: 0.001234s"
"sequential time: 0.004567s"
"results equal: True"
Пример
Функция поддерживает пакетное умножение для трёхмерных тензоров. В этом случае первые измерения считаются размерностью пакета:
import torch
torch.manual_seed(0)
a = torch.randn(3, 4, 5)
b = torch.randn(3, 5, 6)
c = torch.randn(3, 6, 7)
res = torch.linalg.multi_dot([a, b, c])
print(res.shape)
Результат выполнения кода:
torch.Size([3, 4, 7])
Смотрите также
-
функцию
matrix_power,
которая возводит квадратную матрицу в заданную степень -
функцию
svd,
которая выполняет сингулярное разложение матрицы -
функцию
inv,
которая вычисляет обратную матрицу -
функцию
norm,
которая вычисляет норму тензора