| 361 | |
| 362 | class DataModuleFromConfig(pl.LightningDataModule): |
| 363 | def __init__(self, batch_size, train=None, train2=None, validation=None, test=None, predict=None, |
| 364 | wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False, |
| 365 | shuffle_val_dataloader=False, with_prior_preservation=False): |
| 366 | super().__init__() |
| 367 | self.batch_size = batch_size |
| 368 | self.dataset_configs = dict() |
| 369 | self.num_workers = num_workers if num_workers is not None else batch_size * 2 |
| 370 | self.use_worker_init_fn = use_worker_init_fn |
| 371 | self.with_prior_preservation = with_prior_preservation |
| 372 | if train is not None: |
| 373 | self.dataset_configs["train"] = train |
| 374 | self.train_dataloader = self._train_dataloader |
| 375 | if train2 is not None and train2['params']['caption'] != '': |
| 376 | self.dataset_configs["train2"] = train2 |
| 377 | if validation is not None: |
| 378 | self.dataset_configs["validation"] = validation |
| 379 | self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader) |
| 380 | if test is not None: |
| 381 | self.dataset_configs["test"] = test |
| 382 | self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader) |
| 383 | if predict is not None: |
| 384 | self.dataset_configs["predict"] = predict |
| 385 | self.predict_dataloader = self._predict_dataloader |
| 386 | self.wrap = wrap |
| 387 | |
| 388 | def prepare_data(self): |
| 389 | for data_cfg in self.dataset_configs.values(): |