Функция argmin
Функция argmin возвращает индекс первого вхождения минимального значения
в тензоре. Если ось не указана, поиск выполняется по всему тензору.
Параметр dim задаёт ось, вдоль которой производится поиск.
Параметр keepdim определяет, сохранять ли размерность.
Синтаксис
torch.argmin(t, [dim], [keepdim])
Пример
Давайте найдём индекс минимального элемента в одномерном тензоре:
import torch
t = torch.tensor([5, 2, 8, 1, 9])
res = torch.argmin(t)
print(res)
Результат выполнения кода:
tensor(3)
Пример
Найдём индекс минимального элемента вдоль строк (по оси 0)
в двумерном тензоре:
import torch
t = torch.tensor([
[3, 9, 2],
[8, 1, 5],
[4, 7, 6],
])
res = torch.argmin(t, dim=0)
print(res)
Результат выполнения кода:
tensor([2, 1, 0])
Пример
Найдём индекс минимального элемента вдоль столбцов (по оси 1):
import torch
t = torch.tensor([
[3, 9, 2],
[8, 1, 5],
[4, 7, 6],
])
res = torch.argmin(t, dim=1)
print(res)
Результат выполнения кода:
tensor([2, 1, 0])
Пример
Используем параметр keepdim для сохранения размерности:
import torch
t = torch.tensor([
[3, 9, 2],
[8, 1, 5],
[4, 7, 6],
])
res = torch.argmin(t, dim=1, keepdim=True)
print(res)
Результат выполнения кода:
tensor([
[2],
[1],
[0],
])
Пример
При наличии нескольких одинаковых минимальных значений возвращается индекс первого из них:
import torch
t = torch.tensor([4, 1, 7, 1, 3])
res = torch.argmin(t)
print(res)
Результат выполнения кода:
tensor(1)
Пример
Для трёхмерного тензора можно указать любую ось:
import torch
t = torch.tensor([
[
[1, 5, 2],
[4, 9, 3],
],
[
[7, 2, 6],
[0, 8, 1],
],
])
res = torch.argmin(t, dim=2)
print(res)
Результат выполнения кода:
tensor([
[0, 2],
[1, 0],
])