| 172 | return self._seqs_consumed |
| 173 | |
| 174 | def _reset_state(self): |
| 175 | self.epoch += 1 |
| 176 | self.buffer.clear() |
| 177 | self._seqs_consumed = 0 |
| 178 | self._seqs_emitted = 0 |
| 179 | self.np_rng = np.random.default_rng(self.epoch + self.seed) |
| 180 | |
| 181 | # Update the epoch for the sampler |
| 182 | if isinstance(self.src_iterable, torch.utils.data.dataloader.DataLoader): |
| 183 | if isinstance(self.src_iterable.sampler, torch.utils.data.distributed.DistributedSampler): |
| 184 | self.src_iterable.sampler.set_epoch(self.epoch) |
| 185 | |
| 186 | def __iter__(self): |
| 187 | self._reset_state() |