PyTorch Cheatsheet
Datasets and DataLoaders
Use this PyTorch reference while you build software engineering projects, review code, or refresh the syntax you reach for most.
Dataset ABC
Subclass torch.utils.data.Dataset and implement __len__ and __getitem__.
from torch.utils.data import Dataset class MyDataset(Dataset): def __init__(self, data, labels, transform=None): self.data = data self.labels = labels self.transform = transform def __len__(self) -> int: return len(self.data) def __getitem__(self, idx: int): x = self.data[idx] y = self.labels[idx] if self.transform: x = self.transform(x) return x, y dataset = MyDataset(X, y) sample, label = dataset[0] print(len(dataset))
IterableDataset
Use when data is too large to index or comes from a stream.
from torch.utils.data import IterableDataset class StreamDataset(IterableDataset): def __init__(self, filepath): self.filepath = filepath def __iter__(self): with open(self.filepath) as f: for line in f: x, y = parse(line) yield x, y
With
IterableDatasetandnum_workers > 0, handle per-worker splitting yourself:
def __iter__(self): info = torch.utils.data.get_worker_info() if info is None: yield from self.all_records() else: # split records across workers for i, record in enumerate(self.all_records()): if i % info.num_workers == info.id: yield record
Built-in Datasets
from torchvision import datasets, transforms # Image datasets (torchvision) train = datasets.MNIST(root='./data', train=True, download=True, transform=transforms.ToTensor()) datasets.CIFAR10(root='./data', train=True, download=True) datasets.ImageNet(root='/data/imagenet', split='train') datasets.ImageFolder(root='./data/train') # folder-per-class layout datasets.DatasetFolder(root='./data', loader=..., extensions=('.pt',)) # Audio (torchaudio) from torchaudio.datasets import SPEECHCOMMANDS ds = SPEECHCOMMANDS(root='./data', download=True, subset='training')
torchtextis archived and incompatible with torch >= 2.4 — do not use it in new code. For text, use Hugging Facedatasets(works directly withDataLoader):
from datasets import load_dataset # pip install datasets ds = load_dataset('ag_news', split='train') ds = ds.with_format('torch') # returns tensors from __getitem__ loader = torch.utils.data.DataLoader(ds, batch_size=32, shuffle=True)
DataLoader
from torch.utils.data import DataLoader loader = DataLoader( dataset, batch_size=32, shuffle=True, # shuffles every epoch (not for IterableDataset) num_workers=4, # subprocess workers for loading pin_memory=True, # faster host→GPU transfer drop_last=False, # drop last incomplete batch prefetch_factor=2, # batches prefetched per worker (default 2) persistent_workers=True,# keep workers alive between epochs timeout=0, # seconds to wait for a data chunk collate_fn=None, # custom batching function sampler=None, # custom sampling strategy batch_sampler=None, # yields lists of indices (overrides batch_size) worker_init_fn=None, # called at start of each worker process generator=None, # torch.Generator for reproducibility ) # Typical training loop usage for epoch in range(num_epochs): for batch_x, batch_y in loader: batch_x = batch_x.to(device) batch_y = batch_y.to(device) ...
num_workers guidelines
| Scenario | Recommendation |
|---|---|
| Debugging | num_workers=0 (single process, easy to break) |
| CPU-bound loading | num_workers=4–8 |
| GPU-bound training | num_workers=2–4 + pin_memory=True |
| Windows | Keep low (num_workers=0 or 2) — fork not available |
Samplers
from torch.utils.data import ( SequentialSampler, RandomSampler, SubsetRandomSampler, WeightedRandomSampler, BatchSampler, DistributedSampler, ) # Random subset sampler = SubsetRandomSampler(indices=range(1000)) # Class-balanced sampling class_weights = [1.0, 2.0, 0.5] # weight per class sample_weights = [class_weights[label] for label in all_labels] sampler = WeightedRandomSampler(weights=sample_weights, num_samples=len(dataset), replacement=True) # Wrap in DataLoader loader = DataLoader(dataset, batch_size=32, sampler=sampler)
Custom Collation
The default collate_fn stacks tensors, handles None, etc. Override when you need variable-length batches:
def pad_collate(batch): xs, ys = zip(*batch) # pad sequences to max length in batch xs_padded = torch.nn.utils.rnn.pad_sequence(xs, batch_first=True) ys = torch.tensor(ys) return xs_padded, ys loader = DataLoader(dataset, batch_size=32, collate_fn=pad_collate)
RNN padding helpers
from torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence # Pad variable-length sequences padded = pad_sequence(list_of_tensors, batch_first=True, padding_value=0) # Pack for RNN (skip padding in computation) lengths = [len(s) for s in list_of_tensors] packed = pack_padded_sequence(padded, lengths, batch_first=True, enforce_sorted=False) out, hidden = lstm(packed) # Unpack back to padded output, lengths = pad_packed_sequence(out, batch_first=True)
Dataset Utilities
from torch.utils.data import ( random_split, Subset, ConcatDataset, ChainDataset, # for IterableDataset TensorDataset, ) # TensorDataset — wrap tensors directly ds = TensorDataset(X_tensor, y_tensor) # Split into train / val train_ds, val_ds = random_split(dataset, lengths=[0.8, 0.2]) # or fixed sizes: train_ds, val_ds = random_split(dataset, lengths=[800, 200]) # with reproducibility: train_ds, val_ds = random_split(dataset, [800, 200], generator=torch.Generator().manual_seed(42)) # Manual subset by index val_ds = Subset(dataset, indices=range(800, 1000)) # Concatenate datasets combined = ConcatDataset([ds1, ds2, ds3])
Transforms (torchvision)
from torchvision import transforms transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomCrop(224, padding=4), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), # PIL → [0,1] float32 CHW transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) # v2 API (recommended for torchvision >= 0.15) from torchvision.transforms import v2 transform = v2.Compose([ v2.RandomResizedCrop(224), v2.RandomHorizontalFlip(), v2.ToDtype(torch.float32, scale=True), # replaces ToTensor v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])
Reproducibility with DataLoader
def seed_worker(worker_id): import numpy as np, random worker_seed = torch.initial_seed() % 2**32 np.random.seed(worker_seed) random.seed(worker_seed) g = torch.Generator() g.manual_seed(42) loader = DataLoader( dataset, batch_size=32, shuffle=True, worker_init_fn=seed_worker, generator=g, )
Distributed Data Loading
from torch.utils.data.distributed import DistributedSampler sampler = DistributedSampler( dataset, num_replicas=world_size, rank=rank, shuffle=True, drop_last=False, ) loader = DataLoader(dataset, batch_size=32, sampler=sampler, num_workers=4, pin_memory=True) # IMPORTANT: reshuffle each epoch for epoch in range(epochs): sampler.set_epoch(epoch) # ensures different shuffle per epoch for batch in loader: ...