Datasets and DataLoaders

IterableDataset for streams and huge files

IterableDataset serves data as a flowing stream instead of a numbered shelf — the right shape for logs, network feeds and files too big to index, with one sharding trap every user hits.

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.

An IterableDataset hands out samples one after another, like water from a tap — you take what flows, instead of asking for item number i.

A bookshelf and a river. From a bookshelf you can demand book 4,081, any book, any order — that is the regular Dataset, the librarian with numbered shelves.

A river does not do requests. You stand at the bank and take what flows past, in the order it arrives. Some data is a river: a live feed of sensor readings, logs pouring off a server, a file so enormous that jumping to line forty million is itself slow.

IterableDataset is PyTorch's river. It promises one thing only: keep handing me the next sample until you run dry.

Why it exists

The numbered-shelf contract has two costly assumptions: you know how many samples exist, and you can jump to any of them cheaply.

Streams break the first — a live feed has no length yet. Huge compressed files break the second — reading line forty million means decompressing everything before it, so "random access" costs a full read-through. Forcing river-shaped data onto shelf-shaped promises means fake indexes and terrible performance.

So PyTorch offers the second contract. Fewer promises, and in exchange it handles data of any size, arriving at any time, from anywhere.

There is a famous catch. Hire several workers to read the same river, and each reads the whole river — every sample arrives once per worker. Duplicated data, silently. The stream must be deliberately split between them.

How it works

 shelf (Dataset):        river (IterableDataset):
   "give me #4081"          "next, please"
   knows its length          may have no length
   any order                 arrival order

 the trap, with 2 workers:            the fix:
   worker A: reads ALL 20 samples       worker A: takes samples 0, 2, 4, ...
   worker B: reads ALL 20 samples       worker B: takes samples 1, 3, 5, ...
   → model sees everything TWICE        → each sample exactly once

A real example you have seen

Recommendation systems retrain on rivers of clicks that never stop. Nobody numbers those clicks on a shelf first — the training pipeline drinks directly from the stream.

Remember this

  • IterableDataset = next-sample-please; no counting, no jumping.
  • Right shape for streams, logs, and files too big to index.
  • With multiple workers, you must split the stream, or every sample arrives once per worker.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on Windows; the sharding behaviour is identical everywhere.

The trap, demonstrated honestly

duplicate_trap.py
import torch
from torch.utils.data import IterableDataset, DataLoader

class NaiveStream(IterableDataset):
    def __iter__(self):
        for i in range(20):
            yield torch.tensor(i)

if __name__ == "__main__":
    loader = DataLoader(NaiveStream(), batch_size=5, num_workers=2)
    seen = [x for b in loader for x in b.tolist()]
    print("total items:", len(seen), "- expected 20")
    print("every item arrived twice:", sorted(seen) == sorted(list(range(20)) * 2))
Output
total items: 40 - expected 20
every item arrived twice: True

No warning, no error. Each worker got a full copy of the dataset object and iterated it start to finish. Your model trains on every sample twice per epoch, your loss curves look fine, and your epoch takes twice as long as it should. This bug ships to production regularly.

The fix: ask who you are

sharded_stream.py
import torch
from torch.utils.data import IterableDataset, DataLoader, get_worker_info

class LineStream(IterableDataset):
    """Streams numbers 0..19 as if reading lines from a huge file."""
    def __init__(self, n=20):
        self.n = n

    def __iter__(self):
        info = get_worker_info()
        if info is None:                       # single-process loading
            start, step = 0, 1
        else:                                  # each worker takes every k-th line
            start, step = info.id, info.num_workers
        for i in range(start, self.n, step):
            yield torch.tensor(i)

if __name__ == "__main__":
    loader = DataLoader(LineStream(), batch_size=5, num_workers=2)
    seen = []
    for batch in loader:
        seen.extend(batch.tolist())
    print("items seen:", sorted(seen))
    print("total:", len(seen), "- duplicates:", len(seen) - len(set(seen)))
Output
items seen: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19]
total: 20 - duplicates: 0

get_worker_info() is the whole trick. Called inside __iter__, it answers "which worker am I?" — None in single-process mode, otherwise an object carrying id and num_workers. Each worker then takes every k-th item, interleaved, and the union is exactly the stream.

For file-based work, shard at a coarser grain: with 8 log files and 2 workers, worker 0 reads files 0, 2, 4, 6 and worker 1 the rest. Whole files per worker beats interleaved lines — fewer open handles, sequential reads.

What you give up, concretely

No len(loader). Progress bars cannot show totals; code calling len() raises TypeError. Track steps, not epochs.

No shuffle=True. Shuffling means picking indices, and there are none — passing it raises an error. The stream substitute is a shuffle buffer: hold, say, 1,000 samples in a pool, yield one at random as each new one flows in. Partial shuffling, bounded memory.

No samplers. Everything from the sampler lesson is index machinery — none of it applies. Weighting must happen inside your stream logic or the loss.

Common mistakes

The duplicate trap. Above. Any IterableDataset review starts with "where does it call get_worker_info?"

Testing with num_workers=0 only. The trap is invisible in single-process mode — info is None, one full pass, correct results. It appears exactly when you turn workers on. Test both.

An epoch that never ends. A stream reading a live source has no natural end, and for batch in loader: never exits. Cap it: itertools.islice(loader, steps_per_epoch), treating an epoch as a step budget.

Choosing IterableDataset for shelf-shaped data. If your data sits in files you can index, the map-style Dataset keeps shuffling, sampling, and lengths for free. Choose the river only when the data truly flows — or is too big to index, though first read the too-big-for-RAM lesson, which often rescues the shelf.

Try it yourself

Add a shuffle buffer to LineStream: keep a list of up to 8 pending items; on each new item, insert it, then yield a random member (seed a local random.Random(info.id if info else 0)). Verify no duplicates with 2 workers, and observe the order scramble.

What to learn next

Researcher — Mathematics and papers.

The contract, and what the loader does differently

IterableDataset.__iter__ returns an iterator; the fetcher draws batch_size items per batch from it and stops cleanly on StopIteration (partial final batch included, unless drop_last). Under multiprocess loading, each worker constructs its own replica of the dataset and iterates independently — the duplication semantics are documented behaviour, not a bug, because only user code knows the right partition of an opaque stream. get_worker_info() (and worker_init_fn, which can mutate the per-worker replica before iteration) are the two sanctioned sharding hooks. In distributed data parallel, sharding is two-level — by rank, then by worker — and both must be composed by hand or via torchdata's utilities; missing either level silently multiplies data.

Since the loader cannot see indices, in-order delivery guarantees change too: batches from different workers arrive round-robin (worker 0, 1, ..., 0, 1, ...), so any per-worker rate skew shows up as interleaving pattern, not reordering within a worker's stream.

Shuffle buffers and their statistics

A size-$k$ buffer over a stream is a bounded-memory approximate shuffle: each emitted element is drawn uniformly from the current $k$-window. The result is a local permutation — an element cannot appear more than $k$ positions earlier than its arrival, so long-range order (e.g. class-sorted source files) survives shuffling unless $k$ approaches the correlation length of the stream. Practical systems therefore compose two stages: shard-level shuffling (shuffle the file list each epoch) plus sample-level buffering — the design of WebDataset and TFRecord pipelines alike. Buffer size trades memory for decorrelation; the pathology of too-small buffers on sorted data is a measurable accuracy hit, worth an ablation on any new corpus.

The systems argument for streams

Sequential I/O dominates random access on every storage tier: spinning disks by orders of magnitude, object stores (S3-like) by request-latency amortisation, decompression by construction (gzip streams are not seekable; zstd seekable framing is the exception, at a ratio cost). WebDataset formalises the response — shard datasets into tar files, stream shards sequentially, shuffle at both grains — turning web-scale training (LAION-scale image-text) into a bandwidth problem instead of an IOPS problem. FFCV and Mosaic's StreamingDataset occupy the same slot with random-access-friendly layouts and caching, recovering approximate shelf semantics on top of stream economics. The torchdata project's datapipes explored composable stream pipelines in-core; its API has been through deprecation cycles — check current status before adopting, and treat raw IterableDataset as the stable substrate.

Epoch semantics deserve precision in papers: for a true stream, "epoch" is a step budget and sampling is single-pass — the online-learning regime, where each gradient is computed on never-seen data and generalisation-gap arguments simplify (training loss is an estimate of test loss). Mixing that regime's metrics with multi-epoch baselines confounds comparisons; state which regime you are in.

Reading

  • torch.utils.data docs, IterableDataset — the normative multi-worker duplication warning and both sharding hooks.
  • Aizman, Maltby, Breuel (2019), High Performance I/O For Large Scale Deep Learning — the WebDataset design argument.

What to learn next