Ensuring Unique Data Distribution
To prevent multiple workers from yielding the same data shards, you must explicitly partition the dataset within the __iter__ method using torch.utils.data.get_worker_info(). Because IterableDataset does not support the shuffle=True argument in DataLoader, the dataset itself is responsible for determining which subset of data each worker process should stream.
Implementation Steps for Partitioning
- Retrieve Worker Metadata: Call
get_worker_info() inside the __iter__ method to access the id and num_workers of the current process.
- Calculate Shard Offset: Use the worker ID to slice the data source. For example, if you have a list of file paths, each worker should only iterate over paths where
index % num_workers == worker_id.
- Handle Single-Worker Cases: Ensure the logic defaults to processing the full dataset if
num_workers is set to 0 (main process).
Implementing Global Shuffling
True global shuffling is impossible with a pure stream without loading the entire index into memory. The recommended mechanism to approximate this is a Shuffle Buffer. This approach maintains a fixed-size window of samples and yields them randomly, providing local stochasticity that scales with the buffer size.
Recommended Shuffling Workflow
- Buffer Initialization: Create a list (the buffer) and fill it with the first
N samples from the partitioned stream.
- Random Sampling: For every new sample requested, randomly select one element from the buffer to yield, then immediately replace that element with the next sample from the stream.
- Worker Seeding: Use a
worker_init_fn to ensure each worker's random number generator is seeded differently, preventing synchronized shuffling patterns across processes.
# Conceptual implementation within __iter__
worker_info = torch.utils.data.get_worker_info()
if worker_info is None: # single-process loading
iter_start = 0
iter_end = len(self.data)
else:
# Partition data
per_worker = int(math.ceil(len(self.data) / float(worker_info.num_workers)))
worker_id = worker_info.id
iter_start = worker_id * per_worker
iter_end = min(iter_start + per_worker, len(self.data))
# Use a buffer for local shuffling
buffer = []
for i in range(iter_start, iter_end):
item = self.data[i]
if len(buffer) < buffer_size:
buffer.append(item)
else:
idx = random.randint(0, buffer_size - 1)
yield buffer[idx]
buffer[idx] = item
# Drain remaining buffer
random.shuffle(buffer)
for item in buffer: yield item
Verification and Assumptions
This implementation assumes PyTorch 1.7+ where get_worker_info() is stable. To verify the setup, print the worker_id and the first few sample IDs in each worker; they should be distinct and non-overlapping. Warning: Be cautious with buffer_size; since each worker maintains its own buffer, the total memory consumption is buffer_size * num_workers * sample_size.
Diagnostic Detail Needed: Are you streaming from a single large file (like a TFRecord or WebDataset) or a directory of many small files? The partitioning logic differs significantly between these two sources.