IDF / idf /datasets /data_module.py
JK-the-Ko
add IDF codes
e4b5248
Raw
History Blame Contribute Delete
4.87 kB
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