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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 23:47:24