Функция eq
Функция eq выполняет поэлементное сравнение двух тензоров на равенство.
Первым параметром функция принимает тензор, вторым - тензор или скалярное значение.
Результатом является тензор с логическими значениями True и False,
где каждый элемент показывает, равны ли соответствующие элементы входных тензоров.
Функция поддерживает вещание (broadcasting) - автоматическое расширение формы тензоров.
Синтаксис
torch.eq(tensor1, tensor2)
Альтернативный синтаксис с использованием метода eq тензора:
tensor1.eq(tensor2)
Пример
Давайте сравним два тензора с числами и увидим результат сравнения на равенство:
import torch
t1 = torch.tensor([1, 2, 3, 4, 5])
t2 = torch.tensor([1, 2, 0, 4, 8])
res = torch.eq(t1, t2)
print(res)
Результат выполнения кода:
tensor([ True, True, False, True, False])
Пример
Теперь сравним тензор со скалярным значением. Функция применит сравнение к каждому элементу:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
res = torch.eq(t, 3)
print(res)
Результат выполнения кода:
tensor([False, False, True, False, False])
Пример
Сравним двумерные тензоры с разной формой, используя механизм вещания:
import torch
t1 = torch.tensor([
[1, 2, 3],
[4, 5, 6],
])
t2 = torch.tensor([1, 2, 3])
res = t1.eq(t2)
print(res)
Результат выполнения кода:
tensor([
[ True, True, True],
[False, False, False],
])