Implement of DataLoader.
| 45 | |
| 46 | |
| 47 | class DataLoader(object): |
| 48 | """ Implement of DataLoader. """ |
| 49 | |
| 50 | @classmethod |
| 51 | def add_cmdline_argument(cls, group): |
| 52 | group.add_argument('--shuffle', type=str2bool, default=True) |
| 53 | group.add_argument('--sort_pool_size', type=int, default=0) |
| 54 | return group |
| 55 | |
| 56 | def __init__(self, |
| 57 | dataset, |
| 58 | batch_size, |
| 59 | hparams, |
| 60 | collate_fn=None, |
| 61 | sampler=None, |
| 62 | is_test=False): |
| 63 | self.dataset = dataset |
| 64 | self.collate_fn = collate_fn |
| 65 | self.gpu = hparams.gpu |
| 66 | self.sort_pool_size = hparams.sort_pool_size |
| 67 | |
| 68 | if sampler is None: |
| 69 | if hparams.shuffle and not is_test: |
| 70 | sampler = RandomSampler(dataset) |
| 71 | else: |
| 72 | sampler = SequentialSampler(dataset) |
| 73 | |
| 74 | if self.sort_pool_size > 0 and not is_test: |
| 75 | sampler = SortedSampler(sampler, self.sort_pool_size) |
| 76 | |
| 77 | def reader(): |
| 78 | for idx in sampler: |
| 79 | yield idx |
| 80 | |
| 81 | drop_last = False if self.gpu <= 1 or is_test else True |
| 82 | self.reader = batch(reader, batch_size=batch_size, drop_last=drop_last) |
| 83 | self.num_batches = math.floor(len(dataset) / batch_size) if drop_last \ |
| 84 | else math.ceil(len(dataset) / batch_size) |
| 85 | |
| 86 | def __len__(self): |
| 87 | return self.num_batches |
| 88 | |
| 89 | def __iter__(self): |
| 90 | for batch_indices in self.reader(): |
| 91 | samples = [self.dataset[idx] for idx in batch_indices] |
| 92 | yield self.collate_fn(samples) |
| 93 | |
| 94 | |
| 95 | class SequentialDataLoaderWrapper: |
no outgoing calls
searching dependent graphs…