File size: 5,367 Bytes
4947683 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 | import random
from typing import Optional, Sequence
from pathlib import Path
import hydra
import numpy as np
import omegaconf
import pytorch_lightning as pl
import torch
from omegaconf import DictConfig
from torch.utils.data import Dataset
from torch_geometric.loader import DataLoader
from diffcsp.common.utils import PROJECT_ROOT
from diffcsp.common.data_utils import get_scaler_from_data_list
def worker_init_fn(id: int):
"""
DataLoaders workers init function.
Initialize the numpy.random seed correctly for each worker, so that
random augmentations between workers and/or epochs are not identical.
If a global seed is set, the augmentations are deterministic.
https://pytorch.org/docs/stable/notes/randomness.html#dataloader
"""
uint64_seed = torch.initial_seed()
ss = np.random.SeedSequence([uint64_seed])
# More than 128 bits (4 32-bit words) would be overkill.
np.random.seed(ss.generate_state(4))
random.seed(uint64_seed)
class CrystDataModule(pl.LightningDataModule):
def __init__(
self,
datasets: DictConfig,
num_workers: DictConfig,
batch_size: DictConfig,
scaler_path=None,
):
super().__init__()
self.datasets = datasets
self.num_workers = num_workers
self.batch_size = batch_size
self.train_dataset: Optional[Dataset] = None
self.val_datasets: Optional[Sequence[Dataset]] = None
self.test_datasets: Optional[Sequence[Dataset]] = None
self.get_scaler(scaler_path)
def prepare_data(self) -> None:
# download only
pass
def get_scaler(self, scaler_path):
# Load once to compute property scaler
if scaler_path is None:
train_dataset = hydra.utils.instantiate(self.datasets.train)
self.lattice_scaler = get_scaler_from_data_list(
train_dataset.cached_data,
key='scaled_lattice')
self.scaler = get_scaler_from_data_list(
train_dataset.cached_data,
key=train_dataset.prop)
else:
try:
self.lattice_scaler = torch.load(
Path(scaler_path) / 'lattice_scaler.pt')
self.scaler = torch.load(Path(scaler_path) / 'prop_scaler.pt')
except:
train_dataset = hydra.utils.instantiate(self.datasets.train)
self.lattice_scaler = get_scaler_from_data_list(
train_dataset.cached_data,
key='scaled_lattice')
self.scaler = get_scaler_from_data_list(
train_dataset.cached_data,
key=train_dataset.prop)
def setup(self, stage: Optional[str] = None):
"""
construct datasets and assign data scalers.
"""
if stage is None or stage == "fit":
self.train_dataset = hydra.utils.instantiate(self.datasets.train)
self.val_datasets = [
hydra.utils.instantiate(dataset_cfg)
for dataset_cfg in self.datasets.val
]
self.train_dataset.lattice_scaler = self.lattice_scaler
self.train_dataset.scaler = self.scaler
for val_dataset in self.val_datasets:
val_dataset.lattice_scaler = self.lattice_scaler
val_dataset.scaler = self.scaler
if stage is None or stage == "test":
self.test_datasets = [
hydra.utils.instantiate(dataset_cfg)
for dataset_cfg in self.datasets.test
]
for test_dataset in self.test_datasets:
test_dataset.lattice_scaler = self.lattice_scaler
test_dataset.scaler = self.scaler
def train_dataloader(self, shuffle = True) -> DataLoader:
return DataLoader(
self.train_dataset,
shuffle=shuffle,
batch_size=self.batch_size.train,
num_workers=self.num_workers.train,
worker_init_fn=worker_init_fn,
)
def val_dataloader(self) -> Sequence[DataLoader]:
return [
DataLoader(
dataset,
shuffle=False,
batch_size=self.batch_size.val,
num_workers=self.num_workers.val,
worker_init_fn=worker_init_fn,
)
for dataset in self.val_datasets
]
def test_dataloader(self) -> Sequence[DataLoader]:
return [
DataLoader(
dataset,
shuffle=False,
batch_size=self.batch_size.test,
num_workers=self.num_workers.test,
worker_init_fn=worker_init_fn,
)
for dataset in self.test_datasets
]
def __repr__(self) -> str:
return (
f"{self.__class__.__name__}("
f"{self.datasets=}, "
f"{self.num_workers=}, "
f"{self.batch_size=})"
)
@hydra.main(config_path=str(PROJECT_ROOT / "conf"), config_name="default", version_base="1.1")
def main(cfg: omegaconf.DictConfig):
datamodule: pl.LightningDataModule = hydra.utils.instantiate(
cfg.data.datamodule, _recursive_=False
)
datamodule.setup('fit')
import pdb
pdb.set_trace()
if __name__ == "__main__":
main()
|