| 13 | |
| 14 | |
| 15 | class DataModule(LightningDataModule): |
| 16 | def __init__(self, hparams): |
| 17 | super(DataModule, self).__init__() |
| 18 | self.hparams.update(hparams.__dict__) if hasattr(hparams, "__dict__") else self.hparams.update(hparams) |
| 19 | self._mean, self._std = None, None |
| 20 | self._saved_dataloaders = dict() |
| 21 | self.dataset = None |
| 22 | |
| 23 | def prepare_dataset(self): |
| 24 | |
| 25 | assert hasattr(self, f"_prepare_{self.hparams['dataset']}_dataset"), f"Dataset {self.hparams['dataset']} not defined" |
| 26 | dataset_factory = lambda t: getattr(self, f"_prepare_{t}_dataset")() |
| 27 | self.idx_train, self.idx_val, self.idx_test = dataset_factory(self.hparams["dataset"]) |
| 28 | |
| 29 | print(f"train {len(self.idx_train)}, val {len(self.idx_val)}, test {len(self.idx_test)}") |
| 30 | self.train_dataset = Subset(self.dataset, self.idx_train) |
| 31 | self.val_dataset = Subset(self.dataset, self.idx_val) |
| 32 | self.test_dataset = Subset(self.dataset, self.idx_test) |
| 33 | |
| 34 | if self.hparams["standardize"]: |
| 35 | self._standardize() |
| 36 | |
| 37 | def train_dataloader(self): |
| 38 | return self._get_dataloader(self.train_dataset, "train") |
| 39 | |
| 40 | def val_dataloader(self): |
| 41 | loaders = [self._get_dataloader(self.val_dataset, "val")] |
| 42 | delta = 1 if self.hparams['reload'] == 1 else 2 |
| 43 | if ( |
| 44 | len(self.test_dataset) > 0 |
| 45 | and (self.trainer.current_epoch + delta) % self.hparams["test_interval"] == 0 |
| 46 | ): |
| 47 | loaders.append(self._get_dataloader(self.test_dataset, "test")) |
| 48 | return loaders |
| 49 | |
| 50 | def test_dataloader(self): |
| 51 | return self._get_dataloader(self.test_dataset, "test") |
| 52 | |
| 53 | @property |
| 54 | def atomref(self): |
| 55 | if hasattr(self.dataset, "get_atomref"): |
| 56 | return self.dataset.get_atomref() |
| 57 | return None |
| 58 | |
| 59 | @property |
| 60 | def mean(self): |
| 61 | return self._mean |
| 62 | |
| 63 | @property |
| 64 | def std(self): |
| 65 | return self._std |
| 66 | |
| 67 | def _get_dataloader(self, dataset, stage, store_dataloader=True): |
| 68 | store_dataloader = (store_dataloader and not self.hparams["reload"]) |
| 69 | if stage in self._saved_dataloaders and store_dataloader: |
| 70 | return self._saved_dataloaders[stage] |
| 71 | |
| 72 | if stage == "train": |