| 107 | # indefinitely, shuffling items as it goes. |
| 108 | |
| 109 | class InfiniteSampler(torch.utils.data.Sampler): |
| 110 | def __init__(self, dataset, rank=0, num_replicas=1, shuffle=True, seed=0, window_size=0.5): |
| 111 | assert len(dataset) > 0 |
| 112 | assert num_replicas > 0 |
| 113 | assert 0 <= rank < num_replicas |
| 114 | assert 0 <= window_size <= 1 |
| 115 | super().__init__(dataset) |
| 116 | self.dataset = dataset |
| 117 | self.rank = rank |
| 118 | self.num_replicas = num_replicas |
| 119 | self.shuffle = shuffle |
| 120 | self.seed = seed |
| 121 | self.window_size = window_size |
| 122 | |
| 123 | def __iter__(self): |
| 124 | order = np.arange(len(self.dataset)) |
| 125 | rnd = None |
| 126 | window = 0 |
| 127 | if self.shuffle: |
| 128 | rnd = np.random.RandomState(self.seed) |
| 129 | rnd.shuffle(order) |
| 130 | window = int(np.rint(order.size * self.window_size)) |
| 131 | |
| 132 | idx = 0 |
| 133 | while True: |
| 134 | i = idx % order.size |
| 135 | if idx % self.num_replicas == self.rank: |
| 136 | yield order[i] |
| 137 | if window >= 2: |
| 138 | j = (i - rnd.randint(window)) % order.size |
| 139 | order[i], order[j] = order[j], order[i] |
| 140 | idx += 1 |
| 141 | |
| 142 | #---------------------------------------------------------------------------- |
| 143 | # Utilities for operating with torch.nn.Module parameters and buffers. |
nothing calls this directly
no outgoing calls
no test coverage detected