| 146 | |
| 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(): |
| 168 | instantiate_from_config(data_cfg) |
| 169 | |
| 170 | def setup(self, stage=None): |
| 171 | self.datasets = dict( |
| 172 | (k, instantiate_from_config(self.dataset_configs[k])) |
| 173 | for k in self.dataset_configs) |
| 174 | if self.wrap: |
| 175 | for k in self.datasets: |
| 176 | self.datasets[k] = WrappedDataset(self.datasets[k]) |
| 177 | |
| 178 | def _train_dataloader(self): |
| 179 | return DataLoader(self.datasets["train"], batch_size=self.batch_size, |
| 180 | num_workers=self.num_workers, shuffle=True, collate_fn=custom_collate) |
| 181 | |
| 182 | def _val_dataloader(self): |
| 183 | return DataLoader(self.datasets["validation"], |
| 184 | batch_size=self.batch_size, |
| 185 | num_workers=self.num_workers, collate_fn=custom_collate) |
| 186 | |
| 187 | def _test_dataloader(self): |
| 188 | return DataLoader(self.datasets["test"], batch_size=self.batch_size, |
| 189 | num_workers=self.num_workers, collate_fn=custom_collate) |
| 190 | |
| 191 | |
| 192 | class DataModuleTemplateFromConfig(DataModuleFromConfig): |
nothing calls this directly
no outgoing calls
no test coverage detected