PyTorch torch.baddbmm Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.baddbmmIt is a function in PyTorch used to perform batched matrix multiplication and add it to an input batched matrix. It performs matrix multiplication on each pair of matrices in batch1 and batch2, then adds the result to the corresponding matrix in input.

Function Definition

torch.baddbmm(input, batch1, batch2, *, beta=1.0, alpha=1.0, out=None)

Parameters:

  • input(Tensor): Input batched matrix, shape (b, n, p).
  • batch1(Tensor): First batched matrix, shape (b, n, m).
  • batch2(Tensor): Second batched matrix, shape (b, m, p).
  • beta(float, optional): Coefficient multiplied by input, default is 1.0.
  • alpha(float, optional): Coefficient multiplied by the result of batch1 @ batch2, default is 1.0.
  • out(Tensor, optional): Output tensor.

Return Value:

  • torch.Tensor: Returns the sum of the batched matrix multiplication result and the input batched matrix, with shape (b, n, p).

Usage Example

Example

import torch

# Create input batched matrix and batched matrices
input = torch.randn(10, 3, 3)
batch1 = torch.randn(10, 3, 4)
batch2 = torch.randn(10, 4, 3)

# Perform baddbmm
result = torch.baddbmm(input, batch1, batch2)

print("Input batched matrix shape:", input.shape)
print("Batched matrix 1 shape:", batch1.shape)
print("Batched matrix 2 shape:", batch2.shape)
print("Result shape:", result.shape)
print(result[0])  # Print the first result

The output result is:

输入批量矩阵形状: torch.Size([10, 3, 3])
批量矩阵1形状: torch.Size([10, 3, 4])
批量矩阵2形状: torch.Size([10, 4, 3])
结果形状: torch.Size([10, 3, 3])
tensor([[-1.0917,  0.0174, -0.6599],
        [ 1.0403, -0.3905, -0.6314],
        [ 0.0988, -0.3047,  0.4841]])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions