PyTorch torch.mm Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.mmis a function in PyTorch used to perform two-dimensional matrix multiplication.

Function Definition

torch.mm(input, mat2, out)

Usage Example

Example

import torch

A = torch.randn(2, 3)
B = torch.randn(3, 4)

# Matrix multiplication
C = torch.mm(A, B)

print("A shape:", A.shape)
print("B shape:", B.shape)
print("C shape:", C.shape)

The output result is:

A 形状: torch.Size([2, 3])
B 形状: torch.Size([3, 4])
C 形状: torch.Size([2, 4])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions