| 178 | |
| 179 | |
| 180 | class DataModuleFromConfig(pl.LightningDataModule):# batchloader outputshape should be (b,h,w,c) and it will be permuted to (b,c,h,w) in autoencoder.get_input() |
| 181 | def __init__(self, batch_size, train=None, validation=None, test=None, predict=None, |
| 182 | wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False, |
| 183 | shuffle_val_dataloader=False): |
| 184 | super().__init__() |
| 185 | self.batch_size = batch_size |
| 186 | self.dataset_configs = dict() |
| 187 | self.num_workers = num_workers if num_workers is not None else batch_size * 2 |
| 188 | self.use_worker_init_fn = use_worker_init_fn |
| 189 | if train is not None: |
| 190 | self.dataset_configs["train"] = train |
| 191 | self.train_dataloader = self._train_dataloader |
| 192 | if validation is not None: |
| 193 | self.dataset_configs["validation"] = validation |
| 194 | self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader) |
| 195 | if test is not None: |
| 196 | self.dataset_configs["test"] = test |
| 197 | self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader) |
| 198 | if predict is not None: |
| 199 | self.dataset_configs["predict"] = predict |
| 200 | self.predict_dataloader = self._predict_dataloader |
| 201 | self.wrap = wrap |
| 202 | |
| 203 | def prepare_data(self): |
| 204 | for data_cfg in self.dataset_configs.values(): |
| 205 | instantiate_from_config(data_cfg) |
| 206 | |
| 207 | def setup(self, stage=None): |
| 208 | self.datasets = dict( |
| 209 | (k, instantiate_from_config(self.dataset_configs[k])) |
| 210 | for k in self.dataset_configs) |
| 211 | if self.wrap: |
| 212 | for k in self.datasets: |
| 213 | self.datasets[k] = WrappedDataset(self.datasets[k]) |
| 214 | |
| 215 | def _train_dataloader(self): |
| 216 | init_fn = None |
| 217 | return DataLoader(self.datasets["train"], batch_size=self.batch_size ,# sampler=DistributedSampler # np.arange(100), |
| 218 | num_workers=self.num_workers, shuffle=True, |
| 219 | worker_init_fn=init_fn) |
| 220 | |
| 221 | def _val_dataloader(self, shuffle=False): |
| 222 | init_fn = None |
| 223 | return DataLoader(self.datasets["validation"], |
| 224 | batch_size=self.batch_size, |
| 225 | num_workers=self.num_workers, |
| 226 | worker_init_fn=init_fn, |
| 227 | shuffle=shuffle) |
| 228 | |
| 229 | def _test_dataloader(self, shuffle=False): |
| 230 | init_fn = None |
| 231 | # do not shuffle dataloader for iterable dataset |
| 232 | return DataLoader(self.datasets["test"], batch_size=self.batch_size, |
| 233 | num_workers=self.num_workers, worker_init_fn=init_fn, shuffle=shuffle) |
| 234 | |
| 235 | def _predict_dataloader(self, shuffle=False): |
| 236 | init_fn = None |
| 237 | return DataLoader(self.datasets["predict"], batch_size=self.batch_size, |
nothing calls this directly
no outgoing calls
no test coverage detected