РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
278 of 769 menu

Функция 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,
    которая вычисляет норму тензора
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить