(dataset, train=True, is_tensor=True)
| 143 | |
| 144 | |
| 145 | def get_transforms(dataset, train=True, is_tensor=True): |
| 146 | assert (dataset == 'imagenet' or dataset == 'imagenet-mini') |
| 147 | if train: |
| 148 | comp1 = [ |
| 149 | transforms.RandomResizedCrop(224), |
| 150 | transforms.RandomHorizontalFlip(), ] |
| 151 | else: |
| 152 | comp1 = [ |
| 153 | transforms.Resize( [256, 256] ), |
| 154 | transforms.CenterCrop(224), ] |
| 155 | |
| 156 | if is_tensor: |
| 157 | comp2 = [ |
| 158 | torchvision.transforms.Normalize((255*0.5, 255*0.5, 255*0.5), (255., 255., 255.))] |
| 159 | else: |
| 160 | comp2 = [ |
| 161 | transforms.ToTensor(), |
| 162 | transforms.Normalize((0.5, 0.5, 0.5), (1., 1., 1.))] |
| 163 | |
| 164 | trans = transforms.Compose( [*comp1, *comp2] ) |
| 165 | |
| 166 | if is_tensor: trans = ElementWiseTransform(trans) |
| 167 | |
| 168 | return trans |
| 169 | |
| 170 | |
| 171 | def get_filter(fitr): |
no test coverage detected