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
PascalCase.
Discussion
Loading…