| 212 | |
| 213 | class DataModuleFromConfig(pl.LightningDataModule): |
| 214 | def __init__(self, |
| 215 | batch_size: int, |
| 216 | val_batch_size: int = None, |
| 217 | train: dict = None, |
| 218 | validation: dict = None, |
| 219 | test: dict = None, |
| 220 | shuffle_validation: bool = False, |
| 221 | num_workers: int = 0 |
| 222 | ): |
| 223 | super().__init__() |
| 224 | self.batch_size = batch_size |
| 225 | self.train = train |
| 226 | self.validation = validation |
| 227 | self.num_workers = num_workers |
| 228 | self.val_batch_size = val_batch_size if val_batch_size is not None else batch_size |
| 229 | self.shuffle_validation = shuffle_validation |
| 230 | |
| 231 | self.dataset_configs = {} |
| 232 | if train is not None: |
| 233 | self.dataset_configs["train"] = train |
| 234 | self.train_dataloader = self._train_dataloader |
| 235 | if validation is not None: |
| 236 | self.dataset_configs["validation"] = validation |
| 237 | self.val_dataloader = self._val_dataloader |
| 238 | if test is not None: |
| 239 | self.dataset_configs["test"] = test |
| 240 | self.test_dataloader = self._test_dataloader |
| 241 | |
| 242 | def _train_dataloader(self): |
| 243 | return DataLoader(self.datasets["train"], batch_size=self.batch_size, |