| 3 | |
| 4 | |
| 5 | class CustomDatasetDataLoader: |
| 6 | def __init__(self, opt, is_for_train=True): |
| 7 | self._opt = opt |
| 8 | self._is_for_train = is_for_train |
| 9 | self._num_threds = opt.n_threads_train if is_for_train else opt.n_threads_test |
| 10 | self._create_dataset() |
| 11 | |
| 12 | def _create_dataset(self): |
| 13 | self._dataset = DatasetFactory.get_by_name(self._opt.dataset_mode, self._opt, self._is_for_train) |
| 14 | self._dataloader = torch.utils.data.DataLoader( |
| 15 | self._dataset, |
| 16 | batch_size=self._opt.batch_size, |
| 17 | shuffle=not self._opt.serial_batches, |
| 18 | num_workers=int(self._num_threds), |
| 19 | drop_last=True) |
| 20 | |
| 21 | def load_data(self): |
| 22 | return self._dataloader |
| 23 | |
| 24 | def __len__(self): |
| 25 | return len(self._dataset) |