extracts samples only pertaining to this worker's batch
(self, batch)
| 160 | yield idx |
| 161 | |
| 162 | def _batch(self, batch): |
| 163 | """extracts samples only pertaining to this worker's batch""" |
| 164 | start = self.rank*self.batch_size//self.world_size |
| 165 | end = (self.rank+1)*self.batch_size//self.world_size |
| 166 | return batch[start:end] |