| 42 | |
| 43 | |
| 44 | class DataModuleFromConfig(pl.LightningDataModule): |
| 45 | def __init__(self, batch_size, train=None, validation=None, test=None, predict=None, |
| 46 | wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False, |
| 47 | shuffle_val_dataloader=False, train_img=None, |
| 48 | test_max_n_samples=None): |
| 49 | super().__init__() |
| 50 | self.batch_size = batch_size |
| 51 | self.dataset_configs = dict() |
| 52 | self.num_workers = num_workers if num_workers is not None else batch_size * 2 |
| 53 | self.use_worker_init_fn = use_worker_init_fn |
| 54 | if train is not None: |
| 55 | self.dataset_configs["train"] = train |
| 56 | self.train_dataloader = self._train_dataloader |
| 57 | if validation is not None: |
| 58 | self.dataset_configs["validation"] = validation |
| 59 | self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader) |
| 60 | if test is not None: |
| 61 | self.dataset_configs["test"] = test |
| 62 | self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader) |
| 63 | if predict is not None: |
| 64 | self.dataset_configs["predict"] = predict |
| 65 | self.predict_dataloader = self._predict_dataloader |
| 66 | |
| 67 | self.img_loader = None |
| 68 | self.wrap = wrap |
| 69 | self.test_max_n_samples = test_max_n_samples |
| 70 | self.collate_fn = None |
| 71 | |
| 72 | def prepare_data(self): |
| 73 | pass |
| 74 | |
| 75 | def setup(self, stage=None): |
| 76 | self.datasets = dict((k, instantiate_from_config(self.dataset_configs[k])) for k in self.dataset_configs) |
| 77 | if self.wrap: |
| 78 | for k in self.datasets: |
| 79 | self.datasets[k] = WrappedDataset(self.datasets[k]) |
| 80 | |
| 81 | def _train_dataloader(self): |
| 82 | is_iterable_dataset = isinstance(self.datasets['train'], Txt2ImgIterableBaseDataset) |
| 83 | if is_iterable_dataset or self.use_worker_init_fn: |
| 84 | init_fn = worker_init_fn |
| 85 | else: |
| 86 | init_fn = None |
| 87 | loader = DataLoader(self.datasets["train"], batch_size=self.batch_size, |
| 88 | num_workers=self.num_workers, shuffle=False if is_iterable_dataset else True, |
| 89 | worker_init_fn=init_fn, collate_fn=self.collate_fn, |
| 90 | ) |
| 91 | return loader |
| 92 | |
| 93 | def _val_dataloader(self, shuffle=False): |
| 94 | if isinstance(self.datasets['validation'], Txt2ImgIterableBaseDataset) or self.use_worker_init_fn: |
| 95 | init_fn = worker_init_fn |
| 96 | else: |
| 97 | init_fn = None |
| 98 | return DataLoader(self.datasets["validation"], |
| 99 | batch_size=self.batch_size, |
| 100 | num_workers=self.num_workers, |
| 101 | worker_init_fn=init_fn, |
nothing calls this directly
no outgoing calls
no test coverage detected