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.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
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 trainingEach 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
- What DataLoader actually does — the delivery system built on your two answers.
- Writing a collate_fn — when samples refuse to stack into neat batches.
- Datasets that do not fit in memory — lazy loading grown up.
Developer — Code and libraries.
Setup
pip install torchWritten 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
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())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.
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)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
- What DataLoader actually does — the delivery system built on your two answers.
- Writing a collate_fn — when samples refuse to stack into neat batches.
- Datasets that do not fit in memory — lazy loading grown up.
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.datamodule docs — the normative protocol statement, including__getitems__.
What to learn next
- What DataLoader actually does — the delivery system built on your two answers.
- Writing a collate_fn — when samples refuse to stack into neat batches.
- Datasets that do not fit in memory — lazy loading grown up.