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 IterableDataset and num_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')

torchtext is archived and incompatible with torch >= 2.4 — do not use it in new code. For text, use Hugging Face datasets (works directly with DataLoader):

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

ScenarioRecommendation
Debuggingnum_workers=0 (single process, easy to break)
CPU-bound loadingnum_workers=4–8
GPU-bound trainingnum_workers=2–4 + pin_memory=True
WindowsKeep 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:
        ...