| 147 | |
| 148 | class DataModuleFromConfig(pl.LightningDataModule): |
| 149 | def __init__(self, batch_size, train=None, validation=None, test=None, |
| 150 | wrap=False, num_workers=None): |
| 151 | super().__init__() |
| 152 | self.batch_size = batch_size |
| 153 | self.dataset_configs = dict() |
| 154 | self.num_workers = num_workers if num_workers is not None else batch_size*2 |
| 155 | if train is not None: |
| 156 | self.dataset_configs["train"] = train |
| 157 | self.train_dataloader = self._train_dataloader |
| 158 | if validation is not None: |
| 159 | self.dataset_configs["validation"] = validation |
| 160 | self.val_dataloader = self._val_dataloader |
| 161 | if test is not None: |
| 162 | self.dataset_configs["test"] = test |
| 163 | self.test_dataloader = self._test_dataloader |
| 164 | self.wrap = wrap |
| 165 | |
| 166 | def prepare_data(self): |
| 167 | for data_cfg in self.dataset_configs.values(): |