PyTorch torch.broadcast_tensors function
Pytorch torch Reference Manual
torch.broadcast_tensorsIt is a function in PyTorch used to broadcast multiple tensors. It broadcasts the input tensors to the same shape and returns a set of tensors that can be operated on element-wise. The broadcasting rules follow NumPy's broadcasting mechanism.
Function definition
torch.broadcast_tensors(*tensors)
Usage examples
Example
import torch
# Broadcasting tensors of different shapes
x = torch.tensor([[1, 2, 3]]) # Shape (1, 3)
y = torch.tensor([[1], [2], [3]]) # Shape (3, 1)
a, b = torch.broadcast_tensors(x, y)
print("x broadcasted shape:", a.shape)
print("y broadcasted shape:", b.shape)
print("x after broadcasting:")
print(a)
print("y after broadcasting:")
print(b)
print("Element-wise sum:")
print(a + b)
# Multiple tensors
x1 = torch.randn(3, 1, 5)
x2 = torch.randn(1, 4, 5)
x3 = torch.randn(3, 4, 1)
r1, r2, r3 = torch.broadcast_tensors(x1, x2, x3)
print("Broadcasted shape:", r1.shape)
# Broadcasting tensors of different shapes
x = torch.tensor([[1, 2, 3]]) # Shape (1, 3)
y = torch.tensor([[1], [2], [3]]) # Shape (3, 1)
a, b = torch.broadcast_tensors(x, y)
print("x broadcasted shape:", a.shape)
print("y broadcasted shape:", b.shape)
print("x after broadcasting:")
print(a)
print("y after broadcasting:")
print(b)
print("Element-wise sum:")
print(a + b)
# Multiple tensors
x1 = torch.randn(3, 1, 5)
x2 = torch.randn(1, 4, 5)
x3 = torch.randn(3, 4, 1)
r1, r2, r3 = torch.broadcast_tensors(x1, x2, x3)
print("Broadcasted shape:", r1.shape)
Other extensions