espnet3.components.data.epoch_sync_iterator.EpochSyncIterator
espnet3.components.data.epoch_sync_iterator.EpochSyncIterator
class espnet3.components.data.epoch_sync_iterator.EpochSyncIterator(source)
Bases: object
Per-epoch iterator returned by DataLoaderBuilder.
Wraps the iterator produced by an espnet2-style iter factory (iter_factory.build_iter(epoch)) and defines the iterator interface handed to the trainer.
When torch.distributed is initialized, this iterator synchronizes the end of the epoch across ranks. Iter factories like espnet2’s ChunkIterFactory emit a data-dependent number of batches per rank, so under DDP each rank would otherwise end its epoch at a different step. Lightning then moves the exhausted rank into validation while the others are still training, and the two sides issue mismatched collectives that deadlock until the NCCL watchdog kills the job. Before yielding each batch, all ranks all-reduce a 1-element has-next flag with MIN, so the epoch ends on every rank at the same batch index.
Without torch.distributed, batches are yielded unchanged.
- Parameters:source (Union *[*Callable , Iterable ]) – What this rank’s per-epoch passes are built from - not necessarily an iterator itself. Prefer a zero-argument callable returning a fresh iterator, e.g.
partial(iter_factory.build_iter, epoch): it is called once per__iter__, so repeated passes each get their own iterator. A re-iterable container (list,DataLoader) also works. Do not pass a single live generator -build_iterof espnet2’sChunkIterFactoryreturns one, and sharing it across passes empties every pass after the first.__len__is delegated to the source, so it must be sized forlen()to work.
NOTE
Subclasses should override generate() rather than __iter__, so that the epoch-end synchronization stays in place.
NOTE
The synchronization costs one all-reduce of a 1-element tensor per batch, which is negligible next to the gradient synchronization that every distributed training step already performs.
####### Examples
Without torch.distributed, the wrapper is a pass-through:
>>> list(EpochSyncIterator([{"speech": 0}, {"speech": 1}]))
[{'speech': 0}, {'speech': 1}]Under DDP, DataLoaderBuilder wraps the iter factory’s iterator so that every rank stops at the shortest rank’s batch count:
iter_factory = ChunkIterFactory(dataset, batches=batches, ...)
iterator = EpochSyncIterator(
partial(iter_factory.build_iter, epoch) # the call, not its result
)
for batch in iterator: # same number of steps on every rank
...Wrap the source this rank’s per-epoch passes are built from.
generate()
Yield this rank’s batches.
This is the extension point of the class: override it to change what a batch looks like, or which batches are emitted. __iter__ consumes this generator, so an override inherits the epoch-end synchronization for free.
- Yields:Any – The batches of one fresh pass over the source, in order and unchanged.
####### Examples
>>> class NonEmptyIterator(EpochSyncIterator):
... def generate(self):
... for batch in super().generate():
... if len(batch) > 0:
... yield batch
>>> list(NonEmptyIterator([[0], [], [1]]))
[[0], [1]]