MCPcopy Create free account
hub / github.com/Text-to-Audio/Make-An-Audio / DataModuleFromConfig

Class DataModuleFromConfig

main.py:180–238  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

178
179
180class DataModuleFromConfig(pl.LightningDataModule):# batchloader outputshape should be (b,h,w,c) and it will be permuted to (b,c,h,w) in autoencoder.get_input()
181 def __init__(self, batch_size, train=None, validation=None, test=None, predict=None,
182 wrap=False, num_workers=None, shuffle_test_loader=False, use_worker_init_fn=False,
183 shuffle_val_dataloader=False):
184 super().__init__()
185 self.batch_size = batch_size
186 self.dataset_configs = dict()
187 self.num_workers = num_workers if num_workers is not None else batch_size * 2
188 self.use_worker_init_fn = use_worker_init_fn
189 if train is not None:
190 self.dataset_configs["train"] = train
191 self.train_dataloader = self._train_dataloader
192 if validation is not None:
193 self.dataset_configs["validation"] = validation
194 self.val_dataloader = partial(self._val_dataloader, shuffle=shuffle_val_dataloader)
195 if test is not None:
196 self.dataset_configs["test"] = test
197 self.test_dataloader = partial(self._test_dataloader, shuffle=shuffle_test_loader)
198 if predict is not None:
199 self.dataset_configs["predict"] = predict
200 self.predict_dataloader = self._predict_dataloader
201 self.wrap = wrap
202
203 def prepare_data(self):
204 for data_cfg in self.dataset_configs.values():
205 instantiate_from_config(data_cfg)
206
207 def setup(self, stage=None):
208 self.datasets = dict(
209 (k, instantiate_from_config(self.dataset_configs[k]))
210 for k in self.dataset_configs)
211 if self.wrap:
212 for k in self.datasets:
213 self.datasets[k] = WrappedDataset(self.datasets[k])
214
215 def _train_dataloader(self):
216 init_fn = None
217 return DataLoader(self.datasets["train"], batch_size=self.batch_size ,# sampler=DistributedSampler # np.arange(100),
218 num_workers=self.num_workers, shuffle=True,
219 worker_init_fn=init_fn)
220
221 def _val_dataloader(self, shuffle=False):
222 init_fn = None
223 return DataLoader(self.datasets["validation"],
224 batch_size=self.batch_size,
225 num_workers=self.num_workers,
226 worker_init_fn=init_fn,
227 shuffle=shuffle)
228
229 def _test_dataloader(self, shuffle=False):
230 init_fn = None
231 # do not shuffle dataloader for iterable dataset
232 return DataLoader(self.datasets["test"], batch_size=self.batch_size,
233 num_workers=self.num_workers, worker_init_fn=init_fn, shuffle=shuffle)
234
235 def _predict_dataloader(self, shuffle=False):
236 init_fn = None
237 return DataLoader(self.datasets["predict"], batch_size=self.batch_size,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected