PyTorch Custom Datasets and DataLoaders
Build efficient data pipelines — custom Dataset classes, DataLoader configuration, augmentation, and optimized loading for large-scale training.
Real-World Scenario
An ML engineer trains an image classifier on 500,000 product images stored on disk. Loading all images into RAM crashes the machine. A custom Dataset loads images lazily from disk, applies augmentations on-the-fly, and a DataLoader with num_workers=8 prefetches batches in parallel — keeping the GPU busy 95% of the time instead of waiting for data.
Custom Dataset from Files
import torch
from torch.utils.data import Dataset, DataLoader
from pathlib import Path
from PIL import Image
import torchvision.transforms as T
import pandas as pd
import numpy as np
class ImageCSVDataset(Dataset):
"""
Dataset backed by a CSV with columns: image_path, label
Loads images lazily from disk — never loads the full dataset into RAM.
"""
def __init__(
self,
csv_path: str,
image_dir: str,
transform=None,
target_transform=None,
):
self.df = pd.read_csv(csv_path)
self.image_dir = Path(image_dir)
self.transform = transform
self.target_transform = target_transform
# Build label-to-index mapping from unique labels
unique_labels = sorted(self.df["label"].unique())
self.label2idx = {l: i for i, l in enumerate(unique_labels)}
self.classes = unique_labels
def __len__(self) -> int:
return len(self.df)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, int]:
row = self.df.iloc[idx]
image = Image.open(self.image_dir / row["image_path"]).convert("RGB")
label = self.label2idx[row["label"]]
if self.transform:
image = self.transform(image)
if self.target_transform:
label = self.target_transform(label)
return image, label
# Build transforms
train_transform = T.Compose([
T.Resize((224, 224)),
T.RandomHorizontalFlip(p=0.5),
T.RandomRotation(degrees=15),
T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], # ImageNet stats
std=[0.229, 0.224, 0.225]),
])
val_transform = T.Compose([
T.Resize((224, 224)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
# Instantiate (assumes data/train.csv and data/images/ exist)
# train_ds = ImageCSVDataset("data/train.csv", "data/images", transform=train_transform)
# val_ds = ImageCSVDataset("data/val.csv", "data/images", transform=val_transform)
In-Memory Dataset
import torch
from torch.utils.data import Dataset, DataLoader, random_split
import numpy as np
class TabularDataset(Dataset):
"""Wraps numpy arrays or pandas DataFrames for tabular data."""
def __init__(self, X: np.ndarray, y: np.ndarray):
self.X = torch.tensor(X, dtype=torch.float32)
self.y = torch.tensor(y, dtype=torch.long)
def __len__(self):
return len(self.y)
def __getitem__(self, idx):
return self.X[idx], self.y[idx]
# Generate dummy data
rng = np.random.default_rng(42)
X = rng.normal(0, 1, (10_000, 20)).astype(np.float32)
y = rng.integers(0, 5, 10_000)
dataset = TabularDataset(X, y)
# Split with random_split
train_size = int(0.8 * len(dataset))
val_size = len(dataset) - train_size
train_ds, val_ds = random_split(
dataset, [train_size, val_size],
generator=torch.Generator().manual_seed(42)
)
train_loader = DataLoader(
train_ds,
batch_size=256,
shuffle=True,
num_workers=0, # 0 = main process (use 2-8 for disk-bound datasets)
pin_memory=torch.cuda.is_available(),
drop_last=True, # drop the last incomplete batch for stable BatchNorm
)
val_loader = DataLoader(val_ds, batch_size=512, shuffle=False)
print(f"Train batches: {len(train_loader)}")
print(f"Val batches: {len(val_loader)}")
# Verify
X_batch, y_batch = next(iter(train_loader))
print(f"Batch shape: X={X_batch.shape}, y={y_batch.shape}")
Weighted Sampler for Class Imbalance
import torch
from torch.utils.data import DataLoader, WeightedRandomSampler
import numpy as np
# Class imbalance: 90% class 0, 10% class 1
rng = np.random.default_rng(42)
n = 10_000
y_imbal = rng.choice([0, 1], n, p=[0.9, 0.1])
# Compute per-sample weights inversely proportional to class frequency
class_counts = np.bincount(y_imbal)
class_weights = 1.0 / class_counts # weight per class
sample_weights = torch.tensor(class_weights[y_imbal], dtype=torch.float)
# WeightedRandomSampler: samples WITH replacement to balance classes
sampler = WeightedRandomSampler(
weights=sample_weights,
num_samples=len(sample_weights),
replacement=True,
)
X = torch.randn(n, 20)
y = torch.tensor(y_imbal, dtype=torch.long)
dataset = torch.utils.data.TensorDataset(X, y)
loader = DataLoader(dataset, batch_size=64, sampler=sampler)
# Verify balance: the sampled batches should be ~50/50
batch_labels = []
for _, y_batch in loader:
batch_labels.extend(y_batch.tolist())
if len(batch_labels) >= 1000:
break
pos_rate = sum(l == 1 for l in batch_labels[:1000]) / 1000
print(f"Original positive rate: {y_imbal.mean():.1%}")
print(f"Sampled positive rate: {pos_rate:.1%} (should be ~50%)")
Streaming Large Datasets with IterableDataset
import torch
from torch.utils.data import IterableDataset, DataLoader
import numpy as np
class StreamingCSVDataset(IterableDataset):
"""
Streams rows from a large CSV without loading it into memory.
Supports multi-worker loading with worker-level shard assignment.
"""
def __init__(self, file_path: str, chunk_size: int = 1000):
self.file_path = file_path
self.chunk_size = chunk_size
def __iter__(self):
import pandas as pd
# Worker info for multi-process loading — each worker reads a shard
worker_info = torch.utils.data.get_worker_info()
for chunk in pd.read_csv(self.file_path, chunksize=self.chunk_size):
for _, row in chunk.iterrows():
# Parse row into tensors
X = torch.tensor(row[:-1].values.astype(np.float32))
y = torch.tensor(int(row.iloc[-1]), dtype=torch.long)
yield X, y
# Simulated usage (requires a CSV file at data/large_dataset.csv)
# dataset = StreamingCSVDataset("data/large_dataset.csv")
# loader = DataLoader(dataset, batch_size=128, num_workers=4)
# Simpler demonstration with a generator
class SyntheticStreamDataset(IterableDataset):
"""Generate infinite synthetic data for demonstration."""
def __init__(self, n_samples: int = 10_000, n_features: int = 20):
self.n_samples = n_samples
self.n_features = n_features
def __iter__(self):
rng = np.random.default_rng(42)
for _ in range(self.n_samples):
X = torch.tensor(rng.normal(0, 1, self.n_features), dtype=torch.float32)
y = torch.tensor(int(X[0] > 0), dtype=torch.long)
yield X, y
stream_ds = SyntheticStreamDataset(10_000)
stream_loader = DataLoader(stream_ds, batch_size=256)
X_batch, y_batch = next(iter(stream_loader))
print(f"Streaming batch: X={X_batch.shape} y={y_batch.shape}")
DataLoader Performance Tuning
import torch
from torch.utils.data import TensorDataset, DataLoader
import time
X = torch.randn(50_000, 128)
y = torch.randint(0, 10, (50_000,))
dataset = TensorDataset(X, y)
def benchmark_loader(loader: DataLoader, n_batches: int = 50) -> float:
t0 = time.perf_counter()
for i, _ in enumerate(loader):
if i >= n_batches:
break
return (time.perf_counter() - t0) / n_batches * 1000
configs = [
{"batch_size": 256, "num_workers": 0, "pin_memory": False, "prefetch_factor": None},
{"batch_size": 256, "num_workers": 2, "pin_memory": False, "prefetch_factor": 2},
{"batch_size": 256, "num_workers": 4, "pin_memory": True, "prefetch_factor": 2},
{"batch_size": 512, "num_workers": 4, "pin_memory": True, "prefetch_factor": 4},
]
print(f"{'Config':<45} {'ms/batch':>10}")
print("-" * 57)
for cfg in configs:
kw = {k: v for k, v in cfg.items() if v is not None and k != "batch_size"}
loader = DataLoader(dataset, batch_size=cfg["batch_size"], shuffle=True, **kw)
ms = benchmark_loader(loader)
desc = f"bs={cfg['batch_size']} workers={cfg['num_workers']} pin={cfg['pin_memory']}"
print(f"{desc:<45} {ms:>10.2f}")
# Key takeaways:
# - num_workers > 0 is crucial for disk-bound datasets (images)
# - pin_memory=True helps when transferring to GPU
# - For in-memory datasets, num_workers=0 can be fastest (no IPC overhead)
# - prefetch_factor controls how many batches each worker pre-loads Frequently Asked Questions
When should I build a custom Dataset vs use a built-in one?
Build a custom Dataset when: your data is on disk in a custom format, you need on-the-fly augmentation, your dataset doesn't fit in memory, or you need complex sampling logic. PyTorch's built-in datasets (ImageFolder, MNIST, CIFAR) are fine for learning but rarely sufficient for production tasks.
What is pin_memory and when should I enable it?
pin_memory=True allocates data in page-locked (pinned) host memory, which allows faster H2D (host-to-device) transfers. Enable it when training on GPU — it can speed up data loading by 20-30% on typical hardware. Don't use it with CPU training or when memory is very constrained.