Datasets and DataLoaders

Writing your own Dataset

A PyTorch Dataset is two promises — how many samples exist, and how to fetch sample number i — and everything else in the data pipeline is built on those two.

Read these first

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A Dataset is a class that answers two questions: how many samples do you have, and give me sample number i.

Think of a librarian. You do not need to know how the store room is arranged — shelves, boxes, a basement. You ask two things only: "how many books do you hold?" and "bring me book number 4,081." The librarian handles everything behind the counter.

A PyTorch Dataset is that librarian. However your data is stored — one folder of photos, a CSV file, a database — you wrap it in a class that can count its samples and fetch any one of them by number.

Why it exists

Training needs data served in a very particular rhythm: shuffled differently every pass, grouped into batches, fetched fast. That rhythm logic is the same for every project on earth.

What differs is the storage: your files, your formats, your folder layout. PyTorch splits the two cleanly. You write the librarian — the part only you can write. PyTorch provides the delivery system — the DataLoader, which is the next lesson — and the two connect through those two questions.

This split is why the same training loop can feed on photos today and sensor readings tomorrow. Swap librarians; the delivery system never changes.

How it works

your storage          your Dataset               PyTorch's DataLoader
(files, CSV,   →   "I hold 160 samples"    →    shuffles the numbers,
 database)         "sample 42 is: (x, y)"       asks for them in batches,
                                                delivers to training

Each answer to "give me sample i" is one example: the input, and the answer the model should learn — like one photo and its label.

A real example you have seen

Every model you have heard of drank its training data through this pattern or one like it. A photo tagger's librarian read image files and their tags; a music app's read listening histories. Different store rooms, same two questions.

Remember this

  • A Dataset answers "how many?" and "give me number i."
  • You write it because only you know your storage.
  • Everything downstream — shuffling, batching, speed — builds on these two answers.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU. The data is generated with a fixed seed, so this output is exact on 2.5; another version could vary the random values.

A complete Dataset in fifteen lines

mango_dataset.py
import torch
from torch.utils.data import Dataset

class MangoDataset(Dataset):
    """160 fake mangoes: two measurements each, and a ripe/unripe label."""
    def __init__(self):
        g = torch.Generator().manual_seed(0)          # own generator: reproducible data
        self.features = torch.rand(160, 2, generator=g)
        self.labels = (self.features.sum(dim=1) > 1.0).long()

    def __len__(self):
        return len(self.labels)

    def __getitem__(self, idx):
        return self.features[idx], self.labels[idx]

ds = MangoDataset()
print("dataset length:", len(ds))
x, y = ds[0]
print("one sample:", [round(v, 3) for v in x.tolist()], "label", y.item())
print("ripe fraction:", ds.labels.float().mean().item())
Output
dataset length: 160
one sample: [0.496, 0.768] label 1
ripe fraction: 0.40625

The contract, precisely

__len__ returns the sample count. The DataLoader uses it to know the range of valid numbers.

__getitem__(idx) returns sample idx, for any idx from 0 to len - 1, in any order, as many times as asked. Return a tuple; a consistent structure every time. Tensors are ideal, and everything numeric should already be a tensor here, not a Python list.

That is the entire interface. Dataset is a class with no behaviour of its own — subclassing it is a promise that these two methods exist.

Where work should happen: the real design decision

The example above loads everything in __init__ — fine for data that fits in RAM. For 100,000 photos, it is not fine, and the pattern flips to lazy loading: __init__ stores only the list of file paths, and __getitem__ opens one file when asked.

folder_dataset.py
from pathlib import Path
import torch
from torch.utils.data import Dataset

# Stand-in for a real folder: six tiny measurement files.
root = Path("readings")
root.mkdir(exist_ok=True)
for i in range(6):
    (root / f"mango_{i}_{'ripe' if i % 2 else 'unripe'}.txt").write_text(f"{i * 0.1:.1f},{i * 0.2:.1f}")

class FolderDataset(Dataset):
    def __init__(self, root):
        # cheap: collect paths and labels, load NOTHING yet
        self.items = [(p, 1 if "unripe" not in p.name else 0)
                      for p in sorted(Path(root).glob("*.txt"))]

    def __len__(self):
        return len(self.items)

    def __getitem__(self, idx):
        path, label = self.items[idx]
        # expensive work lives here, done for ONE sample per call
        values = [float(v) for v in path.read_text().split(",")]
        return torch.tensor(values), label

ds = FolderDataset(root)
print("samples found:", len(ds))
x, y = ds[3]
print("sample 3:", [round(v, 1) for v in x.tolist()], "label", y)
Output
samples found: 6
sample 3: [0.3, 0.6] label 1

Swap the text files for JPEGs and the file-read for an image decode, and this is every image dataset ever written.

The rule: __init__ should be cheap and hold small things; __getitem__ does the per-sample work. Per-sample work in __getitem__ is what worker processes can parallelise later — num_workers multiplies exactly this method. Image decoding belongs there; so does augmentation, covered in transforms. Datasets larger than memory get their own lesson.

sorted(...) is not decoration. Filesystem listing order is not guaranteed; unsorted, sample 17 is a different file on your laptop and your training server, and your "reproducible" run is not.

Common mistakes

Randomness in the wrong place. Generating data with global randomness in __init__ means a different dataset every run. Use a seeded local Generator as above. Random augmentation is different — that belongs per-call in __getitem__, deliberately fresh each time.

Returning inconsistent shapes. If sample 7 is a 2-vector and sample 8 a 3-vector, batching fails later with a stacking error. Fix the shape here, or read the collate lesson for genuinely variable data.

Labels as Python ints in a tuple of mixed types. Works, but every epoch pays conversion costs. Convert once, in __init__ or on load.

Heavy objects in __init__ you plan to parallelise. Open file handles, database connections, and 4 GB arrays stored on self clash with multiprocessing later — the Windows spawn trap lesson explains the wreckage. Store paths and small metadata; open things lazily.

Try it yourself

Add a transform=None argument to MangoDataset.__init__, and in __getitem__ apply it to x when present. You have reinvented the exact convention torchvision uses — check its source later and find your own design.

What to learn next

Researcher — Mathematics and papers.

The abstraction and its family

torch.utils.data.Dataset[T] is the map-style protocol: __len__ plus __getitem__: int → T, i.e. a finite indexed family ${s_i}_{i=0}^{N-1}$ with random access. Its sibling IterableDataset drops random access for pure iteration — the right shape for streams and sharded archives, with its own worker-sharding obligations (lesson). The split mirrors the storage dichotomy: seekable versus sequential.

Random access is the property doing the heavy lifting: sampling without replacement (a permutation each epoch), weighted sampling, and deterministic resharding all reduce to permuting or reweighting indices — impossible to express over a pure stream without buffering. The composability dividends: Subset (index remapping), ConcatDataset (offset arithmetic over constituent lengths), random_split — all index algebra, no data movement, all O(1) per access.

Two extensions to the base protocol matter in current torch: __getitems__(indices) → list[T] (note the plural), which the default fetcher uses when present to amortise per-call overhead across a batch; and the convention that __getitem__ may raise IndexError to make ConcatDataset bisection safe. Neither is required; both are read reflectively.

Determinism and process semantics

The dataset object is constructed once in the parent process and replicated into workers — by fork inheritance on Linux, by pickling on Windows and macOS (spawn). Consequences: (1) mutable state accumulated in __getitem__ diverges per worker and is lost to the parent; treat the dataset as logically immutable after construction. (2) Anything unpicklable stored on self — open HDF5 handles, lambdas, database connections — breaks spawn-based loading; the standard fix is lazy per-worker initialisation guarded by a None check, or worker_init_fn. (3) Per-sample randomness must key off the DataLoader's per-worker seeding (torch seeds each worker distinctly per epoch) rather than a seed captured in __init__, or augmentations repeat across workers — a real, published class of bug (NumPy generators inherited identical state across forked workers; torch's own generators are per-worker seeded).

Benchmark-scale practice pushes the protocol's limits: when $N$ is $10^9$ (web-scale pretraining), the index list itself is gigabytes and per-sample __getitem__ calls dominate; production systems shift to sharded sequential formats — WebDataset tar shards, Parquet, FFCV — trading random access for throughput, which is exactly the map-to-iterable migration in this section's later lessons.

Reading

  • Paszke et al. (2019), PyTorch: An Imperative Style, High-Performance Deep Learning Library — §data loading, for the design rationale.
  • The torch.utils.data module docs — the normative protocol statement, including __getitems__.

What to learn next

What to learn next

These follow on from what you just read.

  • Datasets and DataLoaders

    What DataLoader actually does

    DataLoader is four small machines in a row — a sampler choosing indices, a batcher grouping them, a fetcher calling your Dataset, and a collator gluing samples into one tensor batch.

  • Datasets and DataLoaders

    Writing a collate_fn

    collate_fn is the function that glues single samples into one batch — write your own when samples have different lengths and the default stacker crashes.

  • Datasets and DataLoaders

    num_workers, prefetching and the Windows spawn trap

    num_workers puts extra processes on data preparation, prefetching keeps batches ready ahead of time — and on Windows, forgetting one guard line crashes training before it starts.