PyTorch로 FashionMNIST 데이터셋 로드 및 레이어 시각화하기

데이터셋 로딩 및 탐색

이 노트북에서는 Fashion-MNIST 데이터셋을 로드하고, 이미지와 레이블의 구조를 분석합니다. 모델 개발 전에 데이터를 직접 확인하는 것은 매우 중요하며, 이는 이미지의 크기, 색상 구성, 레이블 분포 등을 이해하는 데 도움이 됩니다.

PyTorch는 내장된 데이터셋 클래스를 제공하며, 그 중 하나인 FashionMNIST는 이미 data/ 폴더에 다운로드되어 있습니다. 이를 사용해 데이터를 불러오고, DataLoader를 통해 배치 단위로 처리할 수 있습니다.

데이터셋 준비

torch.utils.data.Dataset은 모든 데이터셋의 기본 인터페이스이며, FashionMNIST는 이를 상속하여 이미지와 레이블을 쉽게 로드할 수 있도록 합니다. 또한 데이터 변환(예: 이미지를 텐서로 변환)도 간편하게 적용할 수 있습니다.

다음 코드는 데이터를 텐서 형식으로 변환하고, 학습용 데이터셋을 생성하는 과정입니다.

import torch
import torchvision
from torchvision.datasets import FashionMNIST
from torch.utils.data import DataLoader
from torchvision import transforms

# 이미지를 텐서로 변환 (0~1 범위의 픽셀값 → 텐서)
transform = transforms.ToTensor()

# 학습 데이터셋 로드
train_dataset = FashionMNIST(root='./data', train=True, download=False, transform=transform)

print(f"학습 데이터 수: {len(train_dataset)}")

출력:

학습 데이터 수: 60000

배치 처리 및 데이터 순서 섞기

DataLoader는 데이터를 지정된 배치 크기로 묶고, 각 에포크마다 순서를 무작위로 섞어 줍니다. 이는 모델이 특정 순서에 편향되지 않도록 하는 데 중요합니다.

batch_size = 20
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)

# 클래스 이름 정의
class_names = [
    '티셔츠', '바지', '풀오버', '드레스', '코트',
    '샌달', '셔츠', '스니커즈', '가방', '앵클 부츠'
]

학습 데이터 시각화

이제 배치 단위로 데이터를 가져와, 2행 × 배치 크기/2열의 그리드에 이미지를 출력합니다.

import numpy as np
import matplotlib.pyplot as plt

%matplotlib inline

# 데이터 로더에서 한 배치 가져오기
data_iter = iter(train_loader)
batch_images, batch_labels = next(data_iter)

# 텐서를 넘파이 배열로 변환
batch_images = batch_images.numpy()

# 이미지 시각화
fig = plt.figure(figsize=(25, 4))
for i in range(batch_size):
    ax = fig.add_subplot(2, batch_size//2, i+1, xticks=[], yticks=[])
    ax.imshow(batch_images[i].squeeze(), cmap='gray')
    ax.set_title(class_names[batch_labels[i]])

이미지의 세부 정보 분석

Fashion-MNIST의 각 이미지는 28×28 픽셀의 회색조 이미지이며, 값은 [0, 1] 범위로 정규화되어 있습니다. 이는 신경망이 안정적으로 학습되도록 하기 위한 중요한 조건입니다.

정규화 없이 입력 데이터가 큰 값 범위를 가질 경우, 초기 레이어의 활성화가 과도하게 커져 학습이 불안정해질 수 있습니다. 특히 역전파 과정에서 기울기가 급격히 증가하면서 손실이 발산할 위험이 있습니다.

아래 예시는 특정 이미지의 픽셀 값을 직접 확인하는 방법입니다.

# 특정 이미지 선택 (인덱스 2)
img_idx = 2
image_data = batch_images[img_idx]
image_2d = image_data.squeeze()

# 픽셀 값과 함께 시각화
fig = plt.figure(figsize=(12, 12))
ax = fig.add_subplot(111)
ax.imshow(image_2d, cmap='gray')

width, height = image_2d.shape
threshold = image_2d.max() / 2.5

for x in range(width):
    for y in range(height):
        value = round(image_2d[x][y], 2) if image_2d[x][y] != 0 else 0
        color = 'white' if image_2d[x][y] < threshold else 'black'
        ax.annotate(str(value), xy=(y, x), ha='center', va='center', fontsize=8, color=color)

이러한 시각화는 이미지의 밝기 분포와 특징을 직관적으로 이해하는 데 유용합니다.

태그: PyTorch FashionMNIST DataLoader Data Visualization tensor transformation

7월 31일 22:09에 게시됨