torch.mm
-
torch.mm(input, mat2, *, out=None) → Tensor -
Выполняет матричное умножение матриц
inputиmat2.Если
input— тензор ,mat2— тензор ,outбудет тензором .Примечание
Эта функция не распространяется. Для трансляции матричных произведений см.
torch.matmul().Поддерживает строковые и разреженные 2-мерные тензоры в качестве входных данных, 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/1.13/generated/torch.mm.html