MCPcopy Create free account
hub / github.com/Monalissaa/DisenDiff / DataModuleFromConfig

Class DataModuleFromConfig

train.py:378–463  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

376
377
378class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected