PyTorch torch.combinations Function
PyTorch torch Reference Manual
torch.combinationsIt is a function in PyTorch used to compute all r-element combinations of input tensor elements. It returns all possible combinations of length r from the input tensor.
Function Definition
torch.combinations(input, r=2, with_replacement=False)
Usage Examples
Example
import torch
# Calculate all 2-element combinations
x = torch.tensor([1, 2, 3, 4])
result = torch.combinations(x, r=2)
print("Input:", x)
print("2-element combinations:")
print(result)
# tensor([[1, 2],
# [1, 3],
# [1, 4],
# [2, 3],
# [2, 4],
# [3, 4]])
# 3-element combinations
result3 = torch.combinations(x, r=3)
print("3-element combinations:")
print(result3)
# Combinations with replacement (with_replacement=True)
result_with_replacement = torch.combinations(x, r=2, with_replacement=True)
print("2-element combinations with replacement:")
print(result_with_replacement)
# Calculate all 2-element combinations
x = torch.tensor([1, 2, 3, 4])
result = torch.combinations(x, r=2)
print("Input:", x)
print("2-element combinations:")
print(result)
# tensor([[1, 2],
# [1, 3],
# [1, 4],
# [2, 3],
# [2, 4],
# [3, 4]])
# 3-element combinations
result3 = torch.combinations(x, r=3)
print("3-element combinations:")
print(result3)
# Combinations with replacement (with_replacement=True)
result_with_replacement = torch.combinations(x, r=2, with_replacement=True)
print("2-element combinations with replacement:")
print(result_with_replacement)
Other Extensions