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 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:
        ...