Функция topk
Функция topk возвращает k наибольших или наименьших элементов тензора вдоль указанной размерности. Первым параметром функция принимает тензор, вторым - количество элементов для возврата. Третьим параметром можно указать размерность, по которой выполняется поиск. Четвертым параметром задается направление сортировки: по убыванию (наибольшие) или по возрастанию (наименьшие). Функция возвращает кортеж из двух тензоров: значений и индексов.
Синтаксис
torch.topk(input, k, dim, largest, sorted)
Пример
Давайте найдем три наибольших элемента в одномерном тензоре:
import torch
t = torch.tensor([1, 5, 3, 9, 2])
values, indices = torch.topk(t, 3)
print(values)
print(indices)
Результат выполнения кода:
tensor([9, 5, 3])
tensor([3, 1, 2])
Пример
Давайте найдем два наименьших элемента в одномерном тензоре:
import torch
t = torch.tensor([1, 5, 3, 9, 2])
values, indices = torch.topk(t, 2, largest=False)
print(values)
print(indices)
Результат выполнения кода:
tensor([1, 2])
tensor([0, 4])
Пример
Давайте найдем два наибольших элемента в каждой строке двумерного тензора:
import torch
t = torch.tensor([
[1, 5, 3],
[9, 2, 7],
])
values, indices = torch.topk(t, 2, dim=1)
print(values)
print(indices)
Результат выполнения кода:
tensor([
[5, 3],
[9, 7],
])
tensor([
[1, 2],
[0, 2],
])
Пример
Давайте получим отсортированные значения тензора с помощью функции topk:
import torch
t = torch.tensor([5, 1, 8, 3, 6])
sorted_values, _ = torch.topk(t, t.numel())
print(sorted_values)
Результат выполнения кода:
tensor([8, 6, 5, 3, 1])
Смотрите также
-
функцию
sort,
которая сортирует элементы тензора вдоль указанной размерности -
функцию
argsort,
которая возвращает индексы отсортированных элементов тензора -
функцию
kthvalue,
которая возвращает k-е по величине значение в тензоре -
функцию
max,
которая возвращает максимальный элемент тензора или по размерности