본문 바로가기
C.W.K.
Stream
Lesson 04 of 07 · published

사용자 정의 묶음 구성 함수와 IterableDataset

~12 min · collate, iterable, stream, padding

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

샘플을 그대로 쌓을 수 없을 때

기본 묶음 구성 함수는 각 샘플의 모든 필드를 텐서로 쌓아. 이미지 배치나 고정된 수의 특징을 가진 회귀처럼 모든 샘플의 모양이 같을 때 잘 작동해. 시퀀스 길이가 다르거나, 항목별 메타데이터를 함께 옮기거나, 배치 구조를 직접 정해야 할 때는 별도의 함수가 필요해.

사용자 정의 묶음 구성

묶음 구성 함수는 List[Sample]을 받아 하나의 배치를 반환해. 가장 흔한 예는 가변 길이 시퀀스를 현재 배치에서 가장 긴 시퀀스에 맞춰 패딩하는 거야.

IterableDataset: 스트리밍 데이터

거대한 텍스트 말뭉치, 네트워크로 들어오는 데이터, 온라인 센서 스트림처럼 인덱스로 다루기 어렵거나 디스크에 한 번에 담기 힘든 데이터에는 Dataset 대신 IterableDataset을 구현해. __iter__(self)만 정의하면 PyTorch가 그 반복자에서 샘플을 받아 배치를 구성해.

여러 작업 프로세스로 IterableDataset을 사용할 때는 주의해야 해. 스트림을 직접 나누지 않으면 모든 작업 프로세스가 전체 반복자를 받아 같은 샘플을 중복해서 읽어. __iter__ 안에서 torch.utils.data.get_worker_info()를 사용해 작업 프로세스 ID별로 스트림을 분할해.

Code

가변 길이 시퀀스를 최장 길이에 맞춰 묶기·python
import torch
from torch.utils.data import DataLoader, Dataset
from torch.nn.utils.rnn import pad_sequence

class VarSeqDataset(Dataset):
    def __init__(self):
        self.seqs = [torch.randint(1, 100, (torch.randint(5, 25, (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_pad(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(VarSeqDataset(), batch_size=8, collate_fn=collate_pad)
padded, lengths, labels = next(iter(loader))
print(padded.shape, lengths)   # torch.Size([8, max_len_in_batch]) tensor([...])
샘플별 메타데이터 반환: 사전 형태의 배치·python
import torch
from torch.utils.data import DataLoader, Dataset

class TaggedDataset(Dataset):
    def __init__(self):
        self.X = torch.randn(32, 10)
        self.y = torch.randint(0, 2, (32,))
        self.tags = [f"sample_{i}" for i in range(32)]
    def __len__(self): return 32
    def __getitem__(self, i):
        return {'x': self.X[i], 'y': self.y[i], 'tag': self.tags[i]}

def collate_dict(batch):
    return {
        'x': torch.stack([b['x'] for b in batch]),
        'y': torch.stack([b['y'] for b in batch]),
        'tags': [b['tag'] for b in batch],   # keep as list
    }

loader = DataLoader(TaggedDataset(), batch_size=8, collate_fn=collate_dict)
batch = next(iter(loader))
print(batch['x'].shape, batch['y'].shape, batch['tags'])
작업 프로세스를 올바르게 나누는 IterableDataset·python
import torch
from torch.utils.data import IterableDataset, DataLoader

class StreamingDataset(IterableDataset):
    """Stream samples from a generator. Splits across workers correctly."""
    def __init__(self, n_total=10_000):
        self.n_total = n_total

    def __iter__(self):
        info = torch.utils.data.get_worker_info()
        if info is None:                    # single-process
            start, end = 0, self.n_total
        else:
            per = self.n_total // info.num_workers
            start = info.id * per
            end = self.n_total if info.id == info.num_workers - 1 else start + per

        for i in range(start, end):
            yield torch.randn(10), i % 2

loader = DataLoader(StreamingDataset(), batch_size=32, num_workers=4)
print(sum(b[0].size(0) for b in loader))   # ≈ 10000

External links

Exercise

(image_tensor, list_of_bbox_tensors_per_image, image_id_string)을 반환하는 데이터셋을 구현해 봐. 이미지마다 경계 상자 수는 다르게 해. 이미지는 하나의 배치 텐서로 쌓고, 경계 상자 텐서와 이미지 ID는 각각 Python 목록으로 반환하는 묶음 구성 함수를 작성해. 샘플 몇 개로 모양과 값이 맞는지 검증해.

Progress

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

댓글 0

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

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