PyTorch torch.prod Function
PyTorch torch Reference Manual
torch.prodis a function in PyTorch used to return the product of all elements in a tensor.
Function Definition
torch.prod(input, dim, keepdim=False)
Usage Example
Example
import torch
x = torch.tensor([1, 2, 3, 4])
# Return the product of all elements
print("Product of all elements:", torch.prod(x))
# Product along dim=0
y = torch.tensor([[1, 2, 3], [4, 5, 6]])
print("dim=0 product:", torch.prod(y, dim=0))
print("dim=1 product:", torch.prod(y, dim=1))
x = torch.tensor([1, 2, 3, 4])
# Return the product of all elements
print("Product of all elements:", torch.prod(x))
# Product along dim=0
y = torch.tensor([[1, 2, 3], [4, 5, 6]])
print("dim=0 product:", torch.prod(y, dim=0))
print("dim=1 product:", torch.prod(y, dim=1))
The output result is:
所有元素乘积: tensor(24) dim=0 乘积: tensor([4, 10, 18]) dim=1 乘积: tensor([ 6, 120])
Other Extensions