get_data returns data (images) and targets (labels) shape of data: B, H, W, C shape of labels: B,
(self, svhn_extra=True)
| 220 | self.transform = get_transform(mean[name], std[name], crop_size, train) |
| 221 | |
| 222 | def get_data(self, svhn_extra=True): |
| 223 | """ |
| 224 | get_data returns data (images) and targets (labels) |
| 225 | shape of data: B, H, W, C |
| 226 | shape of labels: B, |
| 227 | """ |
| 228 | dset = getattr(torchvision.datasets, self.name.upper()) |
| 229 | if 'CIFAR' in self.name.upper(): |
| 230 | dset = dset(self.data_dir, train=self.train, download=True) |
| 231 | data, targets = dset.data, dset.targets |
| 232 | return data, targets |
| 233 | elif self.name.upper() == 'SVHN': |
| 234 | if self.train: |
| 235 | if svhn_extra: # train+extra |
| 236 | dset_base = dset(self.data_dir, split='train', download=True) |
| 237 | data_b, targets_b = dset_base.data.transpose([0, 2, 3, 1]), dset_base.labels |
| 238 | dset_extra = dset(self.data_dir, split='extra', download=True) |
| 239 | data_e, targets_e = dset_extra.data.transpose([0, 2, 3, 1]), dset_extra.labels |
| 240 | data = np.concatenate([data_b, data_e]) |
| 241 | targets = np.concatenate([targets_b, targets_e]) |
| 242 | del data_b, data_e |
| 243 | del targets_b, targets_e |
| 244 | else: # train_only |
| 245 | dset = dset(self.data_dir, split='train', download=True) |
| 246 | data, targets = dset.data.transpose([0, 2, 3, 1]), dset.labels |
| 247 | else: # test |
| 248 | dset = dset(self.data_dir, split='test', download=True) |
| 249 | data, targets = dset.data.transpose([0, 2, 3, 1]), dset.labels |
| 250 | return data, targets |
| 251 | elif self.name.upper() == 'STL10': |
| 252 | split = 'train' if self.train else 'test' |
| 253 | dset_lb = dset(self.data_dir, split=split, download=True) |
| 254 | dset_ulb = dset(self.data_dir, split='unlabeled', download=True) |
| 255 | data, targets = dset_lb.data.transpose([0, 2, 3, 1]), dset_lb.labels.astype(np.int64) |
| 256 | ulb_data = dset_ulb.data.transpose([0, 2, 3, 1]) |
| 257 | return data, targets, ulb_data |
| 258 | |
| 259 | def get_dset(self, is_ulb=False, |
| 260 | strong_transform=None, onehot=False): |
no outgoing calls
no test coverage detected