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 (Последовательность[Тензор]) – два или более тензоров для умножения. Первый и последний тензоры могут быть одномерными или двумерными. Все остальные тензоры должны быть двумерными.
- Ключевые аргументы:
-
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/1.13/generated/torch.linalg.multi_dot.html