PyTorch torch.sort Function
PyTorch torch Reference Manual
torch.sortis a function in PyTorch used to sort along a specified dimension.
Function Definition
torch.sort(input, dim, descending, stable, out)
Usage Example
Example
import torch
x = torch.tensor([[3, 1, 2], [6, 4, 5]])
# Sort
values, indices = torch.sort(x)
print("Sorted values:")
print(values)
print("Sorted indices:")
print(indices)
# Sort in descending order
values_desc, _ = torch.sort(x, descending=True)
print("Descending order:")
print(values_desc)
x = torch.tensor([[3, 1, 2], [6, 4, 5]])
# Sort
values, indices = torch.sort(x)
print("Sorted values:")
print(values)
print("Sorted indices:")
print(indices)
# Sort in descending order
values_desc, _ = torch.sort(x, descending=True)
print("Descending order:")
print(values_desc)
Other Extensions