MCPcopy Create free account
hub / github.com/adobe-research/custom-diffusion / DataModuleFromConfig

Class DataModuleFromConfig

train.py:355–440  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

353
354
355class DataModuleFromConfig(pl.LightningDataModule):
356 def __init__(self, batch_size, train=None, train2=None, validation=None, test=None, predict=None,
357 wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False,
358 shuffle_val_dataloader=False):
359 super().__init__()
360 self.batch_size = batch_size
361 self.dataset_configs = dict()
362 self.num_workers = num_workers if num_workers is not None else batch_size * 2
363 self.use_worker_init_fn = use_worker_init_fn
364 if train2 is not None and train2['params']['caption'] != '':
365 self.dataset_configs["train2"] = train2
366 if train is not None:
367 self.dataset_configs["train"] = train
368 self.train_dataloader = self._train_dataloader
369 if validation is not None:
370 self.dataset_configs["validation"] = validation
371 self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader)
372 if test is not None:
373 self.dataset_configs["test"] = test
374 self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader)
375 if predict is not None:
376 self.dataset_configs["predict"] = predict
377 self.predict_dataloader = self._predict_dataloader
378 self.wrap = wrap
379
380 def prepare_data(self):
381 for data_cfg in self.dataset_configs.values():
382 instantiate_from_config(data_cfg)
383
384 def setup(self, stage=None):
385 self.datasets = dict(
386 (k, instantiate_from_config(self.dataset_configs[k]))
387 for k in self.dataset_configs)
388 if self.wrap:
389 for k in self.datasets:
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:
412 init_fn = worker_init_fn

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected