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

PyTorch에서 텐서의 데이터 타입(dtype)을 확인하는 방법

PyTorch 텐서는 동질적(homogeneous)입니다. 즉, 하나의 텐서를 구성하는 모든 요소는 반드시 동일한 데이터 타입을 가집니다. 텐서의 데이터 타입을 확인하려면 '.dtype' 속성에 접근하면 되며, 이 속성은 해당 텐서의 데이터 타입을 반환합니다.

확인 절차

  • 필요한 라이브러리를 임포트합니다. 아래의 모든 Python 예제에서는 torch 라이브러리가 필요하므로, 미리 설치되어 있는지 확인하세요.

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

  • T.dtype을 계산합니다. 여기서 T는 데이터 타입을 확인하려는 텐서입니다.

  • 텐서의 데이터 타입을 출력합니다.

예제 1

다음 Python 프로그램은 텐서의 데이터 타입을 가져오는 기본적인 방법을 보여줍니다.

# 라이브러리 임포트
import torch

# 3x4 크기의 난수 텐서 생성
T = torch.randn(3,4)
print("Original Tensor T:\n", T)

# 위 텐서의 데이터 타입 가져오기
data_type = T.dtype

# 텐서의 데이터 타입 출력
print("Data type of tensor T:\n", data_type)

출력 결과

Original Tensor T:
tensor([[ 2.1768, -0.1328, 0.8155, -0.7967],
        [ 0.1194, 1.0465, 0.0779, 0.9103],
        [-0.1809, 1.8085, 0.8393, -0.2463]])
Data type of tensor T:
torch.float32

예제 2

이번에는 리스트로부터 텐서를 생성한 뒤, 같은 방식으로 데이터 타입을 확인해 보겠습니다.

# 텐서의 데이터 타입을 가져오는 Python 프로그램
# 라이브러리 임포트
import torch

# 리스트로부터 텐서 생성
T = torch.Tensor([1,2,3,4])
print("Original Tensor T:\n", T)

# 위 텐서의 데이터 타입 가져오기
data_type = T.dtype

# 텐서의 데이터 타입 출력
print("Data type of tensor T:\n", data_type)

출력 결과

Original Tensor T:
    tensor([1., 2., 3., 4.])
Data type of tensor T:
    torch.float32

정리

PyTorch에서 텐서의 데이터 타입을 확인하는 것은 매우 간단합니다. .dtype 속성만 호출하면 되며, 예제에서 볼 수 있듯이 torch.randn()으로 생성한 난수 텐서와 torch.Tensor()로 생성한 텐서 모두 기본적으로 torch.float32 타입을 가집니다. 만약 다른 정밀도나 타입(예: int64, float64 등)이 필요하다면, 텐서 생성 시 dtype 인자를 지정하거나 .to(), .type() 메서드를 사용하여 변환할 수 있습니다.