torch.mm
-
torch.mm(input, mat2, *, out=None) → Tensor -
Выполняет матричное умножение матриц
inputиmat2.Если
inputявляется тензором ,mat2является тензором , тоoutбудет тензором .Примечание
Данная функция не выполняет векторное расширение. Для векторного расширения матричных произведений см.
torch.matmul().Поддерживает стэдированные и разреженные двумерные тензоры в качестве входных данных, а также autograd по отношению к стэдированным входным данным.
Эта операция поддерживает аргументы с разреженными макетами. Если
outпредоставлен, его макет будет использован. В противном случае макет результата будет определён на основе макетаinput.Предупреждение
Поддержка разреженных данных — это бета-функция, и некоторые сочетания макета(ов)/типа данных/устройств могут быть не поддерживаемы или не поддерживать autograd. Если вы заметили отсутствие функциональности, пожалуйста, откройте запрос на новую функцию.
Этот оператор поддерживает TensorFloat32.
На некоторых устройствах ROCm при использовании входных данных float16 этот модуль будет использовать разную точность для обратного прохода.
- Параметры
- Ключевые аргументы
-
out (Tensor, необязательно) – выходной тензор.
Пример:
>>> mat1 = torch.randn(2, 3) >>> mat2 = torch.randn(3, 3) >>> torch.mm(mat1, mat2) tensor([[ 0.4851, 0.5037, -0.3633], [-0.0760, -3.6705, 2.4784]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.mm.html