| ''' |
| ----------------------------------------------------------------------------- |
| Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. |
| |
| NVIDIA CORPORATION and its licensors retain all intellectual property |
| and proprietary rights in and to this software, related documentation |
| and any modifications thereto. Any use, reproduction, disclosure or |
| distribution of this software and related documentation without an express |
| license agreement from NVIDIA CORPORATION is strictly prohibited. |
| ----------------------------------------------------------------------------- |
| ''' |
|
|
| import torch |
|
|
|
|
| class MultiEpochsDataLoader(torch.utils.data.DataLoader): |
| """ |
| Relentlessly sample from the dataset. |
| This eliminates the overhead of prefetching data before each epoch. |
| https://github.com/rwightman/pytorch-image-models/blob/master/timm/data/loader.py |
| """ |
|
|
| def __init__(self, *args, **kwargs): |
| super().__init__(*args, **kwargs) |
| self._DataLoader__initialized = False |
| self.batch_sampler = _RepeatSampler(self.batch_sampler) |
| self._DataLoader__initialized = True |
| self.iterator = super().__iter__() |
|
|
| def __len__(self): |
| return len(self.batch_sampler.sampler) |
|
|
| def __iter__(self): |
| for i in range(len(self)): |
| yield next(self.iterator) |
|
|
|
|
| class _RepeatSampler(object): |
| """ Sampler that repeats forever. |
| Args: |
| sampler (Sampler) |
| """ |
|
|
| def __init__(self, sampler): |
| self.sampler = sampler |
|
|
| def __iter__(self): |
| while True: |
| yield from iter(self.sampler) |
|
|