-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdatautils.py
More file actions
83 lines (66 loc) · 2.85 KB
/
Copy pathdatautils.py
File metadata and controls
83 lines (66 loc) · 2.85 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
import torch
import torchvision
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
import numpy as np
from PIL import Image
import random
import os
import scipy.io as sio
import glob
import re
import pickle
def seed_worker(worker_id):
worker_seed = torch.initial_seed() % 2**32
np.random.seed(worker_seed)
random.seed(worker_seed)
g = torch.Generator()
g.manual_seed(0)
def parse_imagenet_val_labels(data_dir):
meta_path = os.path.join(data_dir, 'meta.mat')
meta = sio.loadmat(meta_path, squeeze_me=True)['synsets']
nums_children = list(zip(*meta))[4]
meta = [meta[idx] for idx, num_children in enumerate(nums_children)
if num_children == 0]
idcs, wnids = list(zip(*meta))[:2]
idx_to_wnid = {idx: wnid for idx, wnid in zip(idcs, wnids)}
val_path = os.path.join(data_dir, 'ILSVRC2012_validation_ground_truth.txt')
val_idcs = np.loadtxt(val_path)
val_wnids = [idx_to_wnid[idx] for idx in val_idcs]
label_path = os.path.join(data_dir, 'wnid_to_label.pickle')
with open(label_path, 'rb') as f:
wnid_to_label = pickle.load(f)
val_labels = [wnid_to_label[wnid] for wnid in val_wnids]
return np.array(val_labels)
class Imagenet(Dataset):
"""
Validation dataset of Imagenet
"""
def __init__(self, data_dir, transform):
self.Y = torch.from_numpy(parse_imagenet_val_labels(data_dir)).long()
self.X_path = sorted(glob.glob(os.path.join(data_dir, 'ILSVRC2012_img_val/*.JPEG')),
key=lambda x: re.search('%s(.*)%s' % ('ILSVRC2012_img_val/', '.JPEG'), x).group(1))
self.transform = transform
def __len__(self):
return len(self.X_path)
def __getitem__(self, idx):
img = Image.open(self.X_path[idx]).convert('RGB')
y = self.Y[idx]
if self.transform:
x = self.transform(img)
return x, y
def data_loader(ds_path, batch_size, train_transform, test_transform, num_workers=8):
data_dir = ds_path
if not os.path.isdir(data_dir):
raise Exception('Please download Imagenet2012 dataset!')
train_ds = torchvision.datasets.ImageFolder(os.path.join(data_dir, 'ILSVRC2012_img_train'),
transform=train_transform)
if not os.path.isfile(os.path.join(data_dir, 'wnid_to_label.pickle')):
with open(os.path.join(data_dir, 'wnid_to_label.pickle'), 'wb') as f:
pickle.dump(train_ds.class_to_idx, f)
test_ds = Imagenet(data_dir, test_transform)
train_dl = DataLoader(train_ds, batch_size, shuffle=True, num_workers=num_workers,
worker_init_fn=seed_worker, generator=g)
test_dl = DataLoader(test_ds, min(batch_size, 1024), shuffle=False,
num_workers=num_workers)
return train_dl, test_dl