torch.linalg.multi_dot
-
torch.linalg.multi_dot(tensors, *, out=None) -
Эффективно перемножает две или более матрицы, изменяя порядок перемножения так, чтобы выполнялось наименьшее количество арифметических операций.
Поддерживает входные данные типов float, double, cfloat и cdouble. Данная функция не поддерживает пакетные входные данные.
Каждый тензор в
tensorsдолжен быть двумерным, за исключением первого и последнего, которые могут быть одномерными. Если первый тензор является одномерным вектором формы(n,), он обрабатывается как строчный вектор формы(1, n), аналогично, если последний тензор является одномерным вектором формы(n,), он обрабатывается как столбцовый вектор формы(n, 1).Если первый и последний тензоры являются матрицами, результат будет матрицей. Однако, если один из них является одномерным вектором, то результатом будет одномерный вектор.
Отличия от
numpy.linalg.multi_dot:- В отличие от
numpy.linalg.multi_dot, первый и последний тензоры должны быть либо одномерными, либо двумерными, тогда как NumPy позволяет им быть n-мерными
Предупреждение
Данная функция не выполняет трансляцию.
Примечание
Эта функция реализована путем объединения вызовов
torch.mm()после вычисления оптимального порядка умножения матриц.Примечание
Стоимость перемножения двух матриц с формами
(a, b)и(b, c)составляетa * b * c. Учитывая матрицыA,B,Cс формами(10, 100),(100, 5),(5, 50)соответственно, мы можем вычислить стоимость различных порядков умножения следующим образом:В данном случае перемножение
AиBвначале, а затемCпроисходит в 10 раз быстрее.- Параметры
-
tensors (Sequence[Тензор]) – две или более матриц для перемножения. Первая и последняя матрицы могут быть одномерными или двумерными. Все остальные матрицы должны быть двумерными.
- Ключевые аргументы
-
out (Тензор, необязательно) – выходной тензор. Игнорируется, если
None. По умолчанию:None.
Примеры:
>>> from torch.linalg import multi_dot >>> multi_dot([torch.tensor([1, 2]), torch.tensor([2, 3])]) tensor(8) >>> multi_dot([torch.tensor([[1, 2]]), torch.tensor([2, 3])]) tensor([8]) >>> multi_dot([torch.tensor([[1, 2]]), torch.tensor([[2], [3]])]) tensor([[8]]) >>> A = torch.arange(2 * 3).view(2, 3) >>> B = torch.arange(3 * 2).view(3, 2) >>> C = torch.arange(2 * 2).view(2, 2) >>> multi_dot((A, B, C)) tensor([[ 26, 49], [ 80, 148]]) - В отличие от
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.linalg.multi_dot.html