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