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

Method __init__

train.py:356–378  ·  view source on GitHub ↗
(self, batch_size, train=None, train2=None, validation=None, test=None, predict=None,
                 wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False,
                 shuffle_val_dataloader=False)

Source from the content-addressed store, hash-verified

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():

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected