Функция isclose
Функция isclose в PyTorch используется для поэлементного сравнения двух тензоров.
Она возвращает тензор логических значений, указывающих, являются ли соответствующие
элементы близкими друг к другу с учётом заданных допусков.
Первый параметр - тензор input, второй - тензор other.
Функция также принимает параметры rtol (относительный допуск) и atol (абсолютный допуск),
которые по умолчанию равны 1e-05 и 1e-08 соответственно.
Сравнение выполняется по формуле: abs(a - b) <= (atol + rtol * abs(b)).
Синтаксис
torch.isclose(input, other, rtol=1e-05, atol=1e-08, equal_nan=False)
Пример
Давайте сравним два тензора с числами с плавающей точкой на близость:
import torch
t1 = torch.tensor([1.0, 2.0, 3.0])
t2 = torch.tensor([1.0, 2.0001, 3.0])
res = torch.isclose(t1, t2)
print(res)
Результат выполнения кода:
tensor([ True, False, True])
Пример
Теперь изменим относительный допуск, чтобы второй элемент тоже считался близким:
import torch
t1 = torch.tensor([1.0, 2.0, 3.0])
t2 = torch.tensor([1.0, 2.0001, 3.0])
res = torch.isclose(t1, t2, rtol=1e-3)
print(res)
Результат выполнения кода:
tensor([True, True, True])
Пример
Используем абсолютный допуск для сравнения чисел:
import torch
t1 = torch.tensor([1e-10, 1.0])
t2 = torch.tensor([1.1e-10, 1.0])
res = torch.isclose(t1, t2, atol=2e-10)
print(res)
Результат выполнения кода:
tensor([ True, True])
Пример
Обработаем значения nan с помощью параметра equal_nan:
import torch
t1 = torch.tensor([float('nan'), 1.0])
t2 = torch.tensor([float('nan'), 1.0])
res = torch.isclose(t1, t2, equal_nan=True)
print(res)
Результат выполнения кода:
tensor([ True, True])