Repository navigation
Expand file tree
/
Copy pathdata_sampler.py
More file actions
101 lines (84 loc) · 3.49 KB
/
Copy pathdata_sampler.py
File metadata and controls
101 lines (84 loc) · 3.49 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
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
import os
import torch
import torchvision.transforms as transforms
import torchvision.datasets as datasets
from torch.utils.data.sampler import Sampler
import numpy as np
import distiller
data_transforms = {
'train': transforms.Compose([
transforms.RandomResizedCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
]),
'test': transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
]),
}
def __image_size(dataset):
# un-squeeze is used here to add the batch dimension (value=1), which is missing
return dataset[0][0].unsqueeze(0).size()
def __deterministic_worker_init_fn(worker_id, seed=0):
import random
import numpy
random.seed(seed)
numpy.random.seed(seed)
torch.manual_seed(seed)
def __split_list(l, ratio):
split_idx = int(np.floor(ratio * len(l)))
return l[:split_idx], l[split_idx:]
def _get_subset_length(data_source, effective_size):
if effective_size <= 0 or effective_size > 1:
raise ValueError('effective_size must be in (0..1]')
return int(np.floor(len(data_source) * effective_size))
def _get_sampler(data_source, effective_size, fixed_subset=False, sequential=False):
if fixed_subset:
subset_length = _get_subset_length(data_source, effective_size)
indices = np.random.permutation(len(data_source))
subset_indices = indices[:subset_length]
if sequential:
return SubsetSequentialSampler(subset_indices)
else:
return torch.utils.data.SubsetRandomSampler(subset_indices)
return SwitchingSubsetRandomSampler(data_source, effective_size)
class SwitchingSubsetRandomSampler(Sampler):
"""Samples a random subset of elements from a data source, without replacement.
The subset of elements is re-chosen randomly each time the sampler is enumerated
Args:
data_source (Dataset): dataset to sample from
subset_size (float): value in (0..1], representing the portion of dataset to sample at each enumeration.
"""
def __init__(self, data_source, effective_size):
self.data_source = data_source
self.subset_length = _get_subset_length(data_source, effective_size)
def __iter__(self):
# Randomizing in the same way as in torch.utils.data.sampler.SubsetRandomSampler to maintain
# reproducibility with the previous data loaders implementation
indices = torch.randperm(len(self.data_source))
subset_indices = indices[:self.subset_length]
return (self.data_source[i] for i in subset_indices)
def __len__(self):
return self.subset_length
class SubsetSequentialSampler(torch.utils.data.Sampler):
"""Sequentially samples a subset of the dataset, without replacement.
Arguments:
indices (sequence): a sequence of indices
"""
def __init__(self, indices):
self.indices = indices
def __iter__(self):
return (self.indices[i] for i in range(len(self.indices)))
def __len__(self):
return len(self.indices)
if __name__ = "__main__":
"""
train_dataset = datasets.ImageFolder('/home/swai01/imagenet_datasets/raw-data/train/',
train_transform=data_transforms['train'])
num_train = len(train_dataset)
indices = list(range(num_train))
"""
pass