PyTorch에서 두 텐서(tensors)를 요소 단위(element-wise)로 비교하려면 torch.eq() 메서드를 사용합니다. 이 메서드는 두 텐서의 대응하는 요소들을 하나씩 비교하여, 값이 같으면 True를, 다르면 False를 반환합니다. 차원이 같은 텐서뿐만 아니라 차원이 다른 텐서도 비교할 수 있으며, 단일 요소(singleton)가 아닌 차원에서는 두 텐서의 크기가 서로 일치해야 한다는 점에 유의해야 합니다.
비교 절차
필요한 라이브러리를 임포트합니다. 아래의 모든 Python 예제에서는 torch 라이브러리가 필요하므로, 미리 설치되어 있는지 확인하세요.
PyTorch 텐서를 생성하고 출력합니다.
torch.eq(input1, input2)를 실행합니다. 이 연산은 텐서를 요소 단위로 비교하여 대응하는 요소들이 같으면 True, 다르면 False로 구성된 불리언(Boolean) 텐서를 반환합니다.
반환된 결과 텐서를 출력하여 확인합니다.
예제 1: 1차원 텐서 비교
다음 Python 프로그램은 두 개의 1차원(1-D) 텐서를 요소 단위로 비교하는 방법을 보여줍니다.
# 필요한 라이브러리 임포트
import torch
# 두 개의 텐서 생성
T1 = torch.Tensor([2.4,5.4,-3.44,-5.43,43.5])
T2 = torch.Tensor([2.4,5.5,-3.44,-5.43, 43])
# 생성된 텐서 출력
print("T1:", T1)
print("T2:", T2)
# T1과 T2 텐서를 요소 단위로 비교
print(torch.eq(T1, T2))실행 결과
T1: tensor([ 2.4000, 5.4000, -3.4400, -5.4300, 43.5000]) T2: tensor([ 2.4000, 5.5000, -3.4400, -5.4300, 43.0000]) tensor([ True, False, True, True, False])
예제 2: 2차원 텐서 비교
다음 예제는 두 개의 2차원(2-D) 텐서를 요소 단위로 비교하는 방법을 보여줍니다.
# 필요한 라이브러리 임포트
import torch
# 4x3 크기의 2D 텐서 두 개 생성
T1 = torch.Tensor([[2,3,-32],
[43,4,-53],
[4,37,-4],
[3,75,34]])
T2 = torch.Tensor([[2,3,-32],
[4,4,-53],
[4,37,4],
[3,-75,34]])
# 생성된 텐서 출력
print("T1:", T1)
print("T2:", T2)
# T1과 T2 텐서를 요소 단위로 비교
print(torch.eq(T1, T2))실행 결과
T1: tensor([[ 2., 3., -32.],
[ 43., 4., -53.],
[ 4., 37., -4.],
[ 3., 75., 34.]])
T2: tensor([[ 2., 3., -32.],
[ 4., 4., -53.],
[ 4., 37., 4.],
[ 3., -75., 34.]])
tensor([[ True, True, True],
[False, True, True],
[ True, True, False],
[ True, False, True]])예제 3: 1차원 텐서와 2차원 텐서 비교
torch.eq()는 브로드캐스팅(broadcasting)을 지원하기 때문에, 차원이 다른 텐서 간의 비교도 가능합니다. 다음 프로그램은 1차원 텐서와 2차원 텐서를 요소 단위로 비교하는 방법을 보여줍니다.
# 필요한 라이브러리 임포트
import torch
# 두 개의 텐서 생성
T1 = torch.Tensor([2.4,5.4,-3.44,-5.43,43.5])
T2 = torch.Tensor([[2.4,5.5,-3.44,-5.43, 7],
[1.0,5.4,3.88,4.0,5.78]])
# 생성된 텐서 출력
print("T1:", T1)
print("T2:", T2)
# T1과 T2 텐서를 요소 단위로 비교
print(torch.eq(T1, T2))실행 결과
T1: tensor([ 2.4000, 5.4000, -3.4400, -5.4300, 43.5000])
T2: tensor([[ 2.4000, 5.5000, -3.4400, -5.4300, 7.0000],
[ 1.0000, 5.4000, 3.8800, 4.0000, 5.7800]])
tensor([[ True, False, True, True, False],
[False, True, False, False, False]])마무리
PyTorch에서 텐서를 요소 단위로 비교할 때는 torch.eq() 메서드가 가장 간편한 방법입니다. 동일한 형태(shape)의 텐서뿐 아니라 브로드캐스팅 규칙을 만족하는 서로 다른 차원의 텐서도 비교할 수 있어, 딥러닝 모델 개발 시 마스킹(masking)이나 조건 필터링 등 다양한 용도로 활용됩니다.