(self)
| 128 | self.effective_batch_size = batch_size if gradient_accumulation_steps is None else batch_size * gradient_accumulation_steps |
| 129 | |
| 130 | def __iter__(self): |
| 131 | batch = [] |
| 132 | i = 0 |
| 133 | for idx in self.data_iterator(self.sampler, wrap_around=False): |
| 134 | batch.append(idx) |
| 135 | if len(batch) == self.batch_size: |
| 136 | tbatch = self._batch(batch) |
| 137 | if i >= self.start_iter * self.effective_batch_size: |
| 138 | yield tbatch |
| 139 | self.start_iter = 0 |
| 140 | i += len(batch) |
| 141 | batch = [] |
| 142 | batch_len = len(batch) |
| 143 | if batch_len > 0 and not self.drop_last: |
| 144 | if self.wrap_last: |
| 145 | self.sampler.wrap_around -= (self.batch_size) |
| 146 | self.wrap_around += (len(batch)) |
| 147 | self.wrap_around %= self.batch_size |
| 148 | yield self._batch(batch) |
| 149 | if self.wrap_last: |
| 150 | self.sampler.wrap_around += self.batch_size |
| 151 | |
| 152 | def data_iterator(self, _iter, wrap_around=False): |
| 153 | """iterates through data and handles wrap around""" |
nothing calls this directly
no test coverage detected