extracts samples only pertaining to this worker's batch
(self, batch)
| 103 | return self.train_iters |
| 104 | |
| 105 | def _batch(self, batch): |
| 106 | """extracts samples only pertaining to this worker's batch""" |
| 107 | start = self.rank*self.batch_size//self.world_size |
| 108 | end = (self.rank+1)*self.batch_size//self.world_size |
| 109 | return batch[start:end] |
| 110 | |
| 111 | |
| 112 | class DistributedBatchSampler(data.sampler.BatchSampler): |