如何在Keras Data Generator的on_epoch_end中同步打乱X_train与y_train
解决方法
核心是让打乱后的索引同时映射到图片路径和标签数组,确保两者始终一一对应。你需要补全DataGenerator类中必须的__getitem__和__len__方法,并在获取批次数据时使用打乱后的索引同步提取路径和标签:
import numpy as np from tensorflow import keras from PIL import Image class DataGenerator(keras.utils.Sequence): 'Generates data for Keras' def __init__(self, file_paths, labels, batch_size=32, dim=(240,320), n_channels=3, shuffle=True): self.dim = dim self.batch_size = batch_size self.labels = labels self.file_paths = file_paths self.n_channels = n_channels self.shuffle = shuffle self.on_epoch_end() def __len__(self): 'Denotes the number of batches per epoch' return int(np.floor(len(self.file_paths) / self.batch_size)) def __getitem__(self, index): 'Generate one batch of data' # 获取当前批次的索引范围 indexes = self.indexes[index*self.batch_size:(index+1)*self.batch_size] # 根据索引同步提取图片路径和对应标签 batch_file_paths = [self.file_paths[k] for k in indexes] batch_labels = [self.labels[k] for k in indexes] # 加载并预处理图片 X = self._load_preprocess_images(batch_file_paths) y = np.array(batch_labels) return X, y def on_epoch_end(self): 'Updates indexes after each epoch' self.indexes = np.arange(len(self.file_paths)) if self.shuffle == True: np.random.shuffle(self.indexes) def _load_preprocess_images(self, file_paths): 'Helper function to load and preprocess images' X = np.empty((self.batch_size, *self.dim, self.n_channels)) for i, path in enumerate(file_paths): # 加载图片并调整尺寸 img = Image.open(path).resize(self.dim) # 转换为numpy数组并归一化(根据你的需求调整预处理逻辑) X[i,] = np.array(img) / 255.0 return X
关键说明
__len__:计算每个epoch包含的批次数量,是Sequence类必须实现的方法。__getitem__:核心逻辑在这里,通过indexes(已经打乱的全局索引)同步获取当前批次的图片路径和标签,完全避免了路径和标签不匹配的问题。_load_preprocess_images:封装图片加载和预处理逻辑,可根据你的任务需求(比如灰度图、归一化方式等)调整。
使用方式
初始化生成器时,直接传入拆分好的训练集/验证集路径和标签即可:
# 初始化训练集生成器 train_generator = DataGenerator(X_train, y_train, batch_size=32) # 初始化验证集生成器(验证集通常不需要打乱,可设shuffle=False) val_generator = DataGenerator(X_val, y_val, batch_size=32, shuffle=False) # 模型训练时使用生成器 model.fit(train_generator, validation_data=val_generator, epochs=10)
这样每次epoch结束后,训练集的索引会被打乱,而每个批次的路径和标签始终保持一一对应,不会出现匹配错误。
内容的提问来源于stack exchange,提问作者hels
相关产品推荐
相关产品推荐

