| 376 | |
| 377 | |
| 378 | class DataModuleFromConfig(pl.LightningDataModule): |
| 379 | def __init__(self, batch_size, train=None, train2=None, validation=None, test=None, predict=None, |
| 380 | wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False, |
| 381 | shuffle_val_dataloader=False): |
| 382 | super().__init__() |
| 383 | self.batch_size = batch_size |
| 384 | self.dataset_configs = dict() |
| 385 | self.num_workers = num_workers if num_workers is not None else batch_size * 2 |
| 386 | self.use_worker_init_fn = use_worker_init_fn |
| 387 | if train2 is not None and train2['params']['caption'] != '': |
| 388 | self.dataset_configs["train2"] = train2 |
| 389 | if train is not None: |
| 390 | self.dataset_configs["train"] = train |
| 391 | self.train_dataloader = self._train_dataloader |
| 392 | if validation is not None: |
| 393 | self.dataset_configs["validation"] = validation |
| 394 | self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader) |
| 395 | if test is not None: |
| 396 | self.dataset_configs["test"] = test |
| 397 | self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader) |
| 398 | if predict is not None: |
| 399 | self.dataset_configs["predict"] = predict |
| 400 | self.predict_dataloader = self._predict_dataloader |
| 401 | self.wrap = wrap |
| 402 | |
| 403 | def prepare_data(self): |
| 404 | for data_cfg in self.dataset_configs.values(): |
| 405 | instantiate_from_config(data_cfg) |
| 406 | |
| 407 | def setup(self, stage=None): |
| 408 | self.datasets = dict( |
| 409 | (k, instantiate_from_config(self.dataset_configs[k])) |
| 410 | for k in self.dataset_configs) |
| 411 | if self.wrap: |
| 412 | for k in self.datasets: |
| 413 | self.datasets[k] = WrappedDataset(self.datasets[k]) |
| 414 | |
| 415 | def _train_dataloader(self): |
| 416 | is_iterable_dataset = isinstance(self.datasets['train'], Txt2ImgIterableBaseDataset) |
| 417 | if is_iterable_dataset or self.use_worker_init_fn: |
| 418 | init_fn = worker_init_fn |
| 419 | else: |
| 420 | init_fn = None |
| 421 | if "train2" in self.dataset_configs and self.dataset_configs["train2"]['params']["caption"] != '': |
| 422 | train_set = self.datasets["train"] |
| 423 | train2_set = self.datasets["train2"] |
| 424 | concat_dataset = ConcatDataset(train_set, train2_set) |
| 425 | return DataLoader(concat_dataset, batch_size=self.batch_size // 2, |
| 426 | num_workers=self.num_workers, shuffle=False if is_iterable_dataset else True, |
| 427 | worker_init_fn=init_fn) |
| 428 | else: |
| 429 | return DataLoader(self.datasets["train"], batch_size=self.batch_size, |
| 430 | num_workers=self.num_workers, shuffle=False if is_iterable_dataset else True, |
| 431 | worker_init_fn=init_fn) |
| 432 | |
| 433 | def _val_dataloader(self, shuffle=False): |
| 434 | if isinstance(self.datasets['validation'], Txt2ImgIterableBaseDataset) or self.use_worker_init_fn: |
| 435 | init_fn = worker_init_fn |
nothing calls this directly
no outgoing calls
no test coverage detected