Функция sort
Функция sort сортирует элементы тензора вдоль указанного измерения.
Первым параметром функция принимает тензор.
Вторым параметром можно передать ось для сортировки.
Третьим параметром можно указать сортировку по возрастанию или убыванию.
Функция возвращает кортеж из двух тензоров: отсортированные значения и индексы исходных элементов.
Синтаксис
torch.sort(input, dim=-1, descending=False)
Пример
Давайте отсортируем одномерный тензор по возрастанию:
import torch
t = torch.tensor([5, 2, 8, 1, 9, 3])
sorted_t, indices = torch.sort(t)
print(sorted_t)
print(indices)
Результат выполнения кода:
tensor([1, 2, 3, 5, 8, 9])
tensor([3, 1, 5, 0, 2, 4])
Пример
Отсортируем одномерный тензор по убыванию:
import torch
t = torch.tensor([5, 2, 8, 1, 9, 3])
sorted_t, indices = torch.sort(t, descending=True)
print(sorted_t)
print(indices)
Результат выполнения кода:
tensor([9, 8, 5, 3, 2, 1])
tensor([4, 2, 0, 5, 1, 3])
Пример
Отсортируем двумерный тензор по строкам (по умолчанию):
import torch
t = torch.tensor([
[5, 2, 8],
[1, 9, 3],
[7, 4, 6]
])
sorted_t, indices = torch.sort(t)
print(sorted_t)
print(indices)
Результат выполнения кода:
tensor([
[2, 5, 8],
[1, 3, 9],
[4, 6, 7]
])
tensor([
[1, 0, 2],
[0, 2, 1],
[1, 2, 0]
])
Пример
Отсортируем двумерный тензор по столбцам:
import torch
t = torch.tensor([
[5, 2, 8],
[1, 9, 3],
[7, 4, 6]
])
sorted_t, indices = torch.sort(t, dim=0)
print(sorted_t)
print(indices)
Результат выполнения кода:
tensor([
[1, 2, 3],
[5, 4, 6],
[7, 9, 8]
])
tensor([
[1, 0, 1],
[0, 2, 2],
[2, 1, 0]
])