| 91 | |
| 92 | |
| 93 | def get_cars(augment: bool, train_dir:str, project_dir: str, test_dir:str, img_size = 224): |
| 94 | shape = (3, img_size, img_size) |
| 95 | mean = (0.485, 0.456, 0.406) |
| 96 | std = (0.229, 0.224, 0.225) |
| 97 | |
| 98 | normalize = transforms.Normalize(mean=mean,std=std) |
| 99 | transform_no_augment = transforms.Compose([ |
| 100 | transforms.Resize(size=(img_size, img_size)), |
| 101 | transforms.ToTensor(), |
| 102 | normalize |
| 103 | ]) |
| 104 | |
| 105 | if augment: |
| 106 | transform = transforms.Compose([ |
| 107 | transforms.Resize(size=(img_size+32, img_size+32)), #resize to 256x256 |
| 108 | transforms.RandomOrder([ |
| 109 | transforms.RandomPerspective(distortion_scale=0.5, p = 0.5), |
| 110 | transforms.ColorJitter((0.6,1.4), (0.6,1.4), (0.6,1.4), (-0.4,0.4)), |
| 111 | transforms.RandomHorizontalFlip(), |
| 112 | transforms.RandomAffine(degrees=15,shear=(-2,2)), |
| 113 | ]), |
| 114 | transforms.RandomCrop(size=(img_size, img_size)), #crop to 224x224 |
| 115 | transforms.ToTensor(), |
| 116 | normalize, |
| 117 | ]) |
| 118 | else: |
| 119 | transform = transform_no_augment |
| 120 | |
| 121 | trainset = torchvision.datasets.ImageFolder(train_dir, transform=transform) |
| 122 | projectset = torchvision.datasets.ImageFolder(project_dir, transform=transform_no_augment) |
| 123 | testset = torchvision.datasets.ImageFolder(test_dir, transform=transform_no_augment) |
| 124 | classes = trainset.classes |
| 125 | |
| 126 | return trainset, projectset, testset, classes, shape |
| 127 | |
| 128 | |