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

Keras自定义Sequence类DataGenerator批次乱序、调用次数超量问题求助

问题根因
  • 多worker预取机制导致调用次数超标、顺序混乱:你配置了workers=6,即使use_multiprocessing=False,Keras也会启动多线程预取批次,会提前调用__getitem__加载后续批次,导致调用次数高于预期,且多线程调度下批次返回顺序不固定。
  • Sequence实现存在多处错误:
    1. __init__中调用on_epoch_end()但未初始化self.indexes属性,存在隐患
    2. __init__已经将整个CSV读入self.lines,但__getitem__中又重复调用pandas.read_csv读文件,冗余IO且容易出现行偏移错误
    3. __getitem__调用__data_generation时传入的参数和方法定义不匹配(未定义的self.folder_name、缺少config_file参数)
    4. 未用到self.shuffle对应的索引逻辑,即使后续开启shuffle也不会生效
  • 调用了废弃API:model.fit_generator在TensorFlow 2.x已经被废弃,内部适配逻辑可能存在兼容性问题。
解决步骤
  1. 调整训练配置,关闭多worker预取:将workers设为1,新增max_queue_size=1,禁用预取机制,保证批次按顺序串行调用。
  2. 修正Sequence实现错误,删除冗余IO,补全缺失属性,统一用内存中已加载的self.lines做批次切分,避免重复读文件的偏移错误。
  3. 替换fit_generator为标准model.fit,适配TensorFlow 2.x的官方推荐用法。
修正后的示例代码
import csv
import numpy as np
import pandas as pd
from tensorflow import keras

class DataGenerator(keras.utils.Sequence):
    def __init__(self, file_name, rows_per_batch=50, shuffle=False, config_file=None):
        self.rows_per_batch = rows_per_batch
        self.shuffle = shuffle
        self.file_name = file_name
        self.config_file = config_file # 补全缺失的配置参数
        # 一次性读入所有行(如果文件过大不适合全量读入,可以替换为记录行偏移量的方案,避免全量加载内存)
        with open(self.file_name, 'r') as f:
            reader = csv.reader(f)
            self.lines = list(reader)
        # 初始化索引
        self.indexes = np.arange(len(self.lines))
        self.on_epoch_end()

    def __len__(self):
        # 保证所有样本都被用到可以改成 int(np.ceil(len(self.lines) / self.rows_per_batch))
        return len(self.lines) // self.rows_per_batch

    def __getitem__(self, index):
        # 取当前批次对应的索引范围
        batch_indexes = self.indexes[index*self.rows_per_batch : (index+1)*self.rows_per_batch]
        # 直接从内存的lines里取对应行,不需要重复读文件
        batch_lines = [self.lines[k] for k in batch_indexes]
        X, y = self.__data_generation(batch_lines)
        return X, y

    def on_epoch_end(self):
        if self.shuffle == True:
            np.random.shuffle(self.indexes)

    def __data_generation(self, batch_lines):
        # 直接处理传入的批次行,不需要再读CSV
        df = pd.DataFrame(batch_lines)
        # <这里保留你原有数据处理逻辑>
        X = df.values[:, :-1]
        y = df.values[:, -1]
        return X.astype(np.float32), y.astype(np.float32)

# 主程序
if __name__ == "__main__":
    file_name = "your_data.csv"
    rows_per_batch = 50
    config_file = "your_config.yaml" # 按你的实际需求传入

    train_gen = DataGenerator(file_name, rows_per_batch, shuffle=False, config_file=config_file)
    val_gen = DataGenerator(file_name, rows_per_batch, shuffle=False, config_file=config_file)

    # 替换fit_generator为fit,调整worker和队列配置保证顺序执行
    model.fit(
        x=train_gen,
        validation_data=val_gen,
        use_multiprocessing=False,
        shuffle=False,
        epochs=10,
        workers=1,
        max_queue_size=1 # 关闭预取,保证调用次数和顺序符合预期
    )
补充说明

如果你的CSV文件过大无法全量加载到内存,不要在__init__中全量读入所有行,可以替换为提前统计文件行总数、记录每行的字节偏移量的方案,__getitem__时直接根据偏移量读取对应行,既避免全量内存占用,也不会出现重复读文件的行偏移错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 01:15:03