PyTorch Cheatsheet
Datasets and DataLoaders
Use this PyTorch reference while you build software engineering projects, review code for technical interview prep, or polish examples for a software engineer resume.
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: ...