PyTorch torch.bincount function
Pytorch torch Reference Manual
torch.bincountIt is a function in PyTorch used to count the occurrences of each non-negative integer value. It returns a new 1-dimensional tensor, where the i-th element represents the number of times value i appears in the input. It is commonly used for histogram calculation and group statistics.
Function definition
torch.bincount(input, weights=None, minlength=0)
Usage examples
Example
import torch
# Basic usage: count the occurrences of each value
x = torch.tensor([0, 1, 1, 2, 2, 2, 3, 3, 4])
counts = torch.bincount(x)
print("Input:", x)
print("Counting result:", counts)
# Output: tensor([1, 2, 3, 2, 1])
# Using the weights parameter
x = torch.tensor([0, 1, 1, 2, 2, 2])
weights = torch.tensor([1.0, 2.0, 3.0, 1.0, 2.0, 3.0])
weighted_counts = torch.bincount(x, weights=weights)
print("Weighted count:", weighted_counts)
# Output: tensor([1., 5., 6.])
# Set minimum length
x = torch.tensor([0])
counts = torch.bincount(x, minlength=5)
print("Minimum length 5:", counts)
# Output: tensor([1, 0, 0, 0, 0])
# Basic usage: count the occurrences of each value
x = torch.tensor([0, 1, 1, 2, 2, 2, 3, 3, 4])
counts = torch.bincount(x)
print("Input:", x)
print("Counting result:", counts)
# Output: tensor([1, 2, 3, 2, 1])
# Using the weights parameter
x = torch.tensor([0, 1, 1, 2, 2, 2])
weights = torch.tensor([1.0, 2.0, 3.0, 1.0, 2.0, 3.0])
weighted_counts = torch.bincount(x, weights=weights)
print("Weighted count:", weighted_counts)
# Output: tensor([1., 5., 6.])
# Set minimum length
x = torch.tensor([0])
counts = torch.bincount(x, minlength=5)
print("Minimum length 5:", counts)
# Output: tensor([1, 0, 0, 0, 0])
Other extensions