iwantcoding.com
🔥 Daily 👥 Rooms 🏆 Top Log in Sign up

Datasets & DataLoader

PyTorch’s Dataset + DataLoader handle batching, shuffling, multi-process loading. Subclass Dataset for custom data; DataLoader wires up parallelism + pin_memory.

Dataset, DataLoader, transforms, samplers

EXAMPLE
import torch
from torch.utils.data import Dataset, DataLoader, random_split, Subset, WeightedRandomSampler
from torchvision import transforms
import pandas as pd
from PIL import Image

# 1) Custom Dataset
class ImageCSVDataset(Dataset):
    def __init__(self, csv_path, image_dir, transform=None):
        self.df = pd.read_csv(csv_path)
        self.image_dir = image_dir
        self.transform = transform

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

    def __getitem__(self, idx):
        row = self.df.iloc[idx]
        img = Image.open(f'{self.image_dir}/{row["file"]}').convert('RGB')
        if self.transform:
            img = self.transform(img)
        return img, row['label']

# 2) Compose transforms (vision)
train_tf = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10),
    transforms.ColorJitter(brightness=0.2, contrast=0.2),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])

val_tf = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])

train_ds = ImageCSVDataset('train.csv', 'data/images', transform=train_tf)
val_ds   = ImageCSVDataset('val.csv',   'data/images', transform=val_tf)

# 3) DataLoader
train_dl = DataLoader(
    train_ds,
    batch_size  = 64,
    shuffle     = True,
    num_workers = 4,                # multi-process loading
    pin_memory  = True,             # faster GPU transfer
    persistent_workers = True,      # don't tear down each epoch
    drop_last   = True,             # drop last partial batch (cleaner training)
)
val_dl = DataLoader(val_ds, batch_size=128, shuffle=False, num_workers=4, pin_memory=True)

# 4) Use in a training loop
for epoch in range(epochs):
    for x, y in train_dl:
        x = x.to(device, non_blocking=True)
        y = y.to(device, non_blocking=True)
        loss = model(x).cross_entropy(y).backward()
        optimizer.step()
        optimizer.zero_grad(set_to_none=True)

# 5) Built-in datasets (torchvision)
from torchvision import datasets
train = datasets.CIFAR10('./data', train=True,  download=True, transform=train_tf)
val   = datasets.CIFAR10('./data', train=False, download=True, transform=val_tf)

# Same for torchaudio, torchtext (deprecated; use HF Datasets for NLP now)

# 6) Splitting a dataset
full = ImageCSVDataset('all.csv', 'data/images', transform=val_tf)
train, val = random_split(full, [0.8, 0.2], generator=torch.Generator().manual_seed(42))

# Stratified — use sklearn for the split then Subset
from sklearn.model_selection import StratifiedShuffleSplit
sss = StratifiedShuffleSplit(n_splits=1, test_size=0.2, random_state=42)
idx, _ = next(sss.split(range(len(full)), labels))
train = Subset(full, idx)

# 7) Imbalanced classes — WeightedRandomSampler
class_count = pd.Series(labels).value_counts()
weights = 1.0 / class_count[labels].values
sampler = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)

train_dl = DataLoader(train_ds, batch_size=64, sampler=sampler, num_workers=4, pin_memory=True)

# 8) DistributedSampler — multi-GPU / multi-node
from torch.utils.data.distributed import DistributedSampler
sampler = DistributedSampler(train_ds, shuffle=True)
train_dl = DataLoader(train_ds, batch_size=64, sampler=sampler, num_workers=4, pin_memory=True)
# In your loop: sampler.set_epoch(epoch) before each epoch.

# 9) collate_fn — custom batching (variable-length, mixed dtypes)
def collate_pad(batch):
    sequences, labels = zip(*batch)
    lengths = torch.tensor([len(s) for s in sequences])
    padded  = torch.nn.utils.rnn.pad_sequence(sequences, batch_first=True)
    return padded, lengths, torch.tensor(labels)

dl = DataLoader(text_ds, batch_size=32, collate_fn=collate_pad)

# 10) Streaming / IterableDataset — for huge / generated data
from torch.utils.data import IterableDataset

class S3Stream(IterableDataset):
    def __init__(self, bucket): self.bucket = bucket
    def __iter__(self):
        for obj in s3_client.list_objects(self.bucket):
            yield decode(s3_client.get_object(self.bucket, obj.key))

dl = DataLoader(S3Stream('my-bucket'), batch_size=64, num_workers=8)

# 11) Hugging Face Datasets — modern NLP / multimodal data
from datasets import load_dataset
ds = load_dataset('imdb', split='train')
ds = ds.map(lambda x: tokenizer(x['text'], padding='max_length', truncation=True), batched=True)
ds.set_format(type='torch', columns=['input_ids', 'attention_mask', 'label'])
dl = DataLoader(ds, batch_size=32, shuffle=True)

# 12) Performance tips
# - num_workers: start with 2-4; benchmark; not always more is better
# - pin_memory + non_blocking on .to(device) → overlaps CPU→GPU transfer
# - persistent_workers=True for short epochs
# - Avoid heavy work in __getitem__ (precompute if possible)
# - prefetch_factor=2 (default) — increase if CPU loading is fast and GPU starves

# 13) Profiling
# Use torch.utils.benchmark or PyTorch Profiler to identify if data loading is a bottleneck:
from torch.profiler import profile, ProfilerActivity
with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof:
    for x, y in train_dl:
        x, y = x.cuda(non_blocking=True), y.cuda(non_blocking=True)
        # train step
print(prof.key_averages().table())

# 14) Caching expensive transforms
# Pre-compute and serialise to disk (Parquet, LMDB, WebDataset):
# torch.save(processed_features, 'cache.pt')
# Or use webdataset for tar-based streaming.

# 15) Common pitfalls
#   • num_workers > 0 on Windows without if __name__ == '__main__' → spawn issues
#   • Heavy lambdas in __getitem__ that can't be pickled → workers fail
#   • Reading the same file from many workers → IO bottleneck (use shards)
#   • Forgetting drop_last on small batches → BatchNorm error
#   • set_format('torch') missing → HF datasets return Python lists

# 16) Best practices
#   • Always define a transform pipeline for train + eval (different augmentation)
#   • Use DataLoader's pin_memory + non_blocking transfers
#   • Profile before optimising
#   • For huge data: WebDataset / LMDB / NVIDIA DALI
#   • For HF + transformers: Datasets library is the path of least resistance

Why it matters

A great PyTorch DataLoader keeps every GPU busy: num_workers > 0, pin_memory=True, non_blocking=True on transfers, persistent_workers=True for short epochs. Profile before guessing.

Tip: Tweak the snippet with Try it Yourself », then sit the quiz at the bottom of the page.

Example

Example
from torch.utils.data import Dataset, DataLoader
class Squares(Dataset):
    def __init__(self, n): self.n = n
    def __len__(self): return self.n
    def __getitem__(self, i): return torch.tensor([i]), torch.tensor([i*i])
loader = DataLoader(Squares(1000), batch_size=32, shuffle=True, num_workers=2)
Try it Yourself »

Exercise

Iterate a Dataset in batches.

from torch.utils.data import

Discussion

Loading…