Spaces:
Sleeping
Sleeping
| 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 |