CoolFace
Apppublic

shrimantasatpati/table-extraction

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
data_module.py82 linesDownload Raw Back to components
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