MCPcopy Create free account
hub / github.com/M-Nauta/ProtoTree / get_cars

Function get_cars

util/data.py:93–126  ·  view source on GitHub ↗
(augment: bool, train_dir:str, project_dir: str, test_dir:str, img_size = 224)

Source from the content-addressed store, hash-verified

91
92
93def 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

Callers 1

get_dataFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected