PyTorch torch.broadcast_shapes Function
Pytorch torch Reference Manual
torch.broadcast_shapesIt is a function in PyTorch used to compute the result of broadcasting multiple shapes. It returns the shape obtained after broadcasting the input shapes without actually creating tensors.
Function Definition
torch.broadcast_shapes(*shapes)
Usage Example
Example
import torch
# Compute shape broadcast result
result = torch.broadcast_shapes((3, 1), (1, 4), (3, 4))
print("(3,1) + (1,4) + (3,4) ->", result)
# Output: (3, 4)
# Multiple shapes
result = torch.broadcast_shapes((5,), (1, 5), (3, 1, 5))
print("(5,) + (1,5) + (3,1,5) ->", result)
# Output: (3, 1, 5)
# Single shape
result = torch.broadcast_shapes((2, 3))
print("(2,3) ->", result)
# Output: (2, 3)
# An error is reported when broadcasting is not possible
try:
result = torch.broadcast_shapes((3,), (4,))
except RuntimeError as e:
print("Broadcast error:", e)
# Compute shape broadcast result
result = torch.broadcast_shapes((3, 1), (1, 4), (3, 4))
print("(3,1) + (1,4) + (3,4) ->", result)
# Output: (3, 4)
# Multiple shapes
result = torch.broadcast_shapes((5,), (1, 5), (3, 1, 5))
print("(5,) + (1,5) + (3,1,5) ->", result)
# Output: (3, 1, 5)
# Single shape
result = torch.broadcast_shapes((2, 3))
print("(2,3) ->", result)
# Output: (2, 3)
# An error is reported when broadcasting is not possible
try:
result = torch.broadcast_shapes((3,), (4,))
except RuntimeError as e:
print("Broadcast error:", e)
Other Extensions