MCPcopy Create free account
hub / github.com/dek924/PerX2CT / DataModuleFromConfig

Class DataModuleFromConfig

main.py:148–189  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

146
147
148class DataModuleFromConfig(pl.LightningDataModule):
149 def __init__(self, batch_size, train=None, validation=None, test=None,
150 wrap=False, num_workers=None):
151 super().__init__()
152 self.batch_size = batch_size
153 self.dataset_configs = dict()
154 self.num_workers = num_workers if num_workers is not None else batch_size*2
155 if train is not None:
156 self.dataset_configs["train"] = train
157 self.train_dataloader = self._train_dataloader
158 if validation is not None:
159 self.dataset_configs["validation"] = validation
160 self.val_dataloader = self._val_dataloader
161 if test is not None:
162 self.dataset_configs["test"] = test
163 self.test_dataloader = self._test_dataloader
164 self.wrap = wrap
165
166 def prepare_data(self):
167 for data_cfg in self.dataset_configs.values():
168 instantiate_from_config(data_cfg)
169
170 def setup(self, stage=None):
171 self.datasets = dict(
172 (k, instantiate_from_config(self.dataset_configs[k]))
173 for k in self.dataset_configs)
174 if self.wrap:
175 for k in self.datasets:
176 self.datasets[k] = WrappedDataset(self.datasets[k])
177
178 def _train_dataloader(self):
179 return DataLoader(self.datasets["train"], batch_size=self.batch_size,
180 num_workers=self.num_workers, shuffle=True, collate_fn=custom_collate)
181
182 def _val_dataloader(self):
183 return DataLoader(self.datasets["validation"],
184 batch_size=self.batch_size,
185 num_workers=self.num_workers, collate_fn=custom_collate)
186
187 def _test_dataloader(self):
188 return DataLoader(self.datasets["test"], batch_size=self.batch_size,
189 num_workers=self.num_workers, collate_fn=custom_collate)
190
191
192class DataModuleTemplateFromConfig(DataModuleFromConfig):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected