본문 바로가기
C.W.K.
Stream
Lesson 05 of 08 · published

Dataset, DataLoader, 작업 프로세스 연결

~14 min · dataset, dataloader, num_workers, collate

Level 0텐서 탐구자
0 XP0/62 lessons0/13 achievements
0/120 XP to next level120 XP to go0% complete

GPU가 데이터를 기다리지 않게 하는 두 추상화

PyTorch는 '샘플이 무엇인가'와 '샘플을 어떻게 효율적으로 배치할 것인가'를 분리해:

  • Dataset: __len____getitem__(idx)를 정의해. 샘플이 몇 개이고 i번째 샘플을 어떻게 가져오는지가 최소 계약이야.
  • DataLoader: 데이터셋을 감싸서 배치 구성, 순서 섞기, 작업 프로세스를 이용한 병렬 불러오기, 빠른 GPU 전송을 위한 페이지 고정 메모리를 처리해.

중요한 DataLoader 설정

  • batch_size: 한 번에 묶을 샘플 수야.
  • shuffle=True: 에포크마다 샘플 순서를 무작위로 섞어. 학습에만 쓰고 검증이나 테스트에서는 사용하지 마.
  • num_workers: 데이터를 병렬로 불러올 작업 프로세스 수야. 입출력이 병목인 데이터셋에서는 CPU 코어 수를 기준으로 조정해. 0이면 주 프로세스에서 불러오므로 가장 느리지만 디버깅은 쉬워.
  • pin_memory=True: 배치를 페이지 고정 메모리에 할당해 CPU에서 GPU로 더 빨리 전송할 수 있어. x.to(device, non_blocking=True)와 함께 쓰면 실제 처리량이 좋아질 수 있어.
  • prefetch_factor: 각 작업 프로세스가 미리 불러올 배치 수야. 기본값은 2야. GPU가 다음 배치가 오기 전에 계산을 끝낸다면 값을 올려 봐.
  • persistent_workers=True: 에포크 사이에도 작업 프로세스를 유지해서 매번 다시 시작하는 비용을 줄여.
  • drop_last=True: 마지막의 불완전한 배치를 버려. BatchNorm 통계나 분산 학습처럼 배치 크기를 일정하게 유지해야 할 때 유용해.

macOS의 num_workers 함정

macOS에서 multiprocessing의 기본 시작 방식은 Linux의 'fork'가 아니라 'spawn'이야. 각 작업 프로세스가 모듈을 다시 불러오므로 무거운 데이터셋은 시작이 느릴 수 있고, 피클로 직렬화할 수 없는 객체는 오류를 내. 개발 중에는 num_workers=0으로 디버깅하고, 실제 실행에서 값을 올려 처리량을 맞춰.

Code

사용자 정의 데이터셋: 최소 계약·python
import torch
from torch.utils.data import Dataset, DataLoader

class TensorDataset(Dataset):
    def __init__(self, X, y, transform=None):
        self.X = X
        self.y = y
        self.transform = transform

    def __len__(self):
        return len(self.X)

    def __getitem__(self, idx):
        x = self.X[idx]
        if self.transform is not None:
            x = self.transform(x)
        return x, self.y[idx]

X = torch.randn(1000, 10)
y = torch.randint(0, 5, (1000,))
ds = TensorDataset(X, y)
print(len(ds), ds[0][0].shape, ds[0][1])
운영 환경을 고려한 DataLoader·python
import torch
from torch.utils.data import DataLoader

loader = DataLoader(
    dataset,
    batch_size=64,
    shuffle=True,
    num_workers=8,                # match CPU cores
    pin_memory=True,              # fast CPU→GPU
    prefetch_factor=4,            # batches each worker prefetches ahead
    persistent_workers=True,      # keep workers alive across epochs
    drop_last=True,               # consistent batch shape
)

device = "cuda"
for x, y in loader:
    x = x.to(device, non_blocking=True)
    y = y.to(device, non_blocking=True)
    # ...training step...
random_split: 한 데이터셋에서 학습 / 검증·python
import torch
from torch.utils.data import random_split, DataLoader

dataset = TensorDataset(torch.randn(1000, 10), torch.randint(0, 5, (1000,)))
n_train = int(0.8 * len(dataset))
n_val = len(dataset) - n_train

train_ds, val_ds = random_split(
    dataset, [n_train, n_val],
    generator=torch.Generator().manual_seed(42),   # reproducible split
)

train_loader = DataLoader(train_ds, batch_size=32, shuffle=True)
val_loader = DataLoader(val_ds, batch_size=64, shuffle=False)
사용자 정의 묶음 구성: 가변 길이 시퀀스·python
import torch
from torch.utils.data import DataLoader, Dataset
from torch.nn.utils.rnn import pad_sequence

class VariableSeqDataset(Dataset):
    def __init__(self):
        self.seqs = [torch.randint(0, 100, (torch.randint(5, 20, (1,)).item(),)) for _ in range(64)]
        self.labels = torch.randint(0, 2, (64,))

    def __len__(self): return len(self.seqs)
    def __getitem__(self, i): return self.seqs[i], self.labels[i]

def collate(batch):
    seqs, labels = zip(*batch)
    padded = pad_sequence(seqs, batch_first=True, padding_value=0)
    lengths = torch.tensor([len(s) for s in seqs])
    return padded, lengths, torch.tensor(labels)

loader = DataLoader(VariableSeqDataset(), batch_size=8, collate_fn=collate)
batch = next(iter(loader))
print(batch[0].shape, batch[1], batch[2].shape)
# torch.Size([8, max_len])  tensor([12, 18, ...])  torch.Size([8])

External links

Exercise

루트/class_name/이미지.jpg 형태의 디렉터리 트리에서 이미지를 읽는 데이터셋을 만들어 봐. __len__과 __getitem__을 구현하고, num_workers=4와 pin_memory=True를 설정한 DataLoader로 감싸. 현재 컴퓨터에서 초당 몇 개의 샘플을 처리하는지 재서 기록해 두면 나중에 성능 저하를 발견할 기준선이 돼.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.