import glob import torch from torch.nn.functional import one_hot class DataGenerator: def __init__(self, data_directory, samples): self.data_directory = data_directory self.samples = samples self.files = set(file_name for file_name, index in samples) self.num_samples = None def get_num_samples(self, batch_size=128, num_classes=19 * 19): if self.num_samples is not None: return self.num_samples else: self.num_samples = 0 for X, y in self._generate(batch_size=batch_size, num_classes=num_classes): self.num_samples += X.shape[0] return self.num_samples def _generate(self, batch_size, num_classes): for zip_file_name in self.files: file_name = zip_file_name.replace('.tar.gz', '') + 'train' base = self.data_directory + '/' + file_name + '_features_*.npy' for feature_file in glob.glob(base): label_file = feature_file.replace('features', 'labels') x = torch.from_numpy(np.load(feature_file)).float() y = torch.from_numpy(np.load(label_file)).long() y = one_hot(y, num_classes=num_classes) while x.shape[0] >= batch_size: x_batch, x = x[:batch_size], x[batch_size:] y_batch, y = y[:batch_size], y[batch_size:] yield x_batch, y_batch def generate(self, batch_size=128, num_classes=19 * 19): while True: for item in self._generate(batch_size, num_classes): yield item