Datasets & DataLoaders
PyTorch Dataset and DataLoader feed training loops with batched, shuffled, and optionally augmented data. Efficient input pipelines keep the GPU busy and training fast.
Search across all documentation pages
PyTorch Dataset and DataLoader feed training loops with batched, shuffled, and optionally augmented data. Efficient input pipelines keep the GPU busy and training fast.
Quick-reference recipe card - copy-paste ready.
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])
loader = DataLoader(dataset, batch_size=64, shuffle=True, num_workers=4, pin_memory=True)
for batch_x, batch_y in loader:
batch_x = batch_x.to(device, non_blocking=True)When to reach for this:
collate_fn.pin_memory and non_blocking."""datasets_dataloaders.py - custom Dataset and DataLoader for CSV data."""
from __future__ import annotations
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader, random_split
class TabularDataset(Dataset):
def __init__(self, df: pd.DataFrame, feature_cols: list[str], target_col: str):
self.X = torch.tensor(df[feature_cols].values, dtype=torch.float32)
self.y = torch.tensor(df[target_col].values, dtype=torch.long)
def __len__(self) -> int:
return len(self.y)
def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
return self.X[idx], self.y[idx]
df = pd.read_csv("train.csv")
dataset = TabularDataset(df, feature_cols=["f1", "f2", "f3"], target_col="label")
train_size = int(0.8 * len(dataset))
train_ds, val_ds = random_split(dataset, [train_size, len(dataset) - train_size])
train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=2, pin_memory=True)
val_loader = DataLoader(val_ds, batch_size=256, shuffle=False, num_workers=2)
for epoch in range(3):
for x_batch, y_batch in train_loader:
# x_batch: (128, 3), y_batch: (128,)
passWhat this demonstrates:
Dataset with __len__ and __getitem__.random_split for train/validation partition.DataLoader handles batching and shuffling.pin_memory=True speeds CPU-to-GPU transfer when CUDA is available.Dataset.__getitem__(i) returns one sample; DataLoader collects batch_size samples.shuffle=True randomizes order each epoch.num_workers spawns subprocesses for parallel data loading.collate_fn customizes how samples merge into a batch (padding for sequences).IterableDataset streams data for very large or online sources.| Parameter | Effect | Typical Value |
|---|---|---|
batch_size | Samples per batch | 32-256 (GPU memory dependent) |
shuffle | Randomize order | True for training |
num_workers | Parallel loaders | 4-8 on multi-core |
pin_memory | Page-locked host memory | True with CUDA |
drop_last | Drop incomplete final batch | True for BatchNorm training |
from torchvision import datasets, transforms
# Built-in image dataset with transforms
transform = transforms.Compose([
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
train_set = datasets.CIFAR10(root="./data", train=True, download=True, transform=transform)num_workers=0 in notebooks; use scripts for parallel loading.pin_memory=True and non_blocking=True on .to(device).__getitem__, memory-mapped files, or IterableDataset.__getitem__.shuffle=False for val/test loaders.| Alternative | Use When | Don't Use When |
|---|---|---|
DataLoader | Standard batch training | Streaming infinite data (use IterableDataset) |
torchvision.datasets | Common image benchmarks | Custom data formats |
HuggingFace datasets | NLP/text datasets | Simple image classification |
| WebDataset | Large-scale web-scale training | Small local datasets |
def pad_collate(batch):
xs, ys = zip(*batch)
lengths = [len(x) for x in xs]
padded = torch.nn.utils.rnn.pad_sequence(xs, batch_first=True)
return padded, torch.tensor(ys), torch.tensor(lengths)
loader = DataLoader(ds, collate_fn=pad_collate)num_workers until CPU is saturated.pin_memory=True with CUDA.DataLoader(ds, num_workers=4, persistent_workers=True)__getitem__ or preload in __init__ for small data.__getitem__ call.__getitem__).shuffle - implement shuffling in the iterator or buffer.generator = torch.Generator().manual_seed(42)
loader = DataLoader(ds, shuffle=True, generator=generator)from torch.utils.data import WeightedRandomSampler
sampler = WeightedRandomSampler(weights, num_samples=len(weights))
loader = DataLoader(ds, sampler=sampler)pin_memory has no effect on MPS - use device="mps".num_workers=0 vs 4 to isolate I/O vs compute.__getitem__ does heavy work (decoding images on the fly).Stack versions: This page was written for Python 3.14.0 (stable 3.14, maintenance 3.13), FastAPI 0.115+, Django 5.2, Flask 3.1, Pydantic 2, PyTorch 2.6+, pandas 2.2+, Polars 1.x, ruff 0.9+, and uv 0.6+.
Reviewed by Chris St. John·Last updated Jul 16, 2026