| 26 | return get_tinyImageNet(path) |
| 27 | |
| 28 | def 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 | |
| 52 | def get_tinyImageNet(path): |