PyTorch torch.broadcast_tensors function


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

Pytorch torch 参考手册Pytorch torch Reference Manual

Other extensions