PyTorch torch.bmm Function
Pytorch torch Reference Manual
torch.bmmis a function in PyTorch used to perform batched matrix multiplication.
Function Definition
torch.bmm(input, mat2, out)
Usage Example
Example
import torch
# Batch matrix multiplication
batch_a = torch.randn(10, 3, 4)
batch_b = torch.randn(10, 4, 5)
result = torch.bmm(batch_a, batch_b)
print("Batch result shape:", result.shape)
# Batch matrix multiplication
batch_a = torch.randn(10, 3, 4)
batch_b = torch.randn(10, 4, 5)
result = torch.bmm(batch_a, batch_b)
print("Batch result shape:", result.shape)
The output result is:
批量结果形状: torch.Size([10, 3, 5])
Other Extensions