| 38 | |
| 39 | |
| 40 | def get_transforms(dataset, train=True, is_tensor=True): |
| 41 | if dataset == 'imagenet' or dataset == 'imagenet-mini': |
| 42 | return imagenet_utils.get_transforms(dataset, train, is_tensor) |
| 43 | |
| 44 | if train: |
| 45 | if dataset == 'cifar10' or dataset == 'cifar100': |
| 46 | comp1 = [ |
| 47 | transforms.RandomHorizontalFlip(), |
| 48 | transforms.RandomCrop(32, 4), ] |
| 49 | elif dataset == 'tiny-imagenet': |
| 50 | comp1 = [ |
| 51 | transforms.RandomHorizontalFlip(), |
| 52 | transforms.RandomCrop(64, 8), ] |
| 53 | else: |
| 54 | raise NotImplementedError |
| 55 | else: |
| 56 | comp1 = [] |
| 57 | |
| 58 | if is_tensor: |
| 59 | comp2 = [ |
| 60 | torchvision.transforms.Normalize((255*0.5, 255*0.5, 255*0.5), (255., 255., 255.))] |
| 61 | else: |
| 62 | comp2 = [ |
| 63 | transforms.ToTensor(), |
| 64 | transforms.Normalize((0.5, 0.5, 0.5), (1., 1., 1.))] |
| 65 | |
| 66 | trans = transforms.Compose( [*comp1, *comp2] ) |
| 67 | |
| 68 | if is_tensor: trans = data.ElementWiseTransform(trans) |
| 69 | |
| 70 | return trans |
| 71 | |
| 72 | |
| 73 | def get_filter(fitr): |