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
- Map-style
Dataset: implement__len__and__getitem__(i).IterableDataset: implement__iter__for streams. DataLoaderbatches, shuffles (training), and loads in parallel withnum_workers.- Use
pin_memory=True,persistent_workers=True, andprefetch_factorto overlap loading with GPU compute. - A custom
collate_fnhandles variable-length inputs (padding sequences, building masks). - Apply augmentation only to training data; keep validation transforms deterministic.
- Use samplers (
WeightedRandomSampler,DistributedSampler) for imbalance and multi-GPU training.
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
- torchvision.transforms.v2 handles images, bounding boxes, masks, and videos consistently, and runs on tensors (and on GPU for batched transforms).
- Training gets random augmentations (crops, flips, color jitter, mixup and cutmix); validation and test get deterministic preprocessing (resize, center crop, normalize).
- Heavy augmentation can move to the GPU (Kornia, DALI) when CPU workers bottleneck.
Samplers
WeightedRandomSampleroversamples rare classes for imbalanced datasets (alternatively, weight the loss).DistributedSamplershards data across processes in distributed training; callsampler.set_epoch(epoch)so shuffling differs each epoch.- Custom batch samplers implement length bucketing or grouping.
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:
- Time the loader alone: iterate over it without the model, and compare batches per second with training throughput.
- Use the PyTorch Profiler or
nvidia-smito see GPU idle gaps. - Increase
num_workers, enablepin_memoryandpersistent_workers, and move decoding or augmentation to the GPU. - 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
- PyTorch — The framework overview
- PyTorch Training Loop — Where batches are consumed
- PyTorch Distributed Training — DistributedSampler and sharding
- Feature Engineering — Preparing inputs
- Data Engineering — Offline preprocessing pipelines
- Pandas — Tabular data preparation