English
OneScience
protein structure generation
File size: 8,080 Bytes
8e04e6f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
from abc import ABC, abstractmethod
from typing import Callable, Dict, Iterable, List, Literal, Optional

import lightning as L
from loguru import logger
from torch_geometric import transforms as T
from torch_geometric.data import Dataset
from torch_geometric.loader import DataLoader

from models.utils.cluster_utils import ClusterSampler
from models.utils.dense_padding_data_loader import DensePaddingDataLoader


class BaseLightningDataModule(L.LightningDataModule, ABC):
    """Base class for all datamodules"""

    def __init__(
        self,
        batch_padding: bool = True,
        sampling_mode: Literal["random", "cluster-random", "cluster-reps"] = "random",
        transforms: Optional[List[Callable]] = None,
        pre_transforms: Optional[List[Callable]] = None,
        pre_filters: Optional[List[Callable]] = None,
        batch_size: int = 32,
        num_workers: int = 32,
        pin_memory: bool = False,
    ):
        """Initialising the base data module class.

        Args:
            batch_padding (bool, optional): Whether batches should be padded to a dense representation
                with the length being either a pre-specified max length or the maximum length of the
                sample in the batch (base PyTorch batch) or whether a sparse representation should be
                used (PyG batch). Defaults to True (base PyTorch batch).
            sampling_mode (Literal["random", "cluster-random", "cluster-reps"], optional): How the data should be
                sampled from the dataset later on:
                - "random": Select a random sequence and ignore clusters.
                - "cluster-random": Select a random sequence from each cluster. Keep all samples for each cluster.
                - "cluster-reps": Select the cluster representative from each cluster. Only keep the representative for each cluster.
                  Defaults to "random".
            transforms (List[Callable]): List of transforms applied to each example.
            pre_transforms (List[Callable]): List of transforms applied to each example before processing.
            pre_filters (List[Callable]): List of filters applied to each example before processing.
            batch_size (int, optional): Batch size used for dataloaders. Defaults to 32.
            num_workers (int, optional): Number of workers used for dataloading. Defaults to 32.
            pin_memory (bool, optional): Whether memory should be pinned. Defaults to False.
        """
        super().__init__()
        self.batch_padding = batch_padding
        self.sampling_mode = sampling_mode
        self.transform = (
            self._compose_transforms(transforms) if transforms is not None else None
        )
        self.pre_transform = (
            self._compose_transforms(pre_transforms)
            if pre_transforms is not None
            else None
        )
        self.pre_filter = (
            self._compose_filters(pre_filters) if pre_filters is not None else None
        )
        self.batch_size = batch_size
        self.num_workers = num_workers
        self.pin_memory = pin_memory
        self.train_ds = None
        self.val_ds = None
        self.test_ds = None
        self.clusterid_to_seqid_mappings = None  # for cluster sampling

    def setup(self, stage: Optional[str] = None):
        if stage == "fit" or stage is None:
            self.train_ds = self.train_dataset()
        elif stage == "validation":
            self.val_ds = self.val_dataset()
        elif stage == "test":
            self.test_ds = self.test_dataset()

    def _compose_transforms(self, transforms: Iterable[Callable]) -> T.Compose:
        try:
            return T.Compose(list(transforms.values()))
        except Exception:
            return T.Compose(transforms)

    def _compose_filters(self, filters: Iterable[Callable]) -> T.ComposeFilters:
        try:
            return T.ComposeFilters(list(filters.values()))
        except Exception:
            return T.ComposeFilters(filters)

    @abstractmethod
    def _get_dataset(self, split: str) -> Dataset:
        """Creates a dataset given a split.

        Args:
            split (str): Split for which to get the dataset, with options "train", "val" or "test"

        Returns:
            Dataset: Dataset created for the respective split
        """
        ...

    def train_dataset(self) -> Dataset:
        return self._get_dataset("train")

    def val_dataset(self) -> Dataset:
        return self._get_dataset("val")

    def test_dataset(self) -> Dataset:
        return self._get_dataset("test")

    def _get_dataloader(
        self,
        dataset: Dataset,
        shuffle: bool = False,
        clusterid_to_seqid_mapping: Dict[str, List[str]] = None,
    ) -> DataLoader:
        """Returns the dataloader for the corresponding dataset.

        Args:
            dataset (Dataset): PyG dataset for which the dataloader will be created.
            shuffle (bool, optional): Whether the dataloader should be shuffled. Defaults to False. False when cluster_id mapping is given.
            clusterid_to_seqid_mapping (Dict[str, List[str]], optional): Maps cluster ids to sequence ids. Defaults to None.

        Returns:
            DataLoader: Dataloader to be used by model.
        """
        if self.sampling_mode is None:
            raise ValueError(
                "Sampling mode not set, should be one of 'random', 'cluster-random' or 'cluster-reps'"
            )
        if clusterid_to_seqid_mapping and self.sampling_mode != "random":
            sampler = ClusterSampler(
                dataset=dataset,
                clusterid_to_seqid_mapping=clusterid_to_seqid_mapping,
                sampling_mode=self.sampling_mode,
            )
            shuffle = False
        elif self.sampling_mode == "random":
            sampler = None
            shuffle = shuffle
        else:
            raise ValueError(
                f"Sampling mode is {self.sampling_mode}, but clusterid_to_seqid_mapping is {clusterid_to_seqid_mapping}"
            )

        dataloader_class = DensePaddingDataLoader if self.batch_padding else DataLoader

        return dataloader_class(
            dataset,
            batch_size=self.batch_size,
            sampler=sampler,
            shuffle=shuffle,
            num_workers=self.num_workers,
            pin_memory=self.pin_memory,
            drop_last=True,
        )

    def train_dataloader(self) -> DataLoader:
        if self.train_ds is None:
            self.train_ds = self.train_dataset()
        clusterid_to_seqid_mapping = (
            self.clusterid_to_seqid_mappings["train"]
            if self.clusterid_to_seqid_mappings
            else None
        )
        shuffle = True
        train_dl = self._get_dataloader(
            dataset=self.train_ds,
            shuffle=shuffle,
            clusterid_to_seqid_mapping=clusterid_to_seqid_mapping,
        )
        return train_dl

    def val_dataloader(self) -> DataLoader:
        if self.val_ds is None:
            self.val_ds = self.val_dataset()
        clusterid_to_seqid_mapping = (
            self.clusterid_to_seqid_mappings["val"]
            if self.clusterid_to_seqid_mappings
            else None
        )
        shuffle = False
        logger.info(f"Length of validation set: {len(self.val_ds)}")
        val_dl = self._get_dataloader(
            dataset=self.val_ds,
            shuffle=shuffle,
            clusterid_to_seqid_mapping=clusterid_to_seqid_mapping,
        )
        return val_dl

    def test_dataloader(self) -> DataLoader:
        if self.test_ds is None:
            self.test_ds = self.test_dataset()
        clusterid_to_seqid_mapping = (
            self.clusterid_to_seqid_mappings["test"]
            if self.clusterid_to_seqid_mappings
            else None
        )
        shuffle = False
        test_dl = self._get_dataloader(
            dataset=self.test_ds,
            shuffle=shuffle,
            clusterid_to_seqid_mapping=clusterid_to_seqid_mapping,
        )
        return test_dl