PyTorch torch.broadcast_to function
Pytorch torch Reference Manual
torch.broadcast_toIt is a function in PyTorch used to broadcast a tensor to a specified shape. It returns a view of the input tensor that is broadcast to the target shape. The broadcasting rules follow NumPy's broadcasting mechanism.
Function Definition
torch.broadcast_to(input, shape)
Usage Example
Example
import torch
# Basic usage: broadcast a 1-D tensor to 2-D
x = torch.tensor([1, 2, 3])
y = torch.broadcast_to(x, (3, 3))
print("Original:", x)
print("After broadcasting:")
print(y)
# tensor([[1, 2, 3],
# [1, 2, 3],
# [1, 2, 3]])
# Broadcast a scalar to a larger shape
x = torch.tensor(5)
y = torch.broadcast_to(x, (2, 3, 4))
print("Scalar broadcast to (2,3,4):", y.shape)
# Broadcast a 2-D tensor to 3-D
x = torch.tensor([[1, 2], [3, 4]]) # (2, 2)
y = torch.broadcast_to(x, (3, 2, 2))
print("Broadcast to (3,2,2):", y.shape)
print(y)
# Basic usage: broadcast a 1-D tensor to 2-D
x = torch.tensor([1, 2, 3])
y = torch.broadcast_to(x, (3, 3))
print("Original:", x)
print("After broadcasting:")
print(y)
# tensor([[1, 2, 3],
# [1, 2, 3],
# [1, 2, 3]])
# Broadcast a scalar to a larger shape
x = torch.tensor(5)
y = torch.broadcast_to(x, (2, 3, 4))
print("Scalar broadcast to (2,3,4):", y.shape)
# Broadcast a 2-D tensor to 3-D
x = torch.tensor([[1, 2], [3, 4]]) # (2, 2)
y = torch.broadcast_to(x, (3, 2, 2))
print("Broadcast to (3,2,2):", y.shape)
print(y)
Other Extensions