Функция maximum
Функция maximum вычисляет поэлементный максимум двух тензоров или тензора и скаляра.
Первый аргумент - тензор input.
Второй аргумент - тензор other или скаляр.
Функция возвращает тензор той же формы, что и входные данные.
Поддерживается автоматическое расширение (broadcasting).
Синтаксис
torch.maximum(input, other, *, out=None)
Пример
Давайте сравним два тензора и получим поэлементный максимум:
import torch
t1 = torch.tensor([1, 5, 3, 8])
t2 = torch.tensor([4, 2, 6, 7])
res = torch.maximum(t1, t2)
print(res)
Результат выполнения кода:
tensor([4, 5, 6, 8])
Пример
Давайте сравним тензор со скаляром:
import torch
t = torch.tensor([-2, 5, 0, 3])
res = torch.maximum(t, 1)
print(res)
Результат выполнения кода:
tensor([1, 5, 1, 3])
Пример
Давайте используем broadcasting для двумерных тензоров:
import torch
t1 = torch.tensor([[1, 2], [3, 4]])
t2 = torch.tensor([[0, 5], [2, 1]])
res = torch.maximum(t1, t2)
print(res)
Результат выполнения кода:
tensor([
[1, 5],
[3, 4],
])
Пример
Давайте используем maximum с числами с плавающей запятой:
import torch
t1 = torch.tensor([1.5, -2.3, 3.7])
t2 = torch.tensor([0.2, 4.1, -1.5])
res = torch.maximum(t1, t2)
print(res)
Результат выполнения кода:
tensor([1.5000, 4.1000, 3.7000])
Пример
Давайте используем аргумент out для сохранения результата в заранее созданный тензор:
import torch
t1 = torch.tensor([3, 1, 4])
t2 = torch.tensor([2, 5, 1])
out = torch.empty(3, dtype=torch.long)
torch.maximum(t1, t2, out=out)
print(out)
Результат выполнения кода:
tensor([3, 5, 4])