PyTorch torch.sort Function


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

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions