shrimantasatpati/table-extraction
0
1from typing import Callable, List, Optional, Union2 3import torch4from pytorch_lightning import LightningDataModule5from torch.utils.data import DataLoader, Dataset6 7 8class SampleDataset(Dataset):9 10 def __init__(self,11 x: Union[List, torch.Tensor],12 y: Union[List, torch.Tensor],13 transforms: Optional[Callable] = None) -> None:14 super(SampleDataset, self).__init__()15 self.x = x16 self.y = y17 18 if transforms is None:19 # Replace None with some default transforms20 # If image, could be an Resize and ToTensor21 self.transforms = lambda x: x22 else:23 self.transforms = transforms24 25 def __len__(self):26 return len(self.x)27 28 def __getitem__(self, index: int):29 x = self.x[index]30 y = self.y[index]31 32 x = self.transforms(x)33 return x, y34 35 36class SampleDataModule(LightningDataModule):37 38 def __init__(self,39 x: Union[List, torch.Tensor],40 y: Union[List, torch.Tensor],41 transforms: Optional[Callable] = None,42 val_ratio: float = 0,43 batch_size: int = 32) -> None:44 super(SampleDataModule, self).__init__()45 assert 0 <= val_ratio < 146 assert isinstance(batch_size, int)47 self.x = x48 self.y = y49 50 self.transforms = transforms51 self.val_ratio = val_ratio52 self.batch_size = batch_size53 54 self.setup()55 self.prepare_data()56 57 def setup(self, stage: Optional[str] = None) -> None:58 pass59 60 def prepare_data(self) -> None:61 n_samples: int = len(self.x)62 train_size: int = n_samples - int(n_samples * self.val_ratio)63 64 self.train_dataset = SampleDataset(x=self.x[:train_size],65 y=self.y[:train_size],66 transforms=self.transforms)67 if train_size < n_samples:68 self.val_dataset = SampleDataset(x=self.x[train_size:],69 y=self.y[train_size:],70 transforms=self.transforms)71 else:72 self.val_dataset = SampleDataset(x=self.x[-self.batch_size:],73 y=self.y[-self.batch_size:],74 transforms=self.transforms)75 76 def train_dataloader(self) -> DataLoader:77 return DataLoader(dataset=self.train_dataset,78 batch_size=self.batch_size)79 80 def val_dataloader(self) -> DataLoader:81 return DataLoader(dataset=self.val_dataset, batch_size=self.batch_size)82 