(self)
| 390 | self.datasets[k] = WrappedDataset(self.datasets[k]) |
| 391 | |
| 392 | def _train_dataloader(self): |
| 393 | is_iterable_dataset = isinstance(self.datasets['train'], Txt2ImgIterableBaseDataset) |
| 394 | if is_iterable_dataset or self.use_worker_init_fn: |
| 395 | init_fn = worker_init_fn |
| 396 | else: |
| 397 | init_fn = None |
| 398 | if "train2" in self.dataset_configs and self.dataset_configs["train2"]['params']["caption"] != '': |
| 399 | train_set = self.datasets["train"] |
| 400 | train2_set = self.datasets["train2"] |
| 401 | concat_dataset = ConcatDataset(train_set, train2_set) |
| 402 | return DataLoader(concat_dataset, batch_size=self.batch_size // 2, |
| 403 | num_workers=self.num_workers, shuffle=False if is_iterable_dataset else True, |
| 404 | worker_init_fn=init_fn) |
| 405 | else: |
| 406 | return DataLoader(self.datasets["train"], batch_size=self.batch_size, |
| 407 | num_workers=self.num_workers, shuffle=False if is_iterable_dataset else True, |
| 408 | worker_init_fn=init_fn) |
| 409 | |
| 410 | def _val_dataloader(self, shuffle=False): |
| 411 | if isinstance(self.datasets['validation'], Txt2ImgIterableBaseDataset) or self.use_worker_init_fn: |
nothing calls this directly
no test coverage detected