PyTorch torch.combinations Function


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

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions