MCPcopy Create free account
hub / github.com/cure-lab/deep-active-learning / get_ImageNet

Function get_ImageNet

dataset.py:28–49  ·  view source on GitHub ↗
(path)

Source from the content-addressed store, hash-verified

26 return get_tinyImageNet(path)
27
28def get_ImageNet(path):
29 raw_tr = datasets.ImageFolder(path + '/tinyImageNet/tiny-imagenet-200/train')
30 imagenet_tr_path = path +'imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/'
31 from torchvision import transforms
32 transform = transforms.Compose([transforms.Resize((64, 64))])
33 imagenet_folder = datasets.ImageFolder(imagenet_tr_path, transform=transform)
34 idx_to_class = {}
35 for (class_num, idx) in imagenet_folder.class_to_idx.items():
36 idx_to_class[idx] = class_num
37 X_tr,Y_tr = [], []
38 item_list = imagenet_folder.imgs
39 for (class_num, idx) in raw_tr.class_to_idx.items():
40 new_img_num = 0
41 for ii, (path, target) in enumerate(item_list):
42 if idx_to_class[target] == class_num:
43 X_tr.append(np.array(imagenet_folder[ii][0]))
44 Y_tr.append(idx)
45 new_img_num += 1
46 if new_img_num >= 250:
47 break
48
49 return np.array(X_tr), np.array(Y_tr)
50
51
52def get_tinyImageNet(path):

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected