PyTorch torch.broadcast_shapes Function


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

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions