This commit is contained in:
2023-05-23 15:52:09 +08:00
parent e8f9c8a287
commit 663edbb5e2
88 changed files with 6250 additions and 0 deletions
+40
View File
@@ -0,0 +1,40 @@
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