File size: 4,874 Bytes
e4b5248
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import Any, Tuple, Mapping

from pytorch_lightning.utilities.types import EVAL_DATALOADERS, TRAIN_DATALOADERS
import pytorch_lightning as pl
from torch.utils.data import DataLoader, Dataset
from omegaconf import OmegaConf, DictConfig, ListConfig

from idf.utils.common import instantiate_from_config
from idf.datasets.batch_transform import BatchTransform, IdentityBatchTransform
from idf.datasets.utils import RepeatDataset


class BIRDataModule(pl.LightningDataModule):
    def __init__(
        self,
        train_config: str,
        val_config: str=None
    ) -> "BIRDataModule":
        super().__init__()
        self.train_config = OmegaConf.load(train_config)
        self.val_config = OmegaConf.load(val_config) if val_config else None

    def load_dataset(self, config: Mapping[str, Any]) -> Tuple[Dataset, BatchTransform]:
        dataset = instantiate_from_config(config["dataset"])
        batch_transform = (
            instantiate_from_config(config["batch_transform"])
            if config.get("batch_transform") else IdentityBatchTransform()
        )
        return dataset, batch_transform

    def setup(self, stage: str) -> None:
        if stage == "fit":
            self.train_dataset, self.train_batch_transform = self.load_dataset(self.train_config)
            if self.val_config:
                self.val_dataset, self.val_batch_transform = self.load_dataset(self.val_config)
            else:
                self.val_dataset, self.val_batch_transform = None, None
        else:
            raise NotImplementedError(stage)

    def train_dataloader(self) -> TRAIN_DATALOADERS:
        return DataLoader(
            dataset=self.train_dataset, **self.train_config["data_loader"]
        )

    def val_dataloader(self) -> EVAL_DATALOADERS:
        if self.val_dataset is None:
            return None
        return DataLoader(
            dataset=self.val_dataset, **self.val_config["data_loader"]
        )

    def on_after_batch_transfer(self, batch: Any, dataloader_idx: int) -> Any:
        self.trainer: pl.Trainer
        
        if self.trainer.training:
            return self.train_batch_transform(batch)
        elif self.trainer.validating or self.trainer.sanity_checking:
            return self.val_batch_transform(batch)
        else:
            raise RuntimeError(
                "Trainer state: \n"
                f"training: {self.trainer.training}\n"
                f"validating: {self.trainer.validating}\n"
                f"testing: {self.trainer.testing}\n"
                f"predicting: {self.trainer.predicting}\n"
                f"sanity_checking: {self.trainer.sanity_checking}"
            )

class BaseDataModule(pl.LightningDataModule):
    def __init__(self, train_config: str, val_config: str|ListConfig=None):
        super().__init__()

        self.train_config = OmegaConf.load(train_config)

        if type(val_config) == str:
            self.val_config = OmegaConf.load(val_config)
        elif type(val_config) == ListConfig:
            self.val_config = [OmegaConf.load(vc) for vc in val_config]
        else:
            self.val_config = None
            
    def setup(self, stage: str) -> None:
        if stage == "fit":
            self.train_dataset = instantiate_from_config(self.train_config["dataset"])
            
            repeat_rate = self.train_config["repeat_dataset"]["times"]
            if repeat_rate > 1:
                self.train_dataset = RepeatDataset(self.train_dataset, repeat_rate)

            if type(self.val_config) == DictConfig:
                self.val_dataset = instantiate_from_config(self.val_config["dataset"])
            elif type(self.val_config) == list:
                self.val_dataset = [instantiate_from_config(vc["dataset"]) for vc in self.val_config]
            else:
                self.val_dataset = None
        elif stage == "validate":
            if type(self.val_config) == DictConfig:
                self.val_dataset = instantiate_from_config(self.val_config["dataset"])
            elif type(self.val_config) == list:
                self.val_dataset = [instantiate_from_config(vc["dataset"]) for vc in self.val_config]
            else:
                self.val_dataset = None
        else:
            raise NotImplementedError(stage)

    def train_dataloader(self) -> TRAIN_DATALOADERS:
        return DataLoader(
            dataset=self.train_dataset, **self.train_config["data_loader"]
        )

    def val_dataloader(self) -> EVAL_DATALOADERS:
        if type(self.val_config) == DictConfig:
            return DataLoader(dataset=self.val_dataset, **self.val_config["data_loader"])
        elif type(self.val_config) == list:
            return [DataLoader(dataset=vd, **vc["data_loader"]) 
                    for vd, vc in zip(self.val_dataset, self.val_config)]
        else:
            return None