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