PyTorch Datasets & DataLoaders

A model trains only as fast as data reaches it. PyTorch separates data handling into two pieces: a Dataset, which knows how to load one example, and a DataLoader, which turns a dataset into shuffled batches, using parallel worker processes, pinned memory, and prefetching to keep the GPU fed. Getting this right is often the difference between a GPU at 95% utilization and one idling at 30% while Python decodes JPEGs.

This page covers writing datasets, configuring loaders, batching variable-length data, augmentations, handling class imbalance, and scaling to datasets that don't fit on local disk.

TL;DR

Quick Example

An image dataset with augmentation, and a text dataset with a padding collate function:

Core Concepts

Map-Style vs Iterable Datasets

Keep __getitem__ lightweight and self-contained: open files lazily inside workers (not in __init__ for handles that can't be pickled), and return tensors or NumPy arrays.

DataLoader Parameters

A starting point for num_workers is the number of CPU cores available per GPU, and then measure.

Collate Functions and Padding

The default collate stacks equal-shaped tensors. For variable-length sequences (text, audio, graphs), write a collate_fn that pads to the batch's max length and returns masks. Bucketing (grouping similar lengths into batches) reduces wasted padding and speeds up training for sequence models.

Transforms and Augmentation

Samplers

Large-Scale Data

When data lives in object storage or exceeds local disk: use sharded formats (WebDataset tar shards, Parquet, TFRecord-like formats, or MosaicML Streaming), stream with IterableDataset, cache hot data locally, and preprocess expensive steps (tokenization, decoding) offline in data pipelines. Hugging Face datasets provides memory-mapped Arrow datasets with streaming support.

Diagnosing Input Bottlenecks

Symptoms: low GPU utilization, and step time dominated by data wait. Checks:

  1. Time the loader alone: iterate over it without the model, and compare batches per second with training throughput.
  2. Use the PyTorch Profiler or nvidia-smi to see GPU idle gaps.
  3. Increase num_workers, enable pin_memory and persistent_workers, and move decoding or augmentation to the GPU.
  4. Store data in formats that decode fast (preprocessed tensors, smaller images) on fast storage.

Best Practices

Split Data Before Anything Else

Create train, validation, and test splits early, by the right unit (user, patient, or time period rather than random rows when leakage is possible), and never let augmentation or normalization statistics come from validation or test data. See model evaluation.

Make Datasets Deterministic Given an Index

Randomness belongs in transforms seeded per worker. Deterministic datasets make debugging and reproducibility much easier.

Keep Heavy Work Out of __getitem__ When Possible

Tokenize, resize, or extract features once in preprocessing jobs, and store the results. Recomputing expensive steps every epoch wastes CPU and slows training.

Watch Memory With Many Workers

Each worker is a process. Large Python objects in the dataset (for example lists of millions of strings) can be duplicated through copy-on-access. Store metadata in NumPy arrays, Arrow tables, or memory-mapped files.

Common Mistakes

Shuffling Validation Data or Augmenting It

Random crops or flips on validation data make metrics noisy and optimistic or pessimistic at random. Use deterministic eval transforms, with shuffle=False.

Identical Random Augmentations Across Workers

With older setups, NumPy random states could be duplicated across workers, producing identical "random" augmentations. Use torch's RNG in transforms, or seed per worker via worker_init_fn.

Opening Files or Connections in __init__ With Workers

Database connections, open file handles, and some library objects can't be pickled into worker processes, or become shared incorrectly. Create them lazily inside the worker on first __getitem__ call.

FAQ

How many DataLoader workers should I use?

Start with the number of CPU cores available per GPU (often 4–16) and measure throughput. Too few leaves the GPU waiting, and too many wastes memory and CPU through contention. The right number depends on how expensive each sample is to load and transform.

What does pin_memory do?

It allocates batch tensors in page-locked (pinned) host memory, which allows faster and asynchronous copies to the GPU. Combine it with .to(device, non_blocking=True) to overlap data transfer with computation.

How do I handle variable-length sequences?

Write a collate_fn that pads sequences to the longest in the batch (pad_sequence) and returns an attention mask, or packs them for RNNs. Group similar lengths into batches (bucketing) to minimize padding.

Should I use IterableDataset for large datasets?

When data can't be randomly indexed efficiently (streams, remote shards, generated data), yes. Handle sharding across workers and processes, and approximate shuffling with shuffle buffers or shard-level shuffling. If data fits on local disk with random access, a map-style dataset is simpler.

Related Topics

References