Here's a scene that plays out in every ML team. Someone profiles a training run that feels slow, expecting to find a fat matrix multiply. Instead they find the GPU idle 60% of the time, waiting. The bottleneck isn't the model — it's everything that has to happen before a batch of data reaches the model: reading files off disk, decoding JPEGs, resizing, augmenting, stacking into a tensor, copying to the GPU.
PyTorch splits that work across two abstractions. Dataset knows how to produce one example. DataLoader turns a Dataset into a stream of batches, and — crucially — does the slow parts in parallel worker processes while the GPU is busy with the previous batch.
Dataset: one example at a time
The common kind is a map-style Dataset: implement __len__ and __getitem__(i), and you have something that behaves like a list of (input, target) pairs. Where the data lives — memory, a folder of images, a database, a set of shards — is entirely your business; the DataLoader only ever asks for item i.
from torch.utils.data import Datasetfrom PIL import Imageclass ImageFolderDataset(Dataset):def __init__(self, paths, labels, transform):self.paths = pathsself.labels = labelsself.transform = transform # applied here, so it runs in the workerdef __len__(self):return len(self.paths)def __getitem__(self, i):img = Image.open(self.paths[i]).convert("RGB")return self.transform(img), self.labels[i]
DataLoader: batching, shuffling, parallelism
from torch.utils.data import DataLoaderloader = DataLoader(dataset,batch_size=64,shuffle=True, # reshuffle every epoch (train only)num_workers=8, # worker processes decoding in parallelpin_memory=True, # page-locked memory -> faster host->GPU copypersistent_workers=True,# don't tear workers down between epochsprefetch_factor=2, # batches each worker prepares aheaddrop_last=True, # drop the ragged final batch (stable shapes))for images, labels in loader:images = images.to("cuda", non_blocking=True)labels = labels.to("cuda", non_blocking=True)...
num_workers is the dial that matters most, and more is not always better. Each worker is a process with its own Python interpreter and memory; past the point where workers can keep the GPU fed, you're just paying RAM and startup cost. The curve almost always looks like this:
Illustrative — the plateau point depends on your CPU, storage, and transform cost
Throughput climbs steeply, then flattens once the workers can supply batches as fast as the GPU consumes them. Adding more after the knee costs memory and buys nothing. Measure it for your own pipeline.
| num_workers | images / sec |
|---|---|
| 0 | 850 |
| 1 | 1600 |
| 2 | 2900 |
| 4 | 5200 |
| 6 | 6100 |
| 8 | 6300 |
| 12 | 6250 |
collate_fn: assembling the batch
The default collate_fn takes a list of samples and torch.stacks the tensors — which only works if every sample has the same shape. Images resized to a fixed size are fine. Variable-length sequences are not, and you supply your own:
import torchfrom torch.nn.utils.rnn import pad_sequencedef pad_collate(batch):seqs, labels = zip(*batch)lengths = torch.tensor([len(s) for s in seqs])padded = pad_sequence(seqs, batch_first=True, padding_value=0)return padded, lengths, torch.tensor(labels)loader = DataLoader(dataset, batch_size=32, collate_fn=pad_collate)
Transforms and augmentation
torchvision.transforms.v2 is the current API — it transforms images, bounding boxes, and masks together, and runs on tensors (so it can go on the GPU) as well as PIL images. Compose the pipeline once, hand it to the Dataset, and let it execute in the workers.
from torchvision.transforms import v2train_tf = v2.Compose([v2.RandomResizedCrop(224, antialias=True),v2.RandomHorizontalFlip(),v2.ToDtype(torch.float32, scale=True),v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),])# Validation: deterministic — resize + center crop, no randomness.val_tf = v2.Compose([v2.Resize(256, antialias=True),v2.CenterCrop(224),v2.ToDtype(torch.float32, scale=True),v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),])
Samplers: controlling the order
shuffle=True is shorthand for a RandomSampler. When you need something other than uniform random order — an imbalanced dataset, or one shard per GPU in distributed training — you pass a sampler explicitly. For class imbalance, WeightedRandomSampler oversamples the rare classes:
import torchfrom torch.utils.data import WeightedRandomSampler, DataLoaderclass_count = torch.bincount(torch.tensor(labels))weight_per_class = 1.0 / class_count.float()sample_weights = weight_per_class[torch.tensor(labels)]sampler = WeightedRandomSampler(sample_weights, num_samples=len(labels),replacement=True)# Note: pass a sampler OR shuffle=True, never both.loader = DataLoader(dataset, batch_size=64, sampler=sampler, num_workers=8)
| Knob | Set it to | Why |
|---|---|---|
num_workers | the knee of your throughput curve | parallel decode; measure, don't guess |
pin_memory | True when training on GPU | enables faster async host→device copies |
persistent_workers | True with num_workers > 0 | skip worker re-spawn every epoch |
drop_last | True for training | keeps batch shape constant; avoids a tiny final batch |
prefetch_factor | 2–4 | buffer batches ahead without hoarding RAM |
Datasetyields one example;DataLoaderyields batches and parallelises the slow work.- Put decoding and augmentation in
__getitem__so the workers do it off the critical path. - Tune
num_workersto where throughput plateaus — then stop. pin_memory=True+.to(device, non_blocking=True)overlaps the copy with compute.- Custom
collate_fnfor variable-length data;WeightedRandomSamplerfor imbalance.
References
- [1]torch.utils.data — Dataset, DataLoader, samplers · PyTorch documentationMap- vs iterable-style, multiprocessing behaviour, memory pinning.
- [2]Datasets & DataLoaders — Learn the Basics · PyTorch tutorials
- [3]Transforming and augmenting images (transforms v2) · torchvision documentation
- [4]A performance guide for PyTorch data loading · PyTorch tutorials


