如何实现满足特定采样规则的TensorFlow Keras数据生成器?
针对VoxCeleb数据集的半随机采样Keras生成器修改方案
问题描述
拥有VoxCeleb超大数据集,每条数据包含1-5000范围内的多分类标签与音频录音,无法一次性全量加载,使用TensorFlow 2.8搭配Keras生成器训练。需要生成器遵循以下半随机采样规则:
- 每个批次共包含
n_s个样本; - 批次由
n_c个随机选取的不同类别组成; - 每个类别在批次中包含
n_p个样本(n_s = n_c * n_p); - 每个epoch需覆盖所有标签,确保每个标签至少被遍历一次。
原通用DataGenerator基类代码如下:
from os import path import numpy as np from keras.utils import Sequence from keras.preprocessing.sequence import pad_sequences from pre_processing import load_data # customize function class DataGenerator(Sequence): """Generates data for Keras Sequence based data generator. Suitable for building data generator for training and prediction. """ def __init__(self, list_IDs, labels, n_classes, input_path, target_path, to_fit=True, batch_size=n_s, shuffle=True): """Initialization :param list_IDs: list of all 'label' ids to use in the generator :param to_fit: True to return X and y, False to return X only :param batch_size: batch size at each iteration :param shuffle: True to shuffle label indexes after every epoch """ self.input_path = input_path self.target_path = target_path self.list_IDs = list_IDs self.labels = labels self.n_classes = n_classes self.to_fit = to_fit self.batch_size = batch_size self.shuffle = shuffle self.on_epoch_end() def __len__(self): """Denotes the number of batches per epoch :return: number of batches per epoch """ return int(np.floor(len(self.list_IDs) / self.batch_size)) def __getitem__(self, index): """Generate one batch of data :param index: index of the batch :return: X and y when fitting. X only when predicting """ # Generate indexes of the batch indexes = self.indexes[index * self.batch_size:(index + 1) * self.batch_size] # Find list of IDs list_IDs_temp = [self.list_IDs[k] for k in indexes] # Generate data X = self._generate_X(list_IDs_temp) if self.to_fit: y = self._generate_y(list_IDs_temp) return [X], y else: return [X] def on_epoch_end(self): """ Updates indexes after each epoch """ self.indexes = np.arange(len(self.list_IDs)) if self.shuffle: np.random.shuffle(self.indexes) def _generate_X(self, list_IDs_temp): """Generates data containing batch_size images :param list_IDs_temp: list of label ids to load :return: batch of images """ # Initialization X = [] # Generate data for i, ID in enumerate(list_IDs_temp): # Store sample temp = self._load_input(self.input_path, ID) X.append(temp) X = pad_sequences(X, value=0, padding='post') return X def _generate_y(self, list_IDs_temp): """Generates data containing batch_size masks :param list_IDs_temp: list of label ids to load :return: batch if masks """ # TODO: modify y = [] # Generate data for i, ID in enumerate(list_IDs_temp): # Store sample y.append(self._load_target(self.target_path, ID)) # y = pad_sequences(y, value=0, padding='post') return y def _load_input(self, input_path, ID): feats = load_data(path.join(input_path, ID)) return feats def _load_target(self, target_path, ID): return self.labels[ID]
修改思路
核心是先按类别对样本进行分组,在每个epoch开始时:
- 打乱所有类别顺序,确保每个epoch的类别遍历顺序不同;
- 对每个类别内的样本进行打乱,保证样本选取的随机性;
- 按
n_c个类别为一组,每组每个类别取n_p个样本组成批次; - 处理剩余的类别和样本,确保每个类别至少被遍历一次。
修改后的完整代码
from os import path import numpy as np from keras.utils import Sequence from keras.preprocessing.sequence import pad_sequences from pre_processing import load_data # customize function class SemiRandomDataGenerator(Sequence): """Generates data for Keras with semi-random sampling for VoxCeleb dataset Follows rules: each batch has n_c classes, each class contributes n_p samples; All classes are covered in each epoch. """ def __init__(self, list_IDs, labels, n_classes, input_path, target_path, to_fit=True, n_c=8, n_p=4, shuffle=True): """Initialization :param list_IDs: list of all sample IDs to use in the generator :param labels: dict mapping sample ID to its label :param n_classes: total number of classes :param input_path: path to input features :param target_path: path to target labels (unused here, kept for compatibility) :param to_fit: True to return X and y, False to return X only :param n_c: number of distinct classes per batch :param n_p: number of samples per class in a batch :param shuffle: True to shuffle classes and samples after every epoch """ self.input_path = input_path self.target_path = target_path self.list_IDs = list_IDs self.labels = labels self.n_classes = n_classes self.to_fit = to_fit self.n_c = n_c self.n_p = n_p self.batch_size = n_c * n_p self.shuffle = shuffle # Group sample IDs by their label self.class_to_samples = {} for sample_id in list_IDs: label = self.labels[sample_id] if label not in self.class_to_samples: self.class_to_samples[label] = [] self.class_to_samples[label].append(sample_id) # Filter out classes with insufficient samples (ensure at least n_p samples per class) self.valid_classes = [cls for cls in self.class_to_samples if len(self.class_to_samples[cls]) >= n_p] if len(self.valid_classes) < n_c: raise ValueError(f"Not enough valid classes (need at least {n_c}, got {len(self.valid_classes)})") self.on_epoch_end() def __len__(self): """Denotes the number of batches per epoch Calculated based on total valid classes grouped into n_c-sized chunks, ensuring all classes are covered. """ class_groups = len(self.valid_classes) // self.n_c remaining_classes = len(self.valid_classes) % self.n_c total_batches = class_groups if remaining_classes > 0: total_batches += 1 return total_batches def __getitem__(self, index): """Generate one batch of data :param index: index of the batch :return: X and y when fitting. X only when predicting """ # Get the group of classes for this batch start_idx = index * self.n_c end_idx = start_idx + self.n_c batch_classes = self.shuffled_classes[start_idx:end_idx] # Collect samples for each class in the batch batch_sample_ids = [] for cls in batch_classes: # Get n_p samples from the shuffled list of this class's samples sample_indices = self.class_sample_indices[cls][index * self.n_p : index * self.n_p + self.n_p] # Handle cases where a class has fewer samples than needed for all batches if len(sample_indices) < self.n_p: # Wrap around to the start of the shuffled sample list remaining = self.n_p - len(sample_indices) sample_indices = np.concatenate([sample_indices, self.class_sample_indices[cls][:remaining]]) batch_sample_ids.extend([self.class_to_samples[cls][i] for i in sample_indices]) # Generate data X = self._generate_X(batch_sample_ids) if self.to_fit: y = self._generate_y(batch_sample_ids) return [X], y else: return [X] def on_epoch_end(self): """Updates class order and sample indices after each epoch""" # Shuffle the list of valid classes if self.shuffle: np.random.shuffle(self.valid_classes) self.shuffled_classes = self.valid_classes.copy() # For each class, shuffle its samples and create a repeatable index array self.class_sample_indices = {} for cls in self.valid_classes: samples = self.class_to_samples[cls] if self.shuffle: np.random.shuffle(samples) # Create a tiled index array to handle repeated sampling across batches total_needed_samples = self.__len__() * self.n_p tile_count = int(np.ceil(total_needed_samples / len(samples))) self.class_sample_indices[cls] = np.tile(np.arange(len(samples)), tile_count) def _generate_X(self, list_IDs_temp): """Generates input features batch""" X = [] for ID in list_IDs_temp: temp = self._load_input(self.input_path, ID) X.append(temp) X = pad_sequences(X, value=0, padding='post') return X def _generate_y(self, list_IDs_temp): """Generates label batch""" y = [self.labels[ID] for ID in list_IDs_temp] return y def _load_input(self, input_path, ID): feats = load_data(path.join(input_path, ID)) return feats def _load_target(self, target_path, ID): return self.labels[ID]
关键修改点说明
- 类别分组:在
__init__中按标签将样本ID分组,过滤掉样本数不足n_p的类别,确保每个选中的类别能提供足够样本。 - Epoch初始化:
on_epoch_end中打乱类别顺序和每个类别内的样本顺序,同时为每个类别生成重复的索引数组,避免因样本数不足导致无法生成足够批次。 - 批次计算:
__len__根据可分组的类别数计算总批次,确保所有类别都能被覆盖。 - 批次构建:
__getitem__中按索引选取对应类别的n_p个样本,当类别样本数不足时循环复用已打乱的样本,保证批次符合规则。
内容的提问来源于stack exchange,提问作者Triceratops
相关产品推荐
相关产品推荐

