Skip to main content
PyTorch intermediate Lesson 10 of 11

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.