Функция max
Функция max находит максимальное значение в тензоре.
Если не указывать размерность, она возвращает скалярное значение.
При указании размерности возвращает кортеж из двух тензоров:
максимальные значения и их индексы вдоль этой размерности.
Синтаксис
torch.max(tensor)
torch.max(tensor, dim)
Параметр dim задает размерность, вдоль которой ищется максимум.
Пример
Найдем максимальное значение в одномерном тензоре:
import torch
t = torch.tensor([3, 1, 4, 1, 5])
res = torch.max(t)
print(res)
Результат выполнения кода:
tensor(5)
Пример
Найдем максимальное значение в двумерном тензоре:
import torch
t = torch.tensor([
[1, 8, 3],
[4, 2, 6],
])
res = torch.max(t)
print(res)
Результат выполнения кода:
tensor(8)
Пример
Найдем максимальные значения вдоль столбцов (размерность 0):
import torch
t = torch.tensor([
[1, 8, 3],
[4, 2, 6],
])
res = torch.max(t, dim=0)
print(res)
Результат выполнения кода:
torch.return_types.max(
values=tensor([4, 8, 6]),
indices=tensor([1, 0, 1])
)
В результате получаем максимальные значения по каждому столбцу и их индексы.
Пример
Найдем максимальные значения вдоль строк (размерность 1):
import torch
t = torch.tensor([
[1, 8, 3],
[4, 2, 6],
])
res = torch.max(t, dim=1)
print(res)
Результат выполнения кода:
torch.return_types.max(
values=tensor([8, 6]),
indices=tensor([1, 2])
)
Пример
Используем функцию как метод тензора:
import torch
t = torch.tensor([10, 20, 15])
res = t.max()
print(res)
Результат выполнения кода:
tensor(20)