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

PyTorch에서 텐서 차원 압축(squeeze)과 확장(unsqueeze)하는 방법

텐서 Squeeze와 Unsqueeze란?

PyTorch에서 텐서의 차원을 조작할 때는 torch.squeeze()torch.unsqueeze() 메서드를 사용합니다.

torch.squeeze()는 입력 텐서에서 크기가 1인 차원을 모두 제거한 새로운 텐서를 반환합니다. 예를 들어 입력 텐서의 형태(shape)가 (M × 1 × N × 1 × P)라면, squeeze 후에는 (M × N × P) 형태가 됩니다.

반대로 torch.unsqueeze()는 지정한 위치(dim)에 크기가 1인 새로운 차원을 삽입한 텐서를 반환합니다. 딥러닝에서 배치(batch) 차원을 추가하거나 제거할 때 매우 자주 사용되는 연산입니다.

처리 순서

  • 필요한 라이브러리를 임포트합니다. 아래 모든 Python 예제에서 필요한 라이브러리는 torch이며, 미리 설치되어 있어야 합니다.

  • 텐서를 생성하고 출력합니다.

  • torch.squeeze(input)을 호출합니다. 크기가 1인 차원을 제거하고, 나머지 차원은 그대로 유지한 텐서를 반환합니다.

  • torch.unsqueeze(input, dim)을 호출합니다. 지정한 dim 위치에 크기가 1인 새로운 차원을 삽입한 텐서를 반환합니다.

  • squeeze 또는 unsqueeze된 결과 텐서를 출력합니다.

예제 1

# 텐서를 squeeze/unsqueeze하는 Python 프로그램
# 필요한 라이브러리 임포트
import torch

# 모든 요소가 1인 텐서 생성
T = torch.ones(2,1,2) # 크기 2x1x2
print("Original Tensor T:\n", T )
print("Size of T:", T.size())

# 텐서의 차원 squeeze
squeezed_T = torch.squeeze(T) # 이제 크기 2x2
print("Squeezed_T\n:", squeezed_T )
print("Size of Squeezed_T:", squeezed_T.size())

실행 결과

Original Tensor T:
tensor([[[1., 1.]],
         [[1., 1.]]])
Size of T: torch.Size([2, 1, 2])
Squeezed_T
: tensor([[1., 1.],
         [1., 1.]])
Size of Squeezed_T: torch.Size([2, 2])

위 예제에서 원본 텐서 T는 크기가 (2, 1, 2)였지만, 가운데 크기가 1인 차원이 제거되어 squeeze 후에는 (2, 2)가 된 것을 확인할 수 있습니다.

예제 2

# 텐서를 squeeze/unsqueeze하는 Python 프로그램
# 필요한 라이브러리 임포트
import torch

# 텐서 생성
T = torch.Tensor([1,2,3]) # 크기 3
print("Original Tensor T:\n", T )
print("Size of T:", T.size())

# dim=0 위치에 새로운 차원 삽입 (행 방향)
unsqueezed_T = torch.unsqueeze(T, dim = 0) # 이제 크기 1x3
print("Unsqueezed T\n:", unsqueezed_T )
print("Size of UnSqueezed T:", unsqueezed_T.size())

# dim=1 위치에 새로운 차원 삽입 (열 방향)
unsqueezed_T = torch.unsqueeze(T, dim = 1) # 이제 크기 3x1
print("Unsqueezed T\n:", unsqueezed_T )
print("Size of Unsqueezed T:", unsqueezed_T.size())

실행 결과

Original Tensor T:
   tensor([1., 2., 3.])
Size of T: torch.Size([3])
Unsqueezed T
: tensor([[1., 2., 3.]])
Size of UnSqueezed T: torch.Size([1, 3])
Unsqueezed T
: tensor([[1.],
         [2.],
         [3.]])
Size of Unsqueezed T: torch.Size([3, 1])

이 예제에서 알 수 있듯이, dim 인자에 어떤 값을 지정하느냐에 따라 크기가 1인 차원이 삽입되는 위치가 달라집니다. dim=0을 지정하면 첫 번째 축에 차원이 추가되어 (1, 3) 형태가 되고, dim=1을 지정하면 두 번째 축에 차원이 추가되어 (3, 1) 형태가 됩니다. 이처럼 squeeze와 unsqueeze를 활용하면 모델 입력 형태를 맞추거나 배치 차원을 손쉽게 조절할 수 있습니다.