PyTorch torch.argsort Function
PyTorch torch Reference Manual
torch.argsortIt is a function in PyTorch used to return the indices after sorting. It returns the original position indices of each element after the tensor is sorted.
Function Definition
torch.argsort(input, dim=-1, descending=False, stable=True)
Usage Examples
Example
import torch
# Create tensor
x = torch.tensor([3, 1, 2])
# Return the sorted indices
indices = torch.argsort(x)
print(indices)
# Create tensor
x = torch.tensor([3, 1, 2])
# Return the sorted indices
indices = torch.argsort(x)
print(indices)
The output result is:
tensor([1, 2, 0])
Other Extensions