CoolFace
Apppublic

SV12/ERA_Session13

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
data.py69 linesDownload Raw Back to utils
1import torchvision2import lightning as L3from torch.utils.data import DataLoader4from utils.transforms import train_transform, test_transform5 6 7class Cifar10SearchDataset(torchvision.datasets.CIFAR10):8    def __init__(self, root="~/data", train=True, download=True, transform=None):9        super().__init__(root=root, train=train, download=download, transform=transform)10 11    def __getitem__(self, index):12        image, label = self.data[index], self.targets[index]13        if self.transform is not None:14            transformed = self.transform(image=image)15            image = transformed["image"]16 17        return image, label18 19 20class CIFARDataModule(L.LightningDataModule):21    def __init__(22        self, data_dir="data", batch_size=512, shuffle=True, num_workers=423    ) -> None:24        super().__init__()25        self.data_dir = data_dir26        self.batch_size = batch_size27        self.shuffle = shuffle28        self.num_workers = num_workers29 30    def prepare_data(self) -> None:31        pass32 33    def setup(self, stage=None):34        self.train_dataset = Cifar10SearchDataset(35            root=self.data_dir, train=True, transform=train_transform36        )37 38        self.val_dataset = Cifar10SearchDataset(39            root=self.data_dir, train=False, transform=test_transform40        )41 42        self.test_dataset = Cifar10SearchDataset(43            root=self.data_dir, train=False, transform=test_transform44        )45 46    def train_dataloader(self):47        return DataLoader(48            dataset=self.train_dataset,49            batch_size=self.batch_size,50            shuffle=self.shuffle,51            num_workers=self.num_workers,52        )53 54    def val_dataloader(self):55        return DataLoader(56            dataset=self.val_dataset,57            batch_size=self.batch_size,58            shuffle=self.shuffle,59            num_workers=self.num_workers,60        )61 62    def test_dataloader(self):63        return DataLoader(64            dataset=self.test_dataset,65            batch_size=self.batch_size,66            shuffle=self.shuffle,67            num_workers=self.num_workers,68        )69