TextDataLoader
| 34 | |
| 35 | |
| 36 | class TextDataLoader(object): |
| 37 | """TextDataLoader""" |
| 38 | def __init__(self, dataset, batch_sampler, collate_fn, use_data_parallel=False, use_multiprocess=True): |
| 39 | self.dataset = dataset |
| 40 | self.batch_sampler = batch_sampler |
| 41 | self.fields = self.dataset.fields |
| 42 | self.collate_fn = collate_fn |
| 43 | self.use_data_parallel = use_data_parallel |
| 44 | self.dataloader = io.DataLoader.from_generator(capacity=10, return_list=True, use_multiprocess=use_multiprocess) |
| 45 | self.dataloader.set_batch_generator(self.generator_creator()) |
| 46 | |
| 47 | def __call__(self): |
| 48 | """call""" |
| 49 | return self.dataloader() |
| 50 | |
| 51 | def generator_creator(self): |
| 52 | """Returns a generator, each iteration returns a batch of data""" |
| 53 | def __reader(): |
| 54 | for batch_sample_id in self.batch_sampler: |
| 55 | batch = [] |
| 56 | raw_batch = self.collate_fn([self.dataset[sample_id] for sample_id in batch_sample_id]) |
| 57 | for data, field in zip(raw_batch, self.fields): |
| 58 | if isinstance(data[0], np.ndarray): |
| 59 | data = nn.pad_sequence(data, field.pad_index) |
| 60 | elif isinstance(data[0], Iterable): |
| 61 | data = [nn.pad_sequence(f, field.pad_index) for f in zip(*data)] |
| 62 | batch.append(data) |
| 63 | yield batch |
| 64 | |
| 65 | return __reader |
| 66 | |
| 67 | def __len__(self): |
| 68 | """Returns the number of batches""" |
| 69 | return len(self.batch_sampler) |
| 70 | |
| 71 | |
| 72 | class TextDataset(object): |