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

Функция bmm

Функция bmm (batch matrix multiplication) выполняет пакетное перемножение матриц. Первым параметром функция принимает трёхмерный тензор input, вторым параметром - трёхмерный тензор mat2. Функция возвращает трёхмерный тензор, содержащий результат умножения каждой пары матриц из входных тензоров.

Синтаксис

torch.bmm(input, mat2)

Параметры

Первый параметр input - трёхмерный тензор формы (batch_size, n, m).

Второй параметр mat2 - трёхмерный тензор формы (batch_size, m, p).

Пример

Давайте выполним пакетное умножение двух тензоров размера 2x2:

import torch t1 = torch.tensor([ [[1, 2], [3, 4]], [[5, 6], [7, 8]] ]) t2 = torch.tensor([ [[9, 10], [11, 12]], [[13, 14], [15, 16]] ]) res = torch.bmm(t1, t2) print(res)

Результат выполнения кода:

tensor([ [[ 31, 34], [ 71, 78]], [[155, 166], [211, 226]] ])

Пример

Давайте выполним пакетное умножение тензоров разных размерностей 2x3 и 3x2:

import torch t1 = torch.tensor([ [[1, 2, 3], [4, 5, 6]], [[7, 8, 9], [10, 11, 12]] ]) t2 = torch.tensor([ [[13, 14], [15, 16], [17, 18]], [[19, 20], [21, 22], [23, 24]] ]) res = torch.bmm(t1, t2) print(res)

Результат выполнения кода:

tensor([ [[ 94, 100], [229, 244]], [[508, 532], [733, 766]] ])

Пример

Давайте сравним bmm с обычным умножением через цикл:

import torch t1 = torch.tensor([ [[1, 2], [3, 4]], [[5, 6], [7, 8]] ]) t2 = torch.tensor([ [[9, 10], [11, 12]], [[13, 14], [15, 16]] ]) res1 = torch.bmm(t1, t2) res2 = torch.stack([ torch.mm(t1[0], t2[0]), torch.mm(t1[1], t2[1]) ]) print(res1) print(res2)

Результат выполнения кода:

tensor([ [[ 31, 34], [ 71, 78]], [[155, 166], [211, 226]] ]) tensor([ [[ 31, 34], [ 71, 78]], [[155, 166], [211, 226]] ])

Смотрите также

  • функцию matmul,
    которая выполняет матричное умножение для тензоров любых размерностей
  • функцию mm,
    которая выполняет умножение двумерных матриц
  • функцию addmm,
    которая выполняет умножение матриц и прибавляет результат к другому тензору
  • функцию mv,
    которая выполняет умножение матрицы на вектор
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить