In PyTorch, when is IterableDataset right, and how do you shard it across workers?
answer
- stream, not an index
- every worker starts at the beginning
- who am I, how many of us
- shuffle argument is rejected here
- get_worker_info().id and num_workers
basics
~20 sUse IterableDataset when data arrives as a stream with no cheap random access — a remote log, a compressed shard, a queue. Each worker runs the whole iter, so without sharding by get_worker_info().id every sample is emitted once per worker.
solid answer
~50 sA map-style `Dataset` needs `__len__` and cheap random access by index. When neither exists — you are reading a compressed shard sequentially, consuming a Kafka-like stream, or the corpus is too large to index — subclass `torch.utils.data.IterableDataset` and implement `__iter__`. The trade-off is that `DataLoader` can no longer choose the order: passing `shuffle=True` or a `sampler` with an iterable-style dataset raises a `ValueError`. You get randomness yourself, typically by shuffling the shard order plus a fixed-size shuffle buffer inside `__iter__`. The pitfall is multiprocessing. Each worker gets its own copy of the dataset and runs `__iter__` from the top, so with `num_workers=4` you silently train on four copies of every sample. Fix it by calling `torch.utils.data.get_worker_info()` inside `__iter__`: when it is not `None`, use `info.id` and `info.num_workers` to take a disjoint slice — round-robin over records, or better, assign whole files per worker.
code
python · 21 linesimport torch
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
class RangeStream(IterableDataset):
def __init__(self, start, end):
self.start, self.end = start, end
def __iter__(self):
start, step = self.start, 1
info = get_worker_info()
if info is not None: # shard: worker i takes every num_workers-th item
start = self.start + info.id
step = info.num_workers
for i in range(start, self.end, step):
yield torch.tensor(i)
if __name__ == "__main__":
loader = DataLoader(RangeStream(0, 8), batch_size=4, num_workers=2)
print(torch.cat(list(loader)).sort().values) # 0..7, each exactly oncego deeper
Recognize the two dataset styles and know that an IterableDataset yields samples from __iter__ instead of answering an index, so shuffling is not something DataLoader can do for you.
Be able to state the duplication bug in one sentence — every worker replays the whole stream — and write the get_worker_info() sharding that fixes it.
Discuss the tradeoffs you have actually hit: file-level versus record-level sharding, tail imbalance, buffer-based shuffling quality, and rewriting a loop that no longer knows its step count.
Decide when streaming is warranted at all — shard layout, format and over-sharding factor are platform choices that outlive any one job, and a corpus indexed once may beat a stream that every team re-implements.
## Why the second style exists Map-style datasets assume two things: you know how many samples there are, and you can jump to sample 4,712,003 cheaply. Both fail for genuinely streaming sources — a gzip or tar shard you can only read forward, a network stream, a database cursor, a corpus of billions of tokens where building an index costs more than one training run. `IterableDataset` drops both assumptions: you implement `__iter__` and yield samples, and `DataLoader` batches whatever comes out. ## What DataLoader still does, and what it stops doing Still does: batching into `batch_size` groups, `collate_fn`, `drop_last`, `num_workers`, `pin_memory`. Stops doing: choosing the order. There is no index to permute, so `shuffle=True` is rejected with a `ValueError`, and `sampler` / `batch_sampler` are likewise invalid. Length is also mostly gone — if your dataset implements `__len__`, `DataLoader` will report a length derived from it and warn that the figure can be wrong across multiple workers; if it does not, `len(loader)` raises `TypeError`. Anything in your training loop that assumes a known step count (progress bars, `LambdaLR` schedules by epoch fraction, per-epoch checkpoint math) has to be rewritten to count steps instead. ## The duplication trap This is the interview question. Each worker process receives a copy of the dataset object and, when the main process asks it for a batch, calls `iter(dataset)` on **its own copy**. Nothing coordinates them. With four workers and no sharding logic, worker 0 yields the whole stream, and so do workers 1, 2 and 3: your epoch is four times as long and every sample appears four times, near-adjacently in the batch order. Nothing raises. Loss curves look plausible. It is a textbook silent-wrong-answer bug. `torch.utils.data.get_worker_info()` is the escape hatch. Called inside `__iter__`, it returns `None` in the single-process case and otherwise a `WorkerInfo` carrying `id`, `num_workers`, `seed` and `dataset` (the worker's copy). The two standard sharding patterns: **Record round-robin.** Iterate the whole source but only yield records where `index % info.num_workers == info.id`. Simple and perfectly balanced, but every worker still pays the full read and decompress cost, so it only helps when decode dominates I/O. **Shard assignment.** Split the file list so worker `i` opens only its own subset. This is what real pipelines do: each worker reads distinct bytes, so I/O scales with worker count. The cost is balance — if the number of shards is not a multiple of the worker count, or shards differ in length, some workers finish early and the tail of the epoch runs at reduced parallelism. Keep shards numerous and roughly equal, and over-shard relative to the expected worker count. `worker_init_fn` is the alternative injection point: it receives the worker id and can mutate the worker's dataset copy (reachable via `get_worker_info().dataset`) before iteration starts. Both approaches are equivalent; the `get_worker_info()`-inside-`__iter__` form keeps the logic in one place. ## Shuffling without an index Since the sampler is unavailable, randomness has to come from the stream itself, in two layers: shuffle the shard list per epoch (cheap, coarse), and maintain a reservoir of some thousands of samples in `__iter__` from which you yield at random and refill (fine-grained, bounded memory). Together these approximate a global shuffle well enough for large corpora. Neither is provided out of the box — you write them, or use a library built on `IterableDataset` that does. ## Ordering and epochs Batches from workers are returned in round-robin worker order by default, so the interleaving is deterministic given the shard assignment. "Epoch" becomes a convention: your `__iter__` ends when the assigned shards are exhausted, and the loader restarts it next epoch — so make sure any per-epoch state (shard shuffle seed, buffer) is reset at the top of `__iter__`, not in `__init__`, or every epoch will replay the identical order.
- How do you shuffle an IterableDataset if DataLoader's shuffle argument is rejected?In two layers, both inside `__iter__`. Shuffle the shard or file order each epoch for coarse randomness, then keep a fixed-size buffer of samples and yield a random element from it, refilling from the stream — a reservoir shuffle with bounded memory. Buffer size trades memory for how close you get to a global shuffle; too small and you effectively train on near-sequential data.
- Should you shard by record or by file across workers?By file when I/O matters: each worker opens distinct shards, so read and decompress cost scales with worker count. Record round-robin is perfectly balanced but every worker still reads every byte, which only pays off when decoding dominates. The file approach's weakness is imbalance at the tail, so over-shard relative to your worker count and keep shard sizes similar.
- What does len() return for a DataLoader over an IterableDataset?If the dataset implements `__len__`, DataLoader reports a batch count derived from it and warns that the value may be inaccurate — worker sharding and drop_last can make the real number differ. If the dataset has no `__len__`, `len(loader)` raises TypeError. Training loops over streams should therefore count optimizer steps rather than assume a known number of batches per epoch.
saying these in an interview costs you the question
- Assumes workers automatically split an iterable stream
- Passes shuffle=True with an IterableDataset
- Thinks IterableDataset supports a sampler
- Puts the epoch's shard shuffle in __init__ instead of __iter__
- Expects len(loader) to work for any streaming dataset