You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何实现满足特定采样规则的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开始时:

  1. 打乱所有类别顺序,确保每个epoch的类别遍历顺序不同;
  2. 对每个类别内的样本进行打乱,保证样本选取的随机性;
  3. 按n_c个类别为一组,每组每个类别取n_p个样本组成批次;
  4. 处理剩余的类别和样本,确保每个类别至少被遍历一次。

修改后的完整代码

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]

关键修改点说明

  1. 类别分组:在__init__中按标签将样本ID分组,过滤掉样本数不足n_p的类别,确保每个选中的类别能提供足够样本。
  2. Epoch初始化:on_epoch_end中打乱类别顺序和每个类别内的样本顺序,同时为每个类别生成重复的索引数组,避免因样本数不足导致无法生成足够批次。
  3. 批次计算:__len__根据可分组的类别数计算总批次,确保所有类别都能被覆盖。
  4. 批次构建:__getitem__中按索引选取对应类别的n_p个样本,当类别样本数不足时循环复用已打乱的样本,保证批次符合规则。

内容的提问来源于stack exchange,提问作者Triceratops

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.20 02:01:29