PyTorch에서는 torch.cat()과 torch.stack() 두 가지 함수를 사용하여 두 개 이상의 텐서를 하나로 합칠 수 있습니다. torch.cat()은 여러 텐서를 기존 차원을 따라 이어 붙이는(연결) 함수이고, torch.stack()은 새로운 차원을 만들어 텐서를 쌓는(stack) 함수입니다. 0차원, -1차원 등 다양한 차원을 지정하여 텐서를 원하는 방향으로 결합할 수 있습니다.
두 함수 모두 텐서를 합치는 용도로 사용되지만, 그 동작 방식에는 중요한 차이가 있습니다.
- torch.cat(): 이미 존재하는 차원을 따라 텐서들을 연결하므로, 결과 텐서의 차원 수가 변하지 않습니다.
- torch.stack(): 새로운 차원을 추가하며 텐서들을 쌓기 때문에, 결과 텐서의 차원 수가 1 증가합니다.
진행 단계
- 필요한 라이브러리를 임포트합니다. 아래 모든 예제에서 필요한 파이썬 라이브러리는 torch입니다. 미리 설치되어 있는지 확인하세요.
- 두 개 이상의 PyTorch 텐서를 생성하고 출력합니다.
- torch.cat() 또는 torch.stack()을 사용하여 생성한 텐서들을 결합합니다. 차원 인자(예: 0, -1)를 지정하면 특정 차원을 기준으로 텐서를 합칠 수 있습니다.
- 마지막으로 연결 또는 스택된 최종 텐서를 출력합니다.
예제 1: torch.cat()으로 1D 텐서 연결하기
# PyTorch에서 텐서를 연결하는 파이썬 프로그램
# 필요한 라이브러리 임포트
import torch
# 텐서 생성
T1 = torch.Tensor([1,2,3,4])
T2 = torch.Tensor([0,3,4,1])
T3 = torch.Tensor([4,3,2,5])
# 생성된 텐서 출력
print("T1:", T1)
print("T2:", T2)
print("T3:", T3)
# torch.cat()으로 위 텐서들을 연결
T = torch.cat((T1,T2,T3))
# 연결 후 최종 텐서 출력
print("T:",T)출력 결과
위의 Python 3 코드를 실행하면 다음과 같은 결과가 출력됩니다.
T1: tensor([1., 2., 3., 4.]) T2: tensor([0., 3., 4., 1.]) T3: tensor([4., 3., 2., 5.]) T: tensor([1., 2., 3., 4., 0., 3., 4., 1., 4., 3., 2., 5.])
세 개의 1D 텐서가 순서대로 이어져 하나의 긴 1D 텐서가 된 것을 확인할 수 있습니다.
예제 2: torch.cat()으로 2D 텐서 연결하기
# 필요한 라이브러리 임포트
import torch
# 텐서 생성
T1 = torch.Tensor([[1,2],[3,4]])
T2 = torch.Tensor([[0,3],[4,1]])
T3 = torch.Tensor([[4,3],[2,5]])
# 생성된 텐서 출력
print("T1:\n", T1)
print("T2:\n", T2)
print("T3:\n", T3)
print("0차원으로 텐서 연결(concatenate)")
T = torch.cat((T1,T2,T3), 0)
print("T:\n", T)
print("-1차원으로 텐서 연결(concatenate)")
T = torch.cat((T1,T2,T3), -1)
print("T:\n", T)출력 결과
위의 Python 3 코드를 실행하면 다음과 같은 결과가 출력됩니다.
T1:
tensor([[1., 2.],
[3., 4.]])
T2:
tensor([[0., 3.],
[4., 1.]])
T3:
tensor([[4., 3.],
[2., 5.]])
join(concatenate) tensors in the 0 dimension
T:
tensor([[1., 2.],
[3., 4.],
[0., 3.],
[4., 1.],
[4., 3.],
[2., 5.]])
join(concatenate) tensors in the -1 dimension
T:
tensor([[1., 2., 0., 3., 4., 3.],
[3., 4., 4., 1., 2., 5.]])위 예제에서 2D 텐서들이 0차원과 -1차원을 따라 각각 연결되었습니다. 0차원으로 연결하면 행의 개수가 늘어나고 열의 개수는 그대로 유지됩니다. 반대로 -1차원(마지막 차원)으로 연결하면 열의 개수가 늘어나고 행의 개수는 유지됩니다.
예제 3: torch.stack()으로 1D 텐서 쌓기
# PyTorch에서 텐서를 쌓는 파이썬 프로그램
# 필요한 라이브러리 임포트
import torch
# 텐서 생성
T1 = torch.Tensor([1,2,3,4])
T2 = torch.Tensor([0,3,4,1])
T3 = torch.Tensor([4,3,2,5])
# 생성된 텐서 출력
print("T1:", T1)
print("T2:", T2)
print("T3:", T3)
# "torch.stack()"으로 위 텐서들을 쌓기
print("텐서 쌓기(stack)")
T = torch.stack((T1,T2,T3))
# 결합 후 최종 텐서 출력
print("T:\n",T)
print("0차원으로 텐서 쌓기(stack)")
T = torch.stack((T1,T2,T3), 0)
print("T:\n", T)
print("-1차원으로 텐서 쌓기(stack)")
T = torch.stack((T1,T2,T3), -1)
print("T:\n", T)출력 결과
위의 Python 3 코드를 실행하면 다음과 같은 결과가 출력됩니다.
T1: tensor([1., 2., 3., 4.])
T2: tensor([0., 3., 4., 1.])
T3: tensor([4., 3., 2., 5.])
join(stack) tensors
T:
tensor([[1., 2., 3., 4.],
[0., 3., 4., 1.],
[4., 3., 2., 5.]])
join(stack) tensors in the 0 dimension
T:
tensor([[1., 2., 3., 4.],
[0., 3., 4., 1.],
[4., 3., 2., 5.]])
join(stack) tensors in the -1 dimension
T:
tensor([[1., 0., 4.],
[2., 3., 3.],
[3., 4., 2.],
[4., 1., 5.]])위 예제에서 주목할 점은 1D 텐서들이 stack되어 최종적으로 2D 텐서가 되었다는 것입니다. 이처럼 torch.stack()은 새로운 차원을 추가하기 때문에 입력 텐서보다 차원이 하나 더 높은 결과를 반환합니다.
예제 4: torch.stack()으로 2D 텐서 쌓기
# 필요한 라이브러리 임포트
import torch
# 텐서 생성
T1 = torch.Tensor([[1,2],[3,4]])
T2 = torch.Tensor([[0,3],[4,1]])
T3 = torch.Tensor([[4,3],[2,5]])
# 생성된 텐서 출력
print("T1:\n", T1)
print("T2:\n", T2)
print("T3:\n", T3)
print("0차원으로 텐서 쌓기(stack)")
T = torch.stack((T1,T2,T3), 0)
print("T:\n", T)
print("-1차원으로 텐서 쌓기(stack)")
T = torch.stack((T1,T2,T3), -1)
print("T:\n", T)출력 결과
위의 Python 3 코드를 실행하면 다음과 같은 결과가 출력됩니다.
T1:
tensor([[1., 2.],
[3., 4.]])
T2:
tensor([[0., 3.],
[4., 1.]])
T3:
tensor([[4., 3.],
[2., 5.]])
Join (stack)tensors in the 0 dimension
T:
tensor([[[1., 2.],
[3., 4.]],
[[0., 3.],
[4., 1.]],
[[4., 3.],
[2., 5.]]])
Join(stack) tensors in the -1 dimension
T:
tensor([[[1., 0., 4.],
[2., 3., 3.]],
[[3., 4., 2.],
[4., 1., 5.]]])위 예제에서는 2D 텐서들이 결합(스택)되어 3D 텐서가 생성되었습니다. 이처럼 torch.stack()을 사용하면 입력 텐서의 차원이 하나 증가한다는 점을 명확히 확인할 수 있습니다.
정리
torch.cat()은 기존 차원을 따라 텐서를 이어 붙일 때, torch.stack()은 새로운 차원을 만들어 텐서를 묶을 때 사용합니다. 배치(batch) 데이터를 구성하거나 여러 텐서를 하나의 데이터셋으로 합칠 때 어떤 함수가 적합한지 판단하는 데 이 차이를 반드시 기억해 두세요.