(dataset, root='./data', train=True, fitr=None)
| 84 | |
| 85 | |
| 86 | def get_dataset(dataset, root='./data', train=True, fitr=None): |
| 87 | if dataset == 'imagenet' or dataset == 'imagenet-mini': |
| 88 | return imagenet_utils.get_dataset(dataset, root, train) |
| 89 | |
| 90 | transform = get_transforms(dataset, train=train, is_tensor=False) |
| 91 | lp_fitr = None if fitr is None else get_filter(fitr) |
| 92 | |
| 93 | if dataset == 'cifar10': |
| 94 | target_set = data.datasetCIFAR10(root=root, train=train, transform=transform) |
| 95 | x, y = target_set.data, target_set.targets |
| 96 | elif dataset == 'cifar100': |
| 97 | target_set = data.datasetCIFAR100(root=root, train=train, transform=transform) |
| 98 | x, y = target_set.data, target_set.targets |
| 99 | elif dataset == 'tiny-imagenet': |
| 100 | target_set = data.datasetTinyImageNet(root=root, train=train, transform=transform) |
| 101 | x, y = target_set.x, target_set.y |
| 102 | else: |
| 103 | raise NotImplementedError('dataset {} is not supported'.format(dataset)) |
| 104 | |
| 105 | return data.Dataset(x, y, transform, lp_fitr) |
| 106 | |
| 107 | |
| 108 | def get_indexed_loader(dataset, batch_size, root='./data', train=True): |
no test coverage detected