PyTorch torch.bmm Function


Pytorch torch 参考手册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)

The output result is:

批量结果形状: torch.Size([10, 3, 5])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions