SV12/ERA_Session13
0
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 