| 224 | ) |
| 225 | |
| 226 | class STL10LinearProbeDataModule(_BaseDataModule): |
| 227 | num_classes = 10 |
| 228 | def prepare_data(self): |
| 229 | pass |
| 230 | |
| 231 | def transforms(self): |
| 232 | tform = transforms.Compose([ |
| 233 | transforms.Resize(self.size), |
| 234 | transforms.ToTensor(), |
| 235 | transforms.Normalize((0.4467, 0.4398, 0.4066), (0.2603, 0.2566, 0.2713)) |
| 236 | ]) |
| 237 | return tform |
| 238 | |
| 239 | def setup(self, stage=None): |
| 240 | self.data_train = STL10( |
| 241 | root = self.data_dir, |
| 242 | split = 'train', |
| 243 | transform = self.transforms() |
| 244 | ) |
| 245 | |
| 246 | data_val_test = STL10( |
| 247 | root = self.data_dir, |
| 248 | split = 'test', |
| 249 | transform = self.transforms() |
| 250 | ) |
| 251 | test_size = int(0.7 * len(data_val_test)) |
| 252 | val_size = len(data_val_test) - test_size |
| 253 | self.data_val, self.data_test = random_split(data_val_test, [val_size, test_size]) |
| 254 | |
| 255 | def test_dataloader(self): |
| 256 | return DataLoader( |
| 257 | dataset = self.data_test, |
| 258 | batch_size = self.batch_size, |
| 259 | num_workers = self.num_workers, |
| 260 | pin_memory = self.pin_memory, |
| 261 | shuffle = False, |
| 262 | drop_last = True |
| 263 | ) |
| 264 |
nothing calls this directly
no outgoing calls
no test coverage detected