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