Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

PyTorch에서 두 텐서를 비교하는 방법 – torch.eq() 완벽 가이드

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)이나 조건 필터링 등 다양한 용도로 활용됩니다.